You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

DbContext扩展:优先从本地加载带关联数据的实体

实现优先从本地缓存加载并包含关联数据的DbContext扩展方法

在调用SaveChangesAsync前执行领域事件预处理时,实体已存在于DbContext的本地缓存但未持久化,此时直接查询数据库无意义。需要实现一个扩展方法,优先从本地缓存加载指定ID的实体,本地不存在时再查询数据库,同时支持加载指定的关联数据。

无关联数据的基础实现很简单,但要支持Include关联数据时,本地缓存的实体无法自动加载关联数据,需要手动处理。


实现思路

  1. 支持链式调用:将扩展方法定义在IQueryable<TEntity>上,承接Include后的查询对象,符合常规EF Core调用习惯。
  2. 提取Include导航路径:从传入的IQueryable中解析出所有通过Include指定的导航属性。
  3. 本地实体关联数据加载:如果从本地缓存找到实体,通过DbContext的Entry API手动加载对应的关联数据。
  4. 类型安全的ID匹配:用表达式树实现实体Id属性的访问,兼顾通用性与性能。

完整代码实现

using Microsoft.EntityFrameworkCore;
using Microsoft.EntityFrameworkCore.Query;
using System.Linq.Expressions;
using System.Reflection;

public static class QueryableExtensions
{
    public static async Task<TEntity?> GetByIdLocallyFirstAsync<TEntity, TKey>(
        this IQueryable<TEntity> queryable, 
        TKey id, 
        CancellationToken cancellationToken = default)
        where TEntity : class
    {
        // 从查询对象中获取DbContext实例
        var dbContext = GetDbContext(queryable);
        var dbSet = dbContext.Set<TEntity>();
        
        // 获取实体Id属性的访问器(表达式树编译,避免反射性能损耗)
        var getIdFunc = GetIdPropertyAccessor<TEntity, TKey>();
        
        // 优先从本地缓存查找实体
        var localEntity = dbSet.Local.FirstOrDefault(e => getIdFunc(e).Equals(id));
        
        if (localEntity != null)
        {
            // 解析查询中所有Include的导航属性
            var includedNavigations = GetIncludedNavigations(queryable);
            
            // 手动加载关联数据
            foreach (var navigationPath in includedNavigations)
            {
                await LoadNavigationPropertyAsync(dbContext.Entry(localEntity), navigationPath, cancellationToken);
            }
            
            return localEntity;
        }
        
        // 本地未找到,从数据库执行查询
        return await queryable.FirstOrDefaultAsync(e => getIdFunc(e).Equals(id), cancellationToken);
    }
    
    // 从IQueryable中提取DbContext实例
    private static DbContext GetDbContext<TEntity>(IQueryable<TEntity> queryable)
    {
        var queryCompilerProperty = queryable.Provider.GetType().GetProperty("QueryCompiler", BindingFlags.NonPublic | BindingFlags.Instance);
        if (queryCompilerProperty == null)
            throw new InvalidOperationException("无法从查询提供器中获取DbContext");
            
        var queryCompiler = queryCompilerProperty.GetValue(queryable.Provider);
        var dbContextProperty = queryCompiler.GetType().GetProperty("Context", BindingFlags.NonPublic | BindingFlags.Instance);
        if (dbContextProperty == null)
            throw new InvalidOperationException("无法从查询编译器中获取DbContext");
            
        return (DbContext)dbContextProperty.GetValue(queryCompiler);
    }
    
    // 生成实体Id属性的访问委托
    private static Func<TEntity, TKey> GetIdPropertyAccessor<TEntity, TKey>()
    {
        var idProperty = typeof(TEntity).GetProperty("Id", typeof(TKey));
        if (idProperty == null)
            throw new InvalidOperationException($"实体类型 {typeof(TEntity)} 不存在类型为 {typeof(TKey)} 的Id属性");
            
        var parameter = Expression.Parameter(typeof(TEntity), "e");
        var propertyAccess = Expression.Property(parameter, idProperty);
        return Expression.Lambda<Func<TEntity, TKey>>(propertyAccess, parameter).Compile();
    }
    
    // 解析查询中所有Include的导航路径
    private static List<string> GetIncludedNavigations<TEntity>(IQueryable<TEntity> queryable)
    {
        var includedNavigations = new List<string>();
        var query = queryable.Expression;
        
        while (query is MethodCallExpression methodCall)
        {
            if (methodCall.Method.Name == nameof(EntityFrameworkQueryableExtensions.Include) && methodCall.Arguments.Count == 2)
            {
                var navigationExpr = methodCall.Arguments[1] as LambdaExpression;
                if (navigationExpr != null)
                {
                    includedNavigations.Add(GetNavigationPath(navigationExpr.Body));
                }
            }
            // 如需支持ThenInclude多级关联,可在此扩展解析逻辑
            query = methodCall.Object;
        }
        
        return includedNavigations;
    }
    
    // 从表达式中提取导航属性路径字符串
    private static string GetNavigationPath(Expression expression)
    {
        switch (expression)
        {
            case MemberExpression memberExpr:
                var parentPath = memberExpr.Expression != null ? GetNavigationPath(memberExpr.Expression) : string.Empty;
                return string.IsNullOrEmpty(parentPath) ? memberExpr.Member.Name : $"{parentPath}.{memberExpr.Member.Name}";
            case UnaryExpression unaryExpr:
                return GetNavigationPath(unaryExpr.Operand);
            default:
                throw new InvalidOperationException("无法解析导航属性表达式");
        }
    }
    
    // 异步加载指定的导航属性
    private static async Task LoadNavigationPropertyAsync(EntityEntry entry, string navigationPath, CancellationToken cancellationToken)
    {
        var navigationEntry = entry.Navigation(navigationPath);
        if (!navigationEntry.IsLoaded)
        {
            await navigationEntry.LoadAsync(cancellationToken);
        }
    }
}

使用示例

var offer = await _dbContext.Offers
    .Include(x => x.Commodity)
    .Include(x => x.ContractType)
    .Include(x => x.Customer)
    .Include(x => x.Employee)
    .Include(x => x.DeliveryLocation)
    .Include(x => x.Location)
    .GetByIdLocallyFirstAsync(notification.OfferId, cancellationToken);

关键说明

  • DbContext自动提取:无需手动传入DbContext,通过反射从查询对象中自动获取,简化调用。
  • 关联数据加载逻辑:判断导航属性是否已加载,未加载时调用LoadAsync,如果关联数据存在于本地缓存则直接使用,否则从数据库加载。
  • 性能优化:用表达式树编译Id属性访问委托,比直接反射调用性能更高。

内容的提问来源于stack exchange,提问作者nop

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.14 03:21:15