diff --git a/internal/validation/validate_max_depth_test.go b/internal/validation/validate_max_depth_test.go index 30af940e..8c48fe36 100644 --- a/internal/validation/validate_max_depth_test.go +++ b/internal/validation/validate_max_depth_test.go @@ -320,16 +320,40 @@ func TestMaxDepthFragmentSpreads(t *testing.T) { } query laterDepthValidated { - ...character # depth 1 (+1) - enemies { # depth 1 - friends { # depth 2 - ...character # depth 2 (+1), should error! + characters { # depth 1 + ...character # depth 2 (+1) + enemies { # depth 2 + friends { # depth 3 + ...character # depth 4 (+1), should error! + } } } } `, - depth: 2, - failure: true, + depth: 4, + failure: true, + expectedErrors: []string{"MaxDepthExceeded"}, + }, + { + name: "fragmentChainSpreadShallowFirst", + query: ` + fragment f1 on Character { friends { ...f2 } } + fragment f2 on Character { friends { ...f3 } } + fragment f3 on Character { friends { ...f4 } } + fragment f4 on Character { friends { name } } + + query { + characters { # depth 1 + ...f4 # each fragment is first reached here, where it is shallow + ...f3 + ...f2 + ...f1 # friends nest to depth 5, name to depth 6 + } + } + `, + depth: 4, + failure: true, + expectedErrors: []string{"MaxDepthExceeded"}, }, { name: "spreadAtSameDepth", diff --git a/internal/validation/validation.go b/internal/validation/validation.go index 8715869b..5bed2ce7 100644 --- a/internal/validation/validation.go +++ b/internal/validation/validation.go @@ -309,12 +309,22 @@ func validateValue(c *opContext, v *ast.InputValueDefinition, val any, t ast.Typ } } +// fragmentDepth identifies a fragment spread at a given depth. +type fragmentDepth struct { + frag *ast.FragmentDefinition + depth int +} + // validates the query doesn't go deeper than maxDepth (if set). Returns whether // or not query validated max depth to avoid excessive recursion. // -// The visited map is necessary to ensure that max depth validation does not get stuck in cyclical -// fragment spreads. -func validateMaxDepth(c *opContext, sels []ast.Selection, visited map[*ast.FragmentDefinition]struct{}, depth int) bool { +// The visited map records each fragment together with the depth it was spread +// at. A fragment spread again at a different depth must be checked again, +// because its fields end up at different depths; a spread at a depth already +// checked is skipped, which stops cyclical fragment spreads. Depth never grows +// past maxDepth+1 (a field beyond maxDepth is not descended into), so each +// fragment is walked at most maxDepth+1 times. +func validateMaxDepth(c *opContext, sels []ast.Selection, visited map[fragmentDepth]struct{}, depth int) bool { // maxDepth checking is turned off when maxDepth is 0 if c.maxDepth == 0 { return false @@ -322,7 +332,7 @@ func validateMaxDepth(c *opContext, sels []ast.Selection, visited map[*ast.Fragm exceededMaxDepth := false if visited == nil { - visited = map[*ast.FragmentDefinition]struct{}{} + visited = map[fragmentDepth]struct{}{} } for _, sel := range sels { @@ -348,11 +358,12 @@ func validateMaxDepth(c *opContext, sels []ast.Selection, visited map[*ast.Fragm continue } - if _, ok := visited[frag]; ok { - // we've already seen this fragment, don't check depth again. + key := fragmentDepth{frag: frag, depth: depth} + if _, ok := visited[key]; ok { + // we've already checked this fragment at this depth. continue } - visited[frag] = struct{}{} + visited[key] = struct{}{} // Depth is not incremented because fragments have the same depth as surrounding fields exceededMaxDepth = exceededMaxDepth || validateMaxDepth(c, frag.Selections, visited, depth)