using System; using System.Collections.Generic; using System.Linq; using System.Text; using Microsoft.CodeAnalysis; namespace YesNt.Interpreter.Generator; [Generator] public sealed class StatementRegistryGenerator : ISourceGenerator { private const string StatementAttributeName = "YesNt.Interpreter.Attributes.StatementAttribute"; private const string StaticStatementAttributeName = "YesNt.Interpreter.Attributes.StaticStatementAttribute"; public void Initialize(GeneratorInitializationContext context) { } public void Execute(GeneratorExecutionContext context) { Compilation compilation = context.Compilation; List statementMethods = []; List staticStatementMethods = []; CollectMethods(compilation.Assembly.GlobalNamespace, statementMethods, staticStatementMethods); string source = GenerateRegistrySource(statementMethods, staticStatementMethods); context.AddSource("GeneratedStatementRegistry.g.cs", source); } private static void CollectMethods( INamespaceSymbol namespaceSymbol, List statementMethods, List staticStatementMethods) { foreach (INamespaceSymbol childNamespace in namespaceSymbol.GetNamespaceMembers()) { CollectMethods(childNamespace, statementMethods, staticStatementMethods); } foreach (INamedTypeSymbol type in namespaceSymbol.GetTypeMembers()) { CollectMethods(type, statementMethods, staticStatementMethods); } } private static void CollectMethods( INamedTypeSymbol typeSymbol, List statementMethods, List staticStatementMethods) { foreach (ISymbol member in typeSymbol.GetMembers()) { if (member is IMethodSymbol method && method.MethodKind == MethodKind.Ordinary) { foreach (AttributeData attribute in method.GetAttributes()) { string? attributeName = attribute.AttributeClass?.ToDisplayString(); if (attributeName == StatementAttributeName) { statementMethods.Add(new MethodRegistration(typeSymbol, method, attribute)); } else if (attributeName == StaticStatementAttributeName) { staticStatementMethods.Add(new MethodRegistration(typeSymbol, method, attribute)); } } } } foreach (INamedTypeSymbol nestedType in typeSymbol.GetTypeMembers()) { CollectMethods(nestedType, statementMethods, staticStatementMethods); } } private static string GenerateRegistrySource( List statementMethods, List staticStatementMethods) { StringBuilder sb = new StringBuilder(); _ = sb.AppendLine("// "); _ = sb.AppendLine("#nullable enable"); _ = sb.AppendLine("using System;"); _ = sb.AppendLine("using System.Collections.Generic;"); _ = sb.AppendLine("using System.Linq;"); _ = sb.AppendLine(); _ = sb.AppendLine("namespace YesNt.Interpreter.Runtime;"); _ = sb.AppendLine(); _ = sb.AppendLine("internal static class GeneratedStatementRegistry"); _ = sb.AppendLine("{"); _ = sb.AppendLine(" internal static void Register("); _ = sb.AppendLine(" RuntimeInformation runtimeInfo,"); _ = sb.AppendLine(" out Dictionary> statements,"); _ = sb.AppendLine(" out List> staticStatements)"); _ = sb.AppendLine(" {"); List allTypes = statementMethods .Concat(staticStatementMethods) .Select(x => x.ContainingType) .GroupBy(x => x, SymbolEqualityComparer.Default) .Select(g => g.First()) .OrderBy(x => x.ToDisplayString()) .ToList(); Dictionary instanceNames = new Dictionary(SymbolEqualityComparer.Default); int index = 0; foreach (INamedTypeSymbol type in allTypes) { string instanceName = $"instance{index++}"; instanceNames[type] = instanceName; _ = sb.AppendLine($" var {instanceName} = new global::{type.ToDisplayString()}();"); _ = sb.AppendLine($" {instanceName}.RuntimeInfo = runtimeInfo;"); } _ = sb.AppendLine(" var statementEntries = new List>>();"); foreach (MethodRegistration method in statementMethods .OrderBy(x => x.ContainingType.ToDisplayString()) .ThenBy(x => x.Method.Name)) { string instanceName = instanceNames[method.ContainingType]; string attributeCreation = BuildAttributeCreation("global::YesNt.Interpreter.Attributes.StatementAttribute", method.Attribute); _ = sb.AppendLine($" statementEntries.Add(new({attributeCreation}, {instanceName}.{method.Method.Name}));"); } _ = sb.AppendLine(" var staticEntries = new List>();"); foreach (MethodRegistration method in staticStatementMethods .OrderBy(x => x.ContainingType.ToDisplayString()) .ThenBy(x => x.Method.Name)) { string instanceName = instanceNames[method.ContainingType]; string attributeCreation = BuildAttributeCreation("global::YesNt.Interpreter.Attributes.StaticStatementAttribute", method.Attribute); _ = sb.AppendLine($" staticEntries.Add(new({attributeCreation}, {instanceName}.{method.Method.Name}));"); } _ = sb.AppendLine(" statements = statementEntries"); _ = sb.AppendLine(" .OrderBy(s => s.Key.Priority)"); _ = sb.AppendLine(" .ThenByDescending(s => s.Key.Name.Length)"); _ = sb.AppendLine(" .ToDictionary(x => x.Key, x => x.Value);"); _ = sb.AppendLine(); _ = sb.AppendLine(" staticStatements = staticEntries"); _ = sb.AppendLine(" .OrderBy(s => s.Key.Priority)"); _ = sb.AppendLine(" .ToList();"); _ = sb.AppendLine(" }"); _ = sb.AppendLine("}"); return sb.ToString(); } private static string BuildAttributeCreation(string attributeTypeName, AttributeData attribute) { string ctorArgs = string.Join(", ", attribute.ConstructorArguments.Select(ToLiteral)); string creation = $"new {attributeTypeName}({ctorArgs})"; if (attribute.NamedArguments.Length == 0) { return creation; } string namedArgs = string.Join(", ", attribute.NamedArguments.Select(arg => $"{arg.Key} = {ToLiteral(arg.Value)}")); return $"{creation} {{ {namedArgs} }}"; } private static string ToLiteral(TypedConstant constant) { if (constant.IsNull) { return "null!"; } if (constant.Type is null) { return "null!"; } if (constant.Kind == TypedConstantKind.Enum) { string enumType = $"global::{constant.Type.ToDisplayString()}"; object value = constant.Value!; return $"({enumType}){Convert.ToInt64(value)}"; } return constant.Type.SpecialType switch { SpecialType.System_String => "\"" + EscapeString((string)constant.Value!) + "\"", SpecialType.System_Char => "'" + EscapeChar((char)constant.Value!) + "'", SpecialType.System_Boolean => (bool)constant.Value! ? "true" : "false", SpecialType.System_Int32 => ((int)constant.Value!).ToString(System.Globalization.CultureInfo.InvariantCulture), SpecialType.System_Int64 => ((long)constant.Value!).ToString(System.Globalization.CultureInfo.InvariantCulture) + "L", SpecialType.System_Single => ((float)constant.Value!).ToString(System.Globalization.CultureInfo.InvariantCulture) + "f", SpecialType.System_Double => ((double)constant.Value!).ToString(System.Globalization.CultureInfo.InvariantCulture), _ => constant.Value!.ToString() ?? "null!" }; } private static string EscapeString(string value) { return value .Replace("\\", "\\\\") .Replace("\"", "\\\"") .Replace("\r", "\\r") .Replace("\n", "\\n") .Replace("\t", "\\t"); } private static string EscapeChar(char value) { return value switch { '\\' => "\\\\", '\'' => "\\'", '\r' => "\\r", '\n' => "\\n", '\t' => "\\t", _ => value.ToString() }; } private sealed class MethodRegistration { public MethodRegistration(INamedTypeSymbol containingType, IMethodSymbol method, AttributeData attribute) { ContainingType = containingType; Method = method; Attribute = attribute; } public INamedTypeSymbol ContainingType { get; } public IMethodSymbol Method { get; } public AttributeData Attribute { get; } } }