diff --git a/src/main/kotlin/graphql/kickstart/tools/directive/DirectiveWiringHelper.kt b/src/main/kotlin/graphql/kickstart/tools/directive/DirectiveWiringHelper.kt index b2c69c7c..0631a0d0 100644 --- a/src/main/kotlin/graphql/kickstart/tools/directive/DirectiveWiringHelper.kt +++ b/src/main/kotlin/graphql/kickstart/tools/directive/DirectiveWiringHelper.kt @@ -77,26 +77,30 @@ class DirectiveWiringHelper( var output = wrapper.graphQlType // first the specific named directives wrapper.graphQlType.appliedDirectives.forEach { appliedDirective -> - val env = buildEnvironment(wrapper, appliedDirective) + val env = buildEnvironment(wrapper, output, appliedDirective) val wiring = runtimeWiring.registeredDirectiveWiring[appliedDirective.name] - wiring?.let { output = wrapper.invoker(it, env) } + wiring?.let { output = invokeWiring(wrapper, it, env) } } // now call any statically added to the runtime runtimeWiring.directiveWiring.forEach { staticWiring -> - val env = buildEnvironment(wrapper) - output = wrapper.invoker(staticWiring, env) + val env = buildEnvironment(wrapper, output) + output = invokeWiring(wrapper, staticWiring, env) } // wiring factory is last (if present) - val env = buildEnvironment(wrapper) + val env = buildEnvironment(wrapper, output) if (runtimeWiring.wiringFactory.providesSchemaDirectiveWiring(env)) { val factoryWiring = runtimeWiring.wiringFactory.getSchemaDirectiveWiring(env) - output = wrapper.invoker(factoryWiring, env) + output = invokeWiring(wrapper, factoryWiring, env) } return output } - private fun buildEnvironment(wrapper: WiringWrapper, appliedDirective: GraphQLAppliedDirective? = null): SchemaDirectiveWiringEnvironmentImpl { + private fun invokeWiring(wrapper: WiringWrapper, wiring: SchemaDirectiveWiring, env: SchemaDirectiveWiringEnvironmentImpl): T { + return checkNotNull(wrapper.invoker(wiring, env)) { "The SchemaDirectiveWiring MUST return a non null return value for element '${wrapper.graphQlType.name}'" } + } + + private fun buildEnvironment(wrapper: WiringWrapper, element: T, appliedDirective: GraphQLAppliedDirective? = null): SchemaDirectiveWiringEnvironmentImpl { val type = wrapper.graphQlType val directive = appliedDirective?.let { d -> type.directives.find { it.name == d.name } } val nodeParentTree = buildAstTree(*listOfNotNull( @@ -121,7 +125,7 @@ class DirectiveWiringHelper( is GraphQLFieldsContainer -> schemaDirectiveParameters.newParams(type, nodeParentTree, elementParentTree) else -> schemaDirectiveParameters.newParams(nodeParentTree, elementParentTree) } - return SchemaDirectiveWiringEnvironmentImpl(type, type.directives, type.appliedDirectives, directive, appliedDirective, params) + return SchemaDirectiveWiringEnvironmentImpl(element, type.directives, type.appliedDirectives, directive, appliedDirective, params) } private fun buildAstTree(vararg nodes: NamedNode<*>): NodeParentTree> { @@ -139,7 +143,7 @@ class DirectiveWiringHelper( private data class WiringWrapper( val graphQlType: T, val directiveLocation: Introspection.DirectiveLocation, - val invoker: (SchemaDirectiveWiring, SchemaDirectiveWiringEnvironmentImpl) -> T, + val invoker: (SchemaDirectiveWiring, SchemaDirectiveWiringEnvironmentImpl) -> T?, val fieldsContainer: GraphQLFieldsContainer? = null, val fieldDefinition: GraphQLFieldDefinition? = null, val inputFieldsContainer: GraphQLInputFieldsContainer? = null, diff --git a/src/test/kotlin/graphql/kickstart/tools/DirectiveTest.kt b/src/test/kotlin/graphql/kickstart/tools/DirectiveTest.kt index 5f607a5a..6991bfed 100644 --- a/src/test/kotlin/graphql/kickstart/tools/DirectiveTest.kt +++ b/src/test/kotlin/graphql/kickstart/tools/DirectiveTest.kt @@ -202,6 +202,72 @@ class DirectiveTest { assertEquals(result.getData(), expected) } + @Test + fun `should chain element changes of named directive wirings`() { + val schema = SchemaParser.newParser() + .schemaString( + """ + directive @auth on FIELD_DEFINITION + directive @log on FIELD_DEFINITION + + type Query { + "Name" + name: String @auth @log + } + """) + .resolvers(NameResolver()) + .directive("auth", DescriptionDirective("auth")) + .directive("log", DescriptionDirective("log")) + .build() + .makeExecutableSchema() + + assertEquals(schema.queryType.getField("name").description, "Name +auth +log") + } + + @Test + fun `should chain element changes of named and static directive wirings`() { + val schema = SchemaParser.newParser() + .schemaString( + """ + directive @auth on FIELD_DEFINITION + + type Query { + "Name" + name: String @auth + } + """) + .resolvers(NameResolver()) + .directive("auth", DescriptionDirective("auth")) + .directiveWiring(DescriptionDirective("static")) + .build() + .makeExecutableSchema() + + assertEquals(schema.queryType.getField("name").description, "Name +auth +static") + } + + @Test + fun `should fail when a directive wiring returns null`() { + val error = assertThrows(IllegalStateException::class.java) { + SchemaParser.newParser() + .schemaString( + """ + directive @auth on FIELD_DEFINITION + + type Query { + name: String @auth + } + """) + .resolvers(NameResolver()) + .directive("auth", object : SchemaDirectiveWiring { + override fun onField(environment: SchemaDirectiveWiringEnvironment): GraphQLFieldDefinition? = null + }) + .build() + .makeExecutableSchema() + } + + assertEquals(error.message, "The SchemaDirectiveWiring MUST return a non null return value for element 'name'") + } + @Test fun `should have access to applied directives through the data fetching environment`() { val schema = SchemaParser.newParser() @@ -523,6 +589,13 @@ class DirectiveTest { } } + private class DescriptionDirective(private val name: String) : SchemaDirectiveWiring { + override fun onField(environment: SchemaDirectiveWiringEnvironment): GraphQLFieldDefinition { + val field = environment.element + return field.transform { it.description("${field.description} +$name") } + } + } + private class DoubleDirective : SchemaDirectiveWiring { override fun onField(environment: SchemaDirectiveWiringEnvironment): GraphQLFieldDefinition {