#nullable enable using System; using System.Collections.Concurrent; using System.Collections.Immutable; using System.Linq; using Microsoft.CodeAnalysis; using Microsoft.CodeAnalysis.Diagnostics; using Microsoft.CodeAnalysis.Operations; namespace ShrinkSDK.CodeGen.Analyzers; [DiagnosticAnalyzer(LanguageNames.CSharp)] public sealed class ShrinkCodeGenAnalyzer : DiagnosticAnalyzer { private static readonly DiagnosticDescriptor InvalidEventHandler = new( "SHRINK001", "Invalid event subscriber signature", "Method '{0}' must use void(TEvent), UniTask(TEvent), or UniTask(TEvent, CancellationToken)", "ShrinkSDK.CodeGen", DiagnosticSeverity.Error, true); private static readonly DiagnosticDescriptor DuplicateNetworkContract = new( "SHRINK002", "Duplicate network contract", "Network opcode/route '{0}' is already declared by '{1}'", "ShrinkSDK.CodeGen", DiagnosticSeverity.Error, true); private static readonly DiagnosticDescriptor InvalidInstaller = new( "SHRINK003", "Invalid application installer", "Type '{0}' has ShrinkAppModuleInstaller but does not implement IShrinkAppModuleInstaller", "ShrinkSDK.CodeGen", DiagnosticSeverity.Error, true); private static readonly DiagnosticDescriptor InvalidKey = new( "SHRINK004", "Invalid Context key declaration", "ShrinkKey argument '{0}' is invalid. Use nonempty package/name and a positive major version, e.g. new ShrinkKey(\"game\", \"rage\", 1).", "ShrinkSDK.CodeGen", DiagnosticSeverity.Error, true); public override ImmutableArray SupportedDiagnostics => ImmutableArray.Create(InvalidEventHandler, DuplicateNetworkContract, InvalidInstaller, InvalidKey); public override void Initialize(AnalysisContext context) { context.ConfigureGeneratedCodeAnalysis(GeneratedCodeAnalysisFlags.None); context.EnableConcurrentExecution(); context.RegisterOperationAction(AnalyzeKey, OperationKind.ObjectCreation); context.RegisterCompilationStartAction(start => { var networkContracts = new ConcurrentDictionary(StringComparer.Ordinal); var referencedConflicts = new System.Collections.Generic.List(); foreach (var assembly in start.Compilation.SourceModule.ReferencedAssemblySymbols) foreach (var type in Types(assembly.GlobalNamespace)) { var attribute = type.GetAttributes().FirstOrDefault(a => a.AttributeClass?.ToDisplayString() == "ShrinkNetwork.ShrinkNetworkMessageAttribute"); if (attribute == null || attribute.ConstructorArguments.Length == 0) continue; AddReference("opcode:" + attribute.ConstructorArguments[0].Value, type); var route = attribute.ConstructorArguments.Length > 1 ? attribute.ConstructorArguments[1].Value as string : null; if (!string.IsNullOrWhiteSpace(route)) AddReference("route:" + route!.Trim(), type); } void AddReference(string key, INamedTypeSymbol type) { if (!networkContracts.TryAdd(key, type) && networkContracts.TryGetValue(key, out var previous) && !SymbolEqualityComparer.Default.Equals(type, previous)) referencedConflicts.Add(Diagnostic.Create(DuplicateNetworkContract, Location.None, key, previous.ToDisplayString() + " and " + type.ToDisplayString())); } start.RegisterCompilationEndAction(end => { foreach (var diagnostic in referencedConflicts) end.ReportDiagnostic(diagnostic); }); start.RegisterSymbolAction(symbolContext => AnalyzeMethod(symbolContext), SymbolKind.Method); start.RegisterSymbolAction(symbolContext => AnalyzeType(symbolContext, networkContracts), SymbolKind.NamedType); }); } private static void AnalyzeMethod(SymbolAnalysisContext context) { var method = (IMethodSymbol)context.Symbol; if (!HasAttribute(method, "ShrinkEventBus.ShrinkSubscribeAttribute")) return; var parametersValid = method.Parameters.Length == 1 || method.Parameters.Length == 2 && method.Parameters[1].Type.ToDisplayString() == "System.Threading.CancellationToken"; var returnName = method.ReturnType.ToDisplayString(); var returnValid = method.Parameters.Length == 1 && (method.ReturnsVoid || returnName == "Cysharp.Threading.Tasks.UniTask") || method.Parameters.Length == 2 && returnName == "Cysharp.Threading.Tasks.UniTask"; var eventType = method.Parameters.FirstOrDefault()?.Type; var eventValid = eventType != null && (eventType.ToDisplayString() == "ShrinkEventBus.IShrinkEvent" || eventType.AllInterfaces.Any(i => i.ToDisplayString() == "ShrinkEventBus.IShrinkEvent")); if (!parametersValid || !returnValid || !eventValid || method.IsGenericMethod || method.Parameters.Any(p => p.RefKind != RefKind.None)) context.ReportDiagnostic(Diagnostic.Create(InvalidEventHandler, method.Locations.FirstOrDefault(), method.Name)); } private static void AnalyzeKey(OperationAnalysisContext context) { var creation = (IObjectCreationOperation)context.Operation; if (creation.Type?.OriginalDefinition.ToDisplayString() != "ShrinkContext.ShrinkKey") return; foreach (var argument in creation.Arguments) { if (!argument.Value.ConstantValue.HasValue) continue; var value = argument.Value.ConstantValue.Value; var invalid = argument.Parameter?.Ordinal < 2 ? value == null || value is string text && string.IsNullOrWhiteSpace(text) : value is int version && version <= 0; if (invalid) context.ReportDiagnostic(Diagnostic.Create(InvalidKey, argument.Syntax.GetLocation(), argument.Parameter?.Name)); } } private static void AnalyzeType(SymbolAnalysisContext context, ConcurrentDictionary networkContracts) { var type = (INamedTypeSymbol)context.Symbol; if (HasAttribute(type, "ShrinkApp.ShrinkAppModuleInstallerAttribute") && !type.AllInterfaces.Any(item => item.ToDisplayString() == "ShrinkApp.IShrinkAppModuleInstaller")) context.ReportDiagnostic(Diagnostic.Create(InvalidInstaller, type.Locations.FirstOrDefault(), type.Name)); var network = type.GetAttributes().FirstOrDefault(attribute => attribute.AttributeClass?.ToDisplayString() == "ShrinkNetwork.ShrinkNetworkMessageAttribute"); if (network == null || network.ConstructorArguments.Length == 0) return; var opcode = network.ConstructorArguments[0].Value?.ToString() ?? string.Empty; var route = network.ConstructorArguments.Length > 1 ? network.ConstructorArguments[1].Value as string ?? string.Empty : string.Empty; Check("opcode:" + opcode); if (!string.IsNullOrWhiteSpace(route)) Check("route:" + route.Trim()); void Check(string key) { if (!networkContracts.TryAdd(key, type) && networkContracts.TryGetValue(key, out var previous) && !SymbolEqualityComparer.Default.Equals(previous, type)) context.ReportDiagnostic(Diagnostic.Create(DuplicateNetworkContract, type.Locations.FirstOrDefault(), key, previous.ToDisplayString())); } } private static System.Collections.Generic.IEnumerable Types(INamespaceOrTypeSymbol container) { foreach (var member in container.GetMembers()) { if (member is INamedTypeSymbol type) { yield return type; foreach (var child in Types(type)) yield return child; } else if (member is INamespaceSymbol ns) foreach (var child in Types(ns)) yield return child; } } private static bool HasAttribute(ISymbol symbol, string fullName) => symbol.GetAttributes().Any(attribute => attribute.AttributeClass?.ToDisplayString() == fullName); }