aboutsummaryrefslogtreecommitdiff
path: root/src/Language/GraphQL/Validate/Rules.hs
diff options
context:
space:
mode:
Diffstat (limited to 'src/Language/GraphQL/Validate/Rules.hs')
-rw-r--r--src/Language/GraphQL/Validate/Rules.hs676
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