From db6f58b56b18eddbdb09c53b0ed73101efe50709 Mon Sep 17 00:00:00 2001 From: Andreas Marek Date: Wed, 19 Aug 2026 08:36:06 +1000 Subject: [PATCH] Make fragment cycle validation linear --- .../validation/OperationValidator.java | 84 ++++++++++++------- .../validation/NoFragmentCyclesTest.groovy | 28 +++++++ 2 files changed, 80 insertions(+), 32 deletions(-) diff --git a/src/main/java/graphql/validation/OperationValidator.java b/src/main/java/graphql/validation/OperationValidator.java index fdb0cf04bd..4c1430e504 100644 --- a/src/main/java/graphql/validation/OperationValidator.java +++ b/src/main/java/graphql/validation/OperationValidator.java @@ -66,12 +66,15 @@ import org.jspecify.annotations.NullMarked; import org.jspecify.annotations.Nullable; +import java.util.ArrayDeque; import java.util.ArrayList; import java.util.Collection; import java.util.Collections; +import java.util.Deque; import java.util.EnumSet; import java.util.HashMap; import java.util.HashSet; +import java.util.Iterator; import java.util.LinkedHashMap; import java.util.LinkedHashSet; @@ -301,7 +304,8 @@ public class OperationValidator implements DocumentVisitor { private final Set visitedFragmentSpreads = new HashSet<>(); // --- State: NoFragmentCycles --- - private final Map> fragmentSpreadsMap = new HashMap<>(); + private final Map> fragmentSpreadsMap = new LinkedHashMap<>(); + private final Set fragmentsWithCycleErrors = new HashSet<>(); // --- State: NoUnusedFragments --- private final List allDeclaredFragments = new ArrayList<>(); @@ -368,7 +372,10 @@ public OperationValidator(ValidationContext validationContext, ValidationErrorCo this.rulePredicate = rulePredicate; this.allRulesEnabled = detectAllRulesEnabled(rulePredicate); this.complexityLimits = validationContext.getQueryComplexityLimits(); - prepareFragmentSpreadsMap(); + if (isRuleEnabled(OperationValidationRule.NO_FRAGMENT_CYCLES)) { + prepareFragmentSpreadsMap(); + findFragmentCycles(); + } } private static boolean detectAllRulesEnabled(Predicate predicate) { @@ -1306,7 +1313,7 @@ private void prepareFragmentSpreadsMap() { } private Set gatherSpreads(FragmentDefinition fragmentDefinition) { - final Set spreads = new HashSet<>(); + final Set spreads = new LinkedHashSet<>(); DocumentVisitor visitor = new DocumentVisitor() { @Override public void enter(Node node, List path) { @@ -1324,44 +1331,57 @@ public void leave(Node node, List path) { } private void validateNoFragmentCycles(FragmentDefinition fragmentDefinition) { - ArrayList path = new ArrayList<>(); - path.add(fragmentDefinition.getName()); - Map> transitiveSpreads = buildTransitiveSpreads(path, new HashMap<>()); + if (!fragmentsWithCycleErrors.contains(fragmentDefinition.getName())) { + return; + } + String message = i18n(FragmentCycle, "NoFragmentCycles.cyclesNotAllowed"); + addError(FragmentCycle, Collections.singletonList(fragmentDefinition), message); + } - for (Map.Entry> entry : transitiveSpreads.entrySet()) { - if (entry.getValue().contains(entry.getKey())) { - String message = i18n(FragmentCycle, "NoFragmentCycles.cyclesNotAllowed"); - addError(FragmentCycle, Collections.singletonList(fragmentDefinition), message); + private void findFragmentCycles() { + Set visitedFragments = new HashSet<>(); + for (Map.Entry> entry : fragmentSpreadsMap.entrySet()) { + if (!visitedFragments.add(entry.getKey())) { + continue; } + findFragmentCycles(entry.getKey(), entry.getValue(), visitedFragments); } } - private Map> buildTransitiveSpreads(ArrayList path, Map> transitiveSpreads) { - String name = path.get(path.size() - 1); - if (transitiveSpreads.containsKey(name)) { - return transitiveSpreads; - } - Set spreads = fragmentSpreadsMap.get(name); - if (spreads == null || spreads.isEmpty()) { - return transitiveSpreads; - } - for (String ancestor : path) { - Set ancestorSpreads = transitiveSpreads.get(ancestor); - if (ancestorSpreads == null) { - ancestorSpreads = new HashSet<>(); + private void findFragmentCycles(String firstFragment, Set firstSpreads, Set visitedFragments) { + Set visitingFragments = new HashSet<>(); + Deque fragmentStack = new ArrayDeque<>(); + Deque> spreadIteratorStack = new ArrayDeque<>(); + visitingFragments.add(firstFragment); + fragmentStack.push(firstFragment); + spreadIteratorStack.push(firstSpreads.iterator()); + + while (!fragmentStack.isEmpty()) { + Iterator spreadIterator = spreadIteratorStack.getFirst(); + if (!spreadIterator.hasNext()) { + visitingFragments.remove(fragmentStack.pop()); + spreadIteratorStack.pop(); + continue; } - ancestorSpreads.addAll(spreads); - transitiveSpreads.put(ancestor, ancestorSpreads); - } - for (String child : spreads) { - if (path.contains(child) || transitiveSpreads.containsKey(child)) { + + String childFragment = spreadIterator.next(); + Set childSpreads = fragmentSpreadsMap.get(childFragment); + if (childSpreads == null) { + continue; + } + if (visitingFragments.contains(childFragment)) { + fragmentsWithCycleErrors.add(childFragment); + fragmentsWithCycleErrors.add(fragmentStack.getFirst()); + continue; + } + if (!visitedFragments.add(childFragment)) { continue; } - ArrayList childPath = new ArrayList<>(path); - childPath.add(child); - buildTransitiveSpreads(childPath, transitiveSpreads); + + visitingFragments.add(childFragment); + fragmentStack.push(childFragment); + spreadIteratorStack.push(childSpreads.iterator()); } - return transitiveSpreads; } // --- NoUndefinedVariables --- diff --git a/src/test/groovy/graphql/validation/NoFragmentCyclesTest.groovy b/src/test/groovy/graphql/validation/NoFragmentCyclesTest.groovy index b54ad740bc..876c6ea71c 100644 --- a/src/test/groovy/graphql/validation/NoFragmentCyclesTest.groovy +++ b/src/test/groovy/graphql/validation/NoFragmentCyclesTest.groovy @@ -240,4 +240,32 @@ class NoFragmentCyclesTest extends Specification { errorCollector.containsValidationError(ValidationErrorType.FragmentCycle) errorCollector.getErrors()[0].message == "Validation error (FragmentCycle@[MyFrag]) : Fragment cycles not allowed" } + + def "long acyclic fragment chains are valid"() { + when: + traverse(fragmentChain(1_000, false)) + + then: + errorCollector.getErrors().isEmpty() + } + + def "cycles at the end of long fragment chains are detected"() { + when: + traverse(fragmentChain(1_000, true)) + + then: + errorCollector.containsValidationError(ValidationErrorType.FragmentCycle) + } + + private static String fragmentChain(int fragmentCount, boolean cycle) { + (0.. + String selection = "name" + if (index < fragmentCount - 1) { + selection = "...F${index + 1}" + } else if (cycle) { + selection = "...F${fragmentCount.intdiv(2)}" + } + "fragment F${index} on Dog { ${selection} }" + }.join("\n") + } }