diff --git a/src/main/kotlin/graphql/kickstart/tools/SchemaParser.kt b/src/main/kotlin/graphql/kickstart/tools/SchemaParser.kt index 380db4e1..e8216cff 100644 --- a/src/main/kotlin/graphql/kickstart/tools/SchemaParser.kt +++ b/src/main/kotlin/graphql/kickstart/tools/SchemaParser.kt @@ -380,19 +380,16 @@ class SchemaParser internal constructor( } } .apply { - // a bare @deprecated has no "reason" argument, which makes SchemaPrinter throw a NPE. - // copy the default from the directive definition (for the built-in one: "No longer supported"). - if (directive.name == Directives.DeprecatedDirective.name && directive.arguments.none { it.name == "reason" }) { - val reasonArgument = graphQLArguments["reason"] - if (reasonArgument != null && reasonArgument.hasSetDefaultValue()) { - argument(GraphQLAppliedDirectiveArgument.newArgument() - .name(reasonArgument.name) - .type(reasonArgument.type) - .description(reasonArgument.description) - .inputValueWithState(reasonArgument.argumentDefaultValue) - .build() - ) - } + // arguments that weren't supplied get the default from the directive definition, like graphql-java does. + // this also gives a bare @deprecated its "reason", without which SchemaPrinter throws a NPE. + missingArgumentsWithDefault(directive, graphQLDirective).forEach { graphQLArgument -> + argument(GraphQLAppliedDirectiveArgument.newArgument() + .name(graphQLArgument.name) + .type(graphQLArgument.type) + .description(graphQLArgument.description) + .inputValueWithState(graphQLArgument.argumentDefaultValue) + .build() + ) } } .build() @@ -465,6 +462,15 @@ class SchemaParser internal constructor( .valueLiteral(arg.value) .build()) } + missingArgumentsWithDefault(directive, graphQLDirective).forEach { graphQLArgument -> + val defaultValue = graphQLArgument.argumentDefaultValue + argument(GraphQLArgument.newArgument() + .name(graphQLArgument.name) + .type(graphQLArgument.type) + .description(graphQLArgument.description) + .apply { if (defaultValue.isLiteral) valueLiteral(defaultValue.value as Value<*>) else valueProgrammatic(defaultValue.value) } + .build()) + } } .build() ) @@ -474,6 +480,9 @@ class SchemaParser internal constructor( return output.toTypedArray() } + private fun missingArgumentsWithDefault(directive: Directive, graphQLDirective: GraphQLDirective): List = + graphQLDirective.arguments.filter { it.hasSetDefaultValue() && directive.getArgument(it.name) == null } + private fun determineOutputType(typeDefinition: Type<*>, inputObjects: List) = determineType(GraphQLOutputType::class, typeDefinition, permittedTypesForObject, inputObjects) as GraphQLOutputType diff --git a/src/test/kotlin/graphql/kickstart/tools/DirectiveTest.kt b/src/test/kotlin/graphql/kickstart/tools/DirectiveTest.kt index 5f607a5a..a6d5f139 100644 --- a/src/test/kotlin/graphql/kickstart/tools/DirectiveTest.kt +++ b/src/test/kotlin/graphql/kickstart/tools/DirectiveTest.kt @@ -306,6 +306,54 @@ class DirectiveTest { assert((schema.getType("Book") as GraphQLObjectType).getField("name").isDeprecated) } + @Test + fun `should fill in default values of directive arguments that weren't supplied`() { + val emailDirective = EmailDirective() + val schema = SchemaParser.newParser() + .schemaString( + """ + directive @email(message: String = "{path} must be a valid email") on FIELD_DEFINITION | ARGUMENT_DEFINITION | INPUT_FIELD_DEFINITION + directive @owner(team: String = "books-team") on SCHEMA | ENUM_VALUE + + schema @owner { + query: Query + } + + enum AllowedState { + ALLOWED @owner + DISALLOWED + } + + input PersonInput { + email: String @email + } + + type Query { + contactEmail: String @email + updatePersonEmail(email: String @email, backupEmail: String @email(message: "invalid backup email")): String + updatePerson(person: PersonInput, state: AllowedState): String + } + """) + .resolvers(PersonQueryResolver()) + .directive("email", emailDirective) + .build() + .makeExecutableSchema() + + assertEquals( + emailDirective.messages, + mapOf( + "contactEmail" to ("{path} must be a valid email" to "{path} must be a valid email"), + "email" to ("{path} must be a valid email" to "{path} must be a valid email"), + "backupEmail" to ("invalid backup email" to "invalid backup email") + ) + ) + val inputField = (schema.getType("PersonInput") as GraphQLInputObjectType).getField("email") + assertEquals(inputField.getAppliedDirective("email").getArgument("message")?.getValue(), "{path} must be a valid email") + assertEquals(schema.getSchemaAppliedDirective("owner").getArgument("team")?.getValue(), "books-team") + val enumValue = (schema.getType("AllowedState") as GraphQLEnumType).getValue("ALLOWED")!! + assertEquals(enumValue.getAppliedDirective("owner").getArgument("team")?.getValue(), "books-team") + } + @Test fun `should apply directives on the schema and its extensions`() { val schema = SchemaParser.newParser() @@ -458,6 +506,16 @@ class DirectiveTest { val name: String? ) + private class PersonQueryResolver : GraphQLQueryResolver { + fun contactEmail(): String? = null + fun updatePersonEmail(email: String?, backupEmail: String?): String? = email + fun updatePerson(person: PersonInput?, state: AllowedState?): String? = null + } + + private data class PersonInput( + val email: String? + ) + private class QueryResolver : GraphQLQueryResolver { fun books(): List { return listOf(Book(42L, "Test Book")) @@ -489,6 +547,26 @@ class DirectiveTest { } } + private class EmailDirective : SchemaDirectiveWiring { + val messages = mutableMapOf>() + + override fun onField(environment: SchemaDirectiveWiringEnvironment): GraphQLFieldDefinition { + recordMessage(environment) + return environment.element + } + + override fun onArgument(environment: SchemaDirectiveWiringEnvironment): GraphQLArgument { + recordMessage(environment) + return environment.element + } + + private fun recordMessage(environment: SchemaDirectiveWiringEnvironment<*>) { + val appliedMessage = environment.appliedDirective.getArgument("message")?.getValue() + val legacyMessage = environment.directive.getArgument("message")?.let { GraphQLArgument.getArgumentValue(it) } + messages[environment.element.name] = appliedMessage to legacyMessage + } + } + private class UppercaseDirective : SchemaDirectiveWiring { override fun onObject(environment: SchemaDirectiveWiringEnvironment): GraphQLObjectType { val objectType = environment.element