Moves credential provider registration to new dedicated tool

This commit is contained in:
Ryan Newington
2023-01-31 08:01:58 +11:00
parent 1205f217b3
commit a5a97bdf68
19 changed files with 77 additions and 445 deletions
@@ -1,233 +0,0 @@
using System;
using System.IO;
using System.Linq;
using System.Reflection;
using Microsoft.Win32;
namespace Lithnet.CredentialProvider
{
public static class CredentialProviderRegistrationServices
{
public static void DisableCredentialProvider<T>() where T : CredentialProviderBase => DisableCredentialProvider(typeof(T));
public static void UnregisterCredentialProvider<T>() where T : CredentialProviderBase => UnregisterCredentialProvider(typeof(T));
public static void RegisterCredentialProvider<T>() where T : CredentialProviderBase => RegisterCredentialProvider(typeof(T));
public static void EnableCredentialProvider<T>() where T : CredentialProviderBase => EnableCredentialProvider(typeof(T));
public static void UnregisterCredentialProvider(Type type)
{
DeleteCredentialProviderRegistration(type);
if (IsFrameworkType(type))
{
UnregisterFrameworkAssembly(type);
}
else
{
UnregisterNetCoreAssembly(type);
}
}
public static void RegisterCredentialProvider(Type type)
{
CreateCredentialProviderRegistration(type);
if (IsFrameworkType(type))
{
RegisterFrameworkAssembly(type);
}
else
{
RegisterNetCoreAssembly(type);
}
}
public static void DisableCredentialProvider(Type type)
{
var comGuid = GetComGuid(type);
DisableCredentialProvider(comGuid);
}
public static void DisableCredentialProvider(Guid comGuid)
{
var key = Registry.LocalMachine.OpenSubKey($@"SOFTWARE\Microsoft\Windows\CurrentVersion\Authentication\Credential Providers\{comGuid:B}", true);
key?.SetValue("Disabled", 1);
}
public static void EnableCredentialProvider(Guid comGuid)
{
var key = Registry.LocalMachine.OpenSubKey($@"SOFTWARE\Microsoft\Windows\CurrentVersion\Authentication\Credential Providers\{comGuid:B}", true);
key?.SetValue("Disabled", 0);
}
public static void EnableCredentialProvider(Type type)
{
var comGuid = GetComGuid(type);
EnableCredentialProvider(comGuid);
}
private static void CreateCredentialProviderRegistration(Type t)
{
var comGuid = GetComGuid(t);
var typeName = GetTypeFullName(t);
var key = Registry.LocalMachine.CreateSubKey($@"SOFTWARE\Microsoft\Windows\CurrentVersion\Authentication\Credential Providers\{comGuid:B}", true);
key.SetValue(null, typeName);
}
private static void DeleteCredentialProviderRegistration(Type t)
{
var comGuid = GetComGuid(t);
Registry.LocalMachine.DeleteSubKeyTree($@"SOFTWARE\Microsoft\Windows\CurrentVersion\Authentication\Credential Providers\{comGuid:B}", false);
}
private static void RegisterNetCoreAssembly(Type t)
{
var comGuid = GetComGuid(t);
var typeName = GetTypeFullName(t);
var progId = GetComProgId(t);
var assemblyLocation = GetTypeAssemblyLocation(t);
var dir = Path.GetDirectoryName(assemblyLocation);
var assemblyFile = Path.GetFileNameWithoutExtension(assemblyLocation);
var comHostLocation = Path.Combine(dir, assemblyFile + ".comhost.dll");
var rootClsid = Registry.LocalMachine.CreateSubKey($@"Software\Classes\CLSID\{comGuid:B}", true);
rootClsid.SetValue(null, "CoreCLR COMHost Server");
var inprocKey = rootClsid.CreateSubKey("InprocServer32", true);
inprocKey.SetValue(null, comHostLocation);
inprocKey.SetValue("ThreadingModel", "Both");
var progIdKey = rootClsid.CreateSubKey("ProgId", true);
progIdKey.SetValue(null, progId);
var progIdRoot = Registry.LocalMachine.CreateSubKey($@"Software\Classes\{progId}", true);
progIdRoot.SetValue(null, typeName);
var progIdSubKey = progIdRoot.CreateSubKey("CLSID");
progIdSubKey.SetValue(null, comGuid.ToString("B"));
}
private static void UnregisterNetCoreAssembly(Type t)
{
var comGuid = GetComGuid(t);
var progId = GetComProgId(t);
Registry.LocalMachine.DeleteSubKeyTree($@"Software\Classes\CLSID\{comGuid:B}", false);
Registry.LocalMachine.DeleteSubKeyTree($@"Software\Classes\{progId}", false);
}
private static void RegisterFrameworkAssembly(Type t)
{
var comGuid = GetComGuid(t);
var typeName = GetTypeFullName(t);
var progId = GetComProgId(t);
var rootClsid = Registry.LocalMachine.CreateSubKey($@"Software\Classes\CLSID\{comGuid:B}", true);
rootClsid.SetValue(null, typeName);
rootClsid.CreateSubKey("Implemented Categories");
rootClsid.CreateSubKey(@"Implemented Categories\{62C8FE65-4EBB-45e7-B440-6E39B2CDBF29}");
var inprocKey = rootClsid.CreateSubKey("InprocServer32", true);
inprocKey.SetValue(null, "mscoree.dll");
inprocKey.SetValue("ThreadingModel", "Both");
inprocKey.SetValue("Class", typeName);
inprocKey.SetValue("RuntimeVersion", "v4.0.30319");
inprocKey.SetValue("Assembly", GetTypeAssemblyName(t));
inprocKey.SetValue("CodeBase", GetTypeAssemblyLocation(t));
var progIdKey = rootClsid.CreateSubKey("ProgId", true);
progIdKey.SetValue(null, progId);
var progIdRoot = Registry.LocalMachine.CreateSubKey($@"Software\Classes\{progId}", true);
progIdRoot.SetValue(null, typeName);
var progIdSubKey = progIdRoot.CreateSubKey("CLSID");
progIdSubKey.SetValue(null, comGuid.ToString("B"));
}
private static void UnregisterFrameworkAssembly(Type t)
{
var comGuid = GetComGuid(t);
var progId = GetComProgId(t);
Registry.LocalMachine.DeleteSubKeyTree($@"Software\Classes\CLSID\{comGuid:B}", false);
Registry.LocalMachine.DeleteSubKeyTree($@"Software\Classes\{progId}", false);
}
private static string GetTypeAssemblyLocation(Type type)
{
return type.Assembly.Location;
}
private static string GetTypeAssemblyName(Type type)
{
return type.Assembly.FullName;
}
private static string GetTypeClassName(Type type)
{
return type.Name;
}
private static string GetTypeFullName(Type type)
{
return type.FullName;
}
private static Guid GetComGuid(Type type)
{
var typeId = type.GetCustomAttributeValue("GuidAttribute");
if (typeId == null)
{
throw new ArgumentException($"The type {type.Name} does not have the Guid attribute present");
}
return new Guid(typeId);
}
private static string GetComProgId(Type type)
{
var typeId = type.GetCustomAttributeValue("ProgIdAttribute");
if (typeId == null)
{
throw new ArgumentException($"The type {type.Name} does not have the ProgId attribute present");
}
return typeId;
}
private static bool IsFrameworkType(Type type)
{
var framework = type.Assembly.GetCustomAttributeValue("TargetFrameworkAttribute");
return framework.StartsWith(".NETFramework");
}
private static string GetCustomAttributeValue(this Type type, string attributeName)
{
var cads = type.GetCustomAttributesData();
foreach (CustomAttributeData cad in cads.Where(a => a.AttributeType.Name == attributeName))
{
return cad.ConstructorArguments.FirstOrDefault().Value as string;
}
return String.Empty;
}
private static string GetCustomAttributeValue(this Assembly assembly, string attributeName)
{
foreach (CustomAttributeData cad in assembly.GetCustomAttributesData().Where(a => a.AttributeType.Name == attributeName))
{
return cad.ConstructorArguments.FirstOrDefault().Value as string;
}
return String.Empty;
}
}
}
@@ -53,7 +53,9 @@ namespace Lithnet.CredentialProvider.Interop
public int scode;
[FieldOffset(0)]
#pragma warning disable CS0618 // Type or member is obsolete
[MarshalAs(UnmanagedType.Currency)]
#pragma warning restore CS0618 // Type or member is obsolete
public decimal cyVal;
[FieldOffset(0)]
@@ -1,6 +1,6 @@
<Project Sdk="Microsoft.NET.Sdk">
<PropertyGroup>
<TargetFramework>netstandard2.0</TargetFramework>
<TargetFrameworks>netstandard2.1;net472</TargetFrameworks>
<RegisterForComInterop>false</RegisterForComInterop>
<OutputType>Library</OutputType>
<AllowUnsafeBlocks>true</AllowUnsafeBlocks>
@@ -9,6 +9,7 @@
<AutoGenerateBindingRedirects>true</AutoGenerateBindingRedirects>
<GenerateBindingRedirectsOutputType>true</GenerateBindingRedirectsOutputType>
<LangVersion>9</LangVersion>
<GeneratePackageOnBuild>true</GeneratePackageOnBuild>
</PropertyGroup>
<PropertyGroup>
@@ -17,7 +18,7 @@
<Copyright>Copyright 2023 Lithnet Pty Ltd</Copyright>
<ProductName>Lithnet Windows Credential Provider</ProductName>
<VersionPrefix>1.0.0</VersionPrefix>
<VersionSuffix>beta.26</VersionSuffix>
<VersionSuffix>beta.38</VersionSuffix>
<Authors>Lithnet</Authors>
<AutoGenerateBindingRedirects>true</AutoGenerateBindingRedirects>
<AutoIncrementPackageRevision>true</AutoIncrementPackageRevision>
@@ -30,8 +31,7 @@
</PropertyGroup>
<ItemGroup>
<PackageReference Include="Microsoft.Extensions.Logging.Abstractions" Version="6.0.0" />
<PackageReference Include="Microsoft.Win32.Registry" Version="5.0.0" />
<PackageReference Include="System.Drawing.Common" Version="6.0.0" Condition="$(TargetFramework) == 'netstandard2.0'" />
<PackageReference Include="Microsoft.Extensions.Logging.Abstractions" Version="3.1.32" />
<PackageReference Include="System.Drawing.Common" Version="6.0.0" Condition="$(TargetFramework) == 'netstandard2.1'" />
</ItemGroup>
</Project>
@@ -3,6 +3,7 @@ using System.Diagnostics;
using System.IO;
using System.Reflection;
using System.Runtime.CompilerServices;
using System.Threading;
namespace Lithnet.CredentialProvider.ModuleInit
{
@@ -13,17 +14,47 @@ namespace Lithnet.CredentialProvider.ModuleInit
[ModuleInitializer]
public static void AttachResolver()
{
Trace.WriteLine($"Loaded assembly {Assembly.GetExecutingAssembly().Location}");
Trace.WriteLine("Attaching assembly resolver");
AppDomain.CurrentDomain.AssemblyResolve += CurrentDomain_AssemblyResolve;
AppDomain.CurrentDomain.UnhandledException += CurrentDomain_UnhandledException;
AppDomain.CurrentDomain.FirstChanceException += CurrentDomain_FirstChanceException;
basePath = Path.GetFullPath(Path.GetDirectoryName(Assembly.GetExecutingAssembly().Location));
}
private static void CurrentDomain_FirstChanceException(object sender, System.Runtime.ExceptionServices.FirstChanceExceptionEventArgs e)
{
Trace.WriteLine("First chance exception in credential provider");
Trace.WriteLine((e.Exception)?.ToString() ?? "No exception details provided");
}
private static void CurrentDomain_UnhandledException(object sender, UnhandledExceptionEventArgs e)
{
Trace.WriteLine("Unhandled exception in credential provider");
Trace.WriteLine((e.ExceptionObject as Exception)?.ToString() ?? "No exception details provided");
}
private static Assembly CurrentDomain_AssemblyResolve(object sender, ResolveEventArgs args)
{
var name = new AssemblyName(args.Name);
string assyPath = Path.Combine(basePath, $"{name.Name}.dll");
Trace.WriteLine($"Request for {args.Name}");
#if NETFRAMEWORK
if (name.Name.StartsWith("System.", StringComparison.OrdinalIgnoreCase))
{
string gacPath = $@"C:\Windows\Microsoft.NET\assembly\GAC_MSIL\{name.Name}\v4.0_4.0.0.0__b03f5f7f11d50a3a\{name.Name}.dll";
if (File.Exists(gacPath))
{
Trace.WriteLine($"System assembly found at {gacPath}");
return Assembly.LoadFrom(gacPath);
}
}
#endif
string assyPath = Path.Combine(basePath, $"{name.Name}.dll");
if (File.Exists(assyPath))
{
Trace.WriteLine($"Found at {assyPath}");