diff options
Diffstat (limited to 'src/Language')
| -rw-r--r-- | src/Language/GraphQL/AST/Document.hs | 80 | ||||
| -rw-r--r-- | src/Language/GraphQL/AST/Encoder.hs | 25 | ||||
| -rw-r--r-- | src/Language/GraphQL/Type/Definition.hs | 7 | ||||
| -rw-r--r-- | src/Language/GraphQL/Type/In.hs | 14 | ||||
| -rw-r--r-- | src/Language/GraphQL/Type/Internal.hs | 43 | ||||
| -rw-r--r-- | src/Language/GraphQL/Type/Out.hs | 24 | ||||
| -rw-r--r-- | src/Language/GraphQL/Type/Schema.hs | 22 | ||||
| -rw-r--r-- | src/Language/GraphQL/Validate.hs | 2 | ||||
| -rw-r--r-- | src/Language/GraphQL/Validate/Rules.hs | 676 |
9 files changed, 811 insertions, 82 deletions
diff --git a/src/Language/GraphQL/AST/Document.hs b/src/Language/GraphQL/AST/Document.hs index b30271c..a78b007 100644 --- a/src/Language/GraphQL/AST/Document.hs +++ b/src/Language/GraphQL/AST/Document.hs @@ -1,5 +1,6 @@ {-# LANGUAGE DuplicateRecordFields #-} {-# LANGUAGE ExplicitForAll #-} +{-# LANGUAGE NamedFieldPuns #-} {-# LANGUAGE OverloadedStrings #-} {-# LANGUAGE RecordWildCards #-} {-# LANGUAGE Safe #-} @@ -47,11 +48,15 @@ module Language.GraphQL.AST.Document , UnionMemberTypes(..) , Value(..) , VariableDefinition(..) + , escape ) where +import Data.Char (ord) import Data.Foldable (toList) import Data.Int (Int32) +import Data.List (intercalate) import Data.List.NonEmpty (NonEmpty) +import Numeric (showFloat, showHex) import Data.Text (Text) import qualified Data.Text as Text import Language.GraphQL.AST.DirectiveLocation (DirectiveLocation) @@ -79,7 +84,10 @@ instance Ord Location where data Node a = Node { node :: a , location :: Location - } deriving (Eq, Show) + } deriving Eq + +instance Show a => Show (Node a) where + show Node{ node } = show node instance Functor Node where fmap f Node{..} = Node (f node) location @@ -218,6 +226,28 @@ type TypeCondition = Name -- ** Input Values +escape :: Char -> String +escape char' + | char' == '\\' = "\\\\" + | char' == '\"' = "\\\"" + | char' == '\b' = "\\b" + | char' == '\f' = "\\f" + | char' == '\n' = "\\n" + | char' == '\r' = "\\r" + | char' == '\t' = "\\t" + | char' < '\x0010' = unicode "\\u000" char' + | char' < '\x0020' = unicode "\\u00" char' + | otherwise = [char'] + where + unicode prefix uchar = prefix <> (showHex $ ord uchar) "" + +showList' :: Show a => [a] -> String +showList' list = "[" ++ intercalate ", " (show <$> list) ++ "]" + +showObject :: Show a => [ObjectField a] -> String +showObject fields = + "{ " ++ intercalate ", " (show <$> fields) ++ " }" + -- | Input value (literal or variable). data Value = Variable Name @@ -229,7 +259,19 @@ data Value | Enum Name | List [Value] | Object [ObjectField Value] - deriving (Eq, Show) + deriving Eq + +instance Show Value where + showList = mappend . showList' + show (Variable variableName) = '$' : Text.unpack variableName + show (Int integer) = show integer + show (Float float) = show $ ConstFloat float + show (String text) = show $ ConstString text + show (Boolean boolean) = show boolean + show Null = "null" + show (Enum name) = Text.unpack name + show (List list) = show list + show (Object fields) = showObject fields -- | Constant input value. data ConstValue @@ -241,7 +283,18 @@ data ConstValue | ConstEnum Name | ConstList [ConstValue] | ConstObject [ObjectField ConstValue] - deriving (Eq, Show) + deriving Eq + +instance Show ConstValue where + showList = mappend . showList' + show (ConstInt integer) = show integer + show (ConstFloat float) = showFloat float mempty + show (ConstString text) = "\"" <> Text.foldr (mappend . escape) "\"" text + show (ConstBoolean boolean) = show boolean + show ConstNull = "null" + show (ConstEnum name) = Text.unpack name + show (ConstList list) = show list + show (ConstObject fields) = showObject fields -- | Key-value pair. -- @@ -250,7 +303,13 @@ data ObjectField a = ObjectField { name :: Name , value :: Node a , location :: Location - } deriving (Eq, Show) + } deriving Eq + +instance Show a => Show (ObjectField a) where + show ObjectField{..} = Text.unpack name ++ ": " ++ show value + +instance Functor ObjectField where + fmap f ObjectField{..} = ObjectField name (f <$> value) location -- ** Variables @@ -281,7 +340,12 @@ data Type = TypeNamed Name | TypeList Type | TypeNonNull NonNullType - deriving (Eq, Show) + deriving Eq + +instance Show Type where + show (TypeNamed typeName) = Text.unpack typeName + show (TypeList listType) = concat ["[", show listType, "]"] + show (TypeNonNull nonNullType) = show nonNullType -- | Represents type names. type NamedType = Name @@ -290,7 +354,11 @@ type NamedType = Name data NonNullType = NonNullTypeNamed Name | NonNullTypeList Type - deriving (Eq, Show) + deriving Eq + +instance Show NonNullType where + show (NonNullTypeNamed typeName) = '!' : Text.unpack typeName + show (NonNullTypeList listType) = concat ["![", show listType, "]"] -- ** Directives diff --git a/src/Language/GraphQL/AST/Encoder.hs b/src/Language/GraphQL/AST/Encoder.hs index 9ba51b8..f04f385 100644 --- a/src/Language/GraphQL/AST/Encoder.hs +++ b/src/Language/GraphQL/AST/Encoder.hs @@ -16,7 +16,6 @@ module Language.GraphQL.AST.Encoder , value ) where -import Data.Char (ord) import Data.Foldable (fold) import qualified Data.List.NonEmpty as NonEmpty import Data.Text (Text) @@ -25,7 +24,7 @@ import qualified Data.Text.Lazy as Lazy (Text) import qualified Data.Text.Lazy as Lazy.Text import Data.Text.Lazy.Builder (Builder) import qualified Data.Text.Lazy.Builder as Builder -import Data.Text.Lazy.Builder.Int (decimal, hexadecimal) +import Data.Text.Lazy.Builder.Int (decimal) import Data.Text.Lazy.Builder.RealFloat (realFloat) import qualified Language.GraphQL.AST.Document as Full @@ -234,11 +233,12 @@ quote :: Builder.Builder quote = Builder.singleton '\"' oneLine :: Text -> Builder -oneLine string = quote <> Text.foldr (mappend . escape) quote string +oneLine string = quote <> Text.foldr merge quote string + where + merge = mappend . Builder.fromString . Full.escape stringValue :: Formatter -> Text -> Lazy.Text -stringValue Minified string = Builder.toLazyText - $ quote <> Text.foldr (mappend . escape) quote string +stringValue Minified string = Builder.toLazyText $ oneLine string stringValue (Pretty indentation) string = if hasEscaped string then stringValue Minified string @@ -266,21 +266,6 @@ stringValue (Pretty indentation) string = = Builder.fromLazyText (indent (indentation + 1)) <> line' <> newline <> acc -escape :: Char -> Builder -escape char' - | char' == '\\' = Builder.fromString "\\\\" - | char' == '\"' = Builder.fromString "\\\"" - | char' == '\b' = Builder.fromString "\\b" - | char' == '\f' = Builder.fromString "\\f" - | char' == '\n' = Builder.fromString "\\n" - | char' == '\r' = Builder.fromString "\\r" - | char' == '\t' = Builder.fromString "\\t" - | char' < '\x0010' = unicode "\\u000" char' - | char' < '\x0020' = unicode "\\u00" char' - | otherwise = Builder.singleton char' - where - unicode prefix = mappend (Builder.fromString prefix) . (hexadecimal . ord) - listValue :: Formatter -> [Full.Value] -> Lazy.Text listValue formatter = bracketsCommas formatter $ value formatter diff --git a/src/Language/GraphQL/Type/Definition.hs b/src/Language/GraphQL/Type/Definition.hs index 476fb3a..6ca77aa 100644 --- a/src/Language/GraphQL/Type/Definition.hs +++ b/src/Language/GraphQL/Type/Definition.hs @@ -22,6 +22,7 @@ import Data.HashMap.Strict (HashMap) import qualified Data.HashMap.Strict as HashMap import Data.String (IsString(..)) import Data.Text (Text) +import qualified Data.Text as Text import Language.GraphQL.AST (Name) import Prelude hiding (id) @@ -63,6 +64,9 @@ data ScalarType = ScalarType Name (Maybe Text) instance Eq ScalarType where (ScalarType this _) == (ScalarType that _) = this == that +instance Show ScalarType where + show (ScalarType typeName _) = Text.unpack typeName + -- | Enum type definition. -- -- Some leaf values of requests and input values are Enums. GraphQL serializes @@ -73,6 +77,9 @@ data EnumType = EnumType Name (Maybe Text) (HashMap Name EnumValue) instance Eq EnumType where (EnumType this _ _) == (EnumType that _ _) = this == that +instance Show EnumType where + show (EnumType typeName _ _) = Text.unpack typeName + -- | Enum value is a single member of an 'EnumType'. newtype EnumValue = EnumValue (Maybe Text) diff --git a/src/Language/GraphQL/Type/In.hs b/src/Language/GraphQL/Type/In.hs index 59a6d59..d42599b 100644 --- a/src/Language/GraphQL/Type/In.hs +++ b/src/Language/GraphQL/Type/In.hs @@ -24,6 +24,7 @@ module Language.GraphQL.Type.In import Data.HashMap.Strict (HashMap) import Data.Text (Text) +import qualified Data.Text as Text import Language.GraphQL.AST.Document (Name) import qualified Language.GraphQL.Type.Definition as Definition @@ -40,6 +41,9 @@ data InputObjectType = InputObjectType instance Eq InputObjectType where (InputObjectType this _ _) == (InputObjectType that _ _) = this == that +instance Show InputObjectType where + show (InputObjectType typeName _ _) = Text.unpack typeName + -- | These types may be used as input types for arguments and directives. -- -- GraphQL distinguishes between "wrapping" and "named" types. Each wrapping @@ -56,6 +60,16 @@ data Type | NonNullListType Type deriving Eq +instance Show Type where + show (NamedScalarType scalarType) = show scalarType + show (NamedEnumType enumType) = show enumType + show (NamedInputObjectType inputObjectType) = show inputObjectType + show (ListType baseType) = concat ["[", show baseType, "]"] + show (NonNullScalarType scalarType) = '!' : show scalarType + show (NonNullEnumType enumType) = '!' : show enumType + show (NonNullInputObjectType inputObjectType) = '!' : show inputObjectType + show (NonNullListType baseType) = concat ["![", show baseType, "]"] + -- | Field argument definition. data Argument = Argument (Maybe Text) Type (Maybe Definition.Value) diff --git a/src/Language/GraphQL/Type/Internal.hs b/src/Language/GraphQL/Type/Internal.hs index eb8489c..2081b97 100644 --- a/src/Language/GraphQL/Type/Internal.hs +++ b/src/Language/GraphQL/Type/Internal.hs @@ -14,11 +14,14 @@ module Language.GraphQL.Type.Internal , Type(..) , directives , doesFragmentTypeApply + , implementations , instanceOf + , lookupCompositeField , lookupInputType , lookupTypeCondition , lookupTypeField , mutation + , outToComposite , subscription , query , types @@ -62,26 +65,31 @@ data Schema m = Schema (Maybe (Out.ObjectType m)) Directives (HashMap Full.Name (Type m)) + (HashMap Full.Name [Type m]) -- | Schema query type. query :: forall m. Schema m -> Out.ObjectType m -query (Schema query' _ _ _ _) = query' +query (Schema query' _ _ _ _ _) = query' -- | Schema mutation type. mutation :: forall m. Schema m -> Maybe (Out.ObjectType m) -mutation (Schema _ mutation' _ _ _) = mutation' +mutation (Schema _ mutation' _ _ _ _) = mutation' -- | Schema subscription type. subscription :: forall m. Schema m -> Maybe (Out.ObjectType m) -subscription (Schema _ _ subscription' _ _) = subscription' +subscription (Schema _ _ subscription' _ _ _) = subscription' -- | Schema directive definitions. directives :: forall m. Schema m -> Directives -directives (Schema _ _ _ directives' _) = directives' +directives (Schema _ _ _ directives' _ _) = directives' -- | Types referenced by the schema. types :: forall m. Schema m -> HashMap Full.Name (Type m) -types (Schema _ _ _ _ types') = types' +types (Schema _ _ _ _ types' _) = types' + +-- | Interface implementations. +implementations :: forall m. Schema m -> HashMap Full.Name [Type m] +implementations (Schema _ _ _ _ _ implementations') = implementations' -- | These types may describe the parent context of a selection set. data CompositeType m @@ -160,12 +168,16 @@ lookupInputType (Full.TypeNonNull (Full.NonNullTypeList nonNull)) types' <$> lookupInputType nonNull types' lookupTypeField :: forall a. Full.Name -> Out.Type a -> Maybe (Out.Field a) -lookupTypeField fieldName = \case - Out.ObjectBaseType objectType -> - objectChild objectType - Out.InterfaceBaseType interfaceType -> - interfaceChild interfaceType - Out.ListBaseType listType -> lookupTypeField fieldName listType +lookupTypeField fieldName outputType = + outToComposite outputType >>= lookupCompositeField fieldName + +lookupCompositeField :: forall a + . Full.Name + -> CompositeType a + -> Maybe (Out.Field a) +lookupCompositeField fieldName = \case + CompositeObjectType objectType -> objectChild objectType + CompositeInterfaceType interfaceType -> interfaceChild interfaceType _ -> Nothing where objectChild (Out.ObjectType _ _ _ resolvers) = @@ -174,3 +186,12 @@ lookupTypeField fieldName = \case HashMap.lookup fieldName fields resolverType (Out.ValueResolver objectField _) = objectField resolverType (Out.EventStreamResolver objectField _ _) = objectField + +outToComposite :: forall a. Out.Type a -> Maybe (CompositeType a) +outToComposite = \case + Out.ObjectBaseType objectType -> Just $ CompositeObjectType objectType + Out.InterfaceBaseType interfaceType -> + Just $ CompositeInterfaceType interfaceType + Out.UnionBaseType unionType -> Just $ CompositeUnionType unionType + Out.ListBaseType listType -> outToComposite listType + _ -> Nothing diff --git a/src/Language/GraphQL/Type/Out.hs b/src/Language/GraphQL/Type/Out.hs index b0668f5..847a8a5 100644 --- a/src/Language/GraphQL/Type/Out.hs +++ b/src/Language/GraphQL/Type/Out.hs @@ -38,6 +38,7 @@ import Data.HashMap.Strict (HashMap) import qualified Data.HashMap.Strict as HashMap import Data.Maybe (fromMaybe) import Data.Text (Text) +import qualified Data.Text as Text import Language.GraphQL.AST (Name) import Language.GraphQL.Type.Definition import qualified Language.GraphQL.Type.In as In @@ -52,6 +53,9 @@ data ObjectType m = ObjectType instance forall a. Eq (ObjectType a) where (ObjectType this _ _ _) == (ObjectType that _ _ _) = this == that +instance forall a. Show (ObjectType a) where + show (ObjectType typeName _ _ _) = Text.unpack typeName + -- | Interface Type Definition. -- -- When a field can return one of a heterogeneous set of types, a Interface type @@ -63,6 +67,9 @@ data InterfaceType m = InterfaceType instance forall a. Eq (InterfaceType a) where (InterfaceType this _ _ _) == (InterfaceType that _ _ _) = this == that +instance forall a. Show (InterfaceType a) where + show (InterfaceType typeName _ _ _) = Text.unpack typeName + -- | Union Type Definition. -- -- When a field can return one of a heterogeneous set of types, a Union type is @@ -72,6 +79,9 @@ data UnionType m = UnionType Name (Maybe Text) [ObjectType m] instance forall a. Eq (UnionType a) where (UnionType this _ _) == (UnionType that _ _) = this == that +instance forall a. Show (UnionType a) where + show (UnionType typeName _ _) = Text.unpack typeName + -- | Output object field definition. data Field m = Field (Maybe Text) -- ^ Description. @@ -98,6 +108,20 @@ data Type m | NonNullListType (Type m) deriving Eq +instance forall a. Show (Type a) where + show (NamedScalarType scalarType) = show scalarType + show (NamedEnumType enumType) = show enumType + show (NamedObjectType inputObjectType) = show inputObjectType + show (NamedInterfaceType interfaceType) = show interfaceType + show (NamedUnionType unionType) = show unionType + show (ListType baseType) = concat ["[", show baseType, "]"] + show (NonNullScalarType scalarType) = '!' : show scalarType + show (NonNullEnumType enumType) = '!' : show enumType + show (NonNullObjectType inputObjectType) = '!' : show inputObjectType + show (NonNullInterfaceType interfaceType) = '!' : show interfaceType + show (NonNullUnionType unionType) = '!' : show unionType + show (NonNullListType baseType) = concat ["![", show baseType, "]"] + -- | Matches either 'NamedScalarType' or 'NonNullScalarType'. pattern ScalarBaseType :: forall m. ScalarType -> Type m pattern ScalarBaseType scalarType <- (isScalarType -> Just scalarType) diff --git a/src/Language/GraphQL/Type/Schema.hs b/src/Language/GraphQL/Type/Schema.hs index 099c256..dae8e18 100644 --- a/src/Language/GraphQL/Type/Schema.hs +++ b/src/Language/GraphQL/Type/Schema.hs @@ -23,6 +23,7 @@ import Language.GraphQL.Type.Internal , Schema , Type(..) , directives + , implementations , mutation , subscription , query @@ -41,9 +42,11 @@ schema :: forall m -> Directives -- ^ Directive definitions. -> Schema m -- ^ Schema. schema queryRoot mutationRoot subscriptionRoot directiveDefinitions = - Internal.Schema queryRoot mutationRoot subscriptionRoot allDirectives collectedTypes + Internal.Schema queryRoot mutationRoot subscriptionRoot + allDirectives collectedTypes collectedImplementations where collectedTypes = collectReferencedTypes queryRoot mutationRoot subscriptionRoot + collectedImplementations = collectImplementations collectedTypes allDirectives = HashMap.union directiveDefinitions defaultDirectives defaultDirectives = HashMap.fromList [ ("skip", skipDirective) @@ -153,3 +156,20 @@ collectReferencedTypes queryRoot mutationRoot subscriptionRoot = polymorphicTraverser interfaces fields = flip (foldr visitFields) fields . flip (foldr traverseInterfaceType) interfaces + +-- | Looks for objects and interfaces under the schema types and collects the +-- interfaces they implement. +collectImplementations :: forall m + . HashMap Full.Name (Type m) + -> HashMap Full.Name [Type m] +collectImplementations = HashMap.foldr go HashMap.empty + where + go implementation@(InterfaceType interfaceType) accumulator = + let Out.InterfaceType _ _ interfaces _ = interfaceType + in foldr (add implementation) accumulator interfaces + go implementation@(ObjectType objectType) accumulator = + let Out.ObjectType _ _ interfaces _ = objectType + in foldr (add implementation) accumulator interfaces + go _ accumulator = accumulator + add implementation (Out.InterfaceType typeName _ _ _) accumulator = + HashMap.insertWith (++) typeName [implementation] accumulator diff --git a/src/Language/GraphQL/Validate.hs b/src/Language/GraphQL/Validate.hs index 277f84d..ea72018 100644 --- a/src/Language/GraphQL/Validate.hs +++ b/src/Language/GraphQL/Validate.hs @@ -210,7 +210,7 @@ typeDefinition context rule = \case directives context rule scalarLocation directives' Full.ObjectTypeDefinition _ _ _ directives' fields -> directives context rule objectLocation directives' - >< foldMap (fieldDefinition context rule) fields + >< foldMap (fieldDefinition context rule) fields Full.InterfaceTypeDefinition _ _ directives' fields -> directives context rule interfaceLocation directives' >< foldMap (fieldDefinition context rule) fields 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 |
