Files
YesNt-Interpreter/YesNt.Interpreter.Generator/StatementRegistryGenerator.cs
T
2026-03-04 18:12:08 +01:00

244 lines
9.6 KiB
C#

using Microsoft.CodeAnalysis;
using System;
using System.Collections.Generic;
using System.Linq;
using System.Text;
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<MethodRegistration> statementMethods = [];
List<MethodRegistration> 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<MethodRegistration> statementMethods,
List<MethodRegistration> 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<MethodRegistration> statementMethods,
List<MethodRegistration> 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<MethodRegistration> statementMethods,
List<MethodRegistration> staticStatementMethods)
{
StringBuilder sb = new StringBuilder();
_ = sb.AppendLine("// <auto-generated />");
_ = 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<global::YesNt.Interpreter.Attributes.StatementAttribute, Action<string>> statements,");
_ = sb.AppendLine(" out List<KeyValuePair<global::YesNt.Interpreter.Attributes.StaticStatementAttribute, Action>> staticStatements)");
_ = sb.AppendLine(" {");
List<INamedTypeSymbol> allTypes = statementMethods
.Concat(staticStatementMethods)
.Select(x => x.ContainingType)
.GroupBy(x => x, SymbolEqualityComparer.Default)
.Select(g => g.First())
.OrderBy(x => x.ToDisplayString())
.ToList();
Dictionary<INamedTypeSymbol, string> instanceNames = new Dictionary<INamedTypeSymbol, string>(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<KeyValuePair<global::YesNt.Interpreter.Attributes.StatementAttribute, Action<string>>>();");
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<KeyValuePair<global::YesNt.Interpreter.Attributes.StaticStatementAttribute, Action>>();");
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; }
}
}