310 lines
15 KiB
C#
310 lines
15 KiB
C#
#nullable enable
|
|
|
|
using System;
|
|
using System.Collections.Generic;
|
|
using System.Collections.Immutable;
|
|
using System.Linq;
|
|
using System.Text;
|
|
using Microsoft.CodeAnalysis;
|
|
using Microsoft.CodeAnalysis.CSharp;
|
|
using Microsoft.CodeAnalysis.CSharp.Syntax;
|
|
using Microsoft.CodeAnalysis.Text;
|
|
|
|
namespace ShrinkEventBus.Generator
|
|
{
|
|
[Generator]
|
|
public sealed class ShrinkEventBusGenerator : IIncrementalGenerator
|
|
{
|
|
private const string SubscriberAttribute = "ShrinkEventBus.ShrinkEventSubscriberAttribute";
|
|
private const string SubscribeAttribute = "ShrinkEventBus.ShrinkSubscribeAttribute";
|
|
|
|
private static readonly DiagnosticDescriptor PartialRequired = new DiagnosticDescriptor(
|
|
"SHRINKEVENT001", "Subscriber must be partial",
|
|
"Subscriber type '{0}' must be partial so ShrinkEventBus can generate a reflection-free binding",
|
|
"ShrinkEventBus", DiagnosticSeverity.Error, true);
|
|
|
|
private static readonly DiagnosticDescriptor UnsupportedHandler = new DiagnosticDescriptor(
|
|
"SHRINKEVENT002", "Unsupported handler signature",
|
|
"Handler '{0}' must be void(TEvent), UniTask(TEvent), UniTask(TEvent, CancellationToken), or ValueTask(TEvent, CancellationToken)",
|
|
"ShrinkEventBus", DiagnosticSeverity.Error, true);
|
|
|
|
private static readonly DiagnosticDescriptor TopLevelRequired = new DiagnosticDescriptor(
|
|
"SHRINKEVENT003", "Top-level subscriber required",
|
|
"Subscriber type '{0}' must be top-level in the current generator version",
|
|
"ShrinkEventBus", DiagnosticSeverity.Error, true);
|
|
|
|
public void Initialize(IncrementalGeneratorInitializationContext context)
|
|
{
|
|
context.RegisterSourceOutput(context.CompilationProvider,
|
|
static (sourceContext, compilation) =>
|
|
{
|
|
if (compilation.GetTypeByMetadataName(
|
|
"System.Runtime.CompilerServices.ModuleInitializerAttribute") == null)
|
|
{
|
|
sourceContext.AddSource("ShrinkEventBus.ModuleInitializerAttribute.g.cs",
|
|
SourceText.From(
|
|
"namespace System.Runtime.CompilerServices { [global::System.AttributeUsage(global::System.AttributeTargets.Method, Inherited = false)] internal sealed class ModuleInitializerAttribute : global::System.Attribute { } }",
|
|
Encoding.UTF8));
|
|
}
|
|
});
|
|
|
|
var candidates = context.SyntaxProvider.ForAttributeWithMetadataName(
|
|
SubscriberAttribute,
|
|
static (node, _) => node is ClassDeclarationSyntax,
|
|
static (syntaxContext, _) => (INamedTypeSymbol)syntaxContext.TargetSymbol);
|
|
|
|
context.RegisterSourceOutput(candidates.Collect(), Generate);
|
|
}
|
|
|
|
private static void Generate(SourceProductionContext context,
|
|
ImmutableArray<INamedTypeSymbol> candidates)
|
|
{
|
|
foreach (var type in candidates)
|
|
GenerateType(context, type);
|
|
}
|
|
|
|
private static void GenerateType(SourceProductionContext context, INamedTypeSymbol type)
|
|
{
|
|
if (type.ContainingType != null)
|
|
{
|
|
context.ReportDiagnostic(Diagnostic.Create(TopLevelRequired,
|
|
type.Locations.FirstOrDefault(), type.ToDisplayString()));
|
|
return;
|
|
}
|
|
|
|
var isPartial = type.DeclaringSyntaxReferences
|
|
.Select(reference => reference.GetSyntax())
|
|
.OfType<ClassDeclarationSyntax>()
|
|
.Any(declaration => declaration.Modifiers.Any(SyntaxKind.PartialKeyword));
|
|
if (!isPartial)
|
|
{
|
|
context.ReportDiagnostic(Diagnostic.Create(PartialRequired,
|
|
type.Locations.FirstOrDefault(), type.ToDisplayString()));
|
|
return;
|
|
}
|
|
|
|
var handlers = new List<HandlerModel>();
|
|
foreach (var method in type.GetMembers().OfType<IMethodSymbol>())
|
|
{
|
|
var attribute = method.GetAttributes().FirstOrDefault(item =>
|
|
item.AttributeClass?.ToDisplayString() == SubscribeAttribute);
|
|
if (attribute == null)
|
|
continue;
|
|
|
|
if (!TryCreateHandler(method, attribute, type.IsStatic, out var handler))
|
|
{
|
|
context.ReportDiagnostic(Diagnostic.Create(UnsupportedHandler,
|
|
method.Locations.FirstOrDefault(), method.ToDisplayString()));
|
|
continue;
|
|
}
|
|
handlers.Add(handler);
|
|
}
|
|
|
|
var subscriber = type.GetAttributes().First(item =>
|
|
item.AttributeClass?.ToDisplayString() == SubscriberAttribute);
|
|
var defaultBus = ReadString(subscriber, "DefaultBus");
|
|
var source = type.IsStatic
|
|
? BuildStaticSource(type, handlers, defaultBus)
|
|
: BuildSource(type, handlers, defaultBus);
|
|
var hint = type.ToDisplayString().Replace('.', '_').Replace('+', '_') + ".ShrinkEvents.g.cs";
|
|
context.AddSource(hint, SourceText.From(source, Encoding.UTF8));
|
|
}
|
|
|
|
private static bool TryCreateHandler(IMethodSymbol method, AttributeData attribute,
|
|
bool staticSubscriber,
|
|
out HandlerModel model)
|
|
{
|
|
model = default;
|
|
if (method.IsStatic != staticSubscriber || method.Parameters.Length == 0 || method.Parameters.Length > 2)
|
|
return false;
|
|
|
|
var eventType = method.Parameters[0].Type.ToDisplayString(SymbolDisplayFormat.FullyQualifiedFormat);
|
|
var returnType = method.ReturnType.ToDisplayString();
|
|
var isSync = method.ReturnsVoid && method.Parameters.Length == 1;
|
|
var isValueTask = returnType == "System.Threading.Tasks.ValueTask" &&
|
|
method.Parameters.Length == 2 &&
|
|
method.Parameters[1].Type.ToDisplayString() == "System.Threading.CancellationToken";
|
|
var isUniTask = returnType == "Cysharp.Threading.Tasks.UniTask";
|
|
var isUniTaskLegacy = isUniTask && method.Parameters.Length == 1;
|
|
var isUniTaskCancelable = isUniTask && method.Parameters.Length == 2 &&
|
|
method.Parameters[1].Type.ToDisplayString() == "System.Threading.CancellationToken";
|
|
if (!isSync && !isValueTask && !isUniTaskLegacy && !isUniTaskCancelable)
|
|
return false;
|
|
|
|
var bindingMethod = isSync
|
|
? "Subscribe"
|
|
: isUniTaskLegacy ? "SubscribeAsyncLegacy" : "SubscribeAsync";
|
|
|
|
model = new HandlerModel(
|
|
method.Name,
|
|
eventType,
|
|
bindingMethod,
|
|
ReadString(attribute, "Bus"),
|
|
ReadInt(attribute, "Priority", 2),
|
|
ReadInt(attribute, "NumericPriority", 0),
|
|
ReadBool(attribute, "ReceiveCanceled", false));
|
|
return true;
|
|
}
|
|
|
|
private static string BuildSource(INamedTypeSymbol type, IReadOnlyList<HandlerModel> handlers,
|
|
string defaultBus)
|
|
{
|
|
var builder = new StringBuilder();
|
|
builder.AppendLine("// <auto-generated />");
|
|
builder.AppendLine("#nullable enable");
|
|
if (!type.ContainingNamespace.IsGlobalNamespace)
|
|
{
|
|
builder.Append("namespace ").Append(type.ContainingNamespace.ToDisplayString()).AppendLine();
|
|
builder.AppendLine("{");
|
|
}
|
|
|
|
var accessibility = type.DeclaredAccessibility switch
|
|
{
|
|
Accessibility.Public => "public ",
|
|
Accessibility.Internal => "internal ",
|
|
_ => string.Empty
|
|
};
|
|
builder.Append(" ").Append(accessibility).Append("partial class ")
|
|
.Append(type.Name).AppendLine(" : global::ShrinkEventBus.IShrinkGeneratedSubscriber");
|
|
builder.AppendLine(" {");
|
|
builder.AppendLine(" global::System.IDisposable global::ShrinkEventBus.IShrinkGeneratedSubscriber.AttachGenerated(");
|
|
builder.AppendLine(" global::ShrinkEventBus.IShrinkBusResolver resolver,");
|
|
builder.AppendLine(" global::ShrinkEventBus.ShrinkBusKey? defaultBus)");
|
|
builder.AppendLine(" {");
|
|
builder.AppendLine(" var binding = new global::ShrinkEventBus.ShrinkEventBinding();");
|
|
|
|
foreach (var handler in handlers)
|
|
{
|
|
var bus = string.IsNullOrWhiteSpace(handler.Bus) ? defaultBus : handler.Bus;
|
|
builder.Append(" binding.Add(global::ShrinkEventBus.ShrinkGeneratedBinding.")
|
|
.Append(handler.BindingMethod).Append('<').Append(handler.EventType).AppendLine(">(");
|
|
builder.Append(" resolver, defaultBus, ")
|
|
.Append(ToLiteral(bus)).AppendLine(", this,");
|
|
builder.Append(" this.").Append(handler.MethodName)
|
|
.Append(", (global::ShrinkEventBus.ShrinkEventPriority)")
|
|
.Append(handler.Priority).Append(", ").Append(handler.NumericPriority).Append(", ")
|
|
.Append(handler.ReceiveCanceled ? "true" : "false").AppendLine("));");
|
|
}
|
|
|
|
builder.AppendLine(" return binding;");
|
|
builder.AppendLine(" }");
|
|
builder.AppendLine(" }");
|
|
if (!type.ContainingNamespace.IsGlobalNamespace)
|
|
builder.AppendLine("}");
|
|
return builder.ToString();
|
|
}
|
|
|
|
private static string BuildStaticSource(INamedTypeSymbol type,
|
|
IReadOnlyList<HandlerModel> handlers, string defaultBus)
|
|
{
|
|
var builder = new StringBuilder();
|
|
builder.AppendLine("// <auto-generated />");
|
|
builder.AppendLine("#nullable enable");
|
|
if (!type.ContainingNamespace.IsGlobalNamespace)
|
|
{
|
|
builder.Append("namespace ").Append(type.ContainingNamespace.ToDisplayString()).AppendLine();
|
|
builder.AppendLine("{");
|
|
}
|
|
|
|
var accessibility = type.DeclaredAccessibility switch
|
|
{
|
|
Accessibility.Public => "public ",
|
|
Accessibility.Internal => "internal ",
|
|
_ => string.Empty
|
|
};
|
|
builder.Append(" ").Append(accessibility).Append("static partial class ")
|
|
.Append(type.Name).AppendLine();
|
|
builder.AppendLine(" {");
|
|
builder.AppendLine(" [global::System.Runtime.CompilerServices.ModuleInitializer]");
|
|
builder.AppendLine(" internal static void ShrinkEventBus_RegisterStaticBindings()");
|
|
builder.AppendLine(" {");
|
|
for (var i = 0; i < handlers.Count; i++)
|
|
{
|
|
var bus = string.IsNullOrWhiteSpace(handlers[i].Bus) ? defaultBus : handlers[i].Bus;
|
|
builder.Append(" global::ShrinkEventBus.ShrinkStaticBindingRegistry.Register(")
|
|
.Append("global::ShrinkEventBus.ShrinkBusKey.Parse(")
|
|
.Append(ToLiteral(bus)).Append("), ShrinkEventBus_Bind_").Append(i).AppendLine(");");
|
|
}
|
|
builder.AppendLine(" }");
|
|
|
|
for (var i = 0; i < handlers.Count; i++)
|
|
{
|
|
var handler = handlers[i];
|
|
var bus = string.IsNullOrWhiteSpace(handler.Bus) ? defaultBus : handler.Bus;
|
|
builder.Append(" private static global::System.IDisposable ShrinkEventBus_Bind_")
|
|
.Append(i).AppendLine("(global::ShrinkEventBus.IShrinkBusResolver resolver)");
|
|
builder.AppendLine(" {");
|
|
builder.Append(" return global::ShrinkEventBus.ShrinkGeneratedBinding.")
|
|
.Append(handler.BindingMethod).Append('<').Append(handler.EventType).AppendLine(">(");
|
|
builder.Append(" resolver, null, ").Append(ToLiteral(bus))
|
|
.AppendLine(", null,");
|
|
builder.Append(" ").Append(handler.MethodName)
|
|
.Append(", (global::ShrinkEventBus.ShrinkEventPriority)")
|
|
.Append(handler.Priority).Append(", ").Append(handler.NumericPriority).Append(", ")
|
|
.Append(handler.ReceiveCanceled ? "true" : "false").AppendLine(");");
|
|
builder.AppendLine(" }");
|
|
}
|
|
|
|
builder.AppendLine(" }");
|
|
if (!type.ContainingNamespace.IsGlobalNamespace)
|
|
builder.AppendLine("}");
|
|
return builder.ToString();
|
|
}
|
|
|
|
private static string ReadString(AttributeData attribute, string name)
|
|
{
|
|
foreach (var pair in attribute.NamedArguments)
|
|
{
|
|
if (pair.Key == name)
|
|
return pair.Value.Value as string ?? string.Empty;
|
|
}
|
|
return string.Empty;
|
|
}
|
|
|
|
private static int ReadInt(AttributeData attribute, string name, int defaultValue)
|
|
{
|
|
foreach (var pair in attribute.NamedArguments)
|
|
{
|
|
if (pair.Key == name && pair.Value.Value != null)
|
|
return Convert.ToInt32(pair.Value.Value);
|
|
}
|
|
return defaultValue;
|
|
}
|
|
|
|
private static bool ReadBool(AttributeData attribute, string name, bool defaultValue)
|
|
{
|
|
foreach (var pair in attribute.NamedArguments)
|
|
{
|
|
if (pair.Key == name && pair.Value.Value is bool value)
|
|
return value;
|
|
}
|
|
return defaultValue;
|
|
}
|
|
|
|
private static string ToLiteral(string value) => SymbolDisplay.FormatLiteral(value ?? string.Empty, true);
|
|
|
|
private readonly struct HandlerModel
|
|
{
|
|
public HandlerModel(string methodName, string eventType, string bindingMethod, string bus,
|
|
int priority, int numericPriority, bool receiveCanceled)
|
|
{
|
|
MethodName = methodName;
|
|
EventType = eventType;
|
|
BindingMethod = bindingMethod;
|
|
Bus = bus;
|
|
Priority = priority;
|
|
NumericPriority = numericPriority;
|
|
ReceiveCanceled = receiveCanceled;
|
|
}
|
|
|
|
public string MethodName { get; }
|
|
public string EventType { get; }
|
|
public string BindingMethod { get; }
|
|
public string Bus { get; }
|
|
public int Priority { get; }
|
|
public int NumericPriority { get; }
|
|
public bool ReceiveCanceled { get; }
|
|
}
|
|
}
|
|
}
|