diff --git a/src/Modules/Unix/Microsoft.PowerShell.Security/Microsoft.PowerShell.Security.psd1 b/src/Modules/Unix/Microsoft.PowerShell.Security/Microsoft.PowerShell.Security.psd1 index d5961d1008d..c287a6cef3c 100644 --- a/src/Modules/Unix/Microsoft.PowerShell.Security/Microsoft.PowerShell.Security.psd1 +++ b/src/Modules/Unix/Microsoft.PowerShell.Security/Microsoft.PowerShell.Security.psd1 @@ -7,7 +7,7 @@ ModuleVersion="7.0.0.0" CompatiblePSEditions = @("Core") PowerShellVersion="3.0" FunctionsToExport = @() -CmdletsToExport="Get-Credential", "Get-ExecutionPolicy", "Set-ExecutionPolicy", "ConvertFrom-SecureString", "ConvertTo-SecureString", "Get-PfxCertificate" +CmdletsToExport="Get-Credential", "Get-ExecutionPolicy", "Set-ExecutionPolicy", "ConvertFrom-SecureString", "ConvertTo-SecureString", "Get-PfxCertificate" , "Protect-CmsMessage", "Unprotect-CmsMessage", "Get-CmsMessage" AliasesToExport = @() NestedModules="Microsoft.PowerShell.Security.dll" HelpInfoURI = 'https://go.microsoft.com/fwlink/?linkid=2113533' diff --git a/src/System.Management.Automation/security/SecuritySupport.cs b/src/System.Management.Automation/security/SecuritySupport.cs index ee1058be350..b7b08491e5f 100644 --- a/src/System.Management.Automation/security/SecuritySupport.cs +++ b/src/System.Management.Automation/security/SecuritySupport.cs @@ -142,15 +142,20 @@ internal static void SetExecutionPolicy(ExecutionPolicyScope scope, ExecutionPol switch (policy) { case ExecutionPolicy.Restricted: - executionPolicy = "Restricted"; break; + executionPolicy = "Restricted"; + break; case ExecutionPolicy.AllSigned: - executionPolicy = "AllSigned"; break; + executionPolicy = "AllSigned"; + break; case ExecutionPolicy.RemoteSigned: - executionPolicy = "RemoteSigned"; break; + executionPolicy = "RemoteSigned"; + break; case ExecutionPolicy.Unrestricted: - executionPolicy = "Unrestricted"; break; + executionPolicy = "Unrestricted"; + break; case ExecutionPolicy.Bypass: - executionPolicy = "Bypass"; break; + executionPolicy = "Bypass"; + break; } // Set the execution policy @@ -359,12 +364,18 @@ internal static string GetExecutionPolicy(ExecutionPolicy policy) { switch (policy) { - case ExecutionPolicy.Bypass: return "Bypass"; - case ExecutionPolicy.Unrestricted: return "Unrestricted"; - case ExecutionPolicy.RemoteSigned: return "RemoteSigned"; - case ExecutionPolicy.AllSigned: return "AllSigned"; - case ExecutionPolicy.Restricted: return "Restricted"; - default: return "Restricted"; + case ExecutionPolicy.Bypass: + return "Bypass"; + case ExecutionPolicy.Unrestricted: + return "Unrestricted"; + case ExecutionPolicy.RemoteSigned: + return "RemoteSigned"; + case ExecutionPolicy.AllSigned: + return "AllSigned"; + case ExecutionPolicy.Restricted: + return "Restricted"; + default: + return "Restricted"; } } @@ -595,7 +606,7 @@ internal static void CheckIfFileExists(string filePath) /// True on success, false otherwise. internal static bool CertIsGoodForSigning(X509Certificate2 c) { - if (!CertHasPrivatekey(c)) + if (!c.HasPrivateKey) { return false; } @@ -620,16 +631,20 @@ internal static bool CertIsGoodForEncryption(X509Certificate2 c) private static bool CertHasOid(X509Certificate2 c, string oid) { - Collection ekus = GetCertEKU(c); - - foreach (string testOid in ekus) + foreach (var extension in c.Extensions) { - if (testOid == oid) + if (extension is X509EnhancedKeyUsageExtension ext) { - return true; + foreach (Oid ekuOid in ext.EnhancedKeyUsages) + { + if (ekuOid.Value == oid) + { + return true; + } + } + break; } } - return false; } @@ -644,82 +659,12 @@ private static bool CertHasKeyUsage(X509Certificate2 c, X509KeyUsageFlags keyUsa { return true; } - break; } } - return false; } - /// - /// Check if the specified cert has a private key in it. - /// - /// Certificate object. - /// True on success, false otherwise. - internal static bool CertHasPrivatekey(X509Certificate2 cert) - { - return cert.HasPrivateKey; - } - - /// - /// Get the EKUs of a cert. - /// - /// Certificate object. - /// A collection of cert eku strings. - [ArchitectureSensitive] - internal static Collection GetCertEKU(X509Certificate2 cert) - { - Collection ekus = new Collection(); - IntPtr pCert = cert.Handle; - int structSize = 0; - IntPtr dummy = IntPtr.Zero; - - if (Security.NativeMethods.CertGetEnhancedKeyUsage(pCert, 0, dummy, - out structSize)) - { - if (structSize > 0) - { - IntPtr ekuBuffer = Marshal.AllocHGlobal(structSize); - - try - { - if (Security.NativeMethods.CertGetEnhancedKeyUsage(pCert, 0, - ekuBuffer, - out structSize)) - { - Security.NativeMethods.CERT_ENHKEY_USAGE ekuStruct = - (Security.NativeMethods.CERT_ENHKEY_USAGE) - Marshal.PtrToStructure(ekuBuffer); - IntPtr ep = ekuStruct.rgpszUsageIdentifier; - IntPtr ekuptr; - - for (int i = 0; i < ekuStruct.cUsageIdentifier; i++) - { - ekuptr = Marshal.ReadIntPtr(ep, i * Marshal.SizeOf(ep)); - string eku = Marshal.PtrToStringAnsi(ekuptr); - ekus.Add(eku); - } - } - else - { - throw new System.ComponentModel.Win32Exception(Marshal.GetLastWin32Error()); - } - } - finally - { - Marshal.FreeHGlobal(ekuBuffer); - } - } - } - else - { - throw new System.ComponentModel.Win32Exception(Marshal.GetLastWin32Error()); - } - - return ekus; - } - /// /// Convert an int to a DWORD. /// @@ -1138,8 +1083,10 @@ public void Resolve(SessionState sessionState, ResolutionPurpose purpose, out Er // Process the certificate if that was supplied exactly if (_pendingCertificate != null) { - ProcessResolvedCertificates(purpose, - new List { _pendingCertificate }, out error); + ProcessResolvedCertificates( + purpose, + new X509Certificate2Collection(_pendingCertificate), + out error); if ((error != null) || (Certificates.Count != 0)) { return; @@ -1162,15 +1109,8 @@ public void Resolve(SessionState sessionState, ResolutionPurpose purpose, out Er return; } - // Then by thumbprint - ResolveFromThumbprint(sessionState, purpose, out error); - if ((error != null) || (Certificates.Count != 0)) - { - return; - } - - // Then by Subject Name - ResolveFromSubjectName(sessionState, purpose, out error); + // Then by cert store + ResolveFromStoreById(purpose, out error); if ((error != null) || (Certificates.Count != 0)) { return; @@ -1215,7 +1155,7 @@ private void ResolveFromBase64Encoding(ResolutionPurpose purpose, out ErrorRecor return; } - List certificatesToProcess = new List(); + var certificatesToProcess = new X509Certificate2Collection(); try { X509Certificate2 newCertificate = new X509Certificate2(messageBytes); @@ -1289,7 +1229,7 @@ private void ResolveFromPath(SessionState sessionState, ResolutionPurpose purpos resolvedPaths.Remove(path); } - List certificatesToProcess = new List(); + var certificatesToProcess = new X509Certificate2Collection(); foreach (string path in resolvedPaths) { X509Certificate2 certificate = null; @@ -1311,99 +1251,51 @@ private void ResolveFromPath(SessionState sessionState, ResolutionPurpose purpos } } - private void ResolveFromThumbprint(SessionState sessionState, ResolutionPurpose purpose, out ErrorRecord error) + private void ResolveFromStoreById(ResolutionPurpose purpose, out ErrorRecord error) { - // Quickly check that this is a thumbprint-like pattern (just hex) - if (!System.Text.RegularExpressions.Regex.IsMatch(_identifier, "^[a-f0-9]+$", Text.RegularExpressions.RegexOptions.IgnoreCase)) - { - error = null; - return; - } - - Collection certificates = new Collection(); + error = null; + WildcardPattern subjectNamePattern = WildcardPattern.Get(_identifier, WildcardOptions.IgnoreCase); try { - // Get first from 'My' store - string certificatePath = sessionState.Path.Combine("Microsoft.PowerShell.Security\\Certificate::CurrentUser\\My", _identifier); - if (sessionState.InvokeProvider.Item.Exists(certificatePath)) - { - foreach (PSObject certificateObject in sessionState.InvokeProvider.Item.Get(certificatePath)) - { - certificates.Add(certificateObject); - } - } + var certificatesToProcess = new X509Certificate2Collection(); - // Second from 'LocalMachine' store - certificatePath = sessionState.Path.Combine("Microsoft.PowerShell.Security\\Certificate::LocalMachine\\My", _identifier); - if (sessionState.InvokeProvider.Item.Exists(certificatePath)) + using (var storeCU = new X509Store("my", StoreLocation.CurrentUser)) { - foreach (PSObject certificateObject in sessionState.InvokeProvider.Item.Get(certificatePath)) + storeCU.Open(OpenFlags.ReadOnly); + X509Certificate2Collection storeCerts = storeCU.Certificates; + + if (RuntimeInformation.IsOSPlatform(OSPlatform.Windows)) { - certificates.Add(certificateObject); + using (var storeLM = new X509Store("my", StoreLocation.LocalMachine)) + { + storeLM.Open(OpenFlags.ReadOnly); + storeCerts.AddRange(storeLM.Certificates); + } } - } - } - catch (SessionStateException) - { - // If we got an ItemNotFound / etc., then this didn't represent a valid path. - } - List certificatesToProcess = new List(); - foreach (PSObject certificateObject in certificates) - { - X509Certificate2 certificate = certificateObject.BaseObject as X509Certificate2; - if (certificate != null) - { - certificatesToProcess.Add(certificate); - } - } - - ProcessResolvedCertificates(purpose, certificatesToProcess, out error); - } - - private void ResolveFromSubjectName(SessionState sessionState, ResolutionPurpose purpose, out ErrorRecord error) - { - Collection certificates = new Collection(); - WildcardPattern subjectNamePattern = WildcardPattern.Get(_identifier, WildcardOptions.IgnoreCase); - - try - { - // Get first from 'My' store, then 'LocalMachine' - string[] certificatePaths = new string[] { - "Microsoft.PowerShell.Security\\Certificate::CurrentUser\\My", - "Microsoft.PowerShell.Security\\Certificate::LocalMachine\\My" }; + certificatesToProcess.AddRange(storeCerts.Find(X509FindType.FindByThumbprint, _identifier, validOnly: false)); - foreach (string certificatePath in certificatePaths) - { - foreach (PSObject certificateObject in sessionState.InvokeProvider.ChildItem.Get(certificatePath, false)) + if (certificatesToProcess.Count == 0) { - if (subjectNamePattern.IsMatch(certificateObject.Properties["Subject"].Value.ToString())) + foreach (var cert in storeCerts) { - certificates.Add(certificateObject); + if (subjectNamePattern.IsMatch(cert.Subject) || subjectNamePattern.IsMatch(cert.GetNameInfo(X509NameType.SimpleName, forIssuer: false))) + { + certificatesToProcess.Add(cert); + } } } + + ProcessResolvedCertificates(purpose, certificatesToProcess, out error); } } catch (SessionStateException) { - // If we got an ItemNotFound / etc., then this didn't represent a valid path. } - - List certificatesToProcess = new List(); - foreach (PSObject certificateObject in certificates) - { - X509Certificate2 certificate = certificateObject.BaseObject as X509Certificate2; - if (certificate != null) - { - certificatesToProcess.Add(certificate); - } - } - - ProcessResolvedCertificates(purpose, certificatesToProcess, out error); } - private void ProcessResolvedCertificates(ResolutionPurpose purpose, List certificatesToProcess, out ErrorRecord error) + private void ProcessResolvedCertificates(ResolutionPurpose purpose, X509Certificate2Collection certificatesToProcess, out ErrorRecord error) { error = null; HashSet processedThumbprints = new HashSet(); @@ -1418,9 +1310,14 @@ private void ProcessResolvedCertificates(ResolutionPurpose purpose, List