diff --git a/src/main/kotlin/graphql/kickstart/tools/resolver/MethodFieldResolver.kt b/src/main/kotlin/graphql/kickstart/tools/resolver/MethodFieldResolver.kt index 926ebee4..59ff78a2 100644 --- a/src/main/kotlin/graphql/kickstart/tools/resolver/MethodFieldResolver.kt +++ b/src/main/kotlin/graphql/kickstart/tools/resolver/MethodFieldResolver.kt @@ -15,6 +15,8 @@ import graphql.schema.DataFetchingEnvironment import graphql.schema.GraphQLFieldDefinition import graphql.schema.GraphQLTypeUtil.isScalar import graphql.schema.LightDataFetcher +import kotlinx.coroutines.CoroutineStart +import kotlinx.coroutines.ensureActive import kotlinx.coroutines.future.future import org.apache.commons.lang3.reflect.TypeUtils import org.reactivestreams.Publisher @@ -211,7 +213,10 @@ internal class MethodFieldResolverDataFetcher( val args = this.args.map { it(environment) }.toTypedArray() return if (isSuspendFunction) { - environment.coroutineScope().future(options.coroutineContextProvider.provide()) { + // start undispatched so DataLoader loads are queued before graphql-java dispatches them, + // which runs the block even if the context is already cancelled, hence ensureActive + environment.coroutineScope().future(options.coroutineContextProvider.provide(), CoroutineStart.UNDISPATCHED) { + ensureActive() invokeSuspend(source, method, args)?.transformWithGenericWrapper(options.genericWrappers) { environment } } } else { diff --git a/src/test/kotlin/graphql/kickstart/tools/MethodFieldResolverDataFetcherTest.kt b/src/test/kotlin/graphql/kickstart/tools/MethodFieldResolverDataFetcherTest.kt index 91464566..8247fac6 100644 --- a/src/test/kotlin/graphql/kickstart/tools/MethodFieldResolverDataFetcherTest.kt +++ b/src/test/kotlin/graphql/kickstart/tools/MethodFieldResolverDataFetcherTest.kt @@ -21,6 +21,7 @@ import kotlinx.coroutines.channels.ReceiveChannel import org.junit.Test import org.reactivestreams.Publisher import org.reactivestreams.tck.TestEnvironment +import java.util.concurrent.CancellationException import java.util.concurrent.CompletableFuture import kotlin.coroutines.coroutineContext @@ -57,6 +58,33 @@ class MethodFieldResolverDataFetcherTest { } } + @Test + fun `data fetcher does not invoke suspend function if coroutineContext defined by options is cancelled`() { + // setup + val cancelledClass = CancelledClass() + + val resolver = createFetcher("active", cancelledClass, options = cancelledClass.options) + + // expect + val future = resolver.get(createEnvironment(DataClass())) as CompletableFuture<*> + assert(runCatching { future.get() }.exceptionOrNull() is CancellationException) + assert(!cancelledClass.invoked) + } + + class CancelledClass : GraphQLResolver { + var invoked = false + + val options = SchemaParserOptions.Builder() + .coroutineContext(Dispatchers.Default + Job().apply { cancel() }) + .build() + + @Suppress("UNUSED_PARAMETER") + suspend fun isActive(data: DataClass): Boolean { + invoked = true + return true + } + } + @ExperimentalCoroutinesApi @Test fun `canceling subscription Publisher also cancels underlying Kotlin coroutine channel`() { diff --git a/src/test/kotlin/graphql/kickstart/tools/SuspendFunctionDataLoaderTest.kt b/src/test/kotlin/graphql/kickstart/tools/SuspendFunctionDataLoaderTest.kt new file mode 100644 index 00000000..20b8a4e7 --- /dev/null +++ b/src/test/kotlin/graphql/kickstart/tools/SuspendFunctionDataLoaderTest.kt @@ -0,0 +1,81 @@ +package graphql.kickstart.tools + +import graphql.ExecutionInput +import graphql.GraphQL +import graphql.schema.DataFetchingEnvironment +import kotlinx.coroutines.future.await +import org.dataloader.DataLoaderFactory +import org.dataloader.DataLoaderRegistry +import org.junit.Test +import java.util.concurrent.CompletableFuture +import java.util.concurrent.TimeUnit + +class SuspendFunctionDataLoaderTest { + + private val schema = SchemaParser.newParser() + .schemaString( + """ + type Query { + user(id: Int!): User! + users(ids: [Int!]!): [User!]! + } + + type User { + id: Int! + friend: User! + } + """) + .resolvers(Query(), UserResolver()) + .build() + .makeExecutableSchema() + private val gql = GraphQL.newGraphQL(schema).build() + + @Test + fun `root suspend function can await a data loader`() { + repeat(20) { + val data = execute("{ users(ids: [1, 2, 3]) { id } }") + + assertEquals(data, mapOf("users" to listOf(mapOf("id" to 1), mapOf("id" to 2), mapOf("id" to 3)))) + } + } + + @Test + fun `nested suspend function can await a data loader`() { + repeat(20) { + val data = execute("{ a: user(id: 1) { friend { id } } b: user(id: 2) { friend { id } } }") + + assertEquals(data, mapOf( + "a" to mapOf("friend" to mapOf("id" to 2)), + "b" to mapOf("friend" to mapOf("id" to 3)) + )) + } + } + + private fun execute(query: String): Any? { + val userLoader = DataLoaderFactory.newDataLoader { ids -> + CompletableFuture.supplyAsync { ids.map { User(it) } } + } + val registry = DataLoaderRegistry.newRegistry().register("user", userLoader).build() + + // a resolver that awaits a load before it is dispatched hangs, so don't wait forever + val result = gql.executeAsync(ExecutionInput.newExecutionInput(query).dataLoaderRegistry(registry)) + .get(5, TimeUnit.SECONDS) + + assert(result.errors.isEmpty()) { result.errors.joinToString { it.message } } + return result.getData() + } + + class Query : GraphQLQueryResolver { + fun user(id: Int): User = User(id) + + suspend fun users(ids: List, env: DataFetchingEnvironment): List = + env.getDataLoader("user")!!.loadMany(ids).await() + } + + class UserResolver : GraphQLResolver { + suspend fun friend(user: User, env: DataFetchingEnvironment): User = + env.getDataLoader("user")!!.load(user.id + 1).await() + } + + data class User(val id: Int) +}