diff --git a/src/Authentication/Authentication.Loader/GraphAssemblyLoadContext.cs b/src/Authentication/Authentication.Loader/GraphAssemblyLoadContext.cs index 047356f79e4..d91369525d7 100644 --- a/src/Authentication/Authentication.Loader/GraphAssemblyLoadContext.cs +++ b/src/Authentication/Authentication.Loader/GraphAssemblyLoadContext.cs @@ -35,11 +35,14 @@ public sealed class GraphAssemblyLoadContext : AssemblyLoadContext }; /// - /// Simple names of the assemblies that make up the shared framework (Trusted Platform Assemblies). These must always - /// come from the runtime, never from the module's Dependencies folder: loading e.g. the netstandard build of - /// System.Memory.dll into this context would create a second, incompatible ReadOnlyMemory<T> type. + /// Maps the simple name of each assembly that makes up the shared framework (Trusted Platform Assemblies) to the + /// full path of the runtime copy. These normally come from the runtime, never from the module's Dependencies + /// folder: loading e.g. the netstandard build of System.Memory.dll into this context would create a second, + /// incompatible ReadOnlyMemory<T> type. The path is retained so the runtime version can be compared against + /// the version the module ships, allowing the module's copy to win when it is strictly newer (e.g. the module + /// ships a far newer System.Text.Json than the one bundled with PowerShell 7.2/7.4). /// - private static readonly HashSet s_trustedPlatformAssemblies = GetTrustedPlatformAssemblies(); + private static readonly Dictionary s_trustedPlatformAssemblies = GetTrustedPlatformAssemblies(); private readonly string _dependencyFolder; private readonly string _psEditionDependencyFolder; @@ -66,16 +69,49 @@ public GraphAssemblyLoadContext(string dependencyFolder, string psEditionDepende /// protected override Assembly Load(AssemblyName assemblyName) { - if (s_sharedAssemblyNames.Contains(assemblyName.Name) || s_trustedPlatformAssemblies.Contains(assemblyName.Name)) + if (s_sharedAssemblyNames.Contains(assemblyName.Name)) { // Defer to the default context so type identity is preserved across the boundary. return null; } string path = ResolveManagedPath(assemblyName.Name); + + if (s_trustedPlatformAssemblies.TryGetValue(assemblyName.Name, out string runtimePath)) + { + // This is a shared-framework assembly. Normally the runtime copy must win so that types such as + // ReadOnlyMemory keep a single identity. However, the module may intentionally ship a newer version + // (for example System.Text.Json 10.x against PowerShell 7.2/7.4 which only bundle 6.x/8.x). In that case + // the runtime copy cannot satisfy the reference, so prefer the module's assembly when it is strictly newer. + if (path == null || !ModuleAssemblyIsNewer(path, runtimePath)) + { + return null; + } + } + return path != null ? LoadFromAssemblyPath(path) : null; } + /// + /// Determines whether the assembly the module ships at has a higher assembly version + /// than the runtime copy at . When versions cannot be read the runtime copy is + /// preferred (returns false) to keep the conservative shared-framework behaviour. + /// + private static bool ModuleAssemblyIsNewer(string modulePath, string runtimePath) + { + try + { + Version moduleVersion = AssemblyName.GetAssemblyName(modulePath).Version; + Version runtimeVersion = AssemblyName.GetAssemblyName(runtimePath).Version; + return moduleVersion != null && runtimeVersion != null && moduleVersion > runtimeVersion; + } + catch + { + return false; + } + } + + /// protected override IntPtr LoadUnmanagedDll(string unmanagedDllName) { @@ -115,15 +151,15 @@ private IEnumerable GetNativeCandidates(string unmanagedDllName) yield return Path.Combine(_dependencyFolder, fileName); } - private static HashSet GetTrustedPlatformAssemblies() + private static Dictionary GetTrustedPlatformAssemblies() { - var result = new HashSet(StringComparer.OrdinalIgnoreCase); + var result = new Dictionary(StringComparer.OrdinalIgnoreCase); if (AppContext.GetData("TRUSTED_PLATFORM_ASSEMBLIES") is string tpa) { foreach (string path in tpa.Split(Path.PathSeparator)) { if (!string.IsNullOrEmpty(path)) - result.Add(Path.GetFileNameWithoutExtension(path)); + result[Path.GetFileNameWithoutExtension(path)] = path; } } return result;