From 3f876d3d70870573439e72dd82a4c097fd8744f1 Mon Sep 17 00:00:00 2001 From: Simon Cropp Date: Sat, 26 Sep 2026 21:42:20 +1000 Subject: [PATCH] Treat lookups by key as bounded, and skip Includes EF ignores RejectUnbounded no longer fires for a Where that compares a primary or alternate key with a value, such as Where(_ => _.Id == id), since it returns at most one row. A key looked up in a list, such as ids.Contains(_.Id), is only bounded when MaxInValues is set, and each set of levels decides that with its own MaxInValues. A lookup only counts on the rows of a DbSet, since after Concat, SelectMany, Join, Select or FromSql the same key can be in many rows. MaxSingleQueryCollections no longer counts an Include that a later Select or aggregate makes Entity Framework ignore, since the query returns no entity for it to load into. A projection that could return an entity keeps the Includes counted. --- claude.md | 6 +- readme.md | 20 +- src/EfQueryComplexity/CollectionCounter.cs | 184 +++++++++++- src/EfQueryComplexity/KeyLookup.cs | 241 ++++++++++++++++ .../QueryComplexityLimits.cs | 5 +- .../QueryComplexityOverride.cs | 4 +- src/EfQueryComplexity/QueryShape.cs | 3 +- src/EfQueryComplexity/RowOperators.cs | 48 +++ src/EfQueryComplexity/ShapeAnalyzer.cs | 5 +- src/EfQueryComplexity/UnboundedDetector.cs | 48 ++- src/EfQueryComplexity/UnboundedEntities.cs | 10 +- src/EfQueryComplexity/Violations.cs | 5 +- src/Tests/KeyContext.cs | 51 ++++ src/Tests/KeyLookupTests.cs | 273 ++++++++++++++++++ src/Tests/ShapeTests.cs | 96 ++++++ src/Tests/UnboundedTests.cs | 26 +- 16 files changed, 979 insertions(+), 46 deletions(-) create mode 100644 src/EfQueryComplexity/KeyLookup.cs create mode 100644 src/EfQueryComplexity/RowOperators.cs create mode 100644 src/Tests/KeyContext.cs create mode 100644 src/Tests/KeyLookupTests.cs diff --git a/claude.md b/claude.md index 1e6856f..56a827f 100644 --- a/claude.md +++ b/claude.md @@ -40,8 +40,10 @@ The service provider alone is not enough. Entity Framework keys a compiled query | `QueryComplexityOptionsExtension.cs` | Holds the levels, and keys the internal service provider | | `QueryInterceptor.cs` | `IQueryExpressionInterceptor`: measures shape and strips markers, once per compiled shape | | `ShapeAnalyzer.cs` | One pass measuring nodes, depth, operators, navigations and includes | -| `CollectionCounter.cs` | The collections one SQL query loads, through collection includes and projections. A split query counts none | -| `UnboundedDetector.cs` | The types of the rows a query can return without a limit, found once per compiled shape | +| `CollectionCounter.cs` | The collections one SQL query loads, through collection includes and projections. A split query counts none, and nor does an Include that Entity Framework ignores, because the query returns no entity for it to load into | +| `UnboundedDetector.cs` | The types of the rows a query can return without a limit, found once per compiled shape, with lists limited and without | +| `KeyLookup.cs` | Whether a `Where` looks its rows up by primary or alternate key, which bounds it like a `Take`. A key in a list is only bounded when `MaxInValues` is set | +| `RowOperators.cs` | The operators that return some of their source's rows, unchanged: the ones allowed between a `DbSet` and a lookup by key, and between an ignored Include and what ignores it | | `UnboundedEntities.cs` | Which of those types `RejectUnbounded` checks: `All`, `None`, `AllExcept`, `Only` | | `Sequences.cs` | Whether a type is a sequence, and what it holds | | `Markers.cs`, `MarkerReader.cs` | The per query marker calls, and reading (`Read`) and removing (`Strip`) them | diff --git a/readme.md b/readme.md index bc31c57..b807c27 100644 --- a/readme.md +++ b/readme.md @@ -126,7 +126,7 @@ A query is checked against the throw levels before the log levels, so a throw le | `MaxSingleQueryCollections` | Collections one SQL query loads | While compiled | 1 | | `MaxTake` | The value passed to `Take` | Every execution | 1000 | | `MaxInValues` | Values in the largest list the query sends | Every execution | 1000 | -| `RejectUnbounded` | A query returning rows with no `Take` | While compiled | `All` | +| `RejectUnbounded` | A query returning rows with no `Take` or lookup by key | While compiled | `All` | A check fires when the measured value is greater than the level. A level of `null` turns that check off. @@ -145,23 +145,35 @@ It does not count: * Reference navigations. * A collection only read by an aggregate, like `_.Employees.Count()` or `_.Employees.Any()`, which is a subquery rather than a join. * Any collection in a split query, from `AsSplitQuery()` or `UseQuerySplittingBehavior(QuerySplittingBehavior.SplitQuery)`, since each collection is then loaded by its own query. `AsSingleQuery()` overrides the default. + * An `Include` that Entity Framework ignores, because the query returns no entity for it to load into. A `Select` returning no entity, like `Select(_ => _.Name)` or `Select(_ => new { _.Name })`, or an aggregate like `Count()`, makes Entity Framework ignore the Includes before it. A projection that could return an entity keeps them counted, including one that passes the entity to a method. The log default of 1 matches the point where Entity Framework logs `MultipleCollectionIncludeWarning`. That warning only covers `Include`, and only when no splitting behavior is configured. ### Unbounded queries -A query is bounded when it cannot return more rows than a `Take` allows: +A query is bounded when it cannot return more rows than a `Take`, or a [lookup by key](#lookups-by-key), allows: - A query that returns one row, an aggregate or a count is bounded, so `First`, `Single`, `Count`, `Any`, `Sum` and friends never fire. -- `Take` bounds everything below it. -- `SelectMany`, `Join`, `GroupJoin`, `LeftJoin`, `RightJoin` and `Zip` return more rows than their source, so a `Take` below one of them bounds the source rather than the query. +- `Take` bounds everything below it, and so does a lookup by key. +- `SelectMany`, `Join`, `GroupJoin`, `LeftJoin`, `RightJoin` and `Zip` return more rows than their source, so a `Take` or a lookup below one of them bounds the source rather than the query. - `Concat` and `Union` are bounded only when both sides are. - Every other operator returns no more rows than its source. The message names the types of the rows returned without a `Take`, and so does `QueryComplexityViolation.RowTypes`. A row type is the entity a query reads, not what it projects to, so `Employees.Select(_ => _.Name)` returns `Employee` rows. A query that joins in another sequence returns rows of both types: `Departments.SelectMany(_ => _.Employees)` returns `Department` and `Employee` rows. +### Lookups by key + +A `Where` that compares the key of its rows with a value, like `Employees.Where(_ => _.Id == id)`, returns at most one row, so it needs no `Take`: + +- Other conditions can be added with `&&`. With `||`, each side has to be a lookup. +- A composite key needs every part compared. +- The key is the primary key or an alternate key. A unique index does not count: a filter can make it unique among only some of the rows, and a column that allows null can hold null in many rows. +- `ids.Contains(_.Id)` returns a row for each value in `ids`, so it is only bounded when `MaxInValues` is set, which limits that list. The log and throw levels each decide this with their own `MaxInValues`. For a composite key, one part can be looked up in a list and the rest compared. +- Only the rows of a `DbSet` can be looked up by their key. Between the `DbSet` and the `Where` there can be filters, ordering, `Skip`, `Take`, `Distinct`, `OfType`, and options such as `Include` and `AsNoTracking`. After a `Select`, a `SelectMany`, a `Join` or a `Concat`, the same key can be in many rows, as it can in the rows `FromSql` returns. + + ### Choosing the types to check Some apps have no large table at all. An admin or workflow app where every table holds hundreds or thousands of rows can return all of them, and this check only reports queries that are fine. Turn it off, and keep the rest: diff --git a/src/EfQueryComplexity/CollectionCounter.cs b/src/EfQueryComplexity/CollectionCounter.cs index 7d90964..f6971f6 100644 --- a/src/EfQueryComplexity/CollectionCounter.cs +++ b/src/EfQueryComplexity/CollectionCounter.cs @@ -6,7 +6,8 @@ /// A single query joins every collection it loads, so each multiplies the rows returned for the /// others, a cartesian explosion. A split query loads each collection in its own query, so counts /// none. A collection only read by an aggregate, like _.Employees.Count(), is a subquery rather -/// than a join, so is not counted. +/// than a join, so is not counted. Nor is an Include that Entity Framework ignores, since the query +/// returns no entity for it to load into. /// sealed class CollectionCounter(IModel model) : ExpressionVisitor @@ -14,6 +15,7 @@ sealed class CollectionCounter(IModel model) : // Keyed on the full path from the root, so an Include chain that restates a collection, to // ThenInclude something else below it, counts that collection once HashSet includePaths = []; + HashSet ignoredIncludes = []; int projected; public static int Count(Expression query, IModel model, bool splitByDefault) @@ -56,24 +58,101 @@ protected override Expression VisitMethodCall(MethodCallExpression node) var method = node.Method; var declaringType = method.DeclaringType; - if (declaringType == typeof(EntityFrameworkQueryableExtensions) && - method.Name is "Include" or "ThenInclude") + if (IsInclude(method)) + { + if (!ignoredIncludes.Contains(node)) + { + AddIncludePaths(node); + } + + return base.VisitMethodCall(node); + } + + if (declaringType != typeof(Queryable) && + declaringType != typeof(Enumerable)) { - AddIncludePaths(node); return base.VisitMethodCall(node); } - if ((declaringType == typeof(Queryable) || declaringType == typeof(Enumerable)) && - method.Name == "Select") + if (method.Name == "Select") { + var selector = node.Arguments[1]; + if (!ReturnsEntities(selector)) + { + IgnoreIncludes(node.Arguments[0]); + } + Visit(node.Arguments[0]); - new ProjectionCounter(this).Visit(node.Arguments[1]); + new ProjectionCounter(this).Visit(selector); return node; } + // An aggregate, like Count or Any, returns a value rather than entities + if (node.Arguments is [var source, ..] && + !Sequences.IsSequence(node.Type) && + !HoldsEntities(node.Type)) + { + IgnoreIncludes(source); + } + return base.VisitMethodCall(node); } + static bool IsInclude(MethodInfo method) => + method.DeclaringType == typeof(EntityFrameworkQueryableExtensions) && + method.Name is "Include" or "ThenInclude"; + + // An Include only loads into the entities a query returns, so Entity Framework ignores the ones + // before an operator that returns none. Walking stops at an operator that changes the rows, since + // an Include below it can load into what that operator returns. + void IgnoreIncludes(Expression source) + { + var current = source; + while (current is MethodCallExpression call && + RowOperators.KeepsRows(call.Method)) + { + if (IsInclude(call.Method)) + { + ignoredIncludes.Add(call); + } + + current = call.Arguments[0]; + } + } + + bool ReturnsEntities(Expression selector) + { + // Queryable takes the selector as an expression, and Enumerable, inside a lambda, as a delegate + if (selector is UnaryExpression {NodeType: ExpressionType.Quote} quote) + { + selector = quote.Operand; + } + + // A method group cannot be looked into, so it could return anything + if (selector is not LambdaExpression lambda) + { + return true; + } + + var finder = new EntityFinder(this); + finder.Visit(lambda.Body); + return finder.Found; + } + + // Whether a value is, or holds, entities. A shared type counts, since any entity could use it. + bool HoldsEntities(Type type) + { + var element = Sequences.ElementType(type); + if (element.IsValueType || + element == typeof(string)) + { + return false; + } + + return model.IsShared(element) || + model.FindEntityType(element) != null; + } + void AddIncludePaths(MethodCallExpression node) { var path = new List(); @@ -295,4 +374,95 @@ void VisitSelectors(Expression node) } } } + + /// + /// Finds an entity a projection can return, which the Includes before it would load into. + /// + /// + /// An entity is not returned when a member is read from it, when it is compared, or when a LINQ + /// operator reads it. What those return is checked where it is used. An entity anywhere else, + /// such as in a constructor or passed to a method, could be returned. + /// + sealed class EntityFinder(CollectionCounter counter) : + ExpressionVisitor + { + public bool Found { get; private set; } + + public override Expression? Visit(Expression? node) + { + if (node != null && + counter.HoldsEntities(node.Type)) + { + Found = true; + return node; + } + + return base.Visit(node); + } + + // The parameters are rows coming in, not what the lambda returns + protected override Expression VisitLambda(Expression node) + { + Visit(node.Body); + return node; + } + + protected override Expression VisitMember(MemberExpression node) + { + if (node.Expression != null) + { + Read(node.Expression); + } + + return node; + } + + protected override Expression VisitBinary(BinaryExpression node) + { + if (node.NodeType is ExpressionType.Equal or ExpressionType.NotEqual) + { + Read(node.Left); + Read(node.Right); + return node; + } + + return base.VisitBinary(node); + } + + // A LINQ operator, or EF.Property, reads its source rather than returning it + protected override Expression VisitMethodCall(MethodCallExpression node) + { + var declaringType = node.Method.DeclaringType; + if (node.Arguments is [var source, ..] && + (declaringType == typeof(Queryable) || + declaringType == typeof(Enumerable) || + ShapeAnalyzer.IsProperty(node))) + { + Read(source); + foreach (var argument in node.Arguments.Skip(1)) + { + Visit(argument); + } + + return node; + } + + return base.VisitMethodCall(node); + } + + // Looks into a value that is read rather than returned. A cast, such as the one reaching a + // member of a derived type, is read along with it. + void Read(Expression value) + { + while (value is UnaryExpression + { + NodeType: ExpressionType.Convert or ExpressionType.ConvertChecked or ExpressionType.TypeAs + } cast) + { + value = cast.Operand; + } + + base.Visit(value); + } + } } diff --git a/src/EfQueryComplexity/KeyLookup.cs b/src/EfQueryComplexity/KeyLookup.cs new file mode 100644 index 0000000..47630b9 --- /dev/null +++ b/src/EfQueryComplexity/KeyLookup.cs @@ -0,0 +1,241 @@ +/// +/// Whether a Where looks its rows up by key, so returns no more rows than the values it compares the +/// key with. +/// +/// +/// A key is the primary key or an alternate key. A unique index is not used, since a filter can make +/// it unique among only some of the rows, and a column that allows null can hold null in many rows. +/// +static class KeyLookup +{ + /// A call to Where. + /// + /// Whether the size of a list the query sends is limited. A key looked up in a list returns a row + /// for each value in it, so only bounds the query when the list is limited. + /// + public static bool Bounds(MethodCallExpression where, bool listsLimited) + { + // Queryable takes the predicate as an expression, and Enumerable, inside a lambda, as a delegate + var predicate = where.Arguments[1]; + if (predicate is UnaryExpression {NodeType: ExpressionType.Quote} quote) + { + predicate = quote.Operand; + } + + if (predicate is not LambdaExpression {Parameters: [var row]} lambda) + { + return false; + } + + // A key picks one row only when each row appears once, so between the Where and the DbSet there + // can only be operators that return some of its rows, unchanged + var source = where.Arguments[0]; + while (source is MethodCallExpression call && + RowOperators.KeepsRows(call.Method)) + { + source = call.Arguments[0]; + } + + // FromSql, and the other roots derived from this one, can return a key more than once + if (source is not EntityQueryRootExpression root || + root.GetType() != typeof(EntityQueryRootExpression)) + { + return false; + } + + return Bounds(lambda.Body, row, root.EntityType, listsLimited); + } + + static bool Bounds(Expression predicate, ParameterExpression row, IEntityType entityType, bool listsLimited) + { + var conditions = new List(); + AddConditions(predicate, conditions); + + var compared = new List(); + var listed = new List(); + foreach (var condition in conditions) + { + // Each side of an or returns its own rows, so both have to be bounded + if (condition is BinaryExpression {NodeType: ExpressionType.OrElse} or) + { + if (Bounds(or.Left, row, entityType, listsLimited) && + Bounds(or.Right, row, entityType, listsLimited)) + { + return true; + } + + continue; + } + + if (Compared(condition, row, entityType) is { } comparedProperty) + { + compared.Add(comparedProperty); + continue; + } + + if (listsLimited && + Listed(condition, row, entityType) is { } listedProperty) + { + listed.Add(listedProperty); + } + } + + foreach (var key in entityType.GetKeys()) + { + if (Covers(key, compared, listed)) + { + return true; + } + } + + return false; + } + + // The conditions every row the predicate returns meets + static void AddConditions(Expression predicate, List conditions) + { + if (predicate is BinaryExpression {NodeType: ExpressionType.AndAlso} and) + { + AddConditions(and.Left, conditions); + AddConditions(and.Right, conditions); + return; + } + + conditions.Add(predicate); + } + + // Every part of the key is compared with one value, or every part but one is, and that one is + // looked up in a list + static bool Covers(IKey key, List compared, List listed) + { + IProperty? remaining = null; + foreach (var property in key.Properties) + { + if (compared.Contains(property)) + { + continue; + } + + if (remaining != null) + { + return false; + } + + remaining = property; + } + + return remaining == null || + listed.Contains(remaining); + } + + // The property a condition compares with one value, as in _.Id == id + static IProperty? Compared(Expression condition, ParameterExpression row, IEntityType entityType) + { + if (condition is not BinaryExpression {NodeType: ExpressionType.Equal} equal) + { + return null; + } + + if (IsValue(equal.Right)) + { + return Property(equal.Left, row, entityType); + } + + if (IsValue(equal.Left)) + { + return Property(equal.Right, row, entityType); + } + + return null; + } + + // One value for the whole query, rather than one read from each row. Entity Framework has already + // turned every value it can work out itself into a parameter or a constant. + static bool IsValue(Expression expression) => + Uncast(expression) is QueryParameterExpression or ConstantExpression; + + // The property a condition looks up in a list the query sends, as in ids.Contains(_.Id) + static IProperty? Listed(Expression condition, ParameterExpression row, IEntityType entityType) + { + if (condition is not MethodCallExpression {Method.Name: "Contains"} call) + { + return null; + } + + Expression list; + Expression item; + + // Enumerable.Contains(ids, _.Id), or the Contains of a list type, such as List + if (call is {Object: null, Arguments: [var source, var value]}) + { + list = source; + item = value; + } + else if (call is {Object: { } instance, Arguments: [var only]}) + { + list = instance; + item = only; + } + else + { + return null; + } + + // Only a list the query sends is limited by MaxInValues. A subquery can return any number of + // values. + if (Uncast(list) is not (QueryParameterExpression or ConstantExpression) || + !ValuePlan.IsValueList(list.Type)) + { + return null; + } + + return Property(item, row, entityType); + } + + // The property of the row an expression reads, as in _.Id or EF.Property(_, "Id") + static IProperty? Property(Expression expression, ParameterExpression row, IEntityType entityType) + { + switch (Uncast(expression)) + { + case MemberExpression {Expression: { } instance} member when Uncast(instance) == row: + return entityType.FindProperty(member.Member.Name); + + case MethodCallExpression {Arguments: [var instance, ConstantExpression {Value: string name}]} call + when ShapeAnalyzer.IsProperty(call) && Uncast(instance) == row: + return entityType.FindProperty(name); + + default: + return null; + } + } + + // A cast that keeps different values different, such as the one comparing an int key with an int? + // value, or an enum key with a number. Casting a row to a type in its hierarchy keeps the row. + static Expression Uncast(Expression expression) + { + while (expression is UnaryExpression + { + NodeType: ExpressionType.Convert or ExpressionType.ConvertChecked or ExpressionType.TypeAs, + Operand: var operand + } cast && + (!operand.Type.IsValueType || + Underlying(operand.Type) == Underlying(cast.Type))) + { + expression = operand; + } + + return expression; + } + + // The type that holds the values of a nullable or an enum + static Type Underlying(Type type) + { + type = Nullable.GetUnderlyingType(type) ?? type; + if (type.IsEnum) + { + return Enum.GetUnderlyingType(type); + } + + return type; + } +} diff --git a/src/EfQueryComplexity/QueryComplexityLimits.cs b/src/EfQueryComplexity/QueryComplexityLimits.cs index 63fbcfa..fab126a 100644 --- a/src/EfQueryComplexity/QueryComplexityLimits.cs +++ b/src/EfQueryComplexity/QueryComplexityLimits.cs @@ -21,8 +21,9 @@ /// Maximum value passed to Take. Checked on every execution. /// Maximum number of values in a list the query sends, such as a Contains list. Checked on every execution. /// -/// The types for which a query that returns rows without a Take fires. true checks every -/// type, and false none. +/// The types for which a query that returns rows without a Take, or a lookup by key, fires. +/// true checks every type, and false none. A key looked up in a list, as in +/// ids.Contains(_.Id), is only a limit when MaxInValues is set. /// public sealed record QueryComplexityLimits( int? MaxNodes, diff --git a/src/EfQueryComplexity/QueryComplexityOverride.cs b/src/EfQueryComplexity/QueryComplexityOverride.cs index 5f7c5a2..6ee6bda 100644 --- a/src/EfQueryComplexity/QueryComplexityOverride.cs +++ b/src/EfQueryComplexity/QueryComplexityOverride.cs @@ -39,8 +39,8 @@ public sealed record QueryComplexityOverride public int? MaxInValues { get; init; } /// - /// Whether a query that returns rows without a Take fires. true checks every type, and - /// false none, whichever types the configured levels name. + /// Whether a query that returns rows without a Take, or a lookup by key, fires. true checks + /// every type, and false none, whichever types the configured levels name. /// public bool? RejectUnbounded { get; init; } diff --git a/src/EfQueryComplexity/QueryShape.cs b/src/EfQueryComplexity/QueryShape.cs index 504621a..70c4dd9 100644 --- a/src/EfQueryComplexity/QueryShape.cs +++ b/src/EfQueryComplexity/QueryShape.cs @@ -9,4 +9,5 @@ readonly record struct QueryShape( int Includes, int IncludeDepth, int SingleQueryCollections, - IReadOnlyList UnboundedTypes); + IReadOnlyList UnboundedTypes, + IReadOnlyList UnboundedTypesWhenListsLimited); diff --git a/src/EfQueryComplexity/RowOperators.cs b/src/EfQueryComplexity/RowOperators.cs new file mode 100644 index 0000000..40b8cd7 --- /dev/null +++ b/src/EfQueryComplexity/RowOperators.cs @@ -0,0 +1,48 @@ +/// +/// Which operators return some of their source's rows, each at most once and unchanged. +/// +static class RowOperators +{ + public static bool KeepsRows(MethodInfo method) + { + var declaringType = method.DeclaringType; + var name = method.Name; + + if (declaringType == typeof(Queryable) || + declaringType == typeof(Enumerable)) + { + return name is + nameof(Queryable.Where) or + nameof(Queryable.OrderBy) or + nameof(Queryable.OrderByDescending) or + nameof(Queryable.ThenBy) or + nameof(Queryable.ThenByDescending) or + nameof(Queryable.Skip) or + nameof(Queryable.Take) or + nameof(Queryable.Distinct) or + nameof(Queryable.Reverse) or + nameof(Queryable.OfType) or + nameof(Queryable.Cast); + } + + // These change how the rows are loaded, not which rows are returned + if (declaringType == typeof(EntityFrameworkQueryableExtensions)) + { + return name is + nameof(EntityFrameworkQueryableExtensions.Include) or + nameof(EntityFrameworkQueryableExtensions.ThenInclude) or + nameof(EntityFrameworkQueryableExtensions.AsNoTracking) or + nameof(EntityFrameworkQueryableExtensions.AsNoTrackingWithIdentityResolution) or + nameof(EntityFrameworkQueryableExtensions.AsTracking) or + nameof(EntityFrameworkQueryableExtensions.IgnoreAutoIncludes) or + nameof(EntityFrameworkQueryableExtensions.IgnoreQueryFilters) or + nameof(EntityFrameworkQueryableExtensions.TagWith) or + nameof(EntityFrameworkQueryableExtensions.TagWithCallSite); + } + + return declaringType == typeof(RelationalQueryableExtensions) && + name is + nameof(RelationalQueryableExtensions.AsSplitQuery) or + nameof(RelationalQueryableExtensions.AsSingleQuery); + } +} diff --git a/src/EfQueryComplexity/ShapeAnalyzer.cs b/src/EfQueryComplexity/ShapeAnalyzer.cs index baf7190..f463e04 100644 --- a/src/EfQueryComplexity/ShapeAnalyzer.cs +++ b/src/EfQueryComplexity/ShapeAnalyzer.cs @@ -27,7 +27,8 @@ public static QueryShape Analyze(Expression query, IModel model, bool splitByDef analyzer.includes, analyzer.includeDepth, CollectionCounter.Count(query, model, splitByDefault), - UnboundedDetector.Find(query)); + UnboundedDetector.Find(query, listsLimited: false), + UnboundedDetector.Find(query, listsLimited: true)); } public override Expression? Visit(Expression? node) @@ -189,7 +190,7 @@ int NavigationsInChain(Expression node) } } - static bool IsProperty(MethodCallExpression call) => + public static bool IsProperty(MethodCallExpression call) => call.Method.DeclaringType == typeof(EF) && call.Method.Name == nameof(EF.Property); diff --git a/src/EfQueryComplexity/UnboundedDetector.cs b/src/EfQueryComplexity/UnboundedDetector.cs index 7777f66..83dcae3 100644 --- a/src/EfQueryComplexity/UnboundedDetector.cs +++ b/src/EfQueryComplexity/UnboundedDetector.cs @@ -3,11 +3,17 @@ /// /// /// Whether a query is unbounded for a set of levels is whether those levels check any of these types. -/// So they are found once for each compiled query, rather than once for each set of levels. +/// So they are found once for each compiled query, with lists limited and without, rather than once +/// for each set of levels. /// static class UnboundedDetector { - public static IReadOnlyList Find(Expression query) + /// The query to walk. + /// + /// Whether the size of a list the query sends is limited, so a key looked up in a list bounds the + /// rows. + /// + public static IReadOnlyList Find(Expression query, bool listsLimited) { // A query that does not return a sequence returns one row, an aggregate, or a row count if (!typeof(IQueryable).IsAssignableFrom(query.Type)) @@ -16,14 +22,14 @@ public static IReadOnlyList Find(Expression query) } var types = new List(); - Collect(query, types, takeBounds: true); + Collect(query, types, takeBounds: true, listsLimited); return types; } // Walks from the outermost operator towards the source, adding the row type of every source that - // no Take limits. A sequence that is joined in is walked with takeBounds false, since the operator - // joining it returns more rows than its source whatever limits that sequence. - static void Collect(Expression expression, List types, bool takeBounds) + // no Take or lookup by key limits. A sequence that is joined in is walked with takeBounds false, + // since the operator joining it returns more rows than its source whatever limits that sequence. + static void Collect(Expression expression, List types, bool takeBounds, bool listsLimited) { while (true) { @@ -57,11 +63,21 @@ static void Collect(Expression expression, List types, bool takeBounds) break; + // A lookup by key returns at most one row for each value, so bounds like a Take + case "Where": + if (takeBounds && + KeyLookup.Bounds(call, listsLimited)) + { + return; + } + + break; + // These return more rows than their source, so a Take below one of them bounds the // source rather than the query case "SelectMany": - Collect(source, types, takeBounds); - CollectSelected(call, types); + Collect(source, types, takeBounds, listsLimited); + CollectSelected(call, types, listsLimited); return; case "Join": @@ -69,15 +85,15 @@ static void Collect(Expression expression, List types, bool takeBounds) case "LeftJoin": case "RightJoin": case "Zip": - Collect(source, types, takeBounds); - CollectJoined(call, types); + Collect(source, types, takeBounds, listsLimited); + CollectJoined(call, types, listsLimited); return; case "Concat": case "Union": case "UnionBy": - Collect(source, types, takeBounds); - Collect(call.Arguments[1], types, takeBounds); + Collect(source, types, takeBounds, listsLimited); + Collect(call.Arguments[1], types, takeBounds, listsLimited); return; } } @@ -88,7 +104,7 @@ static void Collect(Expression expression, List types, bool takeBounds) } // The rows a SelectMany joins in are the ones its collection selector returns - static void CollectSelected(MethodCallExpression call, List types) + static void CollectSelected(MethodCallExpression call, List types, bool listsLimited) { var selector = call.Arguments[1]; @@ -100,7 +116,7 @@ static void CollectSelected(MethodCallExpression call, List types) if (selector is LambdaExpression lambda) { - Collect(lambda.Body, types, takeBounds: false); + Collect(lambda.Body, types, takeBounds: false, listsLimited); return; } @@ -111,7 +127,7 @@ static void CollectSelected(MethodCallExpression call, List types) // Join, GroupJoin, LeftJoin and RightJoin take one other sequence, and Zip one or two. Selectors // and comparers are not sequences. - static void CollectJoined(MethodCallExpression call, List types) + static void CollectJoined(MethodCallExpression call, List types, bool listsLimited) { var arguments = call.Arguments; for (var index = 1; index < arguments.Count; index++) @@ -119,7 +135,7 @@ static void CollectJoined(MethodCallExpression call, List types) var argument = arguments[index]; if (Sequences.IsSequence(argument.Type)) { - Collect(argument, types, takeBounds: false); + Collect(argument, types, takeBounds: false, listsLimited); } } } diff --git a/src/EfQueryComplexity/UnboundedEntities.cs b/src/EfQueryComplexity/UnboundedEntities.cs index a05ca4d..16d768c 100644 --- a/src/EfQueryComplexity/UnboundedEntities.cs +++ b/src/EfQueryComplexity/UnboundedEntities.cs @@ -4,11 +4,11 @@ namespace EfQueryComplexity; /// The types RejectUnbounded checks. /// /// -/// A query fires when it returns rows of a checked type without a Take. Naming a type also names -/// the types derived from it, and naming an interface names the types that implement it. Rows that -/// are not entities, such as a list of values, are checked by and -/// , and not by . true converts to -/// , and false to . +/// A query fires when it returns rows of a checked type without a Take, or a lookup by key. Naming a +/// type also names the types derived from it, and naming an interface names the types that +/// implement it. Rows that are not entities, such as a list of values, are checked by +/// and , and not by . true +/// converts to , and false to . /// public sealed class UnboundedEntities : IEquatable diff --git a/src/EfQueryComplexity/Violations.cs b/src/EfQueryComplexity/Violations.cs index ab5704d..062e821 100644 --- a/src/EfQueryComplexity/Violations.cs +++ b/src/EfQueryComplexity/Violations.cs @@ -14,7 +14,10 @@ public static List ForShape(QueryShape shape, QueryCom Add(violations, nameof(QueryComplexityLimits.MaxIncludeDepth), limits.MaxIncludeDepth, shape.IncludeDepth); Add(violations, nameof(QueryComplexityLimits.MaxSingleQueryCollections), limits.MaxSingleQueryCollections, shape.SingleQueryCollections); - var rowTypes = CheckedTypes(shape.UnboundedTypes, limits.RejectUnbounded); + // A key looked up in a list returns a row for each value in it, so only bounds the query when + // MaxInValues limits the list + var unboundedTypes = limits.MaxInValues == null ? shape.UnboundedTypes : shape.UnboundedTypesWhenListsLimited; + var rowTypes = CheckedTypes(unboundedTypes, limits.RejectUnbounded); if (rowTypes != null) { violations.Add( diff --git a/src/Tests/KeyContext.cs b/src/Tests/KeyContext.cs new file mode 100644 index 0000000..5f0430a --- /dev/null +++ b/src/Tests/KeyContext.cs @@ -0,0 +1,51 @@ +// Keys the shared model does not have. Kept out of TestDbContext, since it is also the LocalDB +// schema. +public class KeyContext(DbContextOptions options) : + DbContext(options) +{ + public DbSet Courses => Set(); + public DbSet Enrollments => Set(); + + protected override void OnModelCreating(ModelBuilder builder) + { + var course = builder.Entity(); + course.HasAlternateKey(_ => _.Code); + course.HasIndex(_ => _.Title).IsUnique(); + + builder.Entity(); + + builder.Entity() + .HasKey(_ => new + { + _.StudentId, + _.CourseId + }); + } +} + +public class Course +{ + public int Id { get; set; } + + // An alternate key + public string Code { get; set; } = ""; + + // A unique index + public string Title { get; set; } = ""; + + public int Credits { get; set; } +} + +// A derived type, which shares the key of its base type +public class OnlineCourse : + Course +{ + public string Url { get; set; } = ""; +} + +// A composite key +public class Enrollment +{ + public int StudentId { get; set; } + public int CourseId { get; set; } +} diff --git a/src/Tests/KeyLookupTests.cs b/src/Tests/KeyLookupTests.cs new file mode 100644 index 0000000..a1eefed --- /dev/null +++ b/src/Tests/KeyLookupTests.cs @@ -0,0 +1,273 @@ +public class KeyLookupTests +{ + [Test] + public Task PrimaryKey() + { + var id = 1; + return AssertBounded(context => context.Courses.Where(_ => _.Id == id)); + } + + [Test] + public Task PrimaryKeyConstant() => + AssertBounded(context => context.Courses.Where(_ => _.Id == 1)); + + [Test] + public Task ValueOnTheLeft() + { + var id = 1; + return AssertBounded(context => context.Courses.Where(_ => id == _.Id)); + } + + // The key is cast to int? to compare it + [Test] + public Task NullableValue() + { + int? id = 1; + return AssertBounded(context => context.Courses.Where(_ => _.Id == id)); + } + + [Test] + public Task EfProperty() + { + var id = 1; + return AssertBounded(context => context.Courses.Where(_ => EF.Property(_, nameof(Course.Id)) == id)); + } + + [Test] + public Task OtherConditions() + { + var id = 1; + return AssertBounded(context => context.Courses.Where(_ => _.Credits > 3 && _.Id == id)); + } + + [Test] + public Task EitherKey() + { + var id = 1; + var other = 2; + return AssertBounded(context => context.Courses.Where(_ => _.Id == id || _.Id == other)); + } + + [Test] + public Task KeyOrAnythingElse() + { + var id = 1; + return AssertUnbounded(context => context.Courses.Where(_ => _.Id == id || _.Credits > 3)); + } + + [Test] + public Task AlternateKey() + { + var code = "CS101"; + return AssertBounded(context => context.Courses.Where(_ => _.Code == code)); + } + + // A filter can make an index unique among only some of the rows + [Test] + public Task UniqueIndexIsNotAKey() + { + var title = "Algorithms"; + return AssertUnbounded(context => context.Courses.Where(_ => _.Title == title)); + } + + [Test] + public Task NotAKey() + { + var credits = 3; + return AssertUnbounded(context => context.Courses.Where(_ => _.Credits == credits)); + } + + // Every row can match + [Test] + public Task KeyComparedWithTheRow() => + AssertUnbounded(context => context.Courses.Where(_ => _.Id == _.Credits)); + + [Test] + public Task KeyRange() + { + var id = 1; + return AssertUnbounded(context => context.Courses.Where(_ => _.Id > id)); + } + + [Test] + public Task CompositeKey() + { + var student = 1; + var course = 2; + return AssertBounded(context => context.Enrollments.Where(_ => _.StudentId == student && _.CourseId == course)); + } + + [Test] + public Task PartOfCompositeKey() + { + var student = 1; + return AssertUnbounded(context => context.Enrollments.Where(_ => _.StudentId == student)); + } + + [Test] + public Task OperatorsAround() + { + var id = 1; + return AssertBounded( + context => context.Courses + .AsNoTracking() + .OrderBy(_ => _.Title) + .Where(_ => _.Id == id) + .TagWith("Lookup")); + } + + [Test] + public Task DerivedType() + { + var id = 1; + return AssertBounded(context => context.Courses.OfType().Where(_ => _.Id == id)); + } + + [Test] + public Task DerivedRoot() + { + var id = 1; + return AssertBounded(context => context.Set().Where(_ => _.Id == id)); + } + + // Each course is returned twice + [Test] + public Task RowsRepeated() + { + var id = 1; + return AssertUnbounded(context => context.Courses.Concat(context.Courses).Where(_ => _.Id == id)); + } + + // The raw SQL can return a key more than once + [Test] + public Task FromSql() + { + var id = 1; + return AssertUnbounded(context => context.Courses.FromSql($"select * from Courses").Where(_ => _.Id == id)); + } + + // The Id of a projection is not the key of the rows it was read from + [Test] + public Task KeyOfAProjection() + { + var id = 1; + return AssertUnbounded( + context => context.Courses + .Select(_ => new Course + { + Id = _.Credits + }) + .Where(_ => _.Id == id)); + } + + [Test] + public Task KeyInList() + { + var ids = new List {1, 2, 3}; + return AssertBounded(context => context.Courses.Where(_ => ids.Contains(_.Id)), maxInValues: 10); + } + + [Test] + public Task KeyInArray() + { + int[] ids = [1, 2, 3]; + return AssertBounded(context => context.Courses.Where(_ => ids.Contains(_.Id)), maxInValues: 10); + } + + // Without MaxInValues the list can hold every key + [Test] + public Task KeyInListWithoutMaxInValues() + { + var ids = new List {1, 2, 3}; + return AssertUnbounded(context => context.Courses.Where(_ => ids.Contains(_.Id))); + } + + [Test] + public Task NotAKeyInList() + { + var credits = new List {1, 2, 3}; + return AssertUnbounded(context => context.Courses.Where(_ => credits.Contains(_.Credits)), maxInValues: 10); + } + + [Test] + public Task CompositeKeyWithList() + { + var student = 1; + var courses = new List {1, 2, 3}; + return AssertBounded( + context => context.Enrollments.Where(_ => _.StudentId == student && courses.Contains(_.CourseId)), + maxInValues: 10); + } + + // Each list is limited, but together they allow the product of their sizes + [Test] + public Task CompositeKeyWithTwoLists() + { + var students = new List {1, 2, 3}; + var courses = new List {1, 2, 3}; + return AssertUnbounded( + context => context.Enrollments.Where(_ => students.Contains(_.StudentId) && courses.Contains(_.CourseId)), + maxInValues: 10); + } + + // A subquery is not a list the query sends, so MaxInValues does not limit it + [Test] + public Task KeyInSubquery() => + AssertUnbounded( + context => context.Courses.Where(_ => context.Enrollments.Select(enrollment => enrollment.CourseId).Contains(_.Id)), + maxInValues: 10); + + // Each set of levels decides for itself whether a list is limited + [Test] + public async Task ListBoundsOnlyTheLevelsLimitingIt() + { + var ids = new List {1, 2, 3}; + var (context, logs) = Build( + logAt: Limits.None with + { + RejectUnbounded = true + }, + throwAt: Limits.None with + { + RejectUnbounded = true, + MaxInValues = 10 + }); + + context.Courses.Where(_ => ids.Contains(_.Id)).ToQueryString(); + + await Assert.That(logs.Single()).Contains("RejectUnbounded"); + } + + static async Task AssertBounded(Func query, int? maxInValues = null) + { + var (context, logs) = Build(Limits.None, Checked(maxInValues)); + query(context).ToQueryString(); + await Assert.That(logs.Count).IsEqualTo(0); + } + + static async Task AssertUnbounded(Func query, int? maxInValues = null) + { + var (context, _) = Build(Limits.None, Checked(maxInValues)); + var exception = Assert.Throws(() => query(context).ToQueryString()); + await Assert.That(exception.Violations.Single().Limit).IsEqualTo("RejectUnbounded"); + } + + static QueryComplexityLimits Checked(int? maxInValues) => + Limits.None with + { + RejectUnbounded = true, + MaxInValues = maxInValues + }; + + static (KeyContext context, List logs) Build(QueryComplexityLimits logAt, QueryComplexityLimits throwAt) + { + var logs = new List(); + var options = new DbContextOptionsBuilder() + .UseSqlServer("Server=.;Database=Test;") + .EnableServiceProviderCaching(false) + .LogTo(logs.Add, [QueryComplexityEventId.LimitExceeded], LogLevel.Debug, DbContextLoggerOptions.None) + .UseQueryComplexity(logAt, throwAt) + .Options; + return (new(options), logs); + } +} diff --git a/src/Tests/ShapeTests.cs b/src/Tests/ShapeTests.cs index 17b131a..697c29d 100644 --- a/src/Tests/ShapeTests.cs +++ b/src/Tests/ShapeTests.cs @@ -261,6 +261,102 @@ public Task SingleQueryCollectionsIgnoresSplitByDefault() => .ThenInclude(_ => _.Employees), SplitByDefault); + // The query returns names, so there are no companies for the Includes to load into + [Test] + public Task SingleQueryCollectionsIgnoresIncludesAProjectionDrops() => + AssertNoCollections( + context => context.Companies + .Include(_ => _.Departments) + .ThenInclude(_ => _.Employees) + .Where(_ => _.Name != "") + .OrderBy(_ => _.Name) + .Take(5) + .Select(_ => _.Name)); + + [Test] + public Task SingleQueryCollectionsIgnoresIncludesADtoDrops() => + AssertNoCollections( + context => context.Companies + .Include(_ => _.Departments) + .ThenInclude(_ => _.Employees) + .Select( + _ => new + { + _.Name, + Departments = _.Departments.Count(), + Staffed = _.Departments.Any(department => department.Employees.Count > 0) + })); + + [Test] + public async Task SingleQueryCollectionsIgnoresIncludesAnAggregateDrops() + { + var (context, _) = ContextBuilder.Build(); + var query = context.Companies + .Include(_ => _.Departments) + .ThenInclude(_ => _.Employees); + var count = Expression.Call(typeof(Queryable), nameof(Queryable.Count), [typeof(Company)], query.Expression); + + await Assert.That(CollectionCounter.Count(count, context.Model, splitByDefault: false)).IsEqualTo(0); + } + + // The projection loads the department names itself, and the Include is ignored + [Test] + public async Task SingleQueryCollectionsCountsProjectionNotDroppedInclude() => + await Assert.That( + Measure( + context => context.Companies + .Include(_ => _.Departments) + .Select( + _ => new + { + _.Name, + Departments = _.Departments.Select(department => department.Name).ToList() + }), + SingleQueryCollections)) + .IsEqualTo(1); + + [Test] + public async Task SingleQueryCollectionsCountsIncludesOfAProjectedEntity() => + await Assert.That( + Measure( + context => context.Companies + .Include(_ => _.Departments) + .ThenInclude(_ => _.Employees) + .Select( + _ => new + { + Company = _ + }), + SingleQueryCollections)) + .IsEqualTo(2); + + // Entity Framework carries the ThenInclude onto the department the projection returns + [Test] + public async Task SingleQueryCollectionsCountsIncludesOfAProjectedNavigation() => + await Assert.That( + Measure( + context => context.Employees + .Include(_ => _.Department) + .ThenInclude(_ => _.Employees) + .Select(_ => _.Department), + SingleQueryCollections)) + .IsEqualTo(1); + + // The method runs after the company is loaded with its Includes + [Test] + public async Task SingleQueryCollectionsCountsIncludesOfAnEntityPassedToAMethod() => + await Assert.That( + Measure( + context => context.Companies + .Include(_ => _.Departments) + .ThenInclude(_ => _.Employees) + .Select(_ => Describe(_)), + SingleQueryCollections)) + .IsEqualTo(2); + + static string Describe(Company company) => + company.Name; + static void SplitByDefault(DbContextOptionsBuilder builder) => builder.UseSqlServer( "Server=.;Database=Test;", diff --git a/src/Tests/UnboundedTests.cs b/src/Tests/UnboundedTests.cs index 6d1547a..fee4a06 100644 --- a/src/Tests/UnboundedTests.cs +++ b/src/Tests/UnboundedTests.cs @@ -45,16 +45,34 @@ public Task ConcatOfBounded() => public Task ConcatWithUnbounded() => AssertUnbounded(context => context.Employees.Take(5).Concat(context.Employees)); + // A lookup by key returns at most one row. KeyLookupTests covers which lookups count. + [Test] + public Task KeyLookupWithInclude() + { + var id = 1; + return AssertBounded(context => context.Companies.Where(_ => _.Id == id).Include(_ => _.Departments)); + } + + // The lookup bounds the departments, not the employees each has + [Test] + public Task KeyLookupThenSelectMany() + { + var id = 1; + return AssertRowTypes( + context => context.Departments.Where(_ => _.Id == id).SelectMany(_ => _.Employees), + "Employee"); + } + [Test] public async Task ScalarTerminalsAreBounded() { var (context, _) = ContextBuilder.Build(); var employees = context.Employees; - await Assert.That(UnboundedDetector.Find(Terminal(employees, nameof(Queryable.Count)))).IsEmpty(); - await Assert.That(UnboundedDetector.Find(Terminal(employees, nameof(Queryable.First)))).IsEmpty(); - await Assert.That(UnboundedDetector.Find(Terminal(employees, nameof(Queryable.Any)))).IsEmpty(); - await Assert.That(UnboundedDetector.Find(Terminal(employees, nameof(Queryable.LongCount)))).IsEmpty(); + await Assert.That(UnboundedDetector.Find(Terminal(employees, nameof(Queryable.Count)), listsLimited: false)).IsEmpty(); + await Assert.That(UnboundedDetector.Find(Terminal(employees, nameof(Queryable.First)), listsLimited: false)).IsEmpty(); + await Assert.That(UnboundedDetector.Find(Terminal(employees, nameof(Queryable.Any)), listsLimited: false)).IsEmpty(); + await Assert.That(UnboundedDetector.Find(Terminal(employees, nameof(Queryable.LongCount)), listsLimited: false)).IsEmpty(); } // A projection does not change which rows are read