diff --git a/src/main/kotlin/graphql/kickstart/tools/resolver/FieldResolverScanner.kt b/src/main/kotlin/graphql/kickstart/tools/resolver/FieldResolverScanner.kt index 589ebc48..173cc4a3 100644 --- a/src/main/kotlin/graphql/kickstart/tools/resolver/FieldResolverScanner.kt +++ b/src/main/kotlin/graphql/kickstart/tools/resolver/FieldResolverScanner.kt @@ -16,8 +16,7 @@ import org.apache.commons.lang3.reflect.TypeUtils import org.reactivestreams.Publisher import org.slf4j.LoggerFactory import java.lang.reflect.* -import kotlin.reflect.full.valueParameters -import kotlin.reflect.jvm.javaType +import kotlin.reflect.full.extensionReceiverParameter import kotlin.reflect.jvm.kotlinFunction /** @@ -167,31 +166,23 @@ internal class FieldResolverScanner(val options: SchemaParserOptions) { it == search.requiredFirstParameterType || method.declaringClass.typeParameters.contains(it) } ?: false } else { - true + // an extension receiver can only take the source object + !isExtensionFunction(method) } - val methodParameterCount = getMethodParameterCount(method) - val methodLastParameter = getMethodLastParameter(method) + val methodParameterCount = method.parameterCountWithoutContinuation() + val methodLastParameter = method.parameterTypes.getOrNull(methodParameterCount - 1) val correctParameterCount = methodParameterCount == requiredCount || (methodParameterCount == (requiredCount + 1) && allowedLastArgumentTypes.contains(methodLastParameter)) return correctParameterCount && appropriateFirstParameter } - private fun getMethodParameterCount(method: Method): Int { + private fun isExtensionFunction(method: Method): Boolean { return try { - method.kotlinFunction?.valueParameters?.size ?: method.parameterCount + method.kotlinFunction?.extensionReceiverParameter != null } catch (e: InternalError) { - method.parameterCount - } - } - - private fun getMethodLastParameter(method: Method): Type? { - return try { - method.kotlinFunction?.valueParameters?.lastOrNull()?.type?.javaType - ?: method.parameterTypes.lastOrNull() - } catch (e: InternalError) { - method.parameterTypes.lastOrNull() + false } } diff --git a/src/main/kotlin/graphql/kickstart/tools/resolver/MethodFieldResolver.kt b/src/main/kotlin/graphql/kickstart/tools/resolver/MethodFieldResolver.kt index 59ff78a2..b280c463 100644 --- a/src/main/kotlin/graphql/kickstart/tools/resolver/MethodFieldResolver.kt +++ b/src/main/kotlin/graphql/kickstart/tools/resolver/MethodFieldResolver.kt @@ -7,6 +7,8 @@ import graphql.kickstart.tools.SchemaParserOptions.GenericWrapper import graphql.kickstart.tools.util.JavaType import graphql.kickstart.tools.util.coroutineScope import graphql.kickstart.tools.util.futureValueType +import graphql.kickstart.tools.util.isSuspendFunction +import graphql.kickstart.tools.util.parameterCountWithoutContinuation import graphql.kickstart.tools.util.typeArgument import graphql.kickstart.tools.util.unwrap import graphql.language.* @@ -26,7 +28,6 @@ import java.lang.reflect.Method import java.util.* import java.util.function.Supplier import kotlin.coroutines.intrinsics.suspendCoroutineUninterceptedOrReturn -import kotlin.reflect.full.valueParameters import kotlin.reflect.jvm.javaType import kotlin.reflect.jvm.kotlinFunction @@ -43,7 +44,7 @@ internal class MethodFieldResolver( private val log = LoggerFactory.getLogger(javaClass) private val isSuspendFunction = method.isSuspendFunction() - private val numberOfParameters = method.kotlinFunction?.valueParameters?.size ?: method.parameterCount + private val numberOfParameters = method.parameterCountWithoutContinuation() private val hasAdditionalParameter = numberOfParameters == (field.inputValueDefinitions.size + getIndexOffset() + 1) override fun createDataFetcher(): DataFetcher<*> { @@ -102,7 +103,8 @@ internal class MethodFieldResolver( // Add DataFetchingEnvironment/Context argument if (this.hasAdditionalParameter) { - when (this.method.parameterTypes.last()) { + // suspend functions have a trailing Continuation parameter + when (this.method.parameterTypes[numberOfParameters - 1]) { null -> throw ResolverError("Expected at least one argument but got none, this is most likely a bug with graphql-java-tools") options.contextClass -> args.add { environment -> val context: Any? = environment.graphQlContext[options.contextClass] @@ -282,14 +284,6 @@ private class CompareGenericWrappers { } } -private fun Method.isSuspendFunction(): Boolean { - return try { - this.kotlinFunction?.isSuspend == true - } catch (e: InternalError) { - false - } -} - private suspend inline fun invokeSuspend(target: Any, resolverMethod: Method, args: Array): Any? { return suspendCoroutineUninterceptedOrReturn { continuation -> invoke(resolverMethod, target, args + continuation) diff --git a/src/main/kotlin/graphql/kickstart/tools/util/Utils.kt b/src/main/kotlin/graphql/kickstart/tools/util/Utils.kt index 58c6ef40..ae8dd5c4 100644 --- a/src/main/kotlin/graphql/kickstart/tools/util/Utils.kt +++ b/src/main/kotlin/graphql/kickstart/tools/util/Utils.kt @@ -12,6 +12,7 @@ import java.lang.reflect.ParameterizedType import java.lang.reflect.Proxy import java.util.concurrent.CompletionStage import java.util.concurrent.Future +import kotlin.reflect.jvm.kotlinFunction /** * @author Andrew Potter @@ -59,6 +60,18 @@ internal val Class<*>.declaredNonProxyMethods: List } } +internal fun JavaMethod.isSuspendFunction(): Boolean { + return try { + this.kotlinFunction?.isSuspend == true + } catch (e: InternalError) { + false + } +} + +// the trailing Continuation of a suspend function isn't a resolver argument +internal fun JavaMethod.parameterCountWithoutContinuation(): Int = + if (isSuspendFunction()) parameterCount - 1 else parameterCount + internal fun getDocumentation(node: AbstractNode<*>, options: SchemaParserOptions): String? = when { node is AbstractDescribedNode<*> && node.description != null -> node.description?.content diff --git a/src/test/kotlin/graphql/kickstart/tools/FieldResolverScannerTest.kt b/src/test/kotlin/graphql/kickstart/tools/FieldResolverScannerTest.kt index 4ca5fa54..db917658 100644 --- a/src/test/kotlin/graphql/kickstart/tools/FieldResolverScannerTest.kt +++ b/src/test/kotlin/graphql/kickstart/tools/FieldResolverScannerTest.kt @@ -6,10 +6,12 @@ import graphql.kickstart.tools.resolver.FieldResolverScanner import graphql.kickstart.tools.resolver.MethodFieldResolver import graphql.kickstart.tools.resolver.PropertyFieldResolver import graphql.language.FieldDefinition +import graphql.language.InputValueDefinition import graphql.language.TypeName import graphql.relay.Connection import graphql.relay.DefaultConnection import graphql.relay.DefaultPageInfo +import graphql.schema.DataFetchingEnvironment import kotlinx.coroutines.ExperimentalCoroutinesApi import org.junit.Test import java.util.* @@ -98,6 +100,21 @@ class FieldResolverScannerTest { Locale.setDefault(default) } + @Test + fun `scanner ignores extension functions on root resolvers`() { + val resolver = RootResolverInfo(listOf(UserQuery(), ExtensionQuery()), options) + val field = FieldDefinition.newFieldDefinition() + .name("user") + .type(TypeName("String")) + .inputValueDefinition(InputValueDefinition("id", TypeName("ID"))) + .build() + + val user = scanner.findFieldResolver(field, resolver) + + assert(user is MethodFieldResolver) + assertEquals((user as MethodFieldResolver).method.declaringClass, UserQuery::class.java) + } + @Test fun `scanner finds field resolver methods in priority order`() { val resolverInfo = RootResolverInfo(listOf(PriorityQuery()), options) @@ -123,6 +140,14 @@ class FieldResolverScannerTest { fun field1() {} } + class UserQuery : GraphQLQueryResolver { + fun user(id: String): String = id + } + + class ExtensionQuery : GraphQLQueryResolver { + fun DataFetchingEnvironment.user(): String = field.name + } + class CamelCaseQuery1 : GraphQLQueryResolver { fun getHullType(): HullType = HullType() } diff --git a/src/test/kotlin/graphql/kickstart/tools/MethodFieldResolverDataFetcherTest.kt b/src/test/kotlin/graphql/kickstart/tools/MethodFieldResolverDataFetcherTest.kt index 8247fac6..4f6c490d 100644 --- a/src/test/kotlin/graphql/kickstart/tools/MethodFieldResolverDataFetcherTest.kt +++ b/src/test/kotlin/graphql/kickstart/tools/MethodFieldResolverDataFetcherTest.kt @@ -240,6 +240,54 @@ class MethodFieldResolverDataFetcherTest { assertEquals(resolver.get(createEnvironment(DataClass(), context = context)), true) } + @Test + fun `data fetcher passes context to suspend function if method has extra argument and context is specified`() { + val context = GraphQLContext.newContext().build() + val resolver = createFetcher("active", resolver = object : GraphQLResolver { + suspend fun isActive(dataClass: DataClass, ctx: GraphQLContext): Boolean { + return ctx == context + } + }) + + @Suppress("UNCHECKED_CAST") + val future = resolver.get(createEnvironment(DataClass(), context = context)) as CompletableFuture + assert(future.get()) + } + + @Test + fun `data fetcher passes custom context to suspend function if method has extra argument and custom context is specified`() { + val customContext = ContextClass() + val context = GraphQLContext.of(mapOf(ContextClass::class.java to customContext)) + val options = SchemaParserOptions.newOptions().contextClass(ContextClass::class).build() + val resolver = createFetcher("active", options = options, resolver = object : GraphQLResolver { + suspend fun isActive(dataClass: DataClass, ctx: ContextClass): Boolean { + return ctx == customContext + } + }) + + @Suppress("UNCHECKED_CAST") + val future = resolver.get(createEnvironment(DataClass(), context = context)) as CompletableFuture + assert(future.get()) + } + + @Test + fun `data fetcher uses extension function on the data class`() { + val resolver = createFetcher("name", object : GraphQLResolver { + fun DataClass.name(): String = "extension $name" + }) + + assertEquals(resolver.get(createEnvironment(DataClass())), "extension TestName") + } + + @Test + fun `data fetcher passes environment to extension function if method has extra argument`() { + val resolver = createFetcher("active", object : GraphQLResolver { + fun DataClass.isActive(env: DataFetchingEnvironment): Boolean = env is DataFetchingEnvironment + }) + + assertEquals(resolver.get(createEnvironment(DataClass())), true) + } + @Test fun `data fetcher marshalls input object if required`() { val name = "correct name"