From 1d3c9eb509f5de362c5a123dfd9e2c56ef8276ca Mon Sep 17 00:00:00 2001 From: Koen Date: Sun, 26 Jul 2026 01:23:40 +0000 Subject: [PATCH] support ThenBy/ThenByDescending after Include/ThenInclude --- .../Internal/ExpressiveOptionsExtension.cs | 3 +- .../Transformers/RewriteThenByAfterInclude.cs | 71 +++++++++++++++++++ .../PolyfillInterceptorGenerator.cs | 24 ++++--- .../ExpressiveQueryableExtensions.cs | 35 +++++++++ .../Infrastructure/IncludeTestBase.cs | 45 ++++++++++++ ...enBy_CastsToIOrderedQueryable.verified.txt | 2 +- ...scending_GeneratesInterceptor.verified.txt | 2 +- 7 files changed, 168 insertions(+), 14 deletions(-) create mode 100644 src/ExpressiveSharp.EntityFrameworkCore/Transformers/RewriteThenByAfterInclude.cs diff --git a/src/ExpressiveSharp.EntityFrameworkCore/Infrastructure/Internal/ExpressiveOptionsExtension.cs b/src/ExpressiveSharp.EntityFrameworkCore/Infrastructure/Internal/ExpressiveOptionsExtension.cs index 1541123a..5daa5b8c 100644 --- a/src/ExpressiveSharp.EntityFrameworkCore/Infrastructure/Internal/ExpressiveOptionsExtension.cs +++ b/src/ExpressiveSharp.EntityFrameworkCore/Infrastructure/Internal/ExpressiveOptionsExtension.cs @@ -94,7 +94,8 @@ public void ApplyServices(IServiceCollection services) new RemoveNullConditionalPatterns(), new FlattenTupleComparisons(), new FlattenConcatArrayCalls(), - new FlattenBlockExpressions()); + new FlattenBlockExpressions(), + new Transformers.RewriteThenByAfterInclude()); if (extraTransformers.Length > 0) options.AddTransformers(extraTransformers); return options; diff --git a/src/ExpressiveSharp.EntityFrameworkCore/Transformers/RewriteThenByAfterInclude.cs b/src/ExpressiveSharp.EntityFrameworkCore/Transformers/RewriteThenByAfterInclude.cs new file mode 100644 index 00000000..64d75b53 --- /dev/null +++ b/src/ExpressiveSharp.EntityFrameworkCore/Transformers/RewriteThenByAfterInclude.cs @@ -0,0 +1,71 @@ +using System.Linq.Expressions; +using Microsoft.EntityFrameworkCore; + +namespace ExpressiveSharp.EntityFrameworkCore.Transformers; + +/// +/// Rewrites ThenBy/ThenByDescending applied after Include/ThenInclude +/// into the equivalent tree EF Core can translate: the ordering is applied to the ordered source +/// beneath the include chain. The C# type system cannot express this shape directly, so the +/// generated interceptors produce a cast node () +/// that this transformer resolves before EF Core sees the query. +/// +public sealed class RewriteThenByAfterInclude : ExpressionVisitor, IExpressionTreeTransformer +{ + public Expression Transform(Expression expression) => Visit(expression); + + protected override Expression VisitMethodCall(MethodCallExpression node) + { + if (node.Method.DeclaringType != typeof(Queryable) + || node.Method.Name is not (nameof(Queryable.ThenBy) or nameof(Queryable.ThenByDescending))) + return base.VisitMethodCall(node); + + var arguments = new Expression[node.Arguments.Count]; + for (var i = 0; i < node.Arguments.Count; i++) + arguments[i] = Visit(node.Arguments[i]); + + var source = arguments[0]; + if (source is UnaryExpression { NodeType: ExpressionType.Convert } convert + && typeof(IOrderedQueryable).IsAssignableFrom(convert.Type)) + { + source = convert.Operand; + } + + var includes = new List(); + var baseSource = source; + while (baseSource is MethodCallExpression call && IsIncludeCall(call)) + { + includes.Add(call); + baseSource = call.Arguments[0]; + } + + if (includes.Count == 0) + { + return node.Update(node.Object, arguments); + } + + if (!typeof(IOrderedQueryable).IsAssignableFrom(baseSource.Type)) + { + throw new InvalidOperationException( + $"'{node.Method.Name}' was called after 'Include'/'ThenInclude' on a source that " + + "is not ordered. Call 'OrderBy'/'OrderByDescending' before 'Include', or apply " + + "the complete ordering after the includes."); + } + + arguments[0] = baseSource; + Expression current = Expression.Call(node.Method, arguments); + for (var i = includes.Count - 1; i >= 0; i--) + { + var includeArguments = includes[i].Arguments.ToArray(); + includeArguments[0] = current; + current = includes[i].Update(null, includeArguments); + } + + return current; + } + + private static bool IsIncludeCall(MethodCallExpression call) + => call.Method.DeclaringType == typeof(EntityFrameworkQueryableExtensions) + && call.Method.Name is nameof(EntityFrameworkQueryableExtensions.Include) + or nameof(EntityFrameworkQueryableExtensions.ThenInclude); +} diff --git a/src/ExpressiveSharp.Generator/PolyfillInterceptorGenerator.cs b/src/ExpressiveSharp.Generator/PolyfillInterceptorGenerator.cs index bbd2391e..0f33986e 100644 --- a/src/ExpressiveSharp.Generator/PolyfillInterceptorGenerator.cs +++ b/src/ExpressiveSharp.Generator/PolyfillInterceptorGenerator.cs @@ -619,7 +619,7 @@ private static string MethodId(string op, string fileTag, int line, int col) var typeAliases = new Dictionary(SymbolEqualityComparer.Default); var delegateFqns = new string[funcParamIndices.Count]; string elemRef; - string castFqn; + string sourceRef; string typeParams; string returnRef; string interceptorParamList; @@ -667,9 +667,11 @@ private static string MethodId(string op, string fileTag, int line, int col) delegateFqns[fi] = funcFqnGenerics[fi]; } - castFqn = isOrdered - ? $"global::System.Linq.IOrderedQueryable<{elemRef}>" - : $"global::System.Linq.IQueryable<{elemRef}>"; + // ThenBy/ThenByDescending go through AsOrdered: a plain cast fails when the source's + // expression is not statically ordered (e.g. an EF Core Include after OrderBy). + sourceRef = isOrdered + ? $"global::ExpressiveSharp.ExpressiveQueryableExtensions.AsOrdered<{elemRef}>((global::System.Linq.IQueryable<{elemRef}>)(object)source)" + : $"(global::System.Linq.IQueryable<{elemRef}>)(object)source"; if (isRewritableReturn) { @@ -718,9 +720,9 @@ private static string MethodId(string op, string fileTag, int line, int col) t.ToDisplayString(SymbolDisplayFormat.FullyQualifiedFormat))) + ">"; } - castFqn = isOrdered - ? $"global::System.Linq.IOrderedQueryable<{elemFqn}>" - : $"global::System.Linq.IQueryable<{elemFqn}>"; + sourceRef = isOrdered + ? $"global::ExpressiveSharp.ExpressiveQueryableExtensions.AsOrdered<{elemFqn}>(source)" + : $"(global::System.Linq.IQueryable<{elemFqn}>)source"; returnRef = isRewritableReturn ? returnElemType!.ToDisplayString(SymbolDisplayFormat.FullyQualifiedFormat) @@ -781,7 +783,7 @@ private static string MethodId(string op, string fileTag, int line, int col) { {{allBodies}} return global::ExpressiveSharp.ExpressiveQueryableExtensions.AsExpressive( {{targetTypeFqn}}.{{methodName}}( - ({{castFqn}})source, + {{sourceRef}}, {{queryableArgList}})); } @@ -795,7 +797,7 @@ private static string MethodId(string op, string fileTag, int line, int col) {{interceptorParamList}}) { {{allBodies}} return {{targetTypeFqn}}.{{methodName}}( - ({{castFqn}})source, + {{sourceRef}}, {{queryableArgList}}); } @@ -813,7 +815,7 @@ private static string MethodId(string op, string fileTag, int line, int col) {{allBodies}} return (global::ExpressiveSharp.IExpressiveQueryable<{{returnRef}}>)(object) global::ExpressiveSharp.ExpressiveQueryableExtensions.AsExpressive( {{targetTypeFqn}}.{{methodName}}( - ({{castFqn}})(object)source, + {{sourceRef}}, {{queryableArgList}})); } @@ -827,7 +829,7 @@ private static string MethodId(string op, string fileTag, int line, int col) {{interceptorParamList}}) { {{allBodies}} return {{targetTypeFqn}}.{{methodName}}( - ({{castFqn}})(object)source, + {{sourceRef}}, {{queryableArgList}}); } diff --git a/src/ExpressiveSharp/Extensions/ExpressiveQueryableExtensions.cs b/src/ExpressiveSharp/Extensions/ExpressiveQueryableExtensions.cs index c623d678..af313eba 100644 --- a/src/ExpressiveSharp/Extensions/ExpressiveQueryableExtensions.cs +++ b/src/ExpressiveSharp/Extensions/ExpressiveQueryableExtensions.cs @@ -1,4 +1,7 @@ +using System.Collections; +using System.Collections.Generic; using System.Linq; +using System.Linq.Expressions; namespace ExpressiveSharp { @@ -20,5 +23,37 @@ public static IExpressiveQueryable AsExpressive( // etc.) remain observable through the returned reference. => source as IExpressiveQueryable ?? new ExpressiveQueryableWrapper(source); + + /// + /// Presents a queryable as for the generated + /// ThenBy/ThenByDescending interceptors. When the underlying expression's + /// static type is not ordered (e.g. an EF Core Include call following + /// OrderBy), the expression is wrapped in a cast node so + /// can compose a valid tree. + /// + public static IOrderedQueryable AsOrdered(this IQueryable source) + => source is IOrderedQueryable ordered + && typeof(IOrderedQueryable).IsAssignableFrom(source.Expression.Type) + ? ordered + : new OrderedQueryableAdapter(source); + + private sealed class OrderedQueryableAdapter : IOrderedQueryable + { + private readonly IQueryable _source; + + public OrderedQueryableAdapter(IQueryable source) + { + _source = source; + Expression = typeof(IOrderedQueryable).IsAssignableFrom(source.Expression.Type) + ? source.Expression + : Expression.Convert(source.Expression, typeof(IOrderedQueryable)); + } + + public Expression Expression { get; } + public Type ElementType => _source.ElementType; + public IQueryProvider Provider => _source.Provider; + public IEnumerator GetEnumerator() => _source.GetEnumerator(); + IEnumerator IEnumerable.GetEnumerator() => ((IEnumerable)_source).GetEnumerator(); + } } } diff --git a/tests/ExpressiveSharp.EntityFrameworkCore.IntegrationTests/Infrastructure/IncludeTestBase.cs b/tests/ExpressiveSharp.EntityFrameworkCore.IntegrationTests/Infrastructure/IncludeTestBase.cs index 1233f40c..508764bb 100644 --- a/tests/ExpressiveSharp.EntityFrameworkCore.IntegrationTests/Infrastructure/IncludeTestBase.cs +++ b/tests/ExpressiveSharp.EntityFrameworkCore.IntegrationTests/Infrastructure/IncludeTestBase.cs @@ -1,3 +1,4 @@ +using ExpressiveSharp.IntegrationTests.Scenarios.Store; using ExpressiveSharp.IntegrationTests.Scenarios.Store.Models; using Microsoft.EntityFrameworkCore; using Microsoft.VisualStudio.TestTools.UnitTesting; @@ -105,4 +106,48 @@ public async Task ChainedModifiers_Include_Execute() Assert.AreEqual(1, results.Count); Assert.IsNotNull(results[0].Customer); } + + [TestMethod] + public async Task ThenBy_AfterInclude_ExecutesWithComposedOrdering() + { + var results = await Context.Orders.AsExpressiveDbSet() + .OrderBy(o => o.Status) + .Include(o => o.Customer) + .ThenBy(o => o.Id) + .ToListAsync(); + + var expectedIds = SeedData.Orders + .OrderBy(o => o.Status).ThenBy(o => o.Id) + .Select(o => o.Id).ToList(); + CollectionAssert.AreEqual(expectedIds, results.Select(r => r.Id).ToList()); + Assert.IsNotNull(results.Single(r => r.Id == 1).Customer); + } + + [TestMethod] + public async Task ThenBy_AfterIncludeThenInclude_ExecutesWithComposedOrdering() + { + var results = await Context.Orders.AsExpressiveDbSet() + .OrderBy(o => o.Status) + .Include(o => o.Customer) + .ThenInclude(c => c!.Address) + .ThenBy(o => o.Id) + .ToListAsync(); + + var expectedIds = SeedData.Orders + .OrderBy(o => o.Status).ThenBy(o => o.Id) + .Select(o => o.Id).ToList(); + CollectionAssert.AreEqual(expectedIds, results.Select(r => r.Id).ToList()); + Assert.IsNotNull(results.Single(r => r.Id == 1).Customer!.Address); + } + + [TestMethod] + public async Task ThenBy_AfterInclude_WithoutOrderBy_ThrowsActionableError() + { + var query = Context.Orders.AsExpressiveDbSet() + .Include(o => o.Customer) + .ThenBy(o => o.Id); + + var ex = await Assert.ThrowsExactlyAsync(() => query.ToListAsync()); + StringAssert.Contains(ex.Message, "OrderBy"); + } } diff --git a/tests/ExpressiveSharp.Generator.Tests/PolyfillInterceptorGenerator/OrderByTests.ThenBy_CastsToIOrderedQueryable.verified.txt b/tests/ExpressiveSharp.Generator.Tests/PolyfillInterceptorGenerator/OrderByTests.ThenBy_CastsToIOrderedQueryable.verified.txt index fd5dca7a..5cdb70d4 100644 --- a/tests/ExpressiveSharp.Generator.Tests/PolyfillInterceptorGenerator/OrderByTests.ThenBy_CastsToIOrderedQueryable.verified.txt +++ b/tests/ExpressiveSharp.Generator.Tests/PolyfillInterceptorGenerator/OrderByTests.ThenBy_CastsToIOrderedQueryable.verified.txt @@ -16,7 +16,7 @@ namespace ExpressiveSharp.Generated.Interceptors var __lambda = global::System.Linq.Expressions.Expression.Lambda>(i361d11c18_expr_0, i361d11c18_p_o); return global::ExpressiveSharp.ExpressiveQueryableExtensions.AsExpressive( global::System.Linq.Queryable.ThenBy( - (global::System.Linq.IOrderedQueryable)source, + global::ExpressiveSharp.ExpressiveQueryableExtensions.AsOrdered(source), __lambda)); } [global::System.Runtime.CompilerServices.InterceptsLocationAttribute(/* scrubbed */)] diff --git a/tests/ExpressiveSharp.Generator.Tests/PolyfillInterceptorGenerator/SingleLambdaQueryableTests.ThenByDescending_GeneratesInterceptor.verified.txt b/tests/ExpressiveSharp.Generator.Tests/PolyfillInterceptorGenerator/SingleLambdaQueryableTests.ThenByDescending_GeneratesInterceptor.verified.txt index cb54d4b7..04983437 100644 --- a/tests/ExpressiveSharp.Generator.Tests/PolyfillInterceptorGenerator/SingleLambdaQueryableTests.ThenByDescending_GeneratesInterceptor.verified.txt +++ b/tests/ExpressiveSharp.Generator.Tests/PolyfillInterceptorGenerator/SingleLambdaQueryableTests.ThenByDescending_GeneratesInterceptor.verified.txt @@ -16,7 +16,7 @@ namespace ExpressiveSharp.Generated.Interceptors var __lambda = global::System.Linq.Expressions.Expression.Lambda>(i361d11c18_expr_0, i361d11c18_p_o); return global::ExpressiveSharp.ExpressiveQueryableExtensions.AsExpressive( global::System.Linq.Queryable.ThenByDescending( - (global::System.Linq.IOrderedQueryable)source, + global::ExpressiveSharp.ExpressiveQueryableExtensions.AsOrdered(source), __lambda)); } [global::System.Runtime.CompilerServices.InterceptsLocationAttribute(/* scrubbed */)]