DbContext扩展:优先从本地加载带关联数据的实体
实现优先从本地缓存加载并包含关联数据的DbContext扩展方法
在调用SaveChangesAsync前执行领域事件预处理时,实体已存在于DbContext的本地缓存但未持久化,此时直接查询数据库无意义。需要实现一个扩展方法,优先从本地缓存加载指定ID的实体,本地不存在时再查询数据库,同时支持加载指定的关联数据。
无关联数据的基础实现很简单,但要支持Include关联数据时,本地缓存的实体无法自动加载关联数据,需要手动处理。
实现思路
- 支持链式调用:将扩展方法定义在
IQueryable<TEntity>上,承接Include后的查询对象,符合常规EF Core调用习惯。 - 提取Include导航路径:从传入的
IQueryable中解析出所有通过Include指定的导航属性。 - 本地实体关联数据加载:如果从本地缓存找到实体,通过DbContext的
EntryAPI手动加载对应的关联数据。 - 类型安全的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
相关产品推荐
相关产品推荐

