Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
36 changes: 30 additions & 6 deletions internal/validation/validate_max_depth_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
25 changes: 18 additions & 7 deletions internal/validation/validation.go
Original file line number Diff line number Diff line change
Expand Up @@ -309,20 +309,30 @@ 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
}

exceededMaxDepth := false
if visited == nil {
visited = map[*ast.FragmentDefinition]struct{}{}
visited = map[fragmentDepth]struct{}{}
}

for _, sel := range sels {
Expand All @@ -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)
Expand Down
Loading