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 @@ -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

/**
*
Expand All @@ -37,4 +41,26 @@ internal class SynchronizedDataLoader<K : Any, V : Any>(
synchronized(this) {
super.load(key, keyContext)
}

override fun loadMany(keys: List<K>): CompletableFuture<List<V>> =
synchronized(this) {
super.loadMany(keys)
}

override fun loadMany(keys: List<K>, keyContexts: List<Any>): CompletableFuture<List<V>> =
synchronized(this) {
super.loadMany(keys, keyContexts)
}

override fun loadMany(keysAndContexts: Map<K, *>): CompletableFuture<Map<K, V>> =
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<DataLoaderFactory.Builder<K, V>>): DataLoader<K, V> =
SynchronizedDataLoader(super.transform(builderConsumer))
}
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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<String, String> {
override val dataLoaderName = "NonBatchingDataLoader"
override fun getDataLoader(graphQLContext: GraphQLContext): DataLoader<String, String> =
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<Any?>? {
instrumentedLoads.incrementAndGet()
return null
}
}
val registry = KotlinDataLoaderRegistryFactory(kotlinDataLoader).generate(mockk(relaxed = true), countingInstrumentation)

val dataLoader = requireNotNull(registry.getDataLoader<String, String>(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<String, AtomicInteger>()
val kotlinDataLoader = object : KotlinDataLoader<String, String> {
override val dataLoaderName = "NonBatchingDataLoader"
override fun getDataLoader(graphQLContext: GraphQLContext): DataLoader<String, String> =
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<String, String>(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() })
}
}
}
Loading