diff --git a/executions/graphql-kotlin-dataloader/src/main/kotlin/com/expediagroup/graphql/dataloader/SynchronizedDataLoader.kt b/executions/graphql-kotlin-dataloader/src/main/kotlin/com/expediagroup/graphql/dataloader/SynchronizedDataLoader.kt index 96562fa2f9..173c88971f 100644 --- a/executions/graphql-kotlin-dataloader/src/main/kotlin/com/expediagroup/graphql/dataloader/SynchronizedDataLoader.kt +++ b/executions/graphql-kotlin-dataloader/src/main/kotlin/com/expediagroup/graphql/dataloader/SynchronizedDataLoader.kt @@ -17,8 +17,12 @@ package com.expediagroup.graphql.dataloader import org.dataloader.DataLoader +import org.dataloader.DataLoaderFactory +import org.dataloader.DataLoaderRegistry import org.dataloader.DelegatingDataLoader +import org.dataloader.instrumentation.DataLoaderInstrumentation import java.util.concurrent.CompletableFuture +import java.util.function.Consumer /** * @@ -37,4 +41,26 @@ internal class SynchronizedDataLoader( synchronized(this) { super.load(key, keyContext) } + + override fun loadMany(keys: List): CompletableFuture> = + synchronized(this) { + super.loadMany(keys) + } + + override fun loadMany(keys: List, keyContexts: List): CompletableFuture> = + synchronized(this) { + super.loadMany(keys, keyContexts) + } + + override fun loadMany(keysAndContexts: Map): CompletableFuture> = + synchronized(this) { + super.loadMany(keysAndContexts) + } + + /** + * [DataLoaderRegistry] transforms every registered [DataLoader] to add its name and [DataLoaderInstrumentation], + * the transformed [DataLoader] needs to be synchronized as well + */ + override fun transform(builderConsumer: Consumer>): DataLoader = + SynchronizedDataLoader(super.transform(builderConsumer)) } diff --git a/executions/graphql-kotlin-dataloader/src/test/kotlin/com/expediagroup/graphql/dataloader/KotlinDataLoaderRegistryFactoryTest.kt b/executions/graphql-kotlin-dataloader/src/test/kotlin/com/expediagroup/graphql/dataloader/KotlinDataLoaderRegistryFactoryTest.kt index edb347158f..ddd41d9228 100644 --- a/executions/graphql-kotlin-dataloader/src/test/kotlin/com/expediagroup/graphql/dataloader/KotlinDataLoaderRegistryFactoryTest.kt +++ b/executions/graphql-kotlin-dataloader/src/test/kotlin/com/expediagroup/graphql/dataloader/KotlinDataLoaderRegistryFactoryTest.kt @@ -27,9 +27,11 @@ import org.dataloader.DataLoader import org.dataloader.DataLoaderFactory import org.dataloader.DataLoaderOptions import org.dataloader.instrumentation.DataLoaderInstrumentation +import org.dataloader.instrumentation.DataLoaderInstrumentationContext import org.junit.jupiter.api.Test import reactor.kotlin.core.publisher.toFlux import java.util.concurrent.CompletableFuture +import java.util.concurrent.ConcurrentHashMap import java.util.concurrent.atomic.AtomicInteger import kotlin.test.assertEquals import kotlin.test.assertTrue @@ -94,4 +96,81 @@ class KotlinDataLoaderRegistryFactoryTest { assertEquals(1, invocations.get()) } } + + @Test + fun `cached and non-batched data loader is invoked once for concurrent loads of the same key when registry has instrumentation`() { + runBlocking { + val invocations = AtomicInteger() + val instrumentedLoads = AtomicInteger() + val kotlinDataLoader = object : KotlinDataLoader { + override val dataLoaderName = "NonBatchingDataLoader" + override fun getDataLoader(graphQLContext: GraphQLContext): DataLoader = + DataLoaderFactory.newMappedDataLoader( + { keys -> + invocations.incrementAndGet() + Thread.sleep(100) + CompletableFuture.completedFuture(keys.associateWith { it }) + }, + DataLoaderOptions.newOptions() + .setBatchingEnabled(false) + .build() + ) + } + val countingInstrumentation = object : DataLoaderInstrumentation { + override fun beginLoad(dataLoader: DataLoader<*, *>, key: Any, loadContext: Any?): DataLoaderInstrumentationContext? { + instrumentedLoads.incrementAndGet() + return null + } + } + val registry = KotlinDataLoaderRegistryFactory(kotlinDataLoader).generate(mockk(relaxed = true), countingInstrumentation) + + val dataLoader = requireNotNull(registry.getDataLoader(kotlinDataLoader.dataLoaderName)) + + List(8) { + async(Dispatchers.Default) { + dataLoader.load("same-key").await() + } + }.awaitAll() + + assertEquals(1, invocations.get()) + assertEquals(8, instrumentedLoads.get()) + } + } + + @Test + fun `cached and non-batched data loader is invoked once per key for concurrent loadMany of the same keys`() { + runBlocking { + val invocations = ConcurrentHashMap() + val kotlinDataLoader = object : KotlinDataLoader { + override val dataLoaderName = "NonBatchingDataLoader" + override fun getDataLoader(graphQLContext: GraphQLContext): DataLoader = + DataLoaderFactory.newMappedDataLoader( + { keys -> + keys.forEach { key -> invocations.computeIfAbsent(key) { AtomicInteger() }.incrementAndGet() } + Thread.sleep(100) + CompletableFuture.completedFuture(keys.associateWith { it }) + }, + DataLoaderOptions.newOptions() + .setBatchingEnabled(false) + .build() + ) + } + val registry = KotlinDataLoaderRegistryFactory(kotlinDataLoader).generate(mockk(relaxed = true), object : DataLoaderInstrumentation {}) + + val dataLoader = requireNotNull(registry.getDataLoader(kotlinDataLoader.dataLoaderName)) + val keys = listOf("a", "b") + + List(9) { index -> + async(Dispatchers.Default) { + when (index % 3) { + 0 -> dataLoader.loadMany(keys).await() + 1 -> dataLoader.loadMany(keys, keys.map { key -> "context-$key" }).await() + else -> dataLoader.loadMany(keys.associateWith { key -> "context-$key" }).await().values.toList() + } + } + }.awaitAll() + + assertEquals(mapOf("a" to 1, "b" to 1), invocations.mapValues { (_, count) -> count.get() }) + } + } }