Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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

/**
Expand Down Expand Up @@ -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
}
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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.*
Expand All @@ -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

Expand All @@ -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<*> {
Expand Down Expand Up @@ -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]
Expand Down Expand Up @@ -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?>): Any? {
return suspendCoroutineUninterceptedOrReturn { continuation ->
invoke(resolverMethod, target, args + continuation)
Expand Down
13 changes: 13 additions & 0 deletions src/main/kotlin/graphql/kickstart/tools/util/Utils.kt
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -59,6 +60,18 @@ internal val Class<*>.declaredNonProxyMethods: List<JavaMethod>
}
}

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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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.*
Expand Down Expand Up @@ -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)
Expand All @@ -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()
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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<DataClass> {
suspend fun isActive(dataClass: DataClass, ctx: GraphQLContext): Boolean {
return ctx == context
}
})

@Suppress("UNCHECKED_CAST")
val future = resolver.get(createEnvironment(DataClass(), context = context)) as CompletableFuture<Boolean>
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<DataClass> {
suspend fun isActive(dataClass: DataClass, ctx: ContextClass): Boolean {
return ctx == customContext
}
})

@Suppress("UNCHECKED_CAST")
val future = resolver.get(createEnvironment(DataClass(), context = context)) as CompletableFuture<Boolean>
assert(future.get())
}

@Test
fun `data fetcher uses extension function on the data class`() {
val resolver = createFetcher("name", object : GraphQLResolver<DataClass> {
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<DataClass> {
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"
Expand Down
Loading