#nullable enable using System; using System.Collections.Generic; using System.IO; using System.Linq; using Mono.Cecil; using Mono.Cecil.Cil; namespace ShrinkSDK.CodeGen { public static class ShrinkAssemblyWeaver { public const string Version = "0.1.0"; private const string EventRuntime = "ShrinkEventBus.Runtime"; private const string NetworkEventIntegration = "ShrinkNetwork.Integration.EventBus"; public static ShrinkWeaveResult Weave( string assemblyPath, string? pdbPath, IEnumerable referencePaths, string outputAssemblyPath, string? outputPdbPath, ShrinkCodeGenPlatform platform = ShrinkCodeGenPlatform.EngineNeutral, string? strongNameKeyPath = null) { var diagnostics = new List(); var references = referencePaths.Where(File.Exists).Select(Path.GetFullPath).Distinct(StringComparer.OrdinalIgnoreCase).ToArray(); using var resolver = new PathAssemblyResolver(references.Append(assemblyPath)); try { var hasSymbols = !string.IsNullOrWhiteSpace(pdbPath) && File.Exists(pdbPath); var reader = new ReaderParameters { AssemblyResolver = resolver, ReadingMode = ReadingMode.Immediate, ReadSymbols = hasSymbols, SymbolReaderProvider = hasSymbols ? new PortablePdbReaderProvider() : null }; using var assemblyInput = new MemoryStream(File.ReadAllBytes(assemblyPath)); using var symbolInput = hasSymbols ? new MemoryStream(File.ReadAllBytes(pdbPath!)) : null; if (hasSymbols) reader.SymbolStream = symbolInput; using var assembly = AssemblyDefinition.ReadAssembly(assemblyInput, reader); resolver.SetSelf(assembly); byte[]? strongNameKey = null; if (assembly.Name.HasPublicKey) { if (string.IsNullOrWhiteSpace(strongNameKeyPath)) throw new InvalidOperationException("Signed assemblies are not modified unless ShrinkCodeGenStrongNameKeyFile is configured."); if (!File.Exists(strongNameKeyPath)) throw new InvalidOperationException($"Strong-name key file does not exist: {strongNameKeyPath}"); strongNameKey = File.ReadAllBytes(strongNameKeyPath); } var module = assembly.MainModule; var marker = FindType(module, "ShrinkSDK.Runtime.ShrinkCodeGenWovenAttribute", "ShrinkRuntime.Abstractions"); var existingMarker = marker == null ? null : assembly.CustomAttributes.FirstOrDefault(a => a.AttributeType.FullName == marker.FullName); if (existingMarker != null) { var wovenVersion = existingMarker.ConstructorArguments.Count > 0 ? existingMarker.ConstructorArguments[0].Value as string : null; if (wovenVersion == Version) return new ShrinkWeaveResult(false, 0, 0, 0, diagnostics); throw new InvalidOperationException($"Assembly was woven by ShrinkSDK.CodeGen {wovenVersion ?? "unknown"}. Run a clean rebuild before weaving with {Version}."); } var inputMvid = module.Mvid.ToString("D"); var instanceCount = 0; var staticCount = 0; if (References(module, EventRuntime)) WeaveEventBus(module, platform, ref instanceCount, ref staticCount); var registryCount = WeaveRegistries(module); if (References(module, NetworkEventIntegration)) registryCount += WeaveNetworkEventBindings(module); if (marker != null) AddMarker(module, marker, inputMvid); Directory.CreateDirectory(Path.GetDirectoryName(Path.GetFullPath(outputAssemblyPath))!); using var symbolOutput = hasSymbols ? new MemoryStream() : null; var writer = new WriterParameters { WriteSymbols = hasSymbols, SymbolWriterProvider = hasSymbols ? new PortablePdbWriterProvider() : null, SymbolStream = symbolOutput, StrongNameKeyBlob = strongNameKey }; assembly.Write(outputAssemblyPath, writer); if (hasSymbols && !string.IsNullOrWhiteSpace(outputPdbPath)) File.WriteAllBytes(outputPdbPath!, symbolOutput!.ToArray()); return new ShrinkWeaveResult(true, instanceCount, staticCount, registryCount, diagnostics); } catch (Exception exception) { diagnostics.Add(new ShrinkCodeGenDiagnostic(ShrinkCodeGenDiagnosticSeverity.Error, exception.Message)); return new ShrinkWeaveResult(false, 0, 0, 0, diagnostics); } } private static void WeaveEventBus(ModuleDefinition module, ShrinkCodeGenPlatform platform, ref int instanceCount, ref int staticCount) { var subscriberAttribute = RequireType(module, "ShrinkEventBus.ShrinkEventSubscriberAttribute", EventRuntime); var subscribeAttribute = RequireType(module, "ShrinkEventBus.ShrinkSubscribeAttribute", EventRuntime); var allTypes = AllTypes(module.Types).Where(type => !type.IsInterface).ToArray(); foreach (var type in allTypes.Where(type => HasAttribute(type, subscriberAttribute))) { var instanceHandlers = type.Methods.Where(method => !method.IsStatic && HasAttribute(method, subscribeAttribute)).ToArray(); if (instanceHandlers.Length > 0) { var lifetime = ReadInt(type.CustomAttributes.First(a => a.AttributeType.FullName == subscriberAttribute.FullName), "Lifetime", 0); if (InjectInstanceBinding(module, type, instanceHandlers, subscribeAttribute)) instanceCount++; if (lifetime == 1 && platform == ShrinkCodeGenPlatform.Unity) InjectAwakeToDestroyLifetime(type, module); else if (lifetime != 0) throw new InvalidOperationException($"{type.FullName}: engine-neutral weaving supports Manual lifetime only. Godot Nodes should attach from _Ready and dispose from _ExitTree."); } } var staticTypes = allTypes.Where(type => HasAttribute(type, subscriberAttribute)) .Where(type => type.Methods.Any(method => method.IsStatic && HasAttribute(method, subscribeAttribute))) .ToArray(); if (staticTypes.Length > 0) { staticCount = staticTypes.Sum(type => type.Methods.Count(method => method.IsStatic && HasAttribute(method, subscribeAttribute))); InjectStaticBootstrap(module, staticTypes, subscriberAttribute, subscribeAttribute); } } private static void InjectAwakeToDestroyLifetime(TypeDefinition type, ModuleDefinition module) { if (!InheritsFrom(type, "UnityEngine.MonoBehaviour")) throw new InvalidOperationException($"[ShrinkEventSubscriber(Lifetime = AwakeToDestroy)] requires MonoBehaviour: {type.FullName}."); const string bindingFieldName = "__shrinkEventBusAwakeToDestroyBinding"; if (type.Fields.Any(field => field.Name == bindingFieldName)) throw new InvalidOperationException($"Reserved generated field already exists on {type.FullName}: {bindingFieldName}."); var disposableType = module.ImportReference(typeof(IDisposable)); var bindingField = new FieldDefinition(bindingFieldName, FieldAttributes.Private, disposableType); type.Fields.Add(bindingField); var eventBusType = RequireType(module, "ShrinkEventBus.EventBus", EventRuntime); var attachMethod = ImportMethod(module, eventBusType, method => method.Name == "Attach" && method.IsStatic && method.Parameters.Count == 2); var disposeMethod = module.ImportReference(typeof(IDisposable).GetMethod(nameof(IDisposable.Dispose)) ?? throw new InvalidOperationException("IDisposable.Dispose was not found.")); InjectAwake(type, module, bindingField, attachMethod); InjectOnDestroy(type, module, bindingField, disposeMethod); } private static void InjectAwake(TypeDefinition type, ModuleDefinition module, FieldDefinition bindingField, MethodReference attachMethod) { var awake = type.Methods.FirstOrDefault(method => method.Name == "Awake" && !method.IsStatic && method.Parameters.Count == 0); if (awake != null) { InsertAttachAtStart(awake, module, bindingField, attachMethod); return; } var baseAwake = FindBaseMethodReference(type, "Awake", module); awake = new MethodDefinition("Awake", baseAwake != null ? MethodAttributes.Family | MethodAttributes.HideBySig | MethodAttributes.Virtual : MethodAttributes.Family | MethodAttributes.HideBySig | MethodAttributes.Virtual | MethodAttributes.NewSlot, module.TypeSystem.Void); awake.Body.InitLocals = true; var il = awake.Body.GetILProcessor(); if (baseAwake != null) { il.Emit(OpCodes.Ldarg_0); il.Emit(OpCodes.Call, baseAwake); } EmitAttach(il, awake, module, bindingField, attachMethod); il.Emit(OpCodes.Ret); type.Methods.Add(awake); } private static void InjectOnDestroy(TypeDefinition type, ModuleDefinition module, FieldDefinition bindingField, MethodReference disposeMethod) { var onDestroy = type.Methods.FirstOrDefault(method => method.Name == "OnDestroy" && !method.IsStatic && method.Parameters.Count == 0); if (onDestroy != null) { InsertDisposeAtStart(onDestroy, bindingField, disposeMethod); return; } var baseOnDestroy = FindBaseMethodReference(type, "OnDestroy", module); onDestroy = new MethodDefinition("OnDestroy", baseOnDestroy != null ? MethodAttributes.Family | MethodAttributes.HideBySig | MethodAttributes.Virtual : MethodAttributes.Family | MethodAttributes.HideBySig | MethodAttributes.Virtual | MethodAttributes.NewSlot, module.TypeSystem.Void); var il = onDestroy.Body.GetILProcessor(); EmitDispose(il, bindingField, disposeMethod); if (baseOnDestroy != null) { il.Emit(OpCodes.Ldarg_0); il.Emit(OpCodes.Call, baseOnDestroy); } il.Emit(OpCodes.Ret); type.Methods.Add(onDestroy); } private static void InsertAttachAtStart(MethodDefinition method, ModuleDefinition module, FieldDefinition bindingField, MethodReference attachMethod) { if (!method.HasBody || method.Body.Instructions.Count == 0) throw new InvalidOperationException($"Awake has no body: {method.FullName}."); method.Body.InitLocals = true; var processor = method.Body.GetILProcessor(); var instructions = BuildAttachInstructions(processor, method, module, bindingField, attachMethod); var first = method.Body.Instructions[0]; foreach (var instruction in instructions) processor.InsertBefore(first, instruction); } private static void EmitAttach(ILProcessor il, MethodDefinition method, ModuleDefinition module, FieldDefinition bindingField, MethodReference attachMethod) { foreach (var instruction in BuildAttachInstructions(il, method, module, bindingField, attachMethod)) il.Append(instruction); } private static IReadOnlyList BuildAttachInstructions(ILProcessor il, MethodDefinition method, ModuleDefinition module, FieldDefinition bindingField, MethodReference attachMethod) { var defaultBusType = module.ImportReference(attachMethod.Parameters[1].ParameterType); var defaultBus = new VariableDefinition(defaultBusType); method.Body.Variables.Add(defaultBus); var attached = il.Create(OpCodes.Nop); return new[] { il.Create(OpCodes.Ldarg_0), il.Create(OpCodes.Ldfld, bindingField), il.Create(OpCodes.Brtrue_S, attached), il.Create(OpCodes.Ldarg_0), il.Create(OpCodes.Ldarg_0), il.Create(OpCodes.Ldloca_S, defaultBus), il.Create(OpCodes.Initobj, defaultBusType), il.Create(OpCodes.Ldloc, defaultBus), il.Create(OpCodes.Call, attachMethod), il.Create(OpCodes.Stfld, bindingField), attached }; } private static void InsertDisposeAtStart(MethodDefinition method, FieldDefinition bindingField, MethodReference disposeMethod) { if (!method.HasBody || method.Body.Instructions.Count == 0) throw new InvalidOperationException($"OnDestroy has no body: {method.FullName}."); var processor = method.Body.GetILProcessor(); var instructions = BuildDisposeInstructions(processor, bindingField, disposeMethod); var first = method.Body.Instructions[0]; foreach (var instruction in instructions) processor.InsertBefore(first, instruction); } private static void EmitDispose(ILProcessor il, FieldDefinition bindingField, MethodReference disposeMethod) { foreach (var instruction in BuildDisposeInstructions(il, bindingField, disposeMethod)) il.Append(instruction); } private static IReadOnlyList BuildDisposeInstructions(ILProcessor il, FieldDefinition bindingField, MethodReference disposeMethod) { var disposed = il.Create(OpCodes.Nop); return new[] { il.Create(OpCodes.Ldarg_0), il.Create(OpCodes.Ldfld, bindingField), il.Create(OpCodes.Brfalse_S, disposed), il.Create(OpCodes.Ldarg_0), il.Create(OpCodes.Ldfld, bindingField), il.Create(OpCodes.Callvirt, disposeMethod), il.Create(OpCodes.Ldarg_0), il.Create(OpCodes.Ldnull), il.Create(OpCodes.Stfld, bindingField), disposed }; } private static MethodReference? FindBaseMethodReference(TypeDefinition type, string methodName, ModuleDefinition module) { try { var baseTypeReference = type.BaseType; while (baseTypeReference != null) { var baseType = baseTypeReference.Resolve(); if (baseType == null) break; var method = baseType.Methods.FirstOrDefault(candidate => candidate.Name == methodName && !candidate.IsStatic && candidate.IsVirtual && candidate.Parameters.Count == 0); if (method != null) { if (baseTypeReference is GenericInstanceType genericBase) { return new MethodReference(method.Name, module.ImportReference(method.ReturnType), module.ImportReference(genericBase)) { HasThis = method.HasThis, ExplicitThis = method.ExplicitThis, CallingConvention = method.CallingConvention }; } return module.ImportReference(method); } baseTypeReference = baseType.BaseType; } } catch { } return null; } private static bool InheritsFrom(TypeDefinition type, string expectedFullName) { var current = type.BaseType; while (current != null) { if (current.FullName == expectedFullName) return true; try { current = current.Resolve()?.BaseType; } catch { return false; } } return false; } private static bool InjectInstanceBinding(ModuleDefinition module, TypeDefinition owner, IReadOnlyList handlers, TypeReference subscribeAttribute) { var generatedInterface = RequireType(module, "ShrinkEventBus.IShrinkGeneratedSubscriber", EventRuntime); if (owner.Interfaces.Any(item => item.InterfaceType.FullName == generatedInterface.FullName)) return false; var resolverType = RequireType(module, "ShrinkEventBus.IShrinkBusResolver", EventRuntime); var busKeyType = RequireType(module, "ShrinkEventBus.ShrinkBusKey", EventRuntime); var bindingType = RequireType(module, "ShrinkEventBus.ShrinkEventBinding", EventRuntime); var helperType = RequireType(module, "ShrinkEventBus.ShrinkGeneratedBinding", EventRuntime); var disposableType = RequireType(module, "System.IDisposable", module.TypeSystem.CoreLibrary.Name); var nullableBusKey = GenericType(module, "System.Nullable`1", true, busKeyType); var bindingCtor = ImportMethod(module, bindingType, method => method.IsConstructor && method.Parameters.Count == 0); var bindingAdd = ImportMethod(module, bindingType, method => method.Name == "Add" && method.Parameters.Count == 1); var interfaceMethod = ImportMethod(module, generatedInterface, method => method.Name == "AttachGenerated"); var syncSubscribe = ImportMethod(module, helperType, method => method.Name == "Subscribe" && method.HasGenericParameters); var asyncSubscribe = ImportMethod(module, helperType, method => method.Name == "SubscribeAsync" && method.HasGenericParameters); var legacySubscribe = ImportMethod(module, helperType, method => method.Name == "SubscribeAsyncLegacy" && method.HasGenericParameters); var asyncHandlerType = RequireType(module, "ShrinkEventBus.ShrinkAsyncEventHandler`1", EventRuntime); var generated = new MethodDefinition("ShrinkEventBus.IShrinkGeneratedSubscriber.AttachGenerated", MethodAttributes.Private | MethodAttributes.Final | MethodAttributes.HideBySig | MethodAttributes.NewSlot | MethodAttributes.Virtual, disposableType); generated.Parameters.Add(new ParameterDefinition("resolver", ParameterAttributes.None, resolverType)); generated.Parameters.Add(new ParameterDefinition("defaultBus", ParameterAttributes.Optional, nullableBusKey)); generated.Overrides.Add(interfaceMethod); generated.Body.InitLocals = true; var bindingLocal = new VariableDefinition(bindingType); generated.Body.Variables.Add(bindingLocal); var il = generated.Body.GetILProcessor(); il.Emit(OpCodes.Newobj, bindingCtor); il.Emit(OpCodes.Stloc, bindingLocal); var classDefaultBus = ReadString(owner.CustomAttributes.First(a => a.AttributeType.FullName == "ShrinkEventBus.ShrinkEventSubscriberAttribute"), "DefaultBus"); foreach (var handler in handlers) EmitInstanceSubscription(module, il, handler, handler.CustomAttributes.First(a => a.AttributeType.FullName == subscribeAttribute.FullName), classDefaultBus, bindingLocal, bindingAdd, syncSubscribe, asyncSubscribe, legacySubscribe, asyncHandlerType); il.Emit(OpCodes.Ldloc, bindingLocal); il.Emit(OpCodes.Ret); owner.Interfaces.Add(new InterfaceImplementation(generatedInterface)); owner.Methods.Add(generated); return true; } private static void EmitInstanceSubscription(ModuleDefinition module, ILProcessor il, MethodDefinition handler, CustomAttribute attribute, string classDefaultBus, VariableDefinition bindingLocal, MethodReference bindingAdd, MethodReference syncSubscribe, MethodReference asyncSubscribe, MethodReference legacySubscribe, TypeReference asyncHandlerType) { var signature = ResolveEventHandler(module, handler, syncSubscribe, asyncSubscribe, legacySubscribe, asyncHandlerType); var bus = ReadString(attribute, "Bus"); if (string.IsNullOrWhiteSpace(bus)) bus = classDefaultBus; il.Emit(OpCodes.Ldloc, bindingLocal); il.Emit(OpCodes.Ldarg_1); il.Emit(OpCodes.Ldarg_2); il.Emit(OpCodes.Ldstr, bus ?? string.Empty); il.Emit(OpCodes.Ldarg_0); il.Emit(OpCodes.Ldarg_0); il.Emit(OpCodes.Ldftn, module.ImportReference(handler)); il.Emit(OpCodes.Newobj, DelegateConstructor(module, signature.DelegateType)); EmitOptions(il, attribute); il.Emit(OpCodes.Call, signature.Method); il.Emit(OpCodes.Callvirt, bindingAdd); } private static void InjectStaticBootstrap(ModuleDefinition module, IReadOnlyList types, TypeReference subscriberAttribute, TypeReference subscribeAttribute) { if (module.Types.Any(type => type.FullName == "ShrinkEventBus.Generated.ShrinkGeneratedStaticBindings")) return; var registry = RequireType(module, "ShrinkEventBus.ShrinkStaticBindingRegistry", EventRuntime); var syncRegister = ImportMethod(module, registry, method => method.Name == "Register" && method.HasGenericParameters); var asyncRegister = ImportMethod(module, registry, method => method.Name == "RegisterAsync" && method.HasGenericParameters); var legacyRegister = ImportMethod(module, registry, method => method.Name == "RegisterAsyncLegacy" && method.HasGenericParameters); var asyncHandlerType = RequireType(module, "ShrinkEventBus.ShrinkAsyncEventHandler`1", EventRuntime); var bootstrap = new TypeDefinition("ShrinkEventBus.Generated", "ShrinkGeneratedStaticBindings", TypeAttributes.Abstract | TypeAttributes.Sealed | TypeAttributes.NotPublic, module.TypeSystem.Object); module.Types.Add(bootstrap); var register = new MethodDefinition("Register", MethodAttributes.Assembly | MethodAttributes.Static | MethodAttributes.HideBySig, module.TypeSystem.Void); bootstrap.Methods.Add(register); var il = register.Body.GetILProcessor(); foreach (var type in types) { var classBus = ReadString(type.CustomAttributes.First(a => a.AttributeType.FullName == subscriberAttribute.FullName), "DefaultBus"); foreach (var handler in type.Methods.Where(method => method.IsStatic && HasAttribute(method, subscribeAttribute)).ToArray()) { var attribute = handler.CustomAttributes.First(a => a.AttributeType.FullName == subscribeAttribute.FullName); var signature = ResolveEventHandler(module, handler, syncRegister, asyncRegister, legacyRegister, asyncHandlerType); var bus = ReadString(attribute, "Bus"); if (string.IsNullOrWhiteSpace(bus)) bus = classBus; var bridge = CreateStaticBridge(module, handler); il.Emit(OpCodes.Ldstr, bus ?? string.Empty); il.Emit(OpCodes.Ldnull); il.Emit(OpCodes.Ldftn, bridge); il.Emit(OpCodes.Newobj, DelegateConstructor(module, signature.DelegateType)); EmitOptions(il, attribute); il.Emit(OpCodes.Call, signature.Method); } } il.Emit(OpCodes.Ret); var moduleType = module.Types.First(type => type.Name == ""); var initializer = moduleType.Methods.FirstOrDefault(method => method.Name == ".cctor"); if (initializer == null) { initializer = new MethodDefinition(".cctor", MethodAttributes.Private | MethodAttributes.Static | MethodAttributes.HideBySig | MethodAttributes.SpecialName | MethodAttributes.RTSpecialName, module.TypeSystem.Void); initializer.Body.GetILProcessor().Emit(OpCodes.Ret); moduleType.Methods.Add(initializer); } initializer.Body.GetILProcessor().InsertBefore(initializer.Body.Instructions[0], Instruction.Create(OpCodes.Call, register)); } private static (MethodReference Method, TypeReference DelegateType) ResolveEventHandler(ModuleDefinition module, MethodDefinition handler, MethodReference sync, MethodReference async, MethodReference legacy, TypeReference asyncHandlerType) { if (handler.Parameters.Count is < 1 or > 2) throw new InvalidOperationException($"{handler.FullName}: [ShrinkSubscribe] requires one event parameter and an optional CancellationToken."); var eventType = module.ImportReference(handler.Parameters[0].ParameterType); var eventInterface = RequireType(module, "ShrinkEventBus.IShrinkEvent", EventRuntime); if (!Implements(eventType, eventInterface.FullName)) throw new InvalidOperationException($"{handler.FullName}: event parameter must implement IShrinkEvent."); MethodReference open; TypeReference delegateType; if (handler.ReturnType.MetadataType == MetadataType.Void && handler.Parameters.Count == 1) { open = sync; delegateType = GenericType(module, "System.Action`1", false, eventType); } else if (handler.ReturnType.FullName == "Cysharp.Threading.Tasks.UniTask" && handler.Parameters.Count == 1) { open = legacy; delegateType = GenericType(module, "System.Func`2", false, eventType, module.ImportReference(handler.ReturnType)); } else if (handler.ReturnType.FullName == "Cysharp.Threading.Tasks.UniTask" && handler.Parameters.Count == 2 && handler.Parameters[1].ParameterType.FullName == "System.Threading.CancellationToken") { open = async; delegateType = new GenericInstanceType(asyncHandlerType) { GenericArguments = { eventType } }; } else throw new InvalidOperationException($"Unsupported [ShrinkSubscribe] signature: {handler.FullName}."); var closed = new GenericInstanceMethod(open); closed.GenericArguments.Add(eventType); return (closed, delegateType); } private static MethodReference CreateStaticBridge(ModuleDefinition module, MethodDefinition handler) { var bridge = new MethodDefinition($"ShrinkEventBus.GeneratedStaticHandler_{handler.MetadataToken.ToInt32():X8}", MethodAttributes.Assembly | MethodAttributes.Static | MethodAttributes.HideBySig, module.ImportReference(handler.ReturnType)); foreach (var parameter in handler.Parameters) bridge.Parameters.Add(new ParameterDefinition(parameter.Name, parameter.Attributes, module.ImportReference(parameter.ParameterType))); var il = bridge.Body.GetILProcessor(); foreach (var parameter in bridge.Parameters) il.Emit(OpCodes.Ldarg, parameter); il.Emit(OpCodes.Call, module.ImportReference(handler)); il.Emit(OpCodes.Ret); handler.DeclaringType.Methods.Add(bridge); return bridge; } private static void EmitOptions(ILProcessor il, CustomAttribute attribute) { il.Emit(OpCodes.Ldc_I4, ReadInt(attribute, "Priority", 2)); il.Emit(OpCodes.Ldc_I4, ReadInt(attribute, "NumericPriority", 0)); il.Emit(ReadBool(attribute, "ReceiveCanceled", false) ? OpCodes.Ldc_I4_1 : OpCodes.Ldc_I4_0); } private static int WeaveRegistries(ModuleDefinition module) { var count = 0; count += AddRegistry(module, "ShrinkCommand.ShrinkCommandSubscriberAttribute", "ShrinkCommand.Runtime", "ShrinkCommand.ShrinkCommandStaticRegistryAttribute", type => type.Methods.Any(method => method.IsStatic && HasAttribute(method, "ShrinkCommand.ShrinkCommandAttribute"))); count += AddRegistry(module, "ShrinkNetwork.ShrinkNetworkMessageAttribute", "ShrinkNetwork.Runtime", "ShrinkNetwork.ShrinkNetworkMessageRegistryAttribute", _ => true); count += AddRegistry(module, "ShrinkNetwork.ShrinkNetworkSubscriberAttribute", "ShrinkNetwork.Runtime", "ShrinkNetwork.ShrinkNetworkStaticSubscriberRegistryAttribute", type => type.Methods.Any(method => method.IsStatic && HasAttribute(method, "ShrinkNetwork.ShrinkNetworkSubscribeAttribute"))); count += AddAppRegistry(module); ValidateDuplicateNetworkContracts(module); ValidateDuplicateCommandPaths(module); return count; } private static int WeaveNetworkEventBindings(ModuleDefinition module) { if (module.Types.Any(type => type.FullName == "ShrinkNetwork.Integration.Generated.ShrinkGeneratedNetworkEventBindings")) return 0; var networkEventAttribute = FindType(module, "ShrinkNetwork.Integration.ShrinkNetworkEventAttribute", NetworkEventIntegration); var messageAttribute = FindType(module, "ShrinkNetwork.ShrinkNetworkMessageAttribute", "ShrinkNetwork.Runtime"); var registryType = FindType(module, "ShrinkNetwork.Integration.ShrinkNetworkEventRegistry", NetworkEventIntegration); if (networkEventAttribute == null || messageAttribute == null || registryType == null) return 0; var registrations = AllTypes(module.Types) .Where(type => !type.IsAbstract && HasAttribute(type, networkEventAttribute)) .Select(type => new { Type = type, Message = type.CustomAttributes.FirstOrDefault(attribute => attribute.AttributeType.FullName == messageAttribute.FullName) }) .Where(item => item.Message != null) .ToArray(); if (registrations.Length == 0) return 0; var registerOpen = ImportMethod(module, registryType, method => method.Name == "Register" && method.HasGenericParameters); var bootstrap = new TypeDefinition("ShrinkNetwork.Integration.Generated", "ShrinkGeneratedNetworkEventBindings", TypeAttributes.Abstract | TypeAttributes.Sealed | TypeAttributes.NotPublic, module.TypeSystem.Object); module.Types.Add(bootstrap); var register = new MethodDefinition("Register", MethodAttributes.Assembly | MethodAttributes.Static | MethodAttributes.HideBySig, module.TypeSystem.Void); bootstrap.Methods.Add(register); var il = register.Body.GetILProcessor(); foreach (var item in registrations) { var attribute = item.Message!; if (attribute.ConstructorArguments.Count == 0) throw new InvalidOperationException($"[ShrinkNetworkMessage] on {item.Type.FullName} has no opcode."); var opcode = Convert.ToInt32(attribute.ConstructorArguments[0].Value); var route = attribute.ConstructorArguments.Count > 1 ? attribute.ConstructorArguments[1].Value as string : null; var registerClosed = new GenericInstanceMethod(registerOpen); registerClosed.GenericArguments.Add(module.ImportReference(item.Type)); il.Emit(OpCodes.Ldc_I4, opcode); if (route == null) il.Emit(OpCodes.Ldnull); else il.Emit(OpCodes.Ldstr, route); il.Emit(OpCodes.Call, registerClosed); } il.Emit(OpCodes.Ret); var moduleType = module.Types.First(type => type.Name == ""); var initializer = moduleType.Methods.FirstOrDefault(method => method.Name == ".cctor"); if (initializer == null) { initializer = new MethodDefinition(".cctor", MethodAttributes.Private | MethodAttributes.Static | MethodAttributes.HideBySig | MethodAttributes.SpecialName | MethodAttributes.RTSpecialName, module.TypeSystem.Void); initializer.Body.GetILProcessor().Emit(OpCodes.Ret); moduleType.Methods.Add(initializer); } initializer.Body.GetILProcessor().InsertBefore(initializer.Body.Instructions[0], Instruction.Create(OpCodes.Call, register)); return registrations.Length; } private static int AddRegistry(ModuleDefinition module, string markerName, string assemblyName, string registryName, Func predicate) { var marker = FindType(module, markerName, assemblyName); var registry = FindType(module, registryName, assemblyName); if (marker == null || registry == null) return 0; if (module.Assembly.CustomAttributes.Any(attribute => attribute.AttributeType.FullName == registry.FullName)) return 0; var types = AllTypes(module.Types).Where(type => HasAttribute(type, marker)).Where(predicate).ToArray(); if (types.Length == 0) return 0; AddTypeArrayAttribute(module, registry, types); return types.Length; } private static int AddAppRegistry(ModuleDefinition module) { var marker = FindType(module, "ShrinkApp.ShrinkAppModuleInstallerAttribute", "ShrinkApp.Core.Runtime"); var contract = FindType(module, "ShrinkApp.IShrinkAppModuleInstaller", "ShrinkApp.Core.Runtime"); var registry = FindType(module, "ShrinkApp.ShrinkAppInstallerRegistryAttribute", "ShrinkApp.Core.Runtime"); if (marker == null || contract == null || registry == null) return 0; if (module.Assembly.CustomAttributes.Any(attribute => attribute.AttributeType.FullName == registry.FullName)) return 0; var types = AllTypes(module.Types).Where(type => !type.IsAbstract && HasAttribute(type, marker) && Implements(type, contract.FullName)).ToArray(); if (types.Length == 0) return 0; AddTypeArrayAttribute(module, registry, types); return types.Length; } private static void AddTypeArrayAttribute(ModuleDefinition module, TypeReference registry, IReadOnlyList types) { var ctor = ImportMethod(module, registry, method => method.IsConstructor && method.Parameters.Count == 1 && method.Parameters[0].ParameterType.IsArray); var typeRef = new TypeReference("System", "Type", module, module.TypeSystem.CoreLibrary); var typeArray = new ArrayType(typeRef); var attribute = new CustomAttribute(ctor); attribute.ConstructorArguments.Add(new CustomAttributeArgument(typeArray, types.Select(type => new CustomAttributeArgument(typeRef, module.ImportReference(type))).ToArray())); module.Assembly.CustomAttributes.Add(attribute); } private static void ValidateDuplicateNetworkContracts(ModuleDefinition module) { var seen = new Dictionary(StringComparer.Ordinal); foreach (var type in AllTypes(module.Types)) foreach (var attribute in type.CustomAttributes.Where(a => a.AttributeType.FullName == "ShrinkNetwork.ShrinkNetworkMessageAttribute")) { var opcode = Convert.ToInt32(attribute.ConstructorArguments[0].Value); var route = attribute.ConstructorArguments.Count > 1 ? attribute.ConstructorArguments[1].Value as string ?? string.Empty : string.Empty; var key = $"{opcode}:{route}"; if (seen.TryGetValue(key, out var previous)) throw new InvalidOperationException($"Duplicate network opcode/route {key}: {previous} and {type.FullName}."); seen[key] = type.FullName; } } private static void ValidateDuplicateCommandPaths(ModuleDefinition module) { var seen = new Dictionary(StringComparer.OrdinalIgnoreCase); foreach (var type in AllTypes(module.Types)) foreach (var method in type.Methods) foreach (var attribute in method.CustomAttributes.Where(a => a.AttributeType.FullName == "ShrinkCommand.ShrinkCommandAttribute")) { var path = attribute.ConstructorArguments[0].Value as string ?? string.Empty; if (seen.TryGetValue(path, out var previous)) throw new InvalidOperationException($"Duplicate command path '{path}': {previous} and {method.FullName}."); seen[path] = method.FullName; } } private static void AddMarker(ModuleDefinition module, TypeReference marker, string inputMvid) { var ctor = ImportMethod(module, marker, method => method.IsConstructor && method.Parameters.Count == 2); var attribute = new CustomAttribute(ctor); attribute.ConstructorArguments.Add(new CustomAttributeArgument(module.TypeSystem.String, Version)); attribute.ConstructorArguments.Add(new CustomAttributeArgument(module.TypeSystem.String, inputMvid)); module.Assembly.CustomAttributes.Add(attribute); } private static bool References(ModuleDefinition module, string assemblyName) => module.AssemblyReferences.Any(reference => string.Equals(reference.Name, assemblyName, StringComparison.Ordinal)); private static TypeReference RequireType(ModuleDefinition module, string fullName, string assemblyName) => FindType(module, fullName, assemblyName) ?? throw new InvalidOperationException($"Required type {fullName} was not found in {assemblyName}."); private static TypeReference? FindType(ModuleDefinition module, string fullName, string assemblyName) { if (string.Equals(module.Assembly.Name.Name, assemblyName, StringComparison.Ordinal)) return module.GetType(fullName); var reference = module.AssemblyReferences.FirstOrDefault(item => string.Equals(item.Name, assemblyName, StringComparison.Ordinal)); if (reference == null) return null; var assembly = module.AssemblyResolver.Resolve(reference); return assembly?.MainModule.GetType(fullName) is { } definition ? module.ImportReference(definition) : null; } private static MethodReference ImportMethod(ModuleDefinition module, TypeReference type, Func predicate) { var definition = type.Resolve() ?? throw new InvalidOperationException($"Unable to resolve {type.FullName}."); var method = definition.Methods.SingleOrDefault(predicate) ?? throw new InvalidOperationException($"Required method was not found on {type.FullName}."); return module.ImportReference(method); } private static IEnumerable AllTypes(IEnumerable roots) { foreach (var type in roots) { yield return type; foreach (var nested in AllTypes(type.NestedTypes)) yield return nested; } } private static bool HasAttribute(ICustomAttributeProvider provider, TypeReference attribute) => HasAttribute(provider, attribute.FullName); private static bool HasAttribute(ICustomAttributeProvider provider, string fullName) => provider.CustomAttributes.Any(a => a.AttributeType.FullName == fullName); private static bool Implements(TypeReference type, string interfaceFullName) { try { TypeDefinition? current = type.Resolve(); while (current != null) { if (current.Interfaces.Any(item => item.InterfaceType.FullName == interfaceFullName)) return true; current = current.BaseType?.Resolve(); } } catch { } return false; } private static TypeReference GenericType(ModuleDefinition module, string name, bool valueType, params TypeReference[] arguments) { var open = new TypeReference("System", name.Substring(name.LastIndexOf('.') + 1), module, module.TypeSystem.CoreLibrary, valueType); var result = new GenericInstanceType(open); foreach (var argument in arguments) result.GenericArguments.Add(argument); return result; } private static MethodReference DelegateConstructor(ModuleDefinition module, TypeReference delegateType) { var ctor = new MethodReference(".ctor", module.TypeSystem.Void, delegateType) { HasThis = true }; ctor.Parameters.Add(new ParameterDefinition(module.TypeSystem.Object)); ctor.Parameters.Add(new ParameterDefinition(module.TypeSystem.IntPtr)); return ctor; } private static string ReadString(CustomAttribute attribute, string name) => attribute.Properties.FirstOrDefault(item => item.Name == name).Argument.Value as string ?? string.Empty; private static int ReadInt(CustomAttribute attribute, string name, int fallback) => attribute.Properties.FirstOrDefault(item => item.Name == name) is { Name: not null } item ? Convert.ToInt32(item.Argument.Value) : fallback; private static bool ReadBool(CustomAttribute attribute, string name, bool fallback) => attribute.Properties.FirstOrDefault(item => item.Name == name) is { Name: not null } item ? Convert.ToBoolean(item.Argument.Value) : fallback; } }