diff --git a/Directory.Build.props b/Directory.Build.props index 3d819d0..8e45790 100644 --- a/Directory.Build.props +++ b/Directory.Build.props @@ -20,7 +20,7 @@ - 3.11.0 + 4.0.1 diff --git a/src/EntityFrameworkCore.Projectables.Generator/ExpressionSyntaxRewriter.cs b/src/EntityFrameworkCore.Projectables.Generator/ExpressionSyntaxRewriter.cs index 5d79a38..ab00045 100644 --- a/src/EntityFrameworkCore.Projectables.Generator/ExpressionSyntaxRewriter.cs +++ b/src/EntityFrameworkCore.Projectables.Generator/ExpressionSyntaxRewriter.cs @@ -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 _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; } diff --git a/src/EntityFrameworkCore.Projectables.Generator/ProjectableInterpreter.cs b/src/EntityFrameworkCore.Projectables.Generator/ProjectableInterpreter.cs index 919c8d3..3755611 100644 --- a/src/EntityFrameworkCore.Projectables.Generator/ProjectableInterpreter.cs +++ b/src/EntityFrameworkCore.Projectables.Generator/ProjectableInterpreter.cs @@ -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() .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() diff --git a/src/EntityFrameworkCore.Projectables.Generator/ProjectionExpressionGenerator.cs b/src/EntityFrameworkCore.Projectables.Generator/ProjectionExpressionGenerator.cs index 5986343..c295678 100644 --- a/src/EntityFrameworkCore.Projectables.Generator/ProjectionExpressionGenerator.cs +++ b/src/EntityFrameworkCore.Projectables.Generator/ProjectionExpressionGenerator.cs @@ -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 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)> 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 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> 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()); } }