Ability to have projectable methods that accept a DbContext

This commit is contained in:
Koen Bekkenutte
2021-06-02 23:44:40 +08:00
parent 6f2b30f723
commit 3aa2c811df
8 changed files with 100 additions and 46 deletions
@@ -0,0 +1,39 @@
using System;
using System.Collections.Generic;
using System.Linq;
using System.Linq.Expressions;
using System.Text;
using System.Threading;
using System.Threading.Tasks;
using System.Transactions;
using EntityFrameworkCore.Projectables.Services;
using Microsoft.EntityFrameworkCore.Query;
using Microsoft.EntityFrameworkCore.Query.Internal;
namespace EntityFrameworkCore.Projectables.Infrastructure.Internal
{
[System.Diagnostics.CodeAnalysis.SuppressMessage("Usage", "EF1001:Internal EF Core API usage.", Justification = "Needed")]
public sealed class CustomQueryProvider : IQueryCompiler
{
readonly IQueryCompiler _decoratedQueryCompiler;
readonly ProjectableExpressionReplacer _projectableExpressionReplacer;
public CustomQueryProvider(IQueryCompiler decoratedQueryCompiler)
{
_decoratedQueryCompiler = decoratedQueryCompiler;
_projectableExpressionReplacer = new ProjectableExpressionReplacer(new ProjectionExpressionResolver());
}
public Func<QueryContext, TResult> CreateCompiledAsyncQuery<TResult>(Expression query)
=> _decoratedQueryCompiler.CreateCompiledAsyncQuery<TResult>(Expand(query));
public Func<QueryContext, TResult> CreateCompiledQuery<TResult>(Expression query)
=> _decoratedQueryCompiler.CreateCompiledQuery<TResult>(Expand(query));
public TResult Execute<TResult>(Expression query)
=> _decoratedQueryCompiler.Execute<TResult>(Expand(query));
public TResult ExecuteAsync<TResult>(Expression query, CancellationToken cancellationToken)
=> _decoratedQueryCompiler.ExecuteAsync<TResult>(Expand(query), cancellationToken);
Expression Expand(Expression expression)
=> _projectableExpressionReplacer.Visit(expression);
}
}
@@ -1,40 +0,0 @@
using System;
using System.Collections.Generic;
using System.Linq;
using System.Linq.Expressions;
using System.Text;
using System.Threading.Tasks;
using EntityFrameworkCore.Projectables.Services;
using Microsoft.EntityFrameworkCore.Query;
namespace EntityFrameworkCore.Projectables.Infrastructure.Internal
{
public class CustomQueryTranslationPreprocessorFactory : IQueryTranslationPreprocessorFactory
{
readonly IQueryTranslationPreprocessorFactory _decoratedFactory;
readonly QueryTranslationPreprocessorDependencies _queryTranslationPreprocessorDependencies;
readonly ProjectableExpressionReplacer _projectableExpressionReplacer = new(new ProjectionExpressionResolver());
public CustomQueryTranslationPreprocessorFactory(IQueryTranslationPreprocessorFactory decoratedFactory, QueryTranslationPreprocessorDependencies queryTranslationPreprocessorDependencies)
{
_decoratedFactory = decoratedFactory;
_queryTranslationPreprocessorDependencies = queryTranslationPreprocessorDependencies;
}
public QueryTranslationPreprocessor Create(QueryCompilationContext queryCompilationContext)
=> new CustomQueryTranslationPreprocessor(_queryTranslationPreprocessorDependencies, queryCompilationContext, _projectableExpressionReplacer);
}
public class CustomQueryTranslationPreprocessor : QueryTranslationPreprocessor
{
readonly ProjectableExpressionReplacer _projectableExpressionReplacer;
public CustomQueryTranslationPreprocessor(QueryTranslationPreprocessorDependencies dependencies, QueryCompilationContext queryCompilationContext, ProjectableExpressionReplacer projectableExpressionReplacer) : base(dependencies, queryCompilationContext)
{
_projectableExpressionReplacer = projectableExpressionReplacer;
}
public override Expression Process(Expression query)
=> base.Process(_projectableExpressionReplacer.Visit(query));
}
}
@@ -37,13 +37,13 @@ namespace EntityFrameworkCore.Projectables.Infrastructure.Internal
return ActivatorUtilities.GetServiceOrCreateInstance(services, descriptor.ImplementationType!);
}
var targetDescriptor = services.FirstOrDefault(x => x.ServiceType == typeof(IQueryTranslationPreprocessorFactory));
var targetDescriptor = services.FirstOrDefault(x => x.ServiceType == typeof(IQueryCompiler));
if (targetDescriptor is null)
{
throw new InvalidOperationException("No QueryTranslationPreprocessorFactory is configured yet. Please make sure to configure a database provider first"); ;
throw new InvalidOperationException("No QueryProvider is configured yet. Please make sure to configure a database provider first"); ;
}
var decoratorObjectFactory = ActivatorUtilities.CreateFactory(typeof(CustomQueryTranslationPreprocessorFactory), new [] { targetDescriptor.ServiceType });
var decoratorObjectFactory = ActivatorUtilities.CreateFactory(typeof(CustomQueryProvider), new [] { targetDescriptor.ServiceType });
services.Replace(ServiceDescriptor.Describe(
targetDescriptor.ServiceType,
@@ -52,7 +52,6 @@ namespace EntityFrameworkCore.Projectables.Infrastructure.Internal
));
}
public void Validate(IDbContextOptions options)
{
}
@@ -0,0 +1,10 @@
SELECT [t0].[Id], [t0].[RecordDate], [t0].[UserId]
FROM [User] AS [u]
LEFT JOIN (
SELECT [t].[Id], [t].[RecordDate], [t].[UserId]
FROM (
SELECT [o].[Id], [o].[RecordDate], [o].[UserId], ROW_NUMBER() OVER(PARTITION BY [o].[UserId] ORDER BY [o].[RecordDate] DESC) AS [row]
FROM [Order] AS [o]
) AS [t]
WHERE [t].[row] <= 1
) AS [t0] ON [u].[Id] = [t0].[UserId]
@@ -38,12 +38,18 @@ namespace EntityFrameworkCore.Projectables.FunctionalTests
public IEnumerable<EntityFrameworkCore.Projectables.FunctionalTests.ComplexModelTests.Order> Last2Orders =>
Orders.OrderByDescending(x => x.RecordDate).Take(2);
[Projectable]
public EntityFrameworkCore.Projectables.FunctionalTests.ComplexModelTests.Order GetLastOrderFromExternalDbContext(DbContext dbContext)
=> dbContext.Set<EntityFrameworkCore.Projectables.FunctionalTests.ComplexModelTests.Order>().Where(x => x.UserId == Id).OrderByDescending(x => x.RecordDate).FirstOrDefault();
}
public class Order
{
public int Id { get; set; }
public int UserId { get; set; }
public DateTime RecordDate { get; set; }
}
@@ -69,5 +75,16 @@ namespace EntityFrameworkCore.Projectables.FunctionalTests
return Verifier.Verify(query.ToQueryString());
}
[Fact]
public Task ProjectOverMethodTakingDbContext()
{
using var dbContext = new SampleDbContext<User>();
var query = dbContext.Set<User>()
.Select(x => x.GetLastOrderFromExternalDbContext(dbContext));
return Verifier.Verify(query.ToQueryString());
}
}
}
@@ -1,4 +1,7 @@
namespace EntityFrameworkCore.Projectables.FunctionalTests.ExtensionMethods
using System.Linq;
using Microsoft.EntityFrameworkCore;
namespace EntityFrameworkCore.Projectables.FunctionalTests.ExtensionMethods
{
public static class EntityExtensions
{
@@ -7,5 +10,9 @@
[Projectable]
public static int Foo(this Entity entity) => entity.Id + 1;
[Projectable]
public static Entity? LeadingEntity(this Entity entity, DbContext dbContext)
=> dbContext.Set<Entity>().Where(y => y.Id > entity.Id).FirstOrDefault();
}
}
@@ -0,0 +1,7 @@
SELECT [t].[Id]
FROM [Entity] AS [e]
OUTER APPLY (
SELECT TOP(1) [e0].[Id]
FROM [Entity] AS [e0]
WHERE [e0].[Id] > [e].[Id]
) AS [t]
@@ -5,6 +5,7 @@ using System.Text;
using System.Threading.Tasks;
using EntityFrameworkCore.Projectables.FunctionalTests.Helpers;
using Microsoft.EntityFrameworkCore;
using Microsoft.EntityFrameworkCore.Scaffolding.Metadata;
using ScenarioTests;
using VerifyXunit;
using Xunit;
@@ -36,5 +37,19 @@ namespace EntityFrameworkCore.Projectables.FunctionalTests.ExtensionMethods
return Verifier.Verify(query.ToQueryString());
}
[Fact]
public Task ExtensionMethodAcceptingDbContext()
{
using var dbContext = new SampleDbContext<Entity>();
var sampleQuery = dbContext.Set<Entity>()
.Select(x => dbContext.Set<Entity>().Where(y => y.Id > x.Id).FirstOrDefault());
var query = dbContext.Set<Entity>()
.Select(x => x.LeadingEntity(dbContext));
return Verifier.Verify(query.ToQueryString());
}
}
}