Skip to content
Merged
Show file tree
Hide file tree
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
107 changes: 107 additions & 0 deletions S1API.Tests/Internal/Utils/ReflectionUtilsTests.cs
Original file line number Diff line number Diff line change
@@ -1,4 +1,7 @@
using S1API.Internal.Utils;
using S1API.Logging;
using System.Reflection;
using System.Reflection.Emit;

namespace S1API.Tests.Internal.Utils;

Expand Down Expand Up @@ -36,6 +39,98 @@ public void StaticAccessWalksBaseTypesForNonPublicMembers()
ReflectionUtils.TryGetStaticFieldOrProperty(typeof(DerivedStaticShape), "RuntimeMember"));
}

[Fact]
public void DerivedTypeScanIncludesAssembliesThatReferenceTheBaseAssembly()
{
Assembly[] loadedAssemblies = AppDomain.CurrentDomain.GetAssemblies();

Assert.True(ReflectionUtils.CanContainTypesDerivedFrom(
typeof(ReflectionUtilsTests).Assembly,
typeof(ReflectionUtils).Assembly,
loadedAssemblies));
}

[Fact]
public void GetDerivedClassesFindsTypesInReferencingAssemblies()
{
Assert.Contains(
typeof(DerivedLogShape),
ReflectionUtils.GetDerivedClasses<Log>());
}

[Fact]
public void DerivedTypeScanExcludesAssembliesWithoutAReferencePathToTheBaseAssembly()
{
Assembly[] loadedAssemblies = AppDomain.CurrentDomain.GetAssemblies();

Assert.False(ReflectionUtils.CanContainTypesDerivedFrom(
typeof(string).Assembly,
typeof(ReflectionUtils).Assembly,
loadedAssemblies));
}

