diff --git a/S1API.Tests/Internal/Utils/ReflectionUtilsTests.cs b/S1API.Tests/Internal/Utils/ReflectionUtilsTests.cs index 72a49c5c..3dfc4a81 100644 --- a/S1API.Tests/Internal/Utils/ReflectionUtilsTests.cs +++ b/S1API.Tests/Internal/Utils/ReflectionUtilsTests.cs @@ -91,6 +91,46 @@ public void DerivedTypeScanFollowsTransitiveAssemblyReferences() loadedAssemblies)); } + [Fact] + public void DerivedTypeScanDoesNotFollowSameNameAssembliesWithDifferentIdentities() + { + string assemblyName = $"S1API.ReflectionUtilsTests.Duplicate.{Guid.NewGuid():N}"; + AssemblyBuilder unrelatedAssembly = CreateDynamicAssembly(assemblyName, new Version(1, 0, 0, 0)); + Type unrelatedType = unrelatedAssembly + .DefineDynamicModule(assemblyName) + .DefineType("UnrelatedType", TypeAttributes.Public) + .CreateType()!; + + AssemblyBuilder relatedAssembly = CreateDynamicAssembly(assemblyName, new Version(2, 0, 0, 0)); + relatedAssembly + .DefineDynamicModule(assemblyName) + .DefineType("RelatedType", TypeAttributes.Public, typeof(ReflectionCandidateBridge)) + .CreateType(); + + AssemblyBuilder candidateAssembly = CreateDynamicAssembly( + $"S1API.ReflectionUtilsTests.Candidate.{Guid.NewGuid():N}", + new Version(1, 0, 0, 0)); + ModuleBuilder candidateModule = candidateAssembly.DefineDynamicModule(candidateAssembly.GetName().Name!); + candidateModule + .DefineType("CandidateType", TypeAttributes.Public, unrelatedType) + .CreateType(); + + Assert.False(ReflectionUtils.CanContainTypesDerivedFrom( + candidateAssembly, + typeof(ReflectionUtils).Assembly, + AppDomain.CurrentDomain.GetAssemblies())); + } + + private static AssemblyBuilder CreateDynamicAssembly(string name, Version version) + { + var assemblyName = new AssemblyName(name) + { + Version = version + }; + + return AssemblyBuilder.DefineDynamicAssembly(assemblyName, AssemblyBuilderAccess.Run); + } + private sealed class MonoShape { #pragma warning disable CS0169 diff --git a/S1API/Internal/Utils/ReflectionUtils.cs b/S1API/Internal/Utils/ReflectionUtils.cs index 36ddbfff..853db58b 100644 --- a/S1API/Internal/Utils/ReflectionUtils.cs +++ b/S1API/Internal/Utils/ReflectionUtils.cs @@ -106,7 +106,7 @@ private static bool CanContainTypesDerivedFrom( AssemblyName baseAssemblyName, IReadOnlyDictionary assembliesBySimpleName) { - if (AssemblyName.ReferenceMatchesDefinition(candidateAssembly.GetName(), baseAssemblyName)) + if (AssemblyIdentityMatches(candidateAssembly.GetName(), baseAssemblyName)) return true; return ReferencesAssemblyTransitively( @@ -137,7 +137,7 @@ private static bool ReferencesAssemblyTransitively( foreach (AssemblyName referencedAssembly in referencedAssemblies) { - if (AssemblyName.ReferenceMatchesDefinition(referencedAssembly, baseAssemblyName)) + if (AssemblyIdentityMatches(referencedAssembly, baseAssemblyName)) return true; string referencedName = referencedAssembly.Name ?? string.Empty; @@ -147,6 +147,9 @@ private static bool ReferencesAssemblyTransitively( foreach (Assembly loadedReference in loadedReferences) { + if (!AssemblyIdentityMatches(loadedReference.GetName(), referencedAssembly)) + continue; + if (ReferencesAssemblyTransitively( loadedReference, baseAssemblyName, @@ -161,6 +164,14 @@ private static bool ReferencesAssemblyTransitively( return false; } + private static bool AssemblyIdentityMatches( + AssemblyName referenceAssemblyName, + AssemblyName definitionAssemblyName) => + string.Equals( + referenceAssemblyName.FullName, + definitionAssemblyName.FullName, + StringComparison.OrdinalIgnoreCase); + /// /// INTERNAL: Gets all types by their name. ///