Use incremental source generators

This commit is contained in:
Koen Bekkenutte
2022-01-08 19:13:30 +08:00
parent c251d6d4b1
commit e2670c2771
4 changed files with 123 additions and 68 deletions
+1 -1
View File
@@ -20,7 +20,7 @@
</PropertyGroup>
<PropertyGroup>
<MicrosoftCodeAnalysisVersion>3.11.0</MicrosoftCodeAnalysisVersion>
<MicrosoftCodeAnalysisVersion>4.0.1</MicrosoftCodeAnalysisVersion>
</PropertyGroup>
<PropertyGroup>
@@ -15,14 +15,16 @@ namespace EntityFrameworkCore.Projectables.Generator
readonly INamedTypeSymbol _targetTypeSymbol;
readonly SemanticModel _semanticModel;
readonly NullConditionalRewriteSupport _nullConditionalRewriteSupport;
readonly GeneratorExecutionContext _context;
readonly Compilation _compilation;
readonly SourceProductionContext _context;
readonly Stack<ExpressionSyntax> _conditionalAccessExpressionsStack = new();
public ExpressionSyntaxRewriter(INamedTypeSymbol targetTypeSymbol, SemanticModel semanticModel, NullConditionalRewriteSupport nullConditionalRewriteSupport, GeneratorExecutionContext context)
public ExpressionSyntaxRewriter(INamedTypeSymbol targetTypeSymbol, NullConditionalRewriteSupport nullConditionalRewriteSupport, Compilation compilation, SemanticModel semanticModel, SourceProductionContext context)
{
_targetTypeSymbol = targetTypeSymbol;
_semanticModel = semanticModel;
_nullConditionalRewriteSupport = nullConditionalRewriteSupport;
_semanticModel = semanticModel;
_compilation = compilation;
_context = context;
}
@@ -138,7 +140,7 @@ namespace EntityFrameworkCore.Projectables.Generator
{
rewrite = false;
}
if (targetSymbolInfo.Symbol?.ContainingType is not null && !_context.Compilation.HasImplicitConversion(targetSymbolInfo.Symbol.ContainingType, _targetTypeSymbol))
if (targetSymbolInfo.Symbol?.ContainingType is not null && !_compilation.HasImplicitConversion(targetSymbolInfo.Symbol.ContainingType, _targetTypeSymbol))
{
rewrite = false;
}
@@ -22,17 +22,17 @@ namespace EntityFrameworkCore.Projectables.Generator
yield return namedTypeSymbol.Name;
}
public static ProjectableDescriptor? GetDescriptor(MemberDeclarationSyntax memberDeclarationSyntax, GeneratorExecutionContext context)
public static ProjectableDescriptor? GetDescriptor(Compilation compilation, MemberDeclarationSyntax member, SourceProductionContext context)
{
var semanticModel = context.Compilation.GetSemanticModel(memberDeclarationSyntax.SyntaxTree);
var memberSymbol = semanticModel.GetDeclaredSymbol(memberDeclarationSyntax);
var semanticModel = compilation.GetSemanticModel(member.SyntaxTree);
var memberSymbol = semanticModel.GetDeclaredSymbol(member);
if (memberSymbol is null)
{
return null;
}
var projectableAttributeTypeSymbol = context.Compilation.GetTypeByMetadataName("EntityFrameworkCore.Projectables.ProjectableAttribute");
var projectableAttributeTypeSymbol = compilation.GetTypeByMetadataName("EntityFrameworkCore.Projectables.ProjectableAttribute");
var projectableAttributeClass = memberSymbol.GetAttributes()
.Where(x => x.AttributeClass?.Name == "ProjectableAttribute")
@@ -51,7 +51,7 @@ namespace EntityFrameworkCore.Projectables.Generator
.Cast<NullConditionalRewriteSupport>()
.FirstOrDefault();
var expressionSyntaxRewriter = new ExpressionSyntaxRewriter(memberSymbol.ContainingType, semanticModel, nullConditionalRewriteSupport, context);
var expressionSyntaxRewriter = new ExpressionSyntaxRewriter(memberSymbol.ContainingType, nullConditionalRewriteSupport, compilation, semanticModel, context);
var declarationSyntaxRewriter = new DeclarationSyntaxRewriter(semanticModel);
var descriptor = new ProjectableDescriptor {
@@ -63,7 +63,7 @@ namespace EntityFrameworkCore.Projectables.Generator
TypeParameterList = SyntaxFactory.TypeParameterList()
};
if (!memberDeclarationSyntax.Modifiers.Any(SyntaxKind.StaticKeyword))
if (!member.Modifiers.Any(SyntaxKind.StaticKeyword))
{
descriptor.ParametersList = descriptor.ParametersList.AddParameters(
SyntaxFactory.Parameter(
@@ -93,7 +93,7 @@ namespace EntityFrameworkCore.Projectables.Generator
descriptor.TargetNestedInClassNames = descriptor.NestedInClassNames;
}
if (memberDeclarationSyntax is MethodDeclarationSyntax methodDeclarationSyntax)
if (member is MethodDeclarationSyntax methodDeclarationSyntax)
{
if (methodDeclarationSyntax.ExpressionBody is null)
{
@@ -125,7 +125,7 @@ namespace EntityFrameworkCore.Projectables.Generator
.Select(x => (TypeParameterConstraintClauseSyntax)declarationSyntaxRewriter.Visit(x));
}
}
else if (memberDeclarationSyntax is PropertyDeclarationSyntax propertyDeclarationSyntax)
else if (member is PropertyDeclarationSyntax propertyDeclarationSyntax)
{
if (propertyDeclarationSyntax.ExpressionBody is null)
{
@@ -145,7 +145,7 @@ namespace EntityFrameworkCore.Projectables.Generator
}
descriptor.UsingDirectives =
memberDeclarationSyntax.SyntaxTree
member.SyntaxTree
.GetRoot()
.DescendantNodes()
.OfType<UsingDirectiveSyntax>()
@@ -1,9 +1,11 @@
using EntityFrameworkCore.Projectables.Services;
using Microsoft.CodeAnalysis;
using Microsoft.CodeAnalysis.CSharp;
using Microsoft.CodeAnalysis.CSharp.Syntax;
using Microsoft.CodeAnalysis.Text;
using System;
using System.Collections.Generic;
using System.Collections.Immutable;
using System.Diagnostics;
using System.Linq;
using System.Text;
@@ -13,70 +15,124 @@ using System.Threading.Tasks;
namespace EntityFrameworkCore.Projectables.Generator
{
[Generator]
public class ProjectionExpressionGenerator : ISourceGenerator
public class ProjectionExpressionGenerator : IIncrementalGenerator
{
public void Execute(GeneratorExecutionContext context)
private const string ProjectablesAttributeName = "EntityFrameworkCore.Projectables.ProjectableAttribute";
public void Initialize(GeneratorInitializationContext context) =>
context.RegisterForSyntaxNotifications(() => new SyntaxReceiver());
public void Initialize(IncrementalGeneratorInitializationContext context)
{
if (context.SyntaxReceiver is not SyntaxReceiver receiver)
// Do a simple filter for members
IncrementalValuesProvider<MemberDeclarationSyntax> memberDeclarations = context.SyntaxProvider
.CreateSyntaxProvider(
predicate: static (s, _) => s is MemberDeclarationSyntax m && m.AttributeLists.Count > 0,
transform: static (c, _) => GetSemanticTargetForGeneration(c))
.Where(static m => m is not null)!; // filter out attributed enums that we don't care about
// Combine the selected enums with the `Compilation`
IncrementalValueProvider<(Compilation, ImmutableArray<MemberDeclarationSyntax>)> compilationAndEnums
= context.CompilationProvider.Combine(memberDeclarations.Collect());
// Generate the source using the compilation and enums
context.RegisterSourceOutput(compilationAndEnums,
static (spc, source) => Execute(source.Item1, source.Item2, spc));
}
static MemberDeclarationSyntax? GetSemanticTargetForGeneration(GeneratorSyntaxContext context)
{
// we know the node is a MemberDeclarationSyntax
var memberDeclarationSyntax = (MemberDeclarationSyntax)context.Node;
// loop through all the attributes on the method
foreach (var attributeListSyntax in memberDeclarationSyntax.AttributeLists)
{
foreach (var attributeSyntax in attributeListSyntax.Attributes)
{
if (context.SemanticModel.GetSymbolInfo(attributeSyntax).Symbol is not IMethodSymbol attributeSymbol)
{
// weird, we couldn't get the symbol, ignore it
continue;
}
var attributeContainingTypeSymbol = attributeSymbol.ContainingType;
var fullName = attributeContainingTypeSymbol.ToDisplayString();
// Is the attribute the [Projcetable] attribute?
if (fullName == ProjectablesAttributeName)
{
// return the enum
return memberDeclarationSyntax;
}
}
}
// we didn't find the attribute we were looking for
return null;
}
static void Execute(Compilation compilation, ImmutableArray<MemberDeclarationSyntax> members, SourceProductionContext context)
{
if (members.IsDefaultOrEmpty)
{
// nothing to do yet
return;
}
if (receiver.Candidates?.Count > 0)
var projectables = members
.Select(x => ProjectableInterpreter.GetDescriptor(compilation, x, context))
.Where(x => x is not null)
.Select(x => x!);
var resultBuilder = new StringBuilder();
foreach (var projectable in projectables)
{
var projectables = receiver.Candidates
.Select(x => ProjectableInterpreter.GetDescriptor(x, context))
.Where(x => x is not null)
.Select(x => x!);
var resultBuilder = new StringBuilder();
foreach (var projectable in projectables)
if (projectable.MemberName is null)
{
if (projectable.MemberName is null)
throw new InvalidOperationException("Expected a memberName here");
}
resultBuilder.Clear();
if (projectable.UsingDirectives is not null)
{
foreach (var usingDirective in projectable.UsingDirectives.Distinct())
{
throw new InvalidOperationException("Expected a memberName here");
resultBuilder.AppendLine(usingDirective);
}
}
resultBuilder.Clear();
if (projectable.TargetClassNamespace is not null)
{
var targetClassUsingDirective = $"using {projectable.TargetClassNamespace};";
if (projectable.UsingDirectives is not null)
if (!projectable.UsingDirectives.Contains(targetClassUsingDirective))
{
foreach (var usingDirective in projectable.UsingDirectives.Distinct())
{
resultBuilder.AppendLine(usingDirective);
}
resultBuilder.AppendLine(targetClassUsingDirective);
}
}
if (projectable.TargetClassNamespace is not null)
if (projectable.ClassNamespace is not null && projectable.ClassNamespace != projectable.TargetClassNamespace)
{
var classUsingDirective = $"using {projectable.ClassNamespace};";
if (!projectable.UsingDirectives.Contains(classUsingDirective))
{
var targetClassUsingDirective = $"using {projectable.TargetClassNamespace};";
if (!projectable.UsingDirectives.Contains(targetClassUsingDirective))
{
resultBuilder.AppendLine(targetClassUsingDirective);
}
resultBuilder.AppendLine(classUsingDirective);
}
}
if (projectable.ClassNamespace is not null && projectable.ClassNamespace != projectable.TargetClassNamespace)
{
var classUsingDirective = $"using {projectable.ClassNamespace};";
var generatedClassName = ProjectionExpressionClassNameGenerator.GenerateName(projectable.ClassNamespace, projectable.NestedInClassNames, projectable.MemberName);
if (!projectable.UsingDirectives.Contains(classUsingDirective))
{
resultBuilder.AppendLine(classUsingDirective);
}
}
var lambdaTypeArguments = SyntaxFactory.TypeArgumentList(
SyntaxFactory.SeparatedList(
projectable.ParametersList?.Parameters.Where(p => p.Type is not null).Select(p => p.Type!)
)
);
var generatedClassName = ProjectionExpressionClassNameGenerator.GenerateName(projectable.ClassNamespace, projectable.NestedInClassNames, projectable.MemberName);
var lambdaTypeArguments = SyntaxFactory.TypeArgumentList(
SyntaxFactory.SeparatedList(
projectable.ParametersList?.Parameters.Where(p => p.Type is not null).Select(p => p.Type!)
)
);
resultBuilder.Append($@"
resultBuilder.Append($@"
namespace EntityFrameworkCore.Projectables.Generated
#nullable disable
{{
@@ -84,16 +140,16 @@ namespace EntityFrameworkCore.Projectables.Generated
{{
public static System.Linq.Expressions.Expression<System.Func<{lambdaTypeArguments.Arguments}, {projectable.ReturnTypeName}>> Expression{(projectable.TypeParameterList?.Parameters.Any() == true ? projectable.TypeParameterList.ToString() : string.Empty)}()");
if (projectable.ConstraintClauses is not null)
if (projectable.ConstraintClauses is not null)
{
foreach (var constraintClause in projectable.ConstraintClauses)
{
foreach (var constraintClause in projectable.ConstraintClauses)
{
resultBuilder.Append($@"
resultBuilder.Append($@"
{constraintClause}");
}
}
}
resultBuilder.Append($@"
resultBuilder.Append($@"
{{
return {projectable.ParametersList} =>
{projectable.Body};
@@ -102,12 +158,9 @@ namespace EntityFrameworkCore.Projectables.Generated
}}");
context.AddSource($"{generatedClassName}_Generated", SourceText.From(resultBuilder.ToString(), Encoding.UTF8));
}
context.AddSource($"{generatedClassName}_Generated", SourceText.From(resultBuilder.ToString(), Encoding.UTF8));
}
}
public void Initialize(GeneratorInitializationContext context) =>
context.RegisterForSyntaxNotifications(() => new SyntaxReceiver());
}
}