WingetPackageTools.cs (11525B)
1 // ----------------------------------------------------------------------------- 2 // <copyright file="WingetPackageTools.cs" company="Microsoft Corporation"> 3 // Copyright (c) Microsoft Corporation. Licensed under the MIT License. 4 // </copyright> 5 // ----------------------------------------------------------------------------- 6 7 namespace WinGetMCPServer 8 { 9 using System.ComponentModel; 10 using Microsoft.Management.Deployment; 11 using ModelContextProtocol.Protocol; 12 using ModelContextProtocol.Server; 13 using ModelContextProtocol; 14 using Windows.Foundation; 15 using WinGetMCPServer.Extensions; 16 using WinGetMCPServer.Response; 17 using WinGetMCPServer.Exceptions; 18 19 /// <summary> 20 /// WinGet package tools. 21 /// </summary> 22 [McpServerToolType] 23 internal class WingetPackageTools 24 { 25 private PackageManager packageManager; 26 27 public WingetPackageTools() 28 { 29 packageManager = ServerConnection.Instance; 30 } 31 32 [McpServerTool( 33 Name = "find-winget-packages", 34 Title = "Find WinGet Packages", 35 ReadOnly = true, 36 OpenWorld = false)] 37 [Description("Find installed and available packages using WinGet")] 38 public CallToolResult FindPackages( 39 [Description("Find packages identified by this value")] string query) 40 { 41 try 42 { 43 ToolResponse.CheckGroupPolicy(); 44 45 var catalog = ConnectCatalog(); 46 47 // First attempt a more exact match 48 var findResult = FindForQuery(catalog, query, fullStringMatch: true); 49 50 // If nothing is found, expand to a looser search 51 if ((findResult.Matches?.Count ?? 0) == 0) 52 { 53 findResult = FindForQuery(catalog, query, fullStringMatch: false); 54 } 55 56 if (findResult.Status != FindPackagesResultStatus.Ok) 57 { 58 return PackageResponse.ForFindError(findResult); 59 } 60 61 List<FindPackageResult> contents = new List<FindPackageResult>(); 62 contents.AddPackages(findResult); 63 64 return ToolResponse.FromObject(contents); 65 } 66 catch (ToolResponseException e) 67 { 68 return e.Response; 69 } 70 } 71 72 [McpServerTool( 73 Name = "install-winget-package", 74 Title = "Install WinGet Package", 75 ReadOnly = false, 76 Destructive = true, 77 Idempotent = false, 78 OpenWorld = false)] 79 [Description("Install or update a package using WinGet")] 80 public async Task<CallToolResult> InstallPackage( 81 [Description("The identifier of the WinGet package")] string identifier, 82 IProgress<ProgressNotificationValue> progress, 83 CancellationToken cancellationToken, 84 [Description("The source containing the package")] string? source = null) 85 { 86 try 87 { 88 ToolResponse.CheckGroupPolicy(); 89 90 var packageCatalog = ConnectCatalog(source); 91 92 if (cancellationToken.IsCancellationRequested) 93 { 94 return PackageResponse.ForCancelBeforeSystemChange(); 95 } 96 97 // First attempt a more exact match 98 var findResult = FindForIdentifier(packageCatalog, identifier, expandedFields: false); 99 100 if (cancellationToken.IsCancellationRequested) 101 { 102 return PackageResponse.ForCancelBeforeSystemChange(); 103 } 104 105 // If nothing is found, expand to a looser search 106 if ((findResult.Matches?.Count ?? 0) == 0) 107 { 108 findResult = FindForIdentifier(packageCatalog, identifier, expandedFields: true); 109 } 110 111 if (findResult.Status != FindPackagesResultStatus.Ok) 112 { 113 return PackageResponse.ForFindError(findResult); 114 } 115 116 if (findResult.Matches?.Count == 0) 117 { 118 return PackageResponse.ForEmptyFind(identifier, source); 119 } 120 else if (findResult.Matches?.Count > 1) 121 { 122 return PackageResponse.ForMultiFind(identifier, source, findResult); 123 } 124 125 CatalogPackage catalogPackage = findResult.Matches![0].CatalogPackage; 126 InstallOptions options = new InstallOptions(); 127 IAsyncOperationWithProgress<InstallResult, InstallProgress>? operation = null; 128 129 if (cancellationToken.IsCancellationRequested) 130 { 131 return PackageResponse.ForCancelBeforeSystemChange(); 132 } 133 134 if (catalogPackage.InstalledVersion == null) 135 { 136 operation = packageManager.InstallPackageAsync(catalogPackage, options); 137 } 138 else 139 { 140 operation = packageManager.UpgradePackageAsync(catalogPackage, options); 141 } 142 143 operation.Progress = (asyncInfo, progressInfo) => progress.Report(CreateInstallProgressNotification(ref progressInfo)); 144 using CancellationTokenRegistration registration = cancellationToken.Register(() => operation.Cancel()); 145 146 var installResult = await operation; 147 findResult = null; 148 149 if (installResult.Status == InstallResultStatus.Ok) 150 { 151 findResult = ReFindForPackage(catalogPackage.DefaultInstallVersion); 152 } 153 154 return PackageResponse.ForInstallOperation(installResult, findResult); 155 } 156 catch (ToolResponseException e) 157 { 158 return e.Response; 159 } 160 } 161 162 private ConnectResult ConnectCatalogWithResult(string? catalog = null) 163 { 164 CreateCompositePackageCatalogOptions createCompositePackageCatalogOptions = new CreateCompositePackageCatalogOptions(); 165 166 var catalogs = packageManager.GetPackageCatalogs(); 167 for (int i = 0; i < catalogs.Count; ++i) 168 { 169 var catalogRef = catalogs[i]; 170 if (string.IsNullOrEmpty(catalog) || catalogRef?.Info.Id == catalog) 171 { 172 createCompositePackageCatalogOptions.Catalogs.Add(catalogs[i]); 173 } 174 } 175 createCompositePackageCatalogOptions.CompositeSearchBehavior = CompositeSearchBehavior.AllCatalogs; 176 177 var compositeRef = packageManager.CreateCompositePackageCatalog(createCompositePackageCatalogOptions); 178 return compositeRef.Connect(); 179 } 180 181 private PackageCatalog ConnectCatalog(string? catalog = null) 182 { 183 var result = ConnectCatalogWithResult(catalog); 184 if (result.Status != ConnectResultStatus.Ok) 185 { 186 throw new ToolResponseException(PackageResponse.ForConnectError(result)); 187 } 188 return result.PackageCatalog; 189 } 190 191 private FindPackagesResult FindForQuery(PackageCatalog catalog, string query, bool fullStringMatch) 192 { 193 PackageFieldMatchOption fullStringMatchOption = fullStringMatch ? PackageFieldMatchOption.EqualsCaseInsensitive : PackageFieldMatchOption.ContainsCaseInsensitive; 194 195 FindPackagesOptions findPackageOptions = new(); 196 findPackageOptions.Selectors.Add(new PackageMatchFilter() { Field = PackageMatchField.Id, Option = fullStringMatchOption, Value = query }); 197 findPackageOptions.Selectors.Add(new PackageMatchFilter() { Field = PackageMatchField.Name, Option = fullStringMatchOption, Value = query }); 198 findPackageOptions.Selectors.Add(new PackageMatchFilter() { Field = PackageMatchField.Moniker, Option = PackageFieldMatchOption.EqualsCaseInsensitive, Value = query }); 199 200 return catalog!.FindPackages(findPackageOptions); 201 } 202 203 private FindPackagesResult FindForIdentifier(PackageCatalog catalog, string query, bool expandedFields) 204 { 205 FindPackagesOptions findPackageOptions = new(); 206 findPackageOptions.Selectors.Add(new PackageMatchFilter() { Field = PackageMatchField.Id, Option = PackageFieldMatchOption.EqualsCaseInsensitive, Value = query }); 207 208 if (expandedFields) 209 { 210 findPackageOptions.Selectors.Add(new PackageMatchFilter() { Field = PackageMatchField.Name, Option = PackageFieldMatchOption.EqualsCaseInsensitive, Value = query }); 211 findPackageOptions.Selectors.Add(new PackageMatchFilter() { Field = PackageMatchField.Moniker, Option = PackageFieldMatchOption.EqualsCaseInsensitive, Value = query }); 212 } 213 214 return catalog!.FindPackages(findPackageOptions); 215 } 216 217 private FindPackagesResult? ReFindForPackage(PackageVersionInfo packageVersionInfo) 218 { 219 var connectResult = ConnectCatalogWithResult(packageVersionInfo.PackageCatalog.Info.Id); 220 221 if (connectResult.Status != ConnectResultStatus.Ok) 222 { 223 return null; 224 } 225 226 var catalog = connectResult.PackageCatalog; 227 228 FindPackagesOptions findPackageOptions = new(); 229 findPackageOptions.Selectors.Add(new PackageMatchFilter() { Field = PackageMatchField.Id, Option = PackageFieldMatchOption.Equals, Value = packageVersionInfo.Id }); 230 231 return catalog!.FindPackages(findPackageOptions); 232 } 233 234 private static ProgressNotificationValue CreateInstallProgressNotification(ref InstallProgress installProgress) 235 { 236 string? message = null; 237 238 switch (installProgress.State) 239 { 240 case PackageInstallProgressState.Queued: 241 message = "The install operation is queued"; 242 break; 243 case PackageInstallProgressState.Downloading: 244 message = "The package installer is being downloaded"; 245 break; 246 case PackageInstallProgressState.Installing: 247 message = "The package is being installed"; 248 break; 249 case PackageInstallProgressState.PostInstall: 250 message = "The installation operation is wrapping up"; 251 break; 252 case PackageInstallProgressState.Finished: 253 message = "The install is complete"; 254 break; 255 default: 256 message = "Unknown install state"; 257 break; 258 } 259 260 const float downloadPercentage = 0.8f; 261 262 ProgressNotificationValue result = new ProgressNotificationValue() 263 { 264 Progress = (float)((installProgress.DownloadProgress * downloadPercentage) + (installProgress.InstallationProgress * (1.0f - downloadPercentage))), 265 Message = message, 266 }; 267 268 return result; 269 } 270 } 271 }