[Fact]
public void DerivedTypeScanFollowsTransitiveAssemblyReferences()
{
var assemblyName = new AssemblyName($"S1API.ReflectionUtilsTests.Dynamic.{Guid.NewGuid():N}");
AssemblyBuilder assemblyBuilder = AssemblyBuilder.DefineDynamicAssembly(
assemblyName,
AssemblyBuilderAccess.Run);
ModuleBuilder moduleBuilder = assemblyBuilder.DefineDynamicModule(assemblyName.Name!);
moduleBuilder.DefineType(
"DynamicReflectionCandidate",
TypeAttributes.Public,
typeof(ReflectionCandidateBridge))
.CreateType();

Assembly[] loadedAssemblies = AppDomain.CurrentDomain.GetAssemblies();

Assert.True(ReflectionUtils.CanContainTypesDerivedFrom(
assemblyBuilder,
typeof(ReflectionUtils).Assembly,
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
Expand Down Expand Up @@ -70,4 +165,16 @@ private class BaseStaticShape
private sealed class DerivedStaticShape : BaseStaticShape
{
}

public class ReflectionCandidateBridge
{
}

private sealed class DerivedLogShape : Log
{
public DerivedLogShape()
: base(nameof(DerivedLogShape))
{
}
}
}
54 changes: 8 additions & 46 deletions S1API/Entities/NPC.cs
Original file line number Diff line number Diff line change
Expand Up @@ -1265,29 +1265,11 @@ internal static bool TryGetConfiguredNpcId(System.Type npcType, out string id)
return null;

// Find the NPC type in loaded assemblies
System.Type? npcType = null;
var baseType = typeof(NPC);
var asms = AppDomain.CurrentDomain.GetAssemblies();
for (int ai = 0; ai < asms.Length && npcType == null; ai++)
{
var asm = asms[ai];
if (asm == baseType.Assembly)
continue; // Skip S1API assembly (internal wrappers)

System.Type[] types;
try { types = asm.GetTypes(); } catch { continue; }
for (int ti = 0; ti < types.Length; ti++)
{
var t = types[ti];
if (t == null || t.IsAbstract || !baseType.IsAssignableFrom(t))
continue;
if (t.Name == typeName)
{
npcType = t;
break;
}
}
}
var baseAssembly = typeof(NPC).Assembly;
System.Type? npcType = ReflectionUtils.GetDerivedClasses<NPC>()
.FirstOrDefault(type =>
type.Assembly != baseAssembly &&
type.Name == typeName);

if (npcType == null)
return null;
Expand Down Expand Up @@ -1545,29 +1527,9 @@ internal static void PreRegisterAllNpcPrefabsInternal()
if (spawnables == null)
return;

var baseType = typeof(NPC);
var baseAssembly = baseType.Assembly;
var candidateTypes = new System.Collections.Generic.List<System.Type>();
var asms = AppDomain.CurrentDomain.GetAssemblies();
for (int ai = 0; ai < asms.Length; ai++)
{
var asm = asms[ai];
Type[] types;
try { types = asm.GetTypes(); } catch { continue; }
for (int ti = 0; ti < types.Length; ti++)
{
var t = types[ti];
if (t == null || t.IsAbstract)
continue;
if (baseType.IsAssignableFrom(t))
{
// Skip internal S1API NPC wrappers; only pre-register mod-defined types
if (t.Assembly == baseAssembly)
continue;
candidateTypes.Add(t);
}
}
}
var baseAssembly = typeof(NPC).Assembly;
var candidateTypes = ReflectionUtils.GetDerivedClasses<NPC>()
.Where(type => type.Assembly != baseAssembly);

foreach (System.Type type in candidateTypes.OrderBy(
candidate => candidate.FullName,
Expand Down
136 changes: 130 additions & 6 deletions S1API/Internal/Utils/ReflectionUtils.cs
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,8 @@ namespace S1API.Internal.Utils
/// </summary>
internal static class ReflectionUtils
{
private static readonly Log Logger = new Log("ReflectionUtils");

private const BindingFlags InstanceMemberFlags = BindingFlags.Public
| BindingFlags.NonPublic
| BindingFlags.Instance
Expand All @@ -29,18 +31,32 @@ internal static class ReflectionUtils
internal static List<Type> GetDerivedClasses<TBaseClass>()
{
List<Type> derivedClasses = new List<Type>();
Assembly[] applicableAssemblies = AppDomain.CurrentDomain.GetAssemblies()
.Where(assembly => !ShouldSkipAssembly(assembly))
Type baseType = typeof(TBaseClass);
Assembly baseAssembly = baseType.Assembly;
Assembly[] loadedAssemblies = AppDomain.CurrentDomain.GetAssemblies();
IReadOnlyDictionary<string, Assembly[]> assembliesBySimpleName =
IndexAssembliesBySimpleName(loadedAssemblies);
Assembly[] applicableAssemblies = loadedAssemblies
.Where(assembly => assembly == baseAssembly || !ShouldSkipAssembly(assembly))
.Where(assembly => CanContainTypesDerivedFrom(
assembly,
baseAssembly.GetName(),
assembliesBySimpleName))
.ToArray();

Logger.Debug(
$"[S1API][Reflection] Scanning {applicableAssemblies.Length} of {loadedAssemblies.Length} " +
$"loaded assemblies for types derived from '{baseType.FullName}'.");

foreach (Assembly assembly in applicableAssemblies)
foreach (Type type in SafeGetTypes(assembly))
{
try
{
if (type == null)
continue;
if (typeof(TBaseClass).IsAssignableFrom(type)
&& type != typeof(TBaseClass)
if (baseType.IsAssignableFrom(type)
&& type != baseType
&& !type.IsAbstract)
{
derivedClasses.Add(type);
Expand All @@ -58,6 +74,104 @@ internal static List<Type> GetDerivedClasses<TBaseClass>()
return derivedClasses;
}

internal static bool CanContainTypesDerivedFrom(
Assembly candidateAssembly,
Assembly baseAssembly,
IEnumerable<Assembly> loadedAssemblies)
{
if (candidateAssembly == baseAssembly)
return true;

IReadOnlyDictionary<string, Assembly[]> assembliesBySimpleName =
IndexAssembliesBySimpleName(loadedAssemblies);

return ReferencesAssemblyTransitively(
candidateAssembly,
baseAssembly.GetName(),
assembliesBySimpleName,
new HashSet<Assembly>());
}

private static IReadOnlyDictionary<string, Assembly[]> IndexAssembliesBySimpleName(
IEnumerable<Assembly> loadedAssemblies)
{
return loadedAssemblies
.Where(assembly => assembly != null)
.GroupBy(assembly => assembly.GetName().Name ?? string.Empty, StringComparer.OrdinalIgnoreCase)
.ToDictionary(group => group.Key, group => group.ToArray(), StringComparer.OrdinalIgnoreCase);
}

private static bool CanContainTypesDerivedFrom(
Assembly candidateAssembly,
AssemblyName baseAssemblyName,
IReadOnlyDictionary<string, Assembly[]> assembliesBySimpleName)
{
if (AssemblyIdentityMatches(candidateAssembly.GetName(), baseAssemblyName))
return true;

return ReferencesAssemblyTransitively(
candidateAssembly,
baseAssemblyName,
assembliesBySimpleName,
new HashSet<Assembly>());
}

private static bool ReferencesAssemblyTransitively(
Assembly candidateAssembly,
AssemblyName baseAssemblyName,
IReadOnlyDictionary<string, Assembly[]> assembliesBySimpleName,
HashSet<Assembly> visitedAssemblies)
{
if (!visitedAssemblies.Add(candidateAssembly))
return false;

AssemblyName[] referencedAssemblies;
try
{
referencedAssemblies = candidateAssembly.GetReferencedAssemblies();
}
catch
{
return false;
}

foreach (AssemblyName referencedAssembly in referencedAssemblies)
{
if (AssemblyIdentityMatches(referencedAssembly, baseAssemblyName))
return true;

string referencedName = referencedAssembly.Name ?? string.Empty;
if (!assembliesBySimpleName.TryGetValue(referencedName, out Assembly[]? loadedReferences)
|| loadedReferences == null)
continue;

foreach (Assembly loadedReference in loadedReferences)
Comment thread
ifBars marked this conversation as resolved.
{
if (!AssemblyIdentityMatches(loadedReference.GetName(), referencedAssembly))
continue;

if (ReferencesAssemblyTransitively(
loadedReference,
baseAssemblyName,
assembliesBySimpleName,
visitedAssemblies))
Comment thread
coderabbitai[bot] marked this conversation as resolved.
{
return true;
}
}
}

return false;
}

private static bool AssemblyIdentityMatches(
AssemblyName referenceAssemblyName,
AssemblyName definitionAssemblyName) =>
string.Equals(
referenceAssemblyName.FullName,
definitionAssemblyName.FullName,
StringComparison.OrdinalIgnoreCase);

/// <summary>
/// INTERNAL: Gets all types by their name.
/// </summary>
Expand Down Expand Up @@ -137,16 +251,26 @@ private static bool ShouldSkipAssembly(Assembly assembly)
/// <returns>The types that were successfully loaded from the assembly.</returns>
private static IEnumerable<Type> SafeGetTypes(Assembly asm)
{
string assemblyName = asm.FullName ?? asm.GetName().Name ?? "<unknown>";
Logger.Debug($"[S1API][Reflection] About to enumerate types in '{assemblyName}'.");

try
{
return asm.GetTypes();
Type[] types = asm.GetTypes();
Logger.Debug(
$"[S1API][Reflection] Enumerated {types.Length} types in '{assemblyName}'.");
return types;
}
catch (ReflectionTypeLoadException ex)
{
return ex.Types.Where(t => t != null)!.Cast<Type>();
Type[] loadedTypes = ex.Types.Where(type => type != null).Cast<Type>().ToArray();
Logger.Debug(
$"[S1API][Reflection] Partially enumerated {loadedTypes.Length} types in '{assemblyName}'.");
return loadedTypes;
}
catch
{
Logger.Debug($"[S1API][Reflection] Failed to enumerate types in '{assemblyName}'.");
return Array.Empty<Type>();
}
}
Expand Down
Loading