diff options
Diffstat (limited to 'src/Language/GraphQL/Validate')
| -rw-r--r-- | src/Language/GraphQL/Validate/Rules.hs | 676 |
1 files changed, 633 insertions, 43 deletions
diff --git a/src/Language/GraphQL/Validate/Rules.hs b/src/Language/GraphQL/Validate/Rules.hs index c67df1c..71455d3 100644 --- a/src/Language/GraphQL/Validate/Rules.hs +++ b/src/Language/GraphQL/Validate/Rules.hs @@ -3,6 +3,7 @@ obtain one at https://mozilla.org/MPL/2.0/. -} {-# LANGUAGE LambdaCase #-} +{-# LANGUAGE NamedFieldPuns #-} {-# LANGUAGE OverloadedStrings #-} {-# LANGUAGE RecordWildCards #-} {-# LANGUAGE ScopedTypeVariables #-} @@ -24,6 +25,8 @@ module Language.GraphQL.Validate.Rules , noUndefinedVariablesRule , noUnusedFragmentsRule , noUnusedVariablesRule + , overlappingFieldsCanBeMergedRule + , possibleFragmentSpreadsRule , providedRequiredInputFieldsRule , providedRequiredArgumentsRule , scalarLeafsRule @@ -35,22 +38,24 @@ module Language.GraphQL.Validate.Rules , uniqueInputFieldNamesRule , uniqueOperationNamesRule , uniqueVariableNamesRule + , valuesOfCorrectTypeRule + , variablesInAllowedPositionRule , variablesAreInputTypesRule ) where import Control.Monad ((>=>), foldM) import Control.Monad.Trans.Class (MonadTrans(..)) -import Control.Monad.Trans.Reader (ReaderT(..), asks, mapReaderT) +import Control.Monad.Trans.Reader (ReaderT(..), ask, asks, mapReaderT) import Control.Monad.Trans.State (StateT, evalStateT, gets, modify) import Data.Bifunctor (first) -import Data.Foldable (find, toList) +import Data.Foldable (find, fold, foldl', toList) import qualified Data.HashMap.Strict as HashMap import Data.HashMap.Strict (HashMap) import Data.HashSet (HashSet) import qualified Data.HashSet as HashSet import Data.List (groupBy, sortBy, sortOn) -import Data.Maybe (isNothing, mapMaybe) -import Data.List.NonEmpty (NonEmpty) +import Data.Maybe (catMaybes, fromMaybe, isJust, isNothing, mapMaybe) +import Data.List.NonEmpty (NonEmpty(..)) import Data.Ord (comparing) import Data.Sequence (Seq(..), (|>)) import qualified Data.Sequence as Seq @@ -80,6 +85,7 @@ specifiedRules = -- Fields , fieldsOnCorrectTypeRule , scalarLeafsRule + , overlappingFieldsCanBeMergedRule -- Arguments. , knownArgumentNamesRule , uniqueArgumentNamesRule @@ -91,7 +97,9 @@ specifiedRules = , noUnusedFragmentsRule , fragmentSpreadTargetDefinedRule , noFragmentCyclesRule + , possibleFragmentSpreadsRule -- Values + , valuesOfCorrectTypeRule , knownInputFieldNamesRule , uniqueInputFieldNamesRule , providedRequiredInputFieldsRule @@ -104,6 +112,7 @@ specifiedRules = , variablesAreInputTypesRule , noUndefinedVariablesRule , noUnusedVariablesRule + , variablesInAllowedPositionRule ] -- | Definition must be OperationDefinition or FragmentDefinition. @@ -320,10 +329,8 @@ fragmentSpreadTypeExistenceRule :: forall m. Rule m fragmentSpreadTypeExistenceRule = SelectionRule $ const $ \case Full.FragmentSpreadSelection fragmentSelection | Full.FragmentSpread fragmentName _ location' <- fragmentSelection -> do - ast' <- asks ast - let target = find (isSpreadTarget fragmentName) ast' - typeCondition <- lift $ maybeToSeq $ target >>= extractTypeCondition types' <- asks $ Schema.types . schema + typeCondition <- findSpreadTarget fragmentName case HashMap.lookup typeCondition types' of Nothing -> pure $ Error { message = spreadError fragmentName typeCondition @@ -342,10 +349,6 @@ fragmentSpreadTypeExistenceRule = SelectionRule $ const $ \case Just _ -> lift mempty _ -> lift mempty where - extractTypeCondition (viewFragment -> Just fragmentDefinition) = - let Full.FragmentDefinition _ typeCondition _ _ _ = fragmentDefinition - in Just typeCondition - extractTypeCondition _ = Nothing spreadError fragmentName typeCondition = concat [ "Fragment \"" , Text.unpack fragmentName @@ -451,8 +454,7 @@ filterSelections applyFilter selections noFragmentCyclesRule :: forall m. Rule m noFragmentCyclesRule = FragmentDefinitionRule $ \case Full.FragmentDefinition fragmentName _ _ selections location' -> do - state <- evalStateT (collectFields selections) - (0, fragmentName) + state <- evalStateT (collectCycles selections) (0, fragmentName) let spreadPath = fst <$> sortBy (comparing snd) (HashMap.toList state) case reverse spreadPath of x : _ | x == fragmentName -> pure $ Error @@ -467,10 +469,10 @@ noFragmentCyclesRule = FragmentDefinitionRule $ \case } _ -> lift mempty where - collectFields :: Traversable t + collectCycles :: Traversable t => t Full.Selection -> StateT (Int, Full.Name) (ReaderT (Validation m) Seq) (HashMap Full.Name Int) - collectFields selectionSet = foldM forEach HashMap.empty selectionSet + collectCycles selectionSet = foldM forEach HashMap.empty selectionSet forEach accumulator = \case Full.FieldSelection fieldSelection -> forField accumulator fieldSelection Full.InlineFragmentSelection fragmentSelection -> @@ -487,15 +489,15 @@ noFragmentCyclesRule = FragmentDefinitionRule $ \case then pure newAccumulator else collectFromSpread fragmentName newAccumulator forInline accumulator (Full.InlineFragment _ _ selections _) = - (accumulator <>) <$> collectFields selections + (accumulator <>) <$> collectCycles selections forField accumulator (Full.Field _ _ _ _ selections _) = - (accumulator <>) <$> collectFields selections + (accumulator <>) <$> collectCycles selections collectFromSpread fragmentName accumulator = do ast' <- lift $ asks ast case findFragmentDefinition fragmentName ast' of Nothing -> pure accumulator Just (Full.FragmentDefinition _ _ _ selections _) -> - (accumulator <>) <$> collectFields selections + (accumulator <>) <$> collectCycles selections findFragmentDefinition :: Text -> NonEmpty Full.Definition @@ -531,15 +533,22 @@ uniqueDirectiveNamesRule = DirectivesRule extract (Full.Directive directiveName _ location') = (directiveName, location') -filterDuplicates :: (a -> (Text, Full.Location)) -> String -> [a] -> Seq Error +groupSorted :: forall a. (a -> Text) -> [a] -> [[a]] +groupSorted getName = groupBy equalByName . sortOn getName + where + equalByName lhs rhs = getName lhs == getName rhs + +filterDuplicates :: forall a + . (a -> (Text, Full.Location)) + -> String + -> [a] + -> Seq Error filterDuplicates extract nodeType = Seq.fromList . fmap makeError . filter ((> 1) . length) - . groupBy equalByName - . sortOn getName + . groupSorted getName where getName = fst . extract - equalByName lhs rhs = getName lhs == getName rhs makeError directives' = Error { message = makeMessage $ head directives' , locations = snd . extract <$> directives' @@ -647,12 +656,9 @@ variableUsageDifference difference errorMessage = OperationDefinitionRule $ \cas lift $ lift $ mapArguments arguments <> mapDirectives directives' variableFilter (Full.FragmentSpreadSelection spread) | Full.FragmentSpread fragmentName _ _ <- spread = do - definitions <- lift $ asks ast - visited <- gets (HashSet.member fragmentName) - modify (HashSet.insert fragmentName) - case find (isSpreadTarget fragmentName) definitions of - Just (viewFragment -> Just fragmentDefinition) - | not visited -> diveIntoSpread fragmentDefinition + nonVisitedFragmentDefinition <- visitFragmentDefinition fragmentName + case nonVisitedFragmentDefinition of + Just fragmentDefinition -> diveIntoSpread fragmentDefinition _ -> lift $ lift mempty diveIntoSpread (Full.FragmentDefinition _ _ directives' selections _) = filterSelections' selections @@ -710,7 +716,7 @@ fieldsOnCorrectTypeRule = FieldRule fieldRule fieldRule parentType (Full.Field _ fieldName _ _ _ location') | Just objectType <- parentType , Nothing <- Type.lookupTypeField fieldName objectType - , Just typeName <- compositeTypeName objectType = pure $ Error + , Just typeName <- typeNameIfComposite objectType = pure $ Error { message = errorMessage fieldName typeName , locations = [location'] } @@ -723,20 +729,17 @@ fieldsOnCorrectTypeRule = FieldRule fieldRule , "\"." ] -compositeTypeName :: forall m. Out.Type m -> Maybe Full.Name -compositeTypeName (Out.ObjectBaseType (Out.ObjectType typeName _ _ _)) = - Just typeName -compositeTypeName (Out.InterfaceBaseType interfaceType) = +compositeTypeName :: forall m. Type.CompositeType m -> Full.Name +compositeTypeName (Type.CompositeObjectType (Out.ObjectType typeName _ _ _)) = + typeName +compositeTypeName (Type.CompositeInterfaceType interfaceType) = let Out.InterfaceType typeName _ _ _ = interfaceType - in Just typeName -compositeTypeName (Out.UnionBaseType (Out.UnionType typeName _ _)) = - Just typeName -compositeTypeName (Out.ScalarBaseType _) = - Nothing -compositeTypeName (Out.EnumBaseType _) = - Nothing -compositeTypeName (Out.ListBaseType wrappedType) = - compositeTypeName wrappedType + in typeName +compositeTypeName (Type.CompositeUnionType (Out.UnionType typeName _ _)) = + typeName + +typeNameIfComposite :: forall m. Out.Type m -> Maybe Full.Name +typeNameIfComposite = fmap compositeTypeName . Type.outToComposite -- | Field selections on scalars or enums are never allowed, because they are -- the leaf nodes of any GraphQL query. @@ -794,7 +797,7 @@ knownArgumentNamesRule = ArgumentsRule fieldRule directiveRule where fieldRule (Just objectType) (Full.Field _ fieldName arguments _ _ _) | Just typeField <- Type.lookupTypeField fieldName objectType - , Just typeName <- compositeTypeName objectType = + , Just typeName <- typeNameIfComposite objectType = lift $ foldr (go typeName fieldName typeField) Seq.empty arguments fieldRule _ _ = lift mempty go typeName fieldName fieldDefinition (Full.Argument argumentName _ location') errors @@ -1013,3 +1016,590 @@ providedRequiredInputFieldsRule = ValueRule go constGo , Text.unpack typeName , "\" is required, but it was not provided." ] + +-- | If multiple field selections with the same response names are encountered +-- during execution, the field and arguments to execute and the resulting value +-- should be unambiguous. Therefore any two field selections which might both be +-- encountered for the same object are only valid if they are equivalent. +-- +-- For simple hand‐written GraphQL, this rule is obviously a clear developer +-- error, however nested fragments can make this difficult to detect manually. +overlappingFieldsCanBeMergedRule :: Rule m +overlappingFieldsCanBeMergedRule = OperationDefinitionRule $ \case + Full.SelectionSet selectionSet _ -> do + schema' <- asks schema + go (toList selectionSet) + $ Type.CompositeObjectType + $ Schema.query schema' + Full.OperationDefinition operationType _ _ _ selectionSet _ -> do + schema' <- asks schema + let root = go (toList selectionSet) . Type.CompositeObjectType + case operationType of + Full.Query -> root $ Schema.query schema' + Full.Mutation + | Just objectType <- Schema.mutation schema' -> root objectType + Full.Subscription + | Just objectType <- Schema.mutation schema' -> root objectType + _ -> lift mempty + where + go selectionSet selectionType = do + fieldTuples <- evalStateT (collectFields selectionType selectionSet) HashSet.empty + fieldsInSetCanMerge fieldTuples + fieldsInSetCanMerge :: forall m + . HashMap Full.Name (NonEmpty (Full.Field, Type.CompositeType m)) + -> ReaderT (Validation m) Seq Error + fieldsInSetCanMerge fieldTuples = do + validation <- ask + let (lonely, paired) = flattenPairs fieldTuples + let reader = flip runReaderT validation + lift $ foldMap (reader . visitLonelyFields) lonely + <> foldMap (reader . forEachFieldTuple) paired + forEachFieldTuple :: forall m + . (FieldInfo m, FieldInfo m) + -> ReaderT (Validation m) Seq Error + forEachFieldTuple (fieldA, fieldB) = + case (parent fieldA, parent fieldB) of + (parentA@Type.CompositeObjectType{}, parentB@Type.CompositeObjectType{}) + | parentA /= parentB -> sameResponseShape fieldA fieldB + _ -> mapReaderT (checkEquality (node fieldA) (node fieldB)) + $ sameResponseShape fieldA fieldB + checkEquality fieldA fieldB Seq.Empty + | Full.Field _ fieldNameA _ _ _ _ <- fieldA + , Full.Field _ fieldNameB _ _ _ _ <- fieldB + , fieldNameA /= fieldNameB = pure $ makeError fieldA fieldB + | Full.Field _ fieldNameA argumentsA _ _ locationA <- fieldA + , Full.Field _ _ argumentsB _ _ locationB <- fieldB + , argumentsA /= argumentsB = + let message = concat + [ "Fields \"" + , Text.unpack fieldNameA + , "\" conflict because they have different arguments. Use " + , "different aliases on the fields to fetch both if this " + , "was intentional." + ] + in pure $ Error message [locationB, locationA] + checkEquality _ _ previousErrors = previousErrors + visitLonelyFields FieldInfo{..} = + let Full.Field _ _ _ _ subSelections _ = node + compositeFieldType = Type.outToComposite type' + in maybe (lift Seq.empty) (go subSelections) compositeFieldType + sameResponseShape :: forall m + . FieldInfo m + -> FieldInfo m + -> ReaderT (Validation m) Seq Error + sameResponseShape fieldA fieldB = + let Full.Field _ _ _ _ selectionsA _ = node fieldA + Full.Field _ _ _ _ selectionsB _ = node fieldB + in case unwrapTypes (type' fieldA) (type' fieldB) of + Left True -> lift mempty + Right (compositeA, compositeB) -> do + validation <- ask + let collectFields' composite = flip runReaderT validation + . flip evalStateT HashSet.empty + . collectFields composite + let collectA = collectFields' compositeA selectionsA + let collectB = collectFields' compositeB selectionsB + fieldsInSetCanMerge + $ foldl' (HashMap.unionWith (<>)) HashMap.empty + $ collectA <> collectB + _ -> pure $ makeError (node fieldA) (node fieldB) + makeError fieldA fieldB = + let Full.Field aliasA fieldNameA _ _ _ locationA = fieldA + Full.Field _ fieldNameB _ _ _ locationB = fieldB + message = concat + [ "Fields \"" + , Text.unpack (fromMaybe fieldNameA aliasA) + , "\" conflict because \"" + , Text.unpack fieldNameB + , "\" and \"" + , Text.unpack fieldNameA + , "\" are different fields. Use different aliases on the fields " + , "to fetch both if this was intentional." + ] + in Error message [locationB, locationA] + unwrapTypes typeA@Out.ScalarBaseType{} typeB@Out.ScalarBaseType{} = + Left $ typeA == typeB + unwrapTypes typeA@Out.EnumBaseType{} typeB@Out.EnumBaseType{} = + Left $ typeA == typeB + unwrapTypes (Out.ListType listA) (Out.ListType listB) = + unwrapTypes listA listB + unwrapTypes (Out.NonNullListType listA) (Out.NonNullListType listB) = + unwrapTypes listA listB + unwrapTypes typeA typeB + | Out.isNonNullType typeA == Out.isNonNullType typeB + , Just compositeA <- Type.outToComposite typeA + , Just compositeB <- Type.outToComposite typeB = + Right (compositeA, compositeB) + | otherwise = Left False + flattenPairs :: forall m + . HashMap Full.Name (NonEmpty (Full.Field, Type.CompositeType m)) + -> (Seq (FieldInfo m), Seq (FieldInfo m, FieldInfo m)) + flattenPairs xs = HashMap.foldr splitSingleFields (Seq.empty, Seq.empty) + $ foldr lookupTypeField [] <$> xs + splitSingleFields :: forall m + . [FieldInfo m] + -> (Seq (FieldInfo m), Seq (FieldInfo m, FieldInfo m)) + -> (Seq (FieldInfo m), Seq (FieldInfo m, FieldInfo m)) + splitSingleFields [head'] (fields, pairList) = (fields |> head', pairList) + splitSingleFields xs (fields, pairList) = (fields, pairs pairList xs) + lookupTypeField (field, parentType) accumulator = + let Full.Field _ fieldName _ _ _ _ = field + in case Type.lookupCompositeField fieldName parentType of + Nothing -> accumulator + Just (Out.Field _ typeField _) -> + FieldInfo field typeField parentType : accumulator + pairs :: forall m + . Seq (FieldInfo m, FieldInfo m) + -> [FieldInfo m] + -> Seq (FieldInfo m, FieldInfo m) + pairs accumulator [] = accumulator + pairs accumulator (fieldA : fields) = + pair fieldA (pairs accumulator fields) fields + pair _ accumulator [] = accumulator + pair field accumulator (fieldA : fields) = + pair field accumulator fields |> (field, fieldA) + collectFields objectType = accumulateFields objectType mempty + accumulateFields = foldM . forEach + forEach parentType accumulator = \case + Full.FieldSelection fieldSelection -> + forField parentType accumulator fieldSelection + Full.FragmentSpreadSelection fragmentSelection -> + forSpread accumulator fragmentSelection + Full.InlineFragmentSelection fragmentSelection -> + forInline parentType accumulator fragmentSelection + forField parentType accumulator field@(Full.Field alias fieldName _ _ _ _) = + let key = fromMaybe fieldName alias + value = (field, parentType) :| [] + in pure $ HashMap.insertWith (<>) key value accumulator + forSpread accumulator (Full.FragmentSpread fragmentName _ _) = do + inVisitetFragments <- gets $ HashSet.member fragmentName + if inVisitetFragments + then pure accumulator + else collectFromSpread fragmentName accumulator + forInline parentType accumulator = \case + Full.InlineFragment maybeType _ selections _ + | Just typeCondition <- maybeType -> + collectFromFragment typeCondition selections accumulator + | otherwise -> accumulateFields parentType accumulator $ toList selections + collectFromFragment typeCondition selectionSet' accumulator = do + types' <- lift $ asks $ Schema.types . schema + case Type.lookupTypeCondition typeCondition types' of + Nothing -> pure accumulator + Just compositeType -> + accumulateFields compositeType accumulator $ toList selectionSet' + collectFromSpread fragmentName accumulator = do + modify $ HashSet.insert fragmentName + ast' <- lift $ asks ast + case findFragmentDefinition fragmentName ast' of + Nothing -> pure accumulator + Just (Full.FragmentDefinition _ typeCondition _ selectionSet' _) -> + collectFromFragment typeCondition selectionSet' accumulator + +data FieldInfo m = FieldInfo + { node :: Full.Field + , type' :: Out.Type m + , parent :: Type.CompositeType m + } + +-- | Fragments are declared on a type and will only apply when the runtime +-- object type matches the type condition. They also are spread within the +-- context of a parent type. A fragment spread is only valid if its type +-- condition could ever apply within the parent type. +possibleFragmentSpreadsRule :: forall m. Rule m +possibleFragmentSpreadsRule = SelectionRule go + where + go (Just parentType) (Full.InlineFragmentSelection fragmentSelection) + | Full.InlineFragment maybeType _ _ location' <- fragmentSelection + , Just typeCondition <- maybeType = do + (fragmentTypeName, parentTypeName) <- + compareTypes typeCondition parentType + pure $ Error + { message = concat + [ "Fragment cannot be spread here as objects of type \"" + , Text.unpack parentTypeName + , "\" can never be of type \"" + , Text.unpack fragmentTypeName + , "\"." + ] + , locations = [location'] + } + go (Just parentType) (Full.FragmentSpreadSelection fragmentSelection) + | Full.FragmentSpread fragmentName _ location' <- fragmentSelection = do + typeCondition <- findSpreadTarget fragmentName + (fragmentTypeName, parentTypeName) <- + compareTypes typeCondition parentType + pure $ Error + { message = concat + [ "Fragment \"" + , Text.unpack fragmentName + , "\" cannot be spread here as objects of type \"" + , Text.unpack parentTypeName + , "\" can never be of type \"" + , Text.unpack fragmentTypeName + , "\"." + ] + , locations = [location'] + } + go _ _ = lift mempty + compareTypes typeCondition parentType = do + types' <- asks $ Schema.types . schema + fragmentType <- lift + $ maybeToSeq + $ Type.lookupTypeCondition typeCondition types' + parentComposite <- lift + $ maybeToSeq + $ Type.outToComposite parentType + possibleFragments <- getPossibleTypes fragmentType + possibleParents <- getPossibleTypes parentComposite + let fragmentTypeName = compositeTypeName fragmentType + let parentTypeName = compositeTypeName parentComposite + if HashSet.null $ HashSet.intersection possibleFragments possibleParents + then pure (fragmentTypeName, parentTypeName) + else lift mempty + getPossibleTypeList (Type.CompositeObjectType objectType) = + pure [Schema.ObjectType objectType] + getPossibleTypeList (Type.CompositeUnionType unionType) = + let Out.UnionType _ _ members = unionType + in pure $ Schema.ObjectType <$> members + getPossibleTypeList (Type.CompositeInterfaceType interfaceType) = + let Out.InterfaceType typeName _ _ _ = interfaceType + in HashMap.lookupDefault [] typeName + <$> asks (Schema.implementations . schema) + getPossibleTypes compositeType + = foldr (HashSet.insert . internalTypeName) HashSet.empty + <$> getPossibleTypeList compositeType + +internalTypeName :: forall m. Schema.Type m -> Full.Name +internalTypeName (Schema.ScalarType (Definition.ScalarType typeName _)) = + typeName +internalTypeName (Schema.EnumType (Definition.EnumType typeName _ _)) = typeName +internalTypeName (Schema.ObjectType (Out.ObjectType typeName _ _ _)) = typeName +internalTypeName (Schema.InputObjectType (In.InputObjectType typeName _ _)) = + typeName +internalTypeName (Schema.InterfaceType (Out.InterfaceType typeName _ _ _)) = + typeName +internalTypeName (Schema.UnionType (Out.UnionType typeName _ _)) = typeName + +findSpreadTarget :: Full.Name -> ReaderT (Validation m1) Seq Full.TypeCondition +findSpreadTarget fragmentName = do + ast' <- asks ast + let target = find (isSpreadTarget fragmentName) ast' + lift $ maybeToSeq $ target >>= extractTypeCondition + where + extractTypeCondition (viewFragment -> Just fragmentDefinition) = + let Full.FragmentDefinition _ typeCondition _ _ _ = fragmentDefinition + in Just typeCondition + extractTypeCondition _ = Nothing + +visitFragmentDefinition :: forall m + . Text + -> ValidationState m (Maybe Full.FragmentDefinition) +visitFragmentDefinition fragmentName = do + definitions <- lift $ asks ast + visited <- gets (HashSet.member fragmentName) + modify (HashSet.insert fragmentName) + case find (isSpreadTarget fragmentName) definitions of + Just (viewFragment -> Just fragmentDefinition) + | not visited -> pure $ Just fragmentDefinition + _ -> pure Nothing + +-- | Variable usages must be compatible with the arguments they are passed to. +-- +-- Validation failures occur when variables are used in the context of types +-- that are complete mismatches, or if a nullable type in a variable is passed +-- to a non‐null argument type. +variablesInAllowedPositionRule :: forall m. Rule m +variablesInAllowedPositionRule = OperationDefinitionRule $ \case + Full.OperationDefinition operationType _ variables _ selectionSet _ -> do + schema' <- asks schema + let root = go variables (toList selectionSet) . Type.CompositeObjectType + case operationType of + Full.Query -> root $ Schema.query schema' + Full.Mutation + | Just objectType <- Schema.mutation schema' -> root objectType + Full.Subscription + | Just objectType <- Schema.mutation schema' -> root objectType + _ -> lift mempty + _ -> lift mempty + where + go variables selections selectionType = mapReaderT (foldr (<>) Seq.empty) + $ flip evalStateT HashSet.empty + $ visitSelectionSet variables selectionType + $ toList selections + visitSelectionSet :: Foldable t + => [Full.VariableDefinition] + -> Type.CompositeType m + -> t Full.Selection + -> ValidationState m (Seq Error) + visitSelectionSet variables selectionType selections = + foldM (evaluateSelection variables selectionType) mempty selections + evaluateFieldSelection variables selections accumulator = \case + Just newParentType -> do + let folder = evaluateSelection variables newParentType + selectionErrors <- foldM folder accumulator selections + pure $ accumulator <> selectionErrors + Nothing -> pure accumulator + evaluateSelection :: [Full.VariableDefinition] + -> Type.CompositeType m + -> Seq Error + -> Full.Selection + -> ValidationState m (Seq Error) + evaluateSelection variables selectionType accumulator selection + | Full.FragmentSpreadSelection spread <- selection + , Full.FragmentSpread fragmentName _ _ <- spread = do + types' <- lift $ asks $ Schema.types . schema + nonVisitedFragmentDefinition <- visitFragmentDefinition fragmentName + case nonVisitedFragmentDefinition of + Just fragmentDefinition + | Full.FragmentDefinition _ typeCondition _ _ _ <- fragmentDefinition + , Just spreadType <- Type.lookupTypeCondition typeCondition types' -> do + spreadErrors <- spreadVariables variables spread + selectionErrors <- diveIntoSpread variables spreadType fragmentDefinition + pure $ accumulator <> spreadErrors <> selectionErrors + _ -> lift $ lift mempty + | Full.FieldSelection fieldSelection <- selection + , Full.Field _ fieldName _ _ subselections _ <- fieldSelection = + case Type.lookupCompositeField fieldName selectionType of + Just (Out.Field _ typeField argumentTypes) -> do + fieldErrors <- fieldVariables variables argumentTypes fieldSelection + selectionErrors <- evaluateFieldSelection variables subselections accumulator + $ Type.outToComposite typeField + pure $ selectionErrors <> fieldErrors + Nothing -> pure accumulator + | Full.InlineFragmentSelection inlineSelection <- selection + , Full.InlineFragment typeCondition _ subselections _ <- inlineSelection = do + types' <- lift $ asks $ Schema.types . schema + let inlineType = fromMaybe selectionType + $ typeCondition >>= flip Type.lookupTypeCondition types' + fragmentErrors <- inlineVariables variables inlineSelection + let folder = evaluateSelection variables inlineType + selectionErrors <- foldM folder accumulator subselections + pure $ accumulator <> fragmentErrors <> selectionErrors + inlineVariables variables inline + | Full.InlineFragment _ directives' _ _ <- inline = + mapDirectives variables directives' + fieldVariables :: [Full.VariableDefinition] + -> In.Arguments + -> Full.Field + -> ValidationState m (Seq Error) + fieldVariables variables argumentTypes fieldSelection = do + let Full.Field _ _ arguments directives' _ _ = fieldSelection + argumentErrors <- mapArguments variables argumentTypes arguments + directiveErrors <- mapDirectives variables directives' + pure $ argumentErrors <> directiveErrors + spreadVariables variables (Full.FragmentSpread _ directives' _) = + mapDirectives variables directives' + diveIntoSpread variables fieldType fragmentDefinition = do + let Full.FragmentDefinition _ _ directives' selections _ = + fragmentDefinition + selectionErrors <- visitSelectionSet variables fieldType selections + directiveErrors <- mapDirectives variables directives' + pure $ selectionErrors <> directiveErrors + findDirectiveVariables variables directive = do + let Full.Directive directiveName arguments _ = directive + directiveDefinitions <- lift $ asks $ Schema.directives . schema + case HashMap.lookup directiveName directiveDefinitions of + Just (Schema.Directive _ _ directiveArguments) -> + mapArguments variables directiveArguments arguments + Nothing -> pure mempty + mapArguments variables argumentTypes = fmap fold + . traverse (findArgumentVariables variables argumentTypes) + mapDirectives variables = fmap fold + <$> traverse (findDirectiveVariables variables) + lookupInputObject variables objectFieldValue locationInfo + | Full.Node{ node = Full.Object objectFields } <- objectFieldValue + , Just (expectedType, _) <- locationInfo + , In.InputObjectBaseType inputObjectType <- expectedType + , In.InputObjectType _ _ fieldTypes' <- inputObjectType = + fold <$> traverse (traverseObjectField variables fieldTypes') objectFields + | otherwise = pure mempty + maybeUsageAllowed variableName variables locationInfo + | Just (locationType, locationValue) <- locationInfo + , findVariableDefinition' <- findVariableDefinition variableName + , Just variableDefinition <- find findVariableDefinition' variables + = maybeToSeq + <$> isVariableUsageAllowed locationType locationValue variableDefinition + | otherwise = pure mempty + findArgumentVariables :: [Full.VariableDefinition] + -> HashMap Full.Name In.Argument + -> Full.Argument + -> ValidationState m (Seq Error) + findArgumentVariables variables argumentTypes argument + | Full.Argument argumentName argumentValue _ <- argument + , Full.Node{ node = Full.Variable variableName } <- argumentValue + = maybeUsageAllowed variableName variables + $ locationPair extractArgument argumentTypes argumentName + | Full.Argument argumentName argumentValue _ <- argument + = lookupInputObject variables argumentValue + $ locationPair extractArgument argumentTypes argumentName + extractField (In.InputField _ locationType locationValue) = + (locationType, locationValue) + extractArgument (In.Argument _ locationType locationValue) = + (locationType, locationValue) + locationPair extract fieldTypes name = + extract <$> HashMap.lookup name fieldTypes + traverseObjectField variables fieldTypes Full.ObjectField{..} + | Full.Node{ node = Full.Variable variableName } <- value + = maybeUsageAllowed variableName variables + $ locationPair extractField fieldTypes name + | otherwise = lookupInputObject variables value + $ locationPair extractField fieldTypes name + findVariableDefinition variableName variableDefinition = + let Full.VariableDefinition variableName' _ _ _ = variableDefinition + in variableName == variableName' + isVariableUsageAllowed locationType locationDefaultValue variableDefinition + | Full.VariableDefinition _ variableType _ _ <- variableDefinition + , Full.TypeNonNull _ <- variableType = + typesCompatibleOrError variableDefinition locationType + | Just nullableLocationType <- unwrapInType locationType + , Full.VariableDefinition _ variableType variableDefaultValue _ <- + variableDefinition + , hasNonNullVariableDefaultValue' <- + hasNonNullVariableDefaultValue variableDefaultValue + , hasLocationDefaultValue <- isJust locationDefaultValue = + if (hasNonNullVariableDefaultValue' || hasLocationDefaultValue) + && areTypesCompatible variableType nullableLocationType + then pure Nothing + else pure $ makeError variableDefinition locationType + | otherwise = typesCompatibleOrError variableDefinition locationType + typesCompatibleOrError variableDefinition locationType + | Full.VariableDefinition _ variableType _ _ <- variableDefinition + , areTypesCompatible variableType locationType = pure Nothing + | otherwise = pure $ makeError variableDefinition locationType + areTypesCompatible nonNullType (unwrapInType -> Just nullableLocationType) + | Full.TypeNonNull (Full.NonNullTypeNamed namedType) <- nonNullType = + areTypesCompatible (Full.TypeNamed namedType) nullableLocationType + | Full.TypeNonNull (Full.NonNullTypeList namedList) <- nonNullType = + areTypesCompatible (Full.TypeList namedList) nullableLocationType + areTypesCompatible _ (In.isNonNullType -> True) = False + areTypesCompatible (Full.TypeNonNull nonNullType) locationType + | Full.NonNullTypeNamed namedType <- nonNullType = + areTypesCompatible (Full.TypeNamed namedType) locationType + | Full.NonNullTypeList namedType <- nonNullType = + areTypesCompatible (Full.TypeList namedType) locationType + areTypesCompatible variableType locationType + | Full.TypeList itemVariableType <- variableType + , In.ListType itemLocationType <- locationType = + areTypesCompatible itemVariableType itemLocationType + | areIdentical variableType locationType = True + | otherwise = False + areIdentical (Full.TypeList typeList) (In.ListType itemLocationType) = + areIdentical typeList itemLocationType + areIdentical (Full.TypeNonNull nonNullType) locationType + | Full.NonNullTypeList nonNullList <- nonNullType + , In.NonNullListType itemLocationType <- locationType = + areIdentical nonNullList itemLocationType + | Full.NonNullTypeNamed _ <- nonNullType + , In.ListBaseType _ <- locationType = False + | Full.NonNullTypeNamed nonNullList <- nonNullType + , In.isNonNullType locationType = + nonNullList == inputTypeName locationType + areIdentical (Full.TypeNamed _) (In.ListBaseType _) = False + areIdentical (Full.TypeNamed typeNamed) locationType + | not $ In.isNonNullType locationType = + typeNamed == inputTypeName locationType + areIdentical _ _ = False + hasNonNullVariableDefaultValue (Just (Full.Node Full.ConstNull _)) = False + hasNonNullVariableDefaultValue Nothing = False + hasNonNullVariableDefaultValue _ = True + unwrapInType (In.NonNullScalarType nonNullType) = + Just $ In.NamedScalarType nonNullType + unwrapInType (In.NonNullEnumType nonNullType) = + Just $ In.NamedEnumType nonNullType + unwrapInType (In.NonNullInputObjectType nonNullType) = + Just $ In.NamedInputObjectType nonNullType + unwrapInType (In.NonNullListType nonNullType) = + Just $ In.ListType nonNullType + unwrapInType _ = Nothing + makeError variableDefinition expectedType = + let Full.VariableDefinition variableName variableType _ location' = + variableDefinition + in Just $ Error + { message = concat + [ "Variable \"$" + , Text.unpack variableName + , "\" of type \"" + , show variableType + , "\" used in position expecting type \"" + , show expectedType + , "\"." + ] + , locations = [location'] + } + +-- | Literal values must be compatible with the type expected in the position +-- they are found as per the coercion rules. +-- +-- The type expected in a position include the type defined by the argument a +-- value is provided for, the type defined by an input object field a value is +-- provided for, and the type of a variable definition a default value is +-- provided for. +valuesOfCorrectTypeRule :: forall m. Rule m +valuesOfCorrectTypeRule = ValueRule go constGo + where + go (Just inputType) value + | Just constValue <- toConstNode value = + lift $ check inputType constValue + go _ _ = lift mempty + toConstNode Full.Node{..} = flip Full.Node location <$> toConst node + toConst (Full.Variable _) = Nothing + toConst (Full.Int integer) = Just $ Full.ConstInt integer + toConst (Full.Float double) = Just $ Full.ConstFloat double + toConst (Full.String string) = Just $ Full.ConstString string + toConst (Full.Boolean boolean) = Just $ Full.ConstBoolean boolean + toConst Full.Null = Just Full.ConstNull + toConst (Full.Enum enum) = Just $ Full.ConstEnum enum + toConst (Full.List values) = + Just $ Full.ConstList $ catMaybes $ toConst <$> values + toConst (Full.Object fields) = Just $ Full.ConstObject + $ catMaybes $ constObjectField <$> fields + constObjectField Full.ObjectField{..} + | Just constValue <- toConstNode value = + Just $ Full.ObjectField name constValue location + | otherwise = Nothing + constGo Nothing = const $ lift mempty + constGo (Just inputType) = lift . check inputType + check :: In.Type -> Full.Node Full.ConstValue -> Seq Error + check _ Full.Node{ node = Full.ConstNull } = + mempty -- Ignore, required fields are checked elsewhere. + check (In.ScalarBaseType scalarType) Full.Node{ node } + | Definition.ScalarType "Int" _ <- scalarType + , Full.ConstInt _ <- node = mempty + | Definition.ScalarType "Boolean" _ <- scalarType + , Full.ConstBoolean _ <- node = mempty + | Definition.ScalarType "String" _ <- scalarType + , Full.ConstString _ <- node = mempty + | Definition.ScalarType "ID" _ <- scalarType + , Full.ConstString _ <- node = mempty + | Definition.ScalarType "ID" _ <- scalarType + , Full.ConstInt _ <- node = mempty + | Definition.ScalarType "Float" _ <- scalarType + , Full.ConstFloat _ <- node = mempty + | Definition.ScalarType "Float" _ <- scalarType + , Full.ConstInt _ <- node = mempty + check (In.EnumBaseType enumType) Full.Node{ node } + | Definition.EnumType _ _ members <- enumType + , Full.ConstEnum memberValue <- node + , HashMap.member memberValue members = mempty + check (In.InputObjectBaseType objectType) Full.Node{ node } + | In.InputObjectType _ _ typeFields <- objectType + , Full.ConstObject valueFields <- node = + foldMap (checkObjectField typeFields) valueFields + check (In.ListBaseType listType) constValue@Full.Node{ .. } + | Full.ConstList listValues <- node = + foldMap (check listType) $ flip Full.Node location <$> listValues + | otherwise = check listType constValue + check inputType Full.Node{ .. } = pure $ Error + { message = concat + [ "Value " + , show node, " cannot be coerced to type \"" + , show inputType + , "\"." + ] + , locations = [location] + } + checkObjectField typeFields Full.ObjectField{..} + | Just typeField <- HashMap.lookup name typeFields + , In.InputField _ fieldType _ <- typeField = + check fieldType value + checkObjectField _ _ = mempty |
