Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -35,11 +35,14 @@ public sealed class GraphAssemblyLoadContext : AssemblyLoadContext
};

/// <summary>
/// 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&lt;T&gt; 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&lt;T&gt; 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).
/// </summary>
private static readonly HashSet<string> s_trustedPlatformAssemblies = GetTrustedPlatformAssemblies();
private static readonly Dictionary<string, string> s_trustedPlatformAssemblies = GetTrustedPlatformAssemblies();

private readonly string _dependencyFolder;
private readonly string _psEditionDependencyFolder;
Expand All @@ -66,16 +69,49 @@ public GraphAssemblyLoadContext(string dependencyFolder, string psEditionDepende
/// <inheritdoc/>
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<T> 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;
}

/// <summary>
/// Determines whether the assembly the module ships at <paramref name="modulePath"/> has a higher assembly version
/// than the runtime copy at <paramref name="runtimePath"/>. When versions cannot be read the runtime copy is
/// preferred (returns false) to keep the conservative shared-framework behaviour.
/// </summary>
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;
}
}


/// <inheritdoc/>
protected override IntPtr LoadUnmanagedDll(string unmanagedDllName)
{
Expand Down Expand Up @@ -115,15 +151,15 @@ private IEnumerable<string> GetNativeCandidates(string unmanagedDllName)
yield return Path.Combine(_dependencyFolder, fileName);
}

private static HashSet<string> GetTrustedPlatformAssemblies()
private static Dictionary<string, string> GetTrustedPlatformAssemblies()
{
var result = new HashSet<string>(StringComparer.OrdinalIgnoreCase);
var result = new Dictionary<string, string>(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;
Expand Down