diff --git a/src/main/kotlin/graphql/kickstart/tools/SchemaParser.kt b/src/main/kotlin/graphql/kickstart/tools/SchemaParser.kt index 57d872ff..7417e54c 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 6991bfed..bb434364 100644 --- a/src/test/kotlin/graphql/kickstart/tools/DirectiveTest.kt +++ b/src/test/kotlin/graphql/kickstart/tools/DirectiveTest.kt @@ -372,6 +372,53 @@ 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(primaryEmail: String @email, backupEmail: String @email(message: "invalid backup email")): String + updatePerson(person: PersonInput, state: AllowedState): String + } + """) + .resolvers(PersonQueryResolver()) + .directive("email", emailDirective) + .build() + .makeExecutableSchema() + + val expectedMessages = mapOf( + "contactEmail" to "{path} must be a valid email", + "primaryEmail" to "{path} must be a valid email", + "backupEmail" to "invalid backup email" + ) + assertEquals(emailDirective.appliedMessages, expectedMessages) + assertEquals(emailDirective.legacyMessages, expectedMessages) + 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() @@ -524,6 +571,16 @@ class DirectiveTest { val name: String? ) + private class PersonQueryResolver : GraphQLQueryResolver { + fun contactEmail(): String? = null + fun updatePersonEmail(primaryEmail: String?, backupEmail: String?): String? = primaryEmail + 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")) @@ -555,6 +612,27 @@ class DirectiveTest { } } + private class EmailDirective : SchemaDirectiveWiring { + val appliedMessages = mutableMapOf() + val legacyMessages = 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 name = environment.element.name + appliedMessages[name] = environment.appliedDirective.getArgument("message")?.getValue() + legacyMessages[name] = environment.directive.getArgument("message")?.let { GraphQLArgument.getArgumentValue(it) } + } + } + private class UppercaseDirective : SchemaDirectiveWiring { override fun onObject(environment: SchemaDirectiveWiringEnvironment): GraphQLObjectType { val objectType = environment.element