#nullable enable using System; using System.Collections.Generic; using System.IO; using System.Linq; using Mono.Cecil; using Mono.Cecil.Cil; using Mono.Cecil.Rocks; using Mono.Cecil.Pdb; using Unity.CompilationPipeline.Common.Diagnostics; using Unity.CompilationPipeline.Common.ILPostProcessing; namespace ShrinkEventBus.CodeGen { public sealed class EventBusILPostProcessor : ILPostProcessor { private const string RuntimeAssemblyName = "ShrinkEventBus.Runtime"; public override ILPostProcessor GetInstance() => this; public override bool WillProcess(ICompiledAssembly compiledAssembly) { if (!ReferencesAssembly(compiledAssembly, RuntimeAssemblyName)) return false; return true; } public override ILPostProcessResult Process(ICompiledAssembly compiledAssembly) { var diagnostics = new List(); if (!WillProcess(compiledAssembly)) return new ILPostProcessResult(compiledAssembly.InMemoryAssembly, diagnostics); var assemblyDefinition = AssemblyDefinitionFor(compiledAssembly); var module = assemblyDefinition.MainModule; try { var generatedSubscriberType = FindType(module, "ShrinkEventBus.ShrinkEventSubscriberAttribute", RuntimeAssemblyName); var generatedSubscribeType = FindType(module, "ShrinkEventBus.ShrinkSubscribeAttribute", RuntimeAssemblyName); if (generatedSubscriberType == null || generatedSubscribeType == null) return GetResult(assemblyDefinition, diagnostics); foreach (var type in GetAllTypes(module.Types) .Where(type => !type.IsInterface && !type.IsAbstract) .Where(type => HasAttribute(type, generatedSubscriberType)) .Where(type => type.Methods.Any(method => !method.IsStatic && HasAttribute(method, generatedSubscribeType)))) { InjectGeneratedBinding(type, module, generatedSubscribeType); } var staticSubscriberTypes = GetAllTypes(module.Types) .Where(type => HasAttribute(type, generatedSubscriberType)) .Where(type => type.Methods.Any(method => method.IsStatic && HasAttribute(method, generatedSubscribeType))) .ToArray(); if (staticSubscriberTypes.Length > 0) InjectStaticBootstrap(module, staticSubscriberTypes, generatedSubscribeType); } catch (Exception ex) { diagnostics.Add(new DiagnosticMessage { DiagnosticType = DiagnosticType.Error, MessageData = $"[ShrinkEventBus.CodeGen] {ex.Message}" }); } return GetResult(assemblyDefinition, diagnostics); } private static bool ReferencesAssembly(ICompiledAssembly compiledAssembly, string assemblyName) { return compiledAssembly.References.Any(reference => string.Equals(Path.GetFileNameWithoutExtension(reference), assemblyName, StringComparison.Ordinal)); } private static bool HasAttribute(ICustomAttributeProvider provider, TypeReference expectedAttributeType) { return provider.CustomAttributes.Any(attribute => attribute.AttributeType.FullName == expectedAttributeType.FullName); } private static IEnumerable GetAllTypes(IEnumerable roots) { foreach (var type in roots) { yield return type; foreach (var nested in GetAllTypes(type.NestedTypes)) yield return nested; } } private static void InjectGeneratedBinding(TypeDefinition type, ModuleDefinition module, TypeReference subscribeAttributeType) { var generatedInterface = FindType(module, "ShrinkEventBus.IShrinkGeneratedSubscriber", RuntimeAssemblyName) ?? throw new InvalidOperationException("IShrinkGeneratedSubscriber was not found."); if (type.Interfaces.Any(item => item.InterfaceType.FullName == generatedInterface.FullName)) return; var resolverType = FindType(module, "ShrinkEventBus.IShrinkBusResolver", RuntimeAssemblyName) ?? throw new InvalidOperationException("IShrinkBusResolver was not found."); var busKeyType = FindType(module, "ShrinkEventBus.ShrinkBusKey", RuntimeAssemblyName) ?? throw new InvalidOperationException("ShrinkBusKey was not found."); var bindingType = FindType(module, "ShrinkEventBus.ShrinkEventBinding", RuntimeAssemblyName) ?? throw new InvalidOperationException("ShrinkEventBinding was not found."); var bindingHelperType = FindType(module, "ShrinkEventBus.ShrinkGeneratedBinding", RuntimeAssemblyName) ?? throw new InvalidOperationException("ShrinkGeneratedBinding was not found."); var priorityType = FindType(module, "ShrinkEventBus.ShrinkEventPriority", RuntimeAssemblyName) ?? throw new InvalidOperationException("ShrinkEventPriority was not found."); var asyncHandlerType = FindType(module, "ShrinkEventBus.ShrinkAsyncEventHandler`1", RuntimeAssemblyName) ?? throw new InvalidOperationException("ShrinkAsyncEventHandler was not found."); var nullableBusKey = new GenericInstanceType(module.ImportReference(typeof(Nullable<>))); nullableBusKey.GenericArguments.Add(busKeyType); var disposableType = module.ImportReference(typeof(IDisposable)); var bindingDefinition = bindingType.Resolve() ?? throw new InvalidOperationException("ShrinkEventBinding could not be resolved."); var bindingCtor = module.ImportReference(bindingDefinition.Methods.Single(method => method.IsConstructor && method.Parameters.Count == 0)); var bindingAdd = module.ImportReference(bindingDefinition.Methods.Single(method => method.Name == "Add" && method.Parameters.Count == 1)); var helperDefinition = bindingHelperType.Resolve() ?? throw new InvalidOperationException("ShrinkGeneratedBinding could not be resolved."); var syncSubscribe = module.ImportReference(helperDefinition.Methods.Single(method => method.Name == "Subscribe" && method.HasGenericParameters)); var asyncSubscribe = module.ImportReference(helperDefinition.Methods.Single(method => method.Name == "SubscribeAsync" && method.HasGenericParameters)); var legacyAsyncSubscribe = module.ImportReference(helperDefinition.Methods.Single(method => method.Name == "SubscribeAsyncLegacy" && method.HasGenericParameters)); var interfaceDefinition = generatedInterface.Resolve() ?? throw new InvalidOperationException("IShrinkGeneratedSubscriber could not be resolved."); var interfaceMethod = module.ImportReference(interfaceDefinition.Methods.Single(method => method.Name == "AttachGenerated")); var generatedMethod = new MethodDefinition( "ShrinkEventBus.IShrinkGeneratedSubscriber.AttachGenerated", MethodAttributes.Private | MethodAttributes.Final | MethodAttributes.HideBySig | MethodAttributes.NewSlot | MethodAttributes.Virtual, disposableType); generatedMethod.Parameters.Add(new ParameterDefinition("resolver", ParameterAttributes.None, resolverType)); generatedMethod.Parameters.Add(new ParameterDefinition("defaultBus", ParameterAttributes.Optional, nullableBusKey)); generatedMethod.Overrides.Add(interfaceMethod); generatedMethod.Body.InitLocals = true; var bindingLocal = new VariableDefinition(bindingType); generatedMethod.Body.Variables.Add(bindingLocal); var il = generatedMethod.Body.GetILProcessor(); il.Emit(OpCodes.Newobj, bindingCtor); il.Emit(OpCodes.Stloc, bindingLocal); var classDefaultBus = ReadStringProperty( type.CustomAttributes.First(attribute => attribute.AttributeType.FullName == "ShrinkEventBus.ShrinkEventSubscriberAttribute"), "DefaultBus"); foreach (var handler in type.Methods.Where(method => !method.IsStatic && HasAttribute(method, subscribeAttributeType)).ToArray()) { EmitGeneratedSubscription(il, module, type, handler, handler.CustomAttributes.First(attribute => attribute.AttributeType.FullName == subscribeAttributeType.FullName), classDefaultBus, bindingLocal, bindingAdd, syncSubscribe, asyncSubscribe, legacyAsyncSubscribe, asyncHandlerType); } il.Emit(OpCodes.Ldloc, bindingLocal); il.Emit(OpCodes.Ret); type.Interfaces.Add(new InterfaceImplementation(generatedInterface)); type.Methods.Add(generatedMethod); } private static void EmitGeneratedSubscription(ILProcessor il, ModuleDefinition module, TypeDefinition ownerType, MethodDefinition handler, CustomAttribute attribute, string classDefaultBus, VariableDefinition bindingLocal, MethodReference bindingAdd, MethodReference syncSubscribe, MethodReference asyncSubscribe, MethodReference legacyAsyncSubscribe, TypeReference asyncHandlerType) { if (handler.Parameters.Count is < 1 or > 2) throw new InvalidOperationException( $"[ShrinkSubscribe] method {handler.FullName} must have one event parameter and an optional CancellationToken."); var eventType = module.ImportReference(handler.Parameters[0].ParameterType); var eventInterface = FindType(module, "ShrinkEventBus.IShrinkEvent", RuntimeAssemblyName) ?? throw new InvalidOperationException("IShrinkEvent was not found."); if (!ImplementsInterface(eventType, eventInterface.FullName)) throw new InvalidOperationException( $"[ShrinkSubscribe] method {handler.FullName} event parameter must implement IShrinkEvent."); var bus = ReadStringProperty(attribute, "Bus"); if (string.IsNullOrWhiteSpace(bus)) bus = classDefaultBus; var priority = ReadIntProperty(attribute, "Priority", 2); var numericPriority = ReadIntProperty(attribute, "NumericPriority", 0); var receiveCanceled = ReadBoolProperty(attribute, "ReceiveCanceled", false); MethodReference openSubscribe; TypeReference delegateType; if (handler.ReturnType.FullName == module.TypeSystem.Void.FullName && handler.Parameters.Count == 1) { openSubscribe = syncSubscribe; delegateType = MakeGenericType(module, typeof(Action<>), eventType); } else if (handler.ReturnType.FullName == "Cysharp.Threading.Tasks.UniTask" && handler.Parameters.Count == 1) { openSubscribe = legacyAsyncSubscribe; var uniTaskType = FindType(module, "Cysharp.Threading.Tasks.UniTask", "UniTask") ?? module.ImportReference(handler.ReturnType); delegateType = MakeGenericType(module, typeof(Func<,>), eventType, uniTaskType); } else if (handler.ReturnType.FullName == "Cysharp.Threading.Tasks.UniTask" && handler.Parameters.Count == 2 && handler.Parameters[1].ParameterType.FullName == typeof(System.Threading.CancellationToken).FullName) { openSubscribe = asyncSubscribe; var closedAsyncHandler = new GenericInstanceType(asyncHandlerType); closedAsyncHandler.GenericArguments.Add(eventType); delegateType = closedAsyncHandler; } else { throw new InvalidOperationException( $"Unsupported [ShrinkSubscribe] signature: {handler.FullName}. Use void(T), UniTask(T), or UniTask(T, CancellationToken)."); } var closedSubscribe = new GenericInstanceMethod(openSubscribe); closedSubscribe.GenericArguments.Add(eventType); var delegateCtor = MakeDelegateConstructor(module, delegateType); 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, delegateCtor); il.Emit(OpCodes.Ldc_I4, priority); il.Emit(OpCodes.Ldc_I4, numericPriority); il.Emit(receiveCanceled ? OpCodes.Ldc_I4_1 : OpCodes.Ldc_I4_0); il.Emit(OpCodes.Call, closedSubscribe); il.Emit(OpCodes.Callvirt, bindingAdd); } private static void InjectStaticBootstrap(ModuleDefinition module, IReadOnlyList subscriberTypes, TypeReference subscribeAttributeType) { var registryType = FindType(module, "ShrinkEventBus.ShrinkStaticBindingRegistry", RuntimeAssemblyName) ?? throw new InvalidOperationException("ShrinkStaticBindingRegistry was not found."); var asyncHandlerType = FindType(module, "ShrinkEventBus.ShrinkAsyncEventHandler`1", RuntimeAssemblyName) ?? throw new InvalidOperationException("ShrinkAsyncEventHandler was not found."); var registryDefinition = registryType.Resolve() ?? throw new InvalidOperationException("ShrinkStaticBindingRegistry could not be resolved."); var syncRegister = module.ImportReference(registryDefinition.Methods.Single(method => method.Name == "Register" && method.HasGenericParameters)); var asyncRegister = module.ImportReference(registryDefinition.Methods.Single(method => method.Name == "RegisterAsync" && method.HasGenericParameters)); var legacyAsyncRegister = module.ImportReference(registryDefinition.Methods.Single(method => method.Name == "RegisterAsyncLegacy" && method.HasGenericParameters)); 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 subscriberTypes) { var classDefaultBus = ReadStringProperty( type.CustomAttributes.First(attribute => attribute.AttributeType.FullName == "ShrinkEventBus.ShrinkEventSubscriberAttribute"), "DefaultBus"); foreach (var handler in type.Methods.Where(method => method.IsStatic && HasAttribute(method, subscribeAttributeType)).ToArray()) { EmitStaticGeneratedSubscription(il, module, handler, handler.CustomAttributes.First(attribute => attribute.AttributeType.FullName == subscribeAttributeType.FullName), classDefaultBus, syncRegister, asyncRegister, legacyAsyncRegister, asyncHandlerType); } } il.Emit(OpCodes.Ret); InjectModuleInitializer(module, register); } private static void EmitStaticGeneratedSubscription(ILProcessor il, ModuleDefinition module, MethodDefinition handler, CustomAttribute attribute, string classDefaultBus, MethodReference syncRegister, MethodReference asyncRegister, MethodReference legacyAsyncRegister, TypeReference asyncHandlerType) { if (handler.Parameters.Count is < 1 or > 2) throw new InvalidOperationException( $"[ShrinkSubscribe] method {handler.FullName} must have one event parameter and an optional CancellationToken."); var eventType = module.ImportReference(handler.Parameters[0].ParameterType); var eventInterface = FindType(module, "ShrinkEventBus.IShrinkEvent", RuntimeAssemblyName) ?? throw new InvalidOperationException("IShrinkEvent was not found."); if (!ImplementsInterface(eventType, eventInterface.FullName)) throw new InvalidOperationException( $"[ShrinkSubscribe] method {handler.FullName} event parameter must implement IShrinkEvent."); var bus = ReadStringProperty(attribute, "Bus"); if (string.IsNullOrWhiteSpace(bus)) bus = classDefaultBus; var priority = ReadIntProperty(attribute, "Priority", 2); var numericPriority = ReadIntProperty(attribute, "NumericPriority", 0); var receiveCanceled = ReadBoolProperty(attribute, "ReceiveCanceled", false); MethodReference openRegister; TypeReference delegateType; if (handler.ReturnType.FullName == module.TypeSystem.Void.FullName && handler.Parameters.Count == 1) { openRegister = syncRegister; delegateType = MakeGenericType(module, typeof(Action<>), eventType); } else if (handler.ReturnType.FullName == "Cysharp.Threading.Tasks.UniTask" && handler.Parameters.Count == 1) { openRegister = legacyAsyncRegister; var uniTaskType = FindType(module, "Cysharp.Threading.Tasks.UniTask", "UniTask") ?? module.ImportReference(handler.ReturnType); delegateType = MakeGenericType(module, typeof(Func<,>), eventType, uniTaskType); } else if (handler.ReturnType.FullName == "Cysharp.Threading.Tasks.UniTask" && handler.Parameters.Count == 2 && handler.Parameters[1].ParameterType.FullName == typeof(System.Threading.CancellationToken).FullName) { openRegister = asyncRegister; var closedAsyncHandler = new GenericInstanceType(asyncHandlerType); closedAsyncHandler.GenericArguments.Add(eventType); delegateType = closedAsyncHandler; } else { throw new InvalidOperationException( $"Unsupported static [ShrinkSubscribe] signature: {handler.FullName}."); } var closedRegister = new GenericInstanceMethod(openRegister); closedRegister.GenericArguments.Add(eventType); var delegateCtor = MakeDelegateConstructor(module, delegateType); var handlerBridge = CreateStaticHandlerBridge(module, handler); il.Emit(OpCodes.Ldstr, bus ?? string.Empty); il.Emit(OpCodes.Ldnull); il.Emit(OpCodes.Ldftn, module.ImportReference(handlerBridge)); il.Emit(OpCodes.Newobj, delegateCtor); il.Emit(OpCodes.Ldc_I4, priority); il.Emit(OpCodes.Ldc_I4, numericPriority); il.Emit(receiveCanceled ? OpCodes.Ldc_I4_1 : OpCodes.Ldc_I4_0); il.Emit(OpCodes.Call, closedRegister); } private static MethodDefinition CreateStaticHandlerBridge(ModuleDefinition module, MethodDefinition handler) { var bridge = new MethodDefinition( $"ShrinkEventBus.GeneratedStaticHandler_{handler.MetadataToken.ToInt32():X8}", MethodAttributes.Assembly | MethodAttributes.Static | MethodAttributes.HideBySig, module.ImportReference(handler.ReturnType)); for (var i = 0; i < handler.Parameters.Count; i++) { var parameter = handler.Parameters[i]; bridge.Parameters.Add(new ParameterDefinition(parameter.Name, parameter.Attributes, module.ImportReference(parameter.ParameterType))); } var il = bridge.Body.GetILProcessor(); for (var i = 0; i < bridge.Parameters.Count; i++) il.Emit(OpCodes.Ldarg, bridge.Parameters[i]); il.Emit(OpCodes.Call, module.ImportReference(handler)); il.Emit(OpCodes.Ret); handler.DeclaringType.Methods.Add(bridge); return bridge; } private static void InjectModuleInitializer(ModuleDefinition module, MethodReference register) { 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); } var processor = initializer.Body.GetILProcessor(); processor.InsertBefore(initializer.Body.Instructions[0], processor.Create(OpCodes.Call, register)); } private static GenericInstanceType MakeGenericType(ModuleDefinition module, Type openType, params TypeReference[] arguments) { var result = new GenericInstanceType(module.ImportReference(openType)); foreach (var argument in arguments) result.GenericArguments.Add(argument); return result; } private static MethodReference MakeDelegateConstructor(ModuleDefinition module, TypeReference delegateType) { var ctor = new MethodReference(".ctor", module.TypeSystem.Void, delegateType) { HasThis = true, CallingConvention = MethodCallingConvention.Default }; ctor.Parameters.Add(new ParameterDefinition(module.TypeSystem.Object)); ctor.Parameters.Add(new ParameterDefinition(module.TypeSystem.IntPtr)); return ctor; } private static bool ImplementsInterface(TypeReference type, string interfaceFullName) { TypeDefinition? current; try { current = type.Resolve(); } catch { return false; } while (current != null) { if (current.Interfaces.Any(item => item.InterfaceType.FullName == interfaceFullName)) return true; try { current = current.BaseType?.Resolve(); } catch { return false; } } return false; } private static string ReadStringProperty(CustomAttribute attribute, string name) { foreach (var property in attribute.Properties) { if (property.Name == name) return property.Argument.Value as string ?? string.Empty; } return string.Empty; } private static int ReadIntProperty(CustomAttribute attribute, string name, int defaultValue) { foreach (var property in attribute.Properties) { if (property.Name == name) return Convert.ToInt32(property.Argument.Value); } return defaultValue; } private static bool ReadBoolProperty(CustomAttribute attribute, string name, bool defaultValue) { foreach (var property in attribute.Properties) { if (property.Name == name) return Convert.ToBoolean(property.Argument.Value); } return defaultValue; } private static TypeReference? FindType(ModuleDefinition module, string fullName, string assemblyName) { var resolved = Type.GetType($"{fullName}, {assemblyName}", false); return resolved == null ? null : module.ImportReference(resolved); } private static AssemblyDefinition AssemblyDefinitionFor(ICompiledAssembly compiledAssembly) { var assemblyResolver = new PostProcessorAssemblyResolver(compiledAssembly); var readerParameters = new ReaderParameters { SymbolStream = new MemoryStream(compiledAssembly.InMemoryAssembly.PdbData.ToArray()), SymbolReaderProvider = new PdbReaderProvider(), AssemblyResolver = assemblyResolver, ReflectionImporterProvider = new PostProcessorReflectionImporterProvider(), ReadingMode = ReadingMode.Immediate }; var assemblyDefinition = AssemblyDefinition.ReadAssembly( new MemoryStream(compiledAssembly.InMemoryAssembly.PeData.ToArray()), readerParameters); assemblyResolver.AddAssemblyDefinitionBeingOperatedOn(assemblyDefinition); return assemblyDefinition; } private static ILPostProcessResult GetResult(AssemblyDefinition assemblyDefinition, List diagnostics) { var pe = new MemoryStream(); var pdb = new MemoryStream(); assemblyDefinition.Write(pe, new WriterParameters { SymbolWriterProvider = new PdbWriterProvider(), SymbolStream = pdb, WriteSymbols = true }); return new ILPostProcessResult(new InMemoryAssembly(pe.ToArray(), pdb.ToArray()), diagnostics); } } }