diff --git a/src/EntityFrameworkCore.Projectables.Generator/ProjectableInterpreter.cs b/src/EntityFrameworkCore.Projectables.Generator/ProjectableInterpreter.cs index 463ab46..0d68b2d 100644 --- a/src/EntityFrameworkCore.Projectables.Generator/ProjectableInterpreter.cs +++ b/src/EntityFrameworkCore.Projectables.Generator/ProjectableInterpreter.cs @@ -45,6 +45,7 @@ namespace EntityFrameworkCore.Projectables.Generator var expressionSyntaxRewriter = new ExpressionSyntaxRewriter(memberSymbol.ContainingType, semanticModel); var parameterSyntaxRewriter = new ParameterSyntaxRewriter(semanticModel); + var returnTypeSyntaxRewriter = new ReturnTypeSyntaxRewriter(semanticModel); var descriptor = new ProjectableDescriptor { ClassName = memberSymbol.ContainingType.Name, @@ -91,7 +92,7 @@ namespace EntityFrameworkCore.Projectables.Generator return null; } - descriptor.ReturnTypeName = methodDeclarationSyntax.ReturnType.ToString(); + descriptor.ReturnTypeName = returnTypeSyntaxRewriter.Visit(methodDeclarationSyntax.ReturnType).ToString(); descriptor.Body = expressionSyntaxRewriter.Visit(methodDeclarationSyntax.ExpressionBody.Expression); foreach (var additionalParameter in ((ParameterListSyntax)parameterSyntaxRewriter.Visit(methodDeclarationSyntax.ParameterList)).Parameters) { diff --git a/src/EntityFrameworkCore.Projectables.Generator/ReturnTypeSyntaxRewriter.cs b/src/EntityFrameworkCore.Projectables.Generator/ReturnTypeSyntaxRewriter.cs new file mode 100644 index 0000000..2e0c90b --- /dev/null +++ b/src/EntityFrameworkCore.Projectables.Generator/ReturnTypeSyntaxRewriter.cs @@ -0,0 +1,35 @@ +using System; +using System.Collections.Generic; +using System.Linq; +using System.Text; +using System.Threading.Tasks; +using Microsoft.CodeAnalysis; +using Microsoft.CodeAnalysis.CSharp; +using Microsoft.CodeAnalysis.CSharp.Syntax; + +namespace EntityFrameworkCore.Projectables.Generator +{ + public class ReturnTypeSyntaxRewriter : CSharpSyntaxRewriter + { + readonly SemanticModel _semanticModel; + + public ReturnTypeSyntaxRewriter(SemanticModel semanticModel) + { + _semanticModel = semanticModel; + } + + public override SyntaxNode? VisitNullableType(NullableTypeSyntax node) + { + var typeInfo = _semanticModel.GetTypeInfo(node); + if (typeInfo.Type is not null) + { + if (typeInfo.Type.TypeKind is not TypeKind.Struct) + { + return Visit(node.ElementType); + } + } + + return base.VisitNullableType(node); + } + } +} diff --git a/tests/EntityFrameworkCore.Projectables.Generator.Tests/ProjectionExpressionGeneratorTests.NullableReferenceTypesAreBeingEliminated.verified.txt b/tests/EntityFrameworkCore.Projectables.Generator.Tests/ProjectionExpressionGeneratorTests.NullableReferenceTypesAreBeingEliminated.verified.txt new file mode 100644 index 0000000..1ccc9af --- /dev/null +++ b/tests/EntityFrameworkCore.Projectables.Generator.Tests/ProjectionExpressionGeneratorTests.NullableReferenceTypesAreBeingEliminated.verified.txt @@ -0,0 +1,14 @@ +using System; +using System.Linq; +using EntityFrameworkCore.Projectables; +using Foo; + +namespace EntityFrameworkCore.Projectables.Generated +#nullable disable +{ + public static class Foo_C_NextFoo + { + public static System.Linq.Expressions.Expression> Expression => + (object? unusedArgument,int? nullablePrimitiveArgument) => null; + } +} \ No newline at end of file diff --git a/tests/EntityFrameworkCore.Projectables.Generator.Tests/ProjectionExpressionGeneratorTests.cs b/tests/EntityFrameworkCore.Projectables.Generator.Tests/ProjectionExpressionGeneratorTests.cs index 54359c1..b79208d 100644 --- a/tests/EntityFrameworkCore.Projectables.Generator.Tests/ProjectionExpressionGeneratorTests.cs +++ b/tests/EntityFrameworkCore.Projectables.Generator.Tests/ProjectionExpressionGeneratorTests.cs @@ -384,6 +384,32 @@ namespace Foo { Assert.Single(result.Diagnostics); } + [Fact] + public Task NullableReferenceTypesAreBeingEliminated() + { + var compilation = CreateCompilation(@" +using System; +using System.Linq; +using EntityFrameworkCore.Projectables; + +#nullable enable + +namespace Foo { + static class C { + [Projectable] + public static object? NextFoo(this object? unusedArgument, int? nullablePrimitiveArgument) => null; + } +} +"); + + var result = RunGenerator(compilation); + + Assert.Empty(result.Diagnostics); + Assert.Single(result.GeneratedTrees); + + return Verifier.Verify(result.GeneratedTrees[0].ToString()); + } + #region Helpers Compilation CreateCompilation(string source, bool expectedToCompile = true)