diff --git a/Hosting/NetCord.Hosting/Gateway/GatewayClientHostedService.cs b/Hosting/NetCord.Hosting/Gateway/GatewayClientHostedService.cs index e0388bf08..9c72e8352 100644 --- a/Hosting/NetCord.Hosting/Gateway/GatewayClientHostedService.cs +++ b/Hosting/NetCord.Hosting/Gateway/GatewayClientHostedService.cs @@ -12,12 +12,12 @@ public Task StartAsync(CancellationToken cancellationToken) { var client = services.GetRequiredService(); - foreach (var handler in services.GetServices()) + foreach (var handlerMetadata in services.GetServices()) { - if (handler is IDelegateGatewayHandlerBase delegateHandler) - RegisterDelegateHandler(client, delegateHandler); + if (handlerMetadata is ClassGatewayHandlerMetadata classHandlerMetadata) + RegisterClassHandler(services, client, classHandlerMetadata); else - RegisterClassHandler(client, handler); + RegisterDelegateHandler(services, client, (DelegateGatewayHandlerMetadata)handlerMetadata); } var options = services.GetRequiredService>().Value; diff --git a/Hosting/NetCord.Hosting/Gateway/GatewayEvent.cs b/Hosting/NetCord.Hosting/Gateway/GatewayEvent.cs index 93cc2cf1e..f0b600b57 100644 --- a/Hosting/NetCord.Hosting/Gateway/GatewayEvent.cs +++ b/Hosting/NetCord.Hosting/Gateway/GatewayEvent.cs @@ -2,20 +2,20 @@ namespace NetCord.Hosting.Gateway; public partial class GatewayEvent { - internal GatewayEvent(string name) + internal GatewayEvent(GatewayEventId id) { - Name = name; + Id = id; } - internal string Name { get; } + internal GatewayEventId Id { get; } } public class GatewayEvent { - internal GatewayEvent(string name) + internal GatewayEvent(GatewayEventId id) { - Name = name; + Id = id; } - internal string Name { get; } + internal GatewayEventId Id { get; } } diff --git a/Hosting/NetCord.Hosting/Gateway/GatewayHandlerServiceCollectionExtensions.cs b/Hosting/NetCord.Hosting/Gateway/GatewayHandlerServiceCollectionExtensions.cs index 095e22415..55d537eb2 100644 --- a/Hosting/NetCord.Hosting/Gateway/GatewayHandlerServiceCollectionExtensions.cs +++ b/Hosting/NetCord.Hosting/Gateway/GatewayHandlerServiceCollectionExtensions.cs @@ -2,6 +2,9 @@ using System.Reflection; using Microsoft.Extensions.DependencyInjection; +using Microsoft.Extensions.DependencyInjection.Extensions; + +using NetCord.Gateway; namespace NetCord.Hosting.Gateway; @@ -12,10 +15,13 @@ public static class GatewayHandlerServiceCollectionExtensions /// /// The type of the to add. /// The to add the to. + /// The of the . /// A reference to this instance after the operation has completed. - public static IServiceCollection AddGatewayHandler<[DAM(DAMT.PublicConstructors)] T>(this IServiceCollection services) where T : class, IGatewayHandler + public static IServiceCollection AddGatewayHandler<[DAM(DAMT.PublicConstructors)] T>(this IServiceCollection services, ServiceLifetime lifetime = ServiceLifetime.Singleton) where T : class, IGatewayHandler { - services.AddSingleton(); + services.TryAdd(ServiceDescriptor.Describe(typeof(T), typeof(T), lifetime)); + services.AddSingleton(new ClassGatewayHandlerMetadata(typeof(T), lifetime is ServiceLifetime.Singleton)); + return services; } @@ -25,10 +31,13 @@ public static class GatewayHandlerServiceCollectionExtensions /// The type of the to add. /// The to add the to. /// The factory that creates the . + /// The of the . /// A reference to this instance after the operation has completed. - public static IServiceCollection AddGatewayHandler(this IServiceCollection services, Func implementationFactory) where T : class, IGatewayHandler + public static IServiceCollection AddGatewayHandler(this IServiceCollection services, Func implementationFactory, ServiceLifetime lifetime = ServiceLifetime.Singleton) where T : class, IGatewayHandler { - services.AddSingleton(implementationFactory); + services.TryAdd(ServiceDescriptor.Describe(typeof(T), implementationFactory, lifetime)); + services.AddSingleton(new ClassGatewayHandlerMetadata(typeof(T), lifetime is ServiceLifetime.Singleton)); + return services; } @@ -37,10 +46,13 @@ public static IServiceCollection AddGatewayHandler(this IServiceCollection se /// /// The to add the to. /// The type of the to add. + /// The of the . /// A reference to this instance after the operation has completed. - public static IServiceCollection AddGatewayHandler(this IServiceCollection services, [DAM(DAMT.PublicConstructors)] Type handlerType) + public static IServiceCollection AddGatewayHandler(this IServiceCollection services, [DAM(DAMT.PublicConstructors)] Type handlerType, ServiceLifetime lifetime = ServiceLifetime.Singleton) { - services.AddSingleton(typeof(IGatewayHandler), handlerType); + services.TryAdd(ServiceDescriptor.Describe(handlerType, handlerType, lifetime)); + services.AddSingleton(new ClassGatewayHandlerMetadata(handlerType, lifetime is ServiceLifetime.Singleton)); + return services; } @@ -50,10 +62,15 @@ public static IServiceCollection AddGatewayHandler(this IServiceCollection servi /// The to add the to. /// The gateway event. /// The delegate that represents the handler. + /// The of the . /// A reference to this instance after the operation has completed. - public static IServiceCollection AddGatewayHandler(this IServiceCollection services, GatewayEvent gatewayEvent, Delegate handler) + public static IServiceCollection AddGatewayHandler(this IServiceCollection services, GatewayEvent gatewayEvent, Delegate handler, ServiceLifetime lifetime = ServiceLifetime.Singleton) { - services.AddSingleton(services => new DelegateGatewayHandler(gatewayEvent.Name, services, handler)); + services.AddSingleton(new DelegateGatewayHandlerMetadata( + DelegateHandlerHelper.CreateHandler>(handler, []), + gatewayEvent.Id, + lifetime is ServiceLifetime.Singleton)); + return services; } @@ -64,10 +81,15 @@ public static IServiceCollection AddGatewayHandler(this IServiceCollection servi /// The to add the to. /// The gateway event. /// The delegate that represents the handler. + /// The of the . /// A reference to this instance after the operation has completed. - public static IServiceCollection AddGatewayHandler(this IServiceCollection services, GatewayEvent gatewayEvent, Delegate handler) + public static IServiceCollection AddGatewayHandler(this IServiceCollection services, GatewayEvent gatewayEvent, Delegate handler, ServiceLifetime lifetime = ServiceLifetime.Singleton) { - services.AddSingleton(services => new DelegateGatewayHandler(gatewayEvent.Name, services, handler)); + services.AddSingleton(new DelegateGatewayHandlerMetadata( + DelegateHandlerHelper.CreateHandler>(handler, [typeof(T)]), + gatewayEvent.Id, + lifetime is ServiceLifetime.Singleton)); + return services; } @@ -76,11 +98,14 @@ public static IServiceCollection AddGatewayHandler(this IServiceCollection se /// /// The to add the implementations to. /// The assembly to scan for implementations. + /// The of the implementations. /// A reference to this instance after the operation has completed. [RequiresUnreferencedCode("Types might be removed")] - public static IServiceCollection AddGatewayHandlers(this IServiceCollection services, Assembly assembly) + public static IServiceCollection AddGatewayHandlers(this IServiceCollection services, Assembly assembly, ServiceLifetime lifetime = ServiceLifetime.Singleton) { - AddGatewayHandlers(services, typeof(IGatewayHandler), assembly); + foreach (var type in HandlerHelpers.GetHandlers(typeof(IGatewayHandler), assembly)) + AddGatewayHandler(services, type, lifetime); + return services; } @@ -89,10 +114,13 @@ public static IServiceCollection AddGatewayHandlers(this IServiceCollection serv /// /// The type of the to add. /// The to add the to. + /// The of the . /// A reference to this instance after the operation has completed. - public static IServiceCollection AddShardedGatewayHandler<[DAM(DAMT.PublicConstructors)] T>(this IServiceCollection services) where T : class, IShardedGatewayHandler + public static IServiceCollection AddShardedGatewayHandler<[DAM(DAMT.PublicConstructors)] T>(this IServiceCollection services, ServiceLifetime lifetime = ServiceLifetime.Singleton) where T : class, IShardedGatewayHandler { - services.AddSingleton(); + services.TryAdd(ServiceDescriptor.Describe(typeof(T), typeof(T), lifetime)); + services.AddSingleton(new ClassShardedGatewayHandlerMetadata(typeof(T), lifetime is ServiceLifetime.Singleton)); + return services; } @@ -102,10 +130,13 @@ public static IServiceCollection AddGatewayHandlers(this IServiceCollection serv /// The type of the to add. /// The to add the to. /// The factory that creates the . + /// The of the . /// A reference to this instance after the operation has completed. - public static IServiceCollection AddShardedGatewayHandler(this IServiceCollection services, Func implementationFactory) where T : class, IShardedGatewayHandler + public static IServiceCollection AddShardedGatewayHandler(this IServiceCollection services, Func implementationFactory, ServiceLifetime lifetime = ServiceLifetime.Singleton) where T : class, IShardedGatewayHandler { - services.AddSingleton(implementationFactory); + services.TryAdd(ServiceDescriptor.Describe(typeof(T), implementationFactory, lifetime)); + services.AddSingleton(new ClassShardedGatewayHandlerMetadata(typeof(T), lifetime is ServiceLifetime.Singleton)); + return services; } @@ -114,10 +145,13 @@ public static IServiceCollection AddShardedGatewayHandler(this IServiceCollec /// /// The to add the to. /// The type of the to add. + /// The of the . /// A reference to this instance after the operation has completed. - public static IServiceCollection AddShardedGatewayHandler(this IServiceCollection services, [DAM(DAMT.PublicConstructors)] Type handlerType) + public static IServiceCollection AddShardedGatewayHandler(this IServiceCollection services, [DAM(DAMT.PublicConstructors)] Type handlerType, ServiceLifetime lifetime = ServiceLifetime.Singleton) { - services.AddSingleton(typeof(IShardedGatewayHandler), handlerType); + services.TryAdd(ServiceDescriptor.Describe(handlerType, handlerType, lifetime)); + services.AddSingleton(new ClassGatewayHandlerMetadata(handlerType, lifetime is ServiceLifetime.Singleton)); + return services; } @@ -127,10 +161,15 @@ public static IServiceCollection AddShardedGatewayHandler(this IServiceCollectio /// The to add the to. /// The gateway event. /// The delegate that represents the handler. + /// The of the . /// A reference to this instance after the operation has completed. - public static IServiceCollection AddShardedGatewayHandler(this IServiceCollection services, GatewayEvent gatewayEvent, Delegate handler) + public static IServiceCollection AddShardedGatewayHandler(this IServiceCollection services, GatewayEvent gatewayEvent, Delegate handler, ServiceLifetime lifetime = ServiceLifetime.Singleton) { - services.AddSingleton(services => new DelegateShardedGatewayHandler(gatewayEvent.Name, services, handler)); + services.AddSingleton(new DelegateShardedGatewayHandlerMetadata( + DelegateHandlerHelper.CreateHandler>(handler, [typeof(GatewayClient)]), + gatewayEvent.Id, + lifetime is ServiceLifetime.Singleton)); + return services; } @@ -141,10 +180,15 @@ public static IServiceCollection AddShardedGatewayHandler(this IServiceCollectio /// The to add the to. /// The gateway event. /// The delegate that represents the handler. + /// The of the . /// A reference to this instance after the operation has completed. - public static IServiceCollection AddShardedGatewayHandler(this IServiceCollection services, GatewayEvent gatewayEvent, Delegate handler) + public static IServiceCollection AddShardedGatewayHandler(this IServiceCollection services, GatewayEvent gatewayEvent, Delegate handler, ServiceLifetime lifetime = ServiceLifetime.Singleton) { - services.AddSingleton(services => new DelegateShardedGatewayHandler(gatewayEvent.Name, services, handler)); + services.AddSingleton(new DelegateShardedGatewayHandlerMetadata( + DelegateHandlerHelper.CreateHandler>(handler, [typeof(GatewayClient), typeof(T)]), + gatewayEvent.Id, + lifetime is ServiceLifetime.Singleton)); + return services; } @@ -153,18 +197,14 @@ public static IServiceCollection AddShardedGatewayHandler(this IServiceCollec /// /// The to add the implementations to. /// The assembly to scan for implementations. + /// The of the implementations. /// A reference to this instance after the operation has completed. [RequiresUnreferencedCode("Types might be removed")] - public static IServiceCollection AddShardedGatewayHandlers(this IServiceCollection services, Assembly assembly) + public static IServiceCollection AddShardedGatewayHandlers(this IServiceCollection services, Assembly assembly, ServiceLifetime lifetime = ServiceLifetime.Singleton) { - AddGatewayHandlers(services, typeof(IShardedGatewayHandler), assembly); - return services; - } + foreach (var type in HandlerHelpers.GetHandlers(typeof(IShardedGatewayHandler), assembly)) + AddShardedGatewayHandler(services, type, lifetime); - [RequiresUnreferencedCode("Types might be removed")] - private static void AddGatewayHandlers(IServiceCollection services, Type handlerBase, Assembly assembly) - { - foreach (var type in HandlerHelpers.GetHandlers(handlerBase, assembly)) - services.AddSingleton(handlerBase, type); + return services; } } diff --git a/Hosting/NetCord.Hosting/Gateway/GatewayHandlers.cs b/Hosting/NetCord.Hosting/Gateway/GatewayHandlers.cs index 98672574a..1e7995f78 100644 --- a/Hosting/NetCord.Hosting/Gateway/GatewayHandlers.cs +++ b/Hosting/NetCord.Hosting/Gateway/GatewayHandlers.cs @@ -1,73 +1,39 @@ -using NetCord.Gateway; - namespace NetCord.Hosting.Gateway; public interface IGatewayHandler; -internal interface IDelegateGatewayHandlerBase : IGatewayHandler -{ - internal string Name { get; } -} +public interface IShardedGatewayHandler; -internal interface IDelegateGatewayHandler : IDelegateGatewayHandlerBase +internal abstract class GatewayHandlerMetadata(bool isSingleton) { - public ValueTask HandleAsync(); + public bool IsSingleton => isSingleton; } -internal interface IDelegateGatewayHandler : IDelegateGatewayHandlerBase +internal sealed class ClassGatewayHandlerMetadata(Type handlerType, bool isSingleton) : GatewayHandlerMetadata(isSingleton) { - public ValueTask HandleAsync(T arg); + public Type HandlerType => handlerType; } -internal class DelegateGatewayHandler(string name, IServiceProvider services, Delegate handler) : IDelegateGatewayHandler +internal sealed class DelegateGatewayHandlerMetadata(Delegate handler, GatewayEventId eventId, bool isSingleton) : GatewayHandlerMetadata(isSingleton) { - private readonly Func _handler = DelegateHandlerHelper.CreateHandler>(handler, []); - - string IDelegateGatewayHandlerBase.Name => name; + public Delegate Handler => handler; - public ValueTask HandleAsync() => _handler(services); + public GatewayEventId EventId => eventId; } -internal class DelegateGatewayHandler(string name, IServiceProvider services, Delegate handler) : IDelegateGatewayHandler +internal abstract class ShardedGatewayHandlerMetadata(bool isSingleton) { - private readonly Func _handler = DelegateHandlerHelper.CreateHandler>(handler, [typeof(T)]); - - string IDelegateGatewayHandlerBase.Name => name; - - public ValueTask HandleAsync(T arg) => _handler(arg, services); + public bool IsSingleton => isSingleton; } -public interface IShardedGatewayHandler; - -internal interface IDelegateShardedGatewayHandlerBase : IShardedGatewayHandler +internal sealed class ClassShardedGatewayHandlerMetadata(Type handlerType, bool isSingleton) : ShardedGatewayHandlerMetadata(isSingleton) { - internal string Name { get; } + public Type HandlerType => handlerType; } -internal interface IDelegateShardedGatewayHandler : IDelegateShardedGatewayHandlerBase +internal sealed class DelegateShardedGatewayHandlerMetadata(Delegate handler, GatewayEventId eventId, bool isSingleton) : ShardedGatewayHandlerMetadata(isSingleton) { - public ValueTask HandleAsync(GatewayClient client); -} - -internal interface IDelegateShardedGatewayHandler : IDelegateShardedGatewayHandlerBase -{ - public ValueTask HandleAsync(GatewayClient client, T arg); -} - -internal class DelegateShardedGatewayHandler(string name, IServiceProvider services, Delegate handler) : IDelegateShardedGatewayHandler -{ - private readonly Func _handler = DelegateHandlerHelper.CreateHandler>(handler, [typeof(GatewayClient)]); - - string IDelegateShardedGatewayHandlerBase.Name => name; - - public ValueTask HandleAsync(GatewayClient client) => _handler(client, services); -} - -internal class DelegateShardedGatewayHandler(string name, IServiceProvider services, Delegate handler) : IDelegateShardedGatewayHandler -{ - private readonly Func _handler = DelegateHandlerHelper.CreateHandler>(handler, [typeof(GatewayClient), typeof(T)]); - - string IDelegateShardedGatewayHandlerBase.Name => name; + public Delegate Handler => handler; - public ValueTask HandleAsync(GatewayClient client, T arg) => _handler(client, arg, services); + public GatewayEventId EventId => eventId; } diff --git a/Hosting/NetCord.Hosting/Gateway/ShardedGatewayClientHostedService.cs b/Hosting/NetCord.Hosting/Gateway/ShardedGatewayClientHostedService.cs index f5ebd4cec..5c71d0e09 100644 --- a/Hosting/NetCord.Hosting/Gateway/ShardedGatewayClientHostedService.cs +++ b/Hosting/NetCord.Hosting/Gateway/ShardedGatewayClientHostedService.cs @@ -12,12 +12,12 @@ public Task StartAsync(CancellationToken cancellationToken) { var client = services.GetRequiredService(); - foreach (var handler in services.GetServices()) + foreach (var handlerMetadata in services.GetServices()) { - if (handler is IDelegateShardedGatewayHandlerBase delegateHandler) - RegisterDelegateShardedHandler(client, delegateHandler); + if (handlerMetadata is ClassShardedGatewayHandlerMetadata classHandlerMetadata) + RegisterClassShardedHandler(services, client, classHandlerMetadata); else - RegisterClassShardedHandler(client, handler); + RegisterDelegateShardedHandler(services, client, (DelegateShardedGatewayHandlerMetadata)handlerMetadata); } var options = services.GetRequiredService>().Value; diff --git a/SourceGenerators/HostingGatewayEventsGenerator/HostingGatewayEventsGenerator.cs b/SourceGenerators/HostingGatewayEventsGenerator/HostingGatewayEventsGenerator.cs index b9fdfe818..d4796eba5 100644 --- a/SourceGenerators/HostingGatewayEventsGenerator/HostingGatewayEventsGenerator.cs +++ b/SourceGenerators/HostingGatewayEventsGenerator/HostingGatewayEventsGenerator.cs @@ -52,11 +52,31 @@ private string GenerateEvents(IEventSymbol[] events) StringWriter stringWriter = new(); Setup(stringWriter); + WriteEventsEnum(stringWriter, events); + WriteEvents(stringWriter, events); return stringWriter.ToString(); } + private void WriteEventsEnum(StringWriter stringWriter, IEventSymbol[] events) + { + stringWriter.WriteLine(); + + stringWriter.WriteLine("internal enum GatewayEventId : byte"); + + stringWriter.WriteLine("{"); + + foreach (var eventSymbol in events) + { + stringWriter.WriteIndentation(1); + stringWriter.Write(eventSymbol.Name); + stringWriter.WriteLine(","); + } + + stringWriter.WriteLine("}"); + } + private void WriteEvents(StringWriter stringWriter, IEventSymbol[] events) { stringWriter.WriteLine(); @@ -85,9 +105,9 @@ private void WriteEvents(StringWriter stringWriter, IEventSymbol[] events) stringWriter.Write(" "); stringWriter.Write(eventSymbol.Name); - stringWriter.Write(" => new(\""); + stringWriter.Write(" => new(global::NetCord.Hosting.Gateway.GatewayEventId."); stringWriter.Write(eventSymbol.Name); - stringWriter.WriteLine("\");"); + stringWriter.WriteLine(");"); } stringWriter.WriteLine("}"); @@ -242,73 +262,98 @@ private void WriteRegisterDelegateHandlerMethod(StringWriter stringWriter, IEven stringWriter.WriteLine(); stringWriter.WriteIndentation(1); - stringWriter.WriteLine("private static void RegisterDelegateHandler(global::NetCord.Gateway.GatewayClient client, global::NetCord.Hosting.Gateway.IDelegateGatewayHandlerBase handler)"); + stringWriter.WriteLine("private static void RegisterDelegateHandler(IServiceProvider services, global::NetCord.Gateway.GatewayClient client, global::NetCord.Hosting.Gateway.DelegateGatewayHandlerMetadata handlerMetadata)"); stringWriter.WriteIndentation(1); - stringWriter.Write("{"); + stringWriter.WriteLine("{"); + + stringWriter.WriteIndentation(2); + stringWriter.WriteLine("var isSingleton = handlerMetadata.IsSingleton;"); + + stringWriter.WriteIndentation(2); + stringWriter.WriteLine("switch (handlerMetadata.EventId)"); + + stringWriter.WriteIndentation(2); + stringWriter.WriteLine("{"); - int i = 0; - foreach (var group in events.GroupBy(e => e.Type, SymbolEqualityComparer.Default)) + int eventsLength = events.Length; + for (int i = 0; i < eventsLength; i++) { - var eventType = (INamedTypeSymbol)group.Key!; + var eventSymbol = events[i]; - stringWriter.WriteLine(); + stringWriter.WriteIndentation(3); + stringWriter.Write("case global::NetCord.Hosting.Gateway.GatewayEventId."); + stringWriter.Write(eventSymbol.Name); + stringWriter.WriteLine(":"); - stringWriter.WriteIndentation(2); - stringWriter.Write("if (handler is global::NetCord.Hosting.Gateway.IDelegateGatewayHandler"); + stringWriter.WriteIndentation(4); + stringWriter.Write("var typedHandler"); + stringWriter.Write(i); + stringWriter.Write(" = (global::System.Func<"); - if (eventType.Arity is 2) + var eventType = (INamedTypeSymbol)eventSymbol.Type; + if (eventType.Arity is not 1) { - stringWriter.Write("<"); stringWriter.Write(eventType.TypeArguments[0].ToDisplayString(SymbolDisplayFormat.FullyQualifiedFormat)); - stringWriter.Write(">"); + stringWriter.Write(", "); } - stringWriter.Write(" delegateGatewayHandler"); - stringWriter.Write(i); - stringWriter.WriteLine(")"); + stringWriter.WriteLine("global::System.IServiceProvider, global::System.Threading.Tasks.ValueTask>)handlerMetadata.Handler;"); - stringWriter.WriteIndentation(2); - stringWriter.WriteLine("{"); + stringWriter.WriteIndentation(4); + stringWriter.Write("client."); + stringWriter.Write(eventSymbol.Name); + stringWriter.WriteLine(" += isSingleton"); - stringWriter.WriteIndentation(3); - stringWriter.Write("switch (delegateGatewayHandler"); + stringWriter.WriteIndentation(5); + stringWriter.Write("? ("); + + if (eventType.Arity is not 1) + stringWriter.Write("arg"); + + stringWriter.Write(") => typedHandler"); stringWriter.Write(i); - stringWriter.WriteLine(".Name)"); + stringWriter.Write("("); - stringWriter.WriteIndentation(3); + if (eventType.Arity is not 1) + stringWriter.Write("arg, "); + + stringWriter.WriteLine("services)"); + + stringWriter.WriteIndentation(5); + stringWriter.Write(": async ("); + + if (eventType.Arity is not 1) + stringWriter.Write("arg"); + + stringWriter.WriteLine(") =>"); + + stringWriter.WriteIndentation(5); stringWriter.WriteLine("{"); - foreach (var eventSymbol in group) - { - stringWriter.WriteIndentation(4); - stringWriter.Write("case \""); - stringWriter.Write(eventSymbol.Name); - stringWriter.WriteLine("\":"); - - stringWriter.WriteIndentation(5); - stringWriter.Write("client."); - stringWriter.Write(eventSymbol.Name); - stringWriter.Write(" += delegateGatewayHandler"); - stringWriter.Write(i); - stringWriter.WriteLine(".HandleAsync;"); - - stringWriter.WriteIndentation(5); - stringWriter.WriteLine("break;"); - } + stringWriter.WriteIndentation(6); + stringWriter.WriteLine("await using var scope = global::Microsoft.Extensions.DependencyInjection.ServiceProviderServiceExtensions.CreateAsyncScope(services);"); - stringWriter.WriteIndentation(3); - stringWriter.WriteLine("}"); + stringWriter.WriteIndentation(6); + stringWriter.Write("await typedHandler"); + stringWriter.Write(i); + stringWriter.Write("("); - stringWriter.WriteIndentation(3); - stringWriter.WriteLine("return;"); + if (eventType.Arity is not 1) + stringWriter.Write("arg, "); - stringWriter.WriteIndentation(2); - stringWriter.WriteLine("}"); + stringWriter.WriteLine("scope.ServiceProvider);"); + + stringWriter.WriteIndentation(5); + stringWriter.WriteLine("};"); - i++; + stringWriter.WriteIndentation(4); + stringWriter.WriteLine("break;"); } + stringWriter.WriteIndentation(2); + stringWriter.WriteLine("}"); + stringWriter.WriteIndentation(1); stringWriter.WriteLine("}"); } @@ -318,10 +363,16 @@ private void WriteRegisterClassHandlerMethod(StringWriter stringWriter, IEventSy stringWriter.WriteLine(); stringWriter.WriteIndentation(1); - stringWriter.WriteLine("private static void RegisterClassHandler(global::NetCord.Gateway.GatewayClient client, global::NetCord.Hosting.Gateway.IGatewayHandler handler)"); + stringWriter.WriteLine("private static void RegisterClassHandler(global::System.IServiceProvider services, global::NetCord.Gateway.GatewayClient client, global::NetCord.Hosting.Gateway.ClassGatewayHandlerMetadata handlerMetadata)"); stringWriter.WriteIndentation(1); - stringWriter.Write("{"); + stringWriter.WriteLine("{"); + + stringWriter.WriteIndentation(2); + stringWriter.WriteLine("var handlerType = handlerMetadata.HandlerType;"); + + stringWriter.WriteIndentation(2); + stringWriter.WriteLine("var isSingleton = handlerMetadata.IsSingleton;"); int eventsLength = events.Length; @@ -329,21 +380,57 @@ private void WriteRegisterClassHandlerMethod(StringWriter stringWriter, IEventSy { var eventSymbol = events[i]; + var eventType = (INamedTypeSymbol)eventSymbol.Type; + stringWriter.WriteLine(); stringWriter.WriteIndentation(2); - stringWriter.Write("if (handler is global::NetCord.Hosting.Gateway.I"); + stringWriter.Write("if (typeof(global::NetCord.Hosting.Gateway.I"); stringWriter.Write(eventSymbol.Name); - stringWriter.Write("GatewayHandler gatewayHandler"); - stringWriter.Write(i); - stringWriter.WriteLine(")"); + stringWriter.WriteLine("GatewayHandler).IsAssignableFrom(handlerType))"); + + stringWriter.WriteIndentation(2); + stringWriter.WriteLine("{"); stringWriter.WriteIndentation(3); stringWriter.Write("client."); stringWriter.Write(eventSymbol.Name); - stringWriter.Write(" += gatewayHandler"); - stringWriter.Write(i); - stringWriter.WriteLine(".HandleAsync;"); + stringWriter.WriteLine(" += isSingleton"); + + stringWriter.WriteIndentation(4); + stringWriter.Write("? ((global::NetCord.Hosting.Gateway.I"); + stringWriter.Write(eventSymbol.Name); + stringWriter.WriteLine("GatewayHandler)global::Microsoft.Extensions.DependencyInjection.ServiceProviderServiceExtensions.GetRequiredService(services, handlerMetadata.HandlerType)).HandleAsync"); + + stringWriter.WriteIndentation(4); + stringWriter.Write(": async ("); + + if (eventType.Arity is not 1) + stringWriter.Write("arg"); + + stringWriter.WriteLine(") =>"); + + stringWriter.WriteIndentation(4); + stringWriter.WriteLine("{"); + + stringWriter.WriteIndentation(5); + stringWriter.WriteLine("await using var scope = global::Microsoft.Extensions.DependencyInjection.ServiceProviderServiceExtensions.CreateAsyncScope(services);"); + + stringWriter.WriteIndentation(5); + stringWriter.Write("await ((global::NetCord.Hosting.Gateway.I"); + stringWriter.Write(eventSymbol.Name); + stringWriter.Write("GatewayHandler)global::Microsoft.Extensions.DependencyInjection.ServiceProviderServiceExtensions.GetRequiredService(scope.ServiceProvider, handlerMetadata.HandlerType)).HandleAsync("); + + if (eventType.Arity is not 1) + stringWriter.Write("arg"); + + stringWriter.WriteLine(");"); + + stringWriter.WriteIndentation(4); + stringWriter.WriteLine("};"); + + stringWriter.WriteIndentation(2); + stringWriter.WriteLine("}"); } stringWriter.WriteIndentation(1); @@ -355,73 +442,98 @@ private void WriteRegisterDelegateShardedHandlerMethod(StringWriter stringWriter stringWriter.WriteLine(); stringWriter.WriteIndentation(1); - stringWriter.WriteLine("private static void RegisterDelegateShardedHandler(global::NetCord.Gateway.ShardedGatewayClient client, global::NetCord.Hosting.Gateway.IDelegateShardedGatewayHandlerBase handler)"); + stringWriter.WriteLine("private static void RegisterDelegateShardedHandler(IServiceProvider services, global::NetCord.Gateway.ShardedGatewayClient client, global::NetCord.Hosting.Gateway.DelegateShardedGatewayHandlerMetadata handlerMetadata)"); stringWriter.WriteIndentation(1); - stringWriter.Write("{"); + stringWriter.WriteLine("{"); + + stringWriter.WriteIndentation(2); + stringWriter.WriteLine("var isSingleton = handlerMetadata.IsSingleton;"); - int i = 0; - foreach (var group in events.GroupBy(e => e.Type, SymbolEqualityComparer.Default)) + stringWriter.WriteIndentation(2); + stringWriter.WriteLine("switch (handlerMetadata.EventId)"); + + stringWriter.WriteIndentation(2); + stringWriter.WriteLine("{"); + + int eventsLength = events.Length; + for (int i = 0; i < eventsLength; i++) { - var eventType = (INamedTypeSymbol)group.Key!; + var eventSymbol = events[i]; - stringWriter.WriteLine(); + stringWriter.WriteIndentation(3); + stringWriter.Write("case global::NetCord.Hosting.Gateway.GatewayEventId."); + stringWriter.Write(eventSymbol.Name); + stringWriter.WriteLine(":"); - stringWriter.WriteIndentation(2); - stringWriter.Write("if (handler is global::NetCord.Hosting.Gateway.IDelegateShardedGatewayHandler"); + stringWriter.WriteIndentation(4); + stringWriter.Write("var typedHandler"); + stringWriter.Write(i); + stringWriter.Write(" = (global::System.Func"); + stringWriter.Write(", "); } - stringWriter.Write(" delegateGatewayHandler"); - stringWriter.Write(i); - stringWriter.WriteLine(")"); + stringWriter.WriteLine("global::System.IServiceProvider, global::System.Threading.Tasks.ValueTask>)handlerMetadata.Handler;"); - stringWriter.WriteIndentation(2); - stringWriter.WriteLine("{"); + stringWriter.WriteIndentation(4); + stringWriter.Write("client."); + stringWriter.Write(eventSymbol.Name); + stringWriter.WriteLine(" += isSingleton"); - stringWriter.WriteIndentation(3); - stringWriter.Write("switch (delegateGatewayHandler"); + stringWriter.WriteIndentation(5); + stringWriter.Write("? (client"); + + if (eventType.Arity is not 1) + stringWriter.Write(", arg"); + + stringWriter.Write(") => typedHandler"); stringWriter.Write(i); - stringWriter.WriteLine(".Name)"); + stringWriter.Write("(client, "); - stringWriter.WriteIndentation(3); + if (eventType.Arity is not 1) + stringWriter.Write("arg, "); + + stringWriter.WriteLine("services)"); + + stringWriter.WriteIndentation(5); + stringWriter.Write(": async (client"); + + if (eventType.Arity is not 1) + stringWriter.Write(", arg"); + + stringWriter.WriteLine(") =>"); + + stringWriter.WriteIndentation(5); stringWriter.WriteLine("{"); - foreach (var eventSymbol in group) - { - stringWriter.WriteIndentation(4); - stringWriter.Write("case \""); - stringWriter.Write(eventSymbol.Name); - stringWriter.WriteLine("\":"); - - stringWriter.WriteIndentation(5); - stringWriter.Write("client."); - stringWriter.Write(eventSymbol.Name); - stringWriter.Write(" += delegateGatewayHandler"); - stringWriter.Write(i); - stringWriter.WriteLine(".HandleAsync;"); - - stringWriter.WriteIndentation(5); - stringWriter.WriteLine("break;"); - } + stringWriter.WriteIndentation(6); + stringWriter.WriteLine("await using var scope = global::Microsoft.Extensions.DependencyInjection.ServiceProviderServiceExtensions.CreateAsyncScope(services);"); - stringWriter.WriteIndentation(3); - stringWriter.WriteLine("}"); + stringWriter.WriteIndentation(6); + stringWriter.Write("await typedHandler"); + stringWriter.Write(i); + stringWriter.Write("(client, "); - stringWriter.WriteIndentation(3); - stringWriter.WriteLine("return;"); + if (eventType.Arity is not 1) + stringWriter.Write("arg, "); - stringWriter.WriteIndentation(2); - stringWriter.WriteLine("}"); + stringWriter.WriteLine("scope.ServiceProvider);"); - i++; + stringWriter.WriteIndentation(5); + stringWriter.WriteLine("};"); + + stringWriter.WriteIndentation(4); + stringWriter.WriteLine("break;"); } + stringWriter.WriteIndentation(2); + stringWriter.WriteLine("}"); + stringWriter.WriteIndentation(1); stringWriter.WriteLine("}"); } @@ -431,10 +543,16 @@ private void WriteRegisterClassShardedHandlerMethod(StringWriter stringWriter, I stringWriter.WriteLine(); stringWriter.WriteIndentation(1); - stringWriter.WriteLine("private static void RegisterClassShardedHandler(global::NetCord.Gateway.ShardedGatewayClient client, global::NetCord.Hosting.Gateway.IShardedGatewayHandler handler)"); + stringWriter.WriteLine("private static void RegisterClassShardedHandler(global::System.IServiceProvider services, global::NetCord.Gateway.ShardedGatewayClient client, global::NetCord.Hosting.Gateway.ClassShardedGatewayHandlerMetadata handlerMetadata)"); stringWriter.WriteIndentation(1); - stringWriter.Write("{"); + stringWriter.WriteLine("{"); + + stringWriter.WriteIndentation(2); + stringWriter.WriteLine("var handlerType = handlerMetadata.HandlerType;"); + + stringWriter.WriteIndentation(2); + stringWriter.WriteLine("var isSingleton = handlerMetadata.IsSingleton;"); int eventsLength = events.Length; @@ -442,21 +560,57 @@ private void WriteRegisterClassShardedHandlerMethod(StringWriter stringWriter, I { var eventSymbol = events[i]; + var eventType = (INamedTypeSymbol)eventSymbol.Type; + stringWriter.WriteLine(); stringWriter.WriteIndentation(2); - stringWriter.Write("if (handler is global::NetCord.Hosting.Gateway.I"); + stringWriter.Write("if (typeof(global::NetCord.Hosting.Gateway.I"); stringWriter.Write(eventSymbol.Name); - stringWriter.Write("ShardedGatewayHandler gatewayHandler"); - stringWriter.Write(i); - stringWriter.WriteLine(")"); + stringWriter.WriteLine("ShardedGatewayHandler).IsAssignableFrom(handlerType))"); + + stringWriter.WriteIndentation(2); + stringWriter.WriteLine("{"); stringWriter.WriteIndentation(3); stringWriter.Write("client."); stringWriter.Write(eventSymbol.Name); - stringWriter.Write(" += gatewayHandler"); - stringWriter.Write(i); - stringWriter.WriteLine(".HandleAsync;"); + stringWriter.WriteLine(" += isSingleton"); + + stringWriter.WriteIndentation(4); + stringWriter.Write("? ((global::NetCord.Hosting.Gateway.I"); + stringWriter.Write(eventSymbol.Name); + stringWriter.WriteLine("ShardedGatewayHandler)global::Microsoft.Extensions.DependencyInjection.ServiceProviderServiceExtensions.GetRequiredService(services, handlerMetadata.HandlerType)).HandleAsync"); + + stringWriter.WriteIndentation(4); + stringWriter.Write(": async (client"); + + if (eventType.Arity is not 1) + stringWriter.Write(", arg"); + + stringWriter.WriteLine(") =>"); + + stringWriter.WriteIndentation(4); + stringWriter.WriteLine("{"); + + stringWriter.WriteIndentation(5); + stringWriter.WriteLine("await using var scope = global::Microsoft.Extensions.DependencyInjection.ServiceProviderServiceExtensions.CreateAsyncScope(services);"); + + stringWriter.WriteIndentation(5); + stringWriter.Write("await ((global::NetCord.Hosting.Gateway.I"); + stringWriter.Write(eventSymbol.Name); + stringWriter.Write("ShardedGatewayHandler)global::Microsoft.Extensions.DependencyInjection.ServiceProviderServiceExtensions.GetRequiredService(scope.ServiceProvider, handlerMetadata.HandlerType)).HandleAsync(client"); + + if (eventType.Arity is not 1) + stringWriter.Write(", arg"); + + stringWriter.WriteLine(");"); + + stringWriter.WriteIndentation(4); + stringWriter.WriteLine("};"); + + stringWriter.WriteIndentation(2); + stringWriter.WriteLine("}"); } stringWriter.WriteIndentation(1); diff --git a/Tests/NetCord.Test.Hosting/Program.cs b/Tests/NetCord.Test.Hosting/Program.cs index c51d53a60..cc8d7477e 100644 --- a/Tests/NetCord.Test.Hosting/Program.cs +++ b/Tests/NetCord.Test.Hosting/Program.cs @@ -61,9 +61,9 @@ .AddComponentInteractions() .AddComponentInteractions() .AddCommands() - .AddGatewayHandler(GatewayEvent.MessageCreate, (Message message, ILogger logger) => logger.LogInformation("Content: {}", message.Content)) + .AddGatewayHandler(GatewayEvent.MessageCreate, (Message message, ILogger logger, IServiceProvider p) => logger.LogInformation("Content: {}", message.Content), ServiceLifetime.Scoped) .AddGatewayHandler() - .AddGatewayHandler() + .AddGatewayHandler(ServiceLifetime.Scoped) .AddGatewayHandler() .AddSingleton("Wzium") .AddKeyedSingleton("key", "Wzium2");