diff --git a/.github/workflows/pull_request_event.yml b/.github/workflows/pull_request_event.yml
index 6b2db40..d2d3cef 100644
--- a/.github/workflows/pull_request_event.yml
+++ b/.github/workflows/pull_request_event.yml
@@ -18,6 +18,16 @@ jobs:
if: ${{ github.event_name == 'pull_request'
&& (github.event.action == 'opened' || github.event.action == 'synchronize' ||
github.event.action == 'reopened' || github.event.action == 'ready_for_review') }}
+ services:
+ redis:
+ image: redis:6.2.7-alpine
+ ports:
+ - 6379:6379
+ options: >-
+ --health-cmd "redis-cli ping"
+ --health-interval 10s
+ --health-timeout 5s
+ --health-retries 5
steps:
- uses: actions/checkout@v4
with:
@@ -28,4 +38,7 @@ jobs:
distribution: 'temurin'
java-version: 17
- name: Test
+ env:
+ TEST_REDIS_HOST: localhost
+ TEST_REDIS_PORT: 6379
run: ./gradlew clean test --info
diff --git a/AGENTS.md b/AGENTS.md
new file mode 100644
index 0000000..8696336
--- /dev/null
+++ b/AGENTS.md
@@ -0,0 +1,28 @@
+# Repository Guidelines
+
+## Project Structure & Module Organization
+- Core libraries live in module directories such as `core`, `core-reactor`, and `core-kotlin-coroutine`; each follows the Gradle layout `src/main` and `src/test`.
+- Spring adapters sit under `core-spring*` modules, while runnable samples are in `req-shield-*example` projects.
+- Shared utilities and constants are centralized in `support`. Generated build outputs stay under each module's `build/` folder.
+
+## Build, Test, and Development Commands
+- `./gradlew build` compiles all modules, runs unit tests, and assembles artifacts; pass `--parallel` for faster local feedback.
+- `./gradlew test` executes Kotlin/JVM unit tests across every enabled module.
+- `./gradlew ktlintCheck` enforces the project's formatting contract before you open a PR.
+- Use `./gradlew :req-shield-spring-boot3-example:bootRun` (or another sample module) to manually exercise integration paths.
+
+## Coding Style & Naming Conventions
+- Kotlin sources use 4-space indentation, `UpperCamelCase` for types, and `lowerCamelCase` for functions and properties.
+- Keep package names lowercase and aligned with module boundaries (e.g., `com.linecorp.reqshield.core`).
+- Always add the Apache 2.0 copyright header shown in `CONTRIBUTING.md` to new files.
+- Prefer early-return patterns and meaningful exception messages; align with the `ErrorCode` enums already defined.
+
+## Testing Guidelines
+- Write tests with JUnit 5 (`org.junit.jupiter`) and place them under `src/test/kotlin` mirroring the `src/main` package.
+- Use descriptive method names such as `shouldCollapseConcurrentRequests()` and cover both success and failure paths.
+- When adding integration behaviour, extend the corresponding example module and run `./gradlew test` before submission.
+
+## Commit & Pull Request Guidelines
+- Follow the repository's history of concise, imperative commits (e.g., `Add cache invalidation helper`).
+- Reference related issues in the body, summarise motivation, modifications, and results, and include screenshots/logs for behaviour changes.
+- Verify CLS (Contributor License Agreement) status, ensure CI passes locally, and request review from a maintainer familiar with your module.
diff --git a/CLAUDE.md b/CLAUDE.md
index e31092d..abe76ee 100644
--- a/CLAUDE.md
+++ b/CLAUDE.md
@@ -90,7 +90,7 @@ Contains shared:
- `isLocalLock`: Use local vs distributed locking (default: true)
- `lockTimeoutMillis`: Lock acquisition timeout (default: 3000ms)
- `decisionForUpdate`: Percentage of TTL after which to trigger async cache refresh (default: 80)
-- `maxAttemptGetCache`: Max retry attempts when waiting for cache (default: 10)
+- `maxAttemptGetCache`: Max retry attempts when waiting for cache (default: 60)
- `reqShieldWorkMode`: CREATE_AND_UPDATE_CACHE | ONLY_CREATE_CACHE | ONLY_UPDATE_CACHE
### Work Modes
diff --git a/README.md b/README.md
index 05506bb..aaa56cd 100644
--- a/README.md
+++ b/README.md
@@ -23,6 +23,44 @@ A lib that regulates the cache-based requests an application receives in terms o
`implementation("com.linecorp.cse.reqshield:core-spring-webflux:{version}")`
`implementation("com.linecorp.cse.reqshield:core-spring-webflux-kotlin-coroutine:{version}")`
+## Testing & Integration Tips
+
+### Integration tests with Redis (Testcontainers)
+
+- Redis-backed integration tests using Testcontainers always run as part of module test tasks.
+- Requirements:
+ - A working local Docker daemon with network access to pull `redis:6.2.7-alpine` on first run.
+ - Sufficient permissions to start containers from tests.
+- If you need to temporarily bypass Redis ITs locally (e.g., no Docker), run specific unit-test-only tasks or exclude the example modules when invoking Gradle.
+
+### WebFlux null handling
+
+- `@ReqShieldCacheable(nullHandling = ...)` controls how `null` values are emitted in WebFlux:
+ - `EMIT_EMPTY` (default): map `null` to `Mono.empty()`.
+ - `ERROR`: throw an `IllegalStateException` if a `null` value is produced.
+
+### Global lock guidance
+
+- When `isLocalLock = false`, you must provide real global lock/unlock implementations.
+- Recommended approach with Redis:
+ - Lock: `SETNX lock:{key} 1` + `PEXPIRE lock:{key} {ttlMillis}`
+ - Unlock: `DEL lock:{key}`
+- The provided defaults return `true` and are only suitable for local/dev usage.
+
+### Reactor Scheduler tuning
+
+- Reactor-based modules accept a `Scheduler` (e.g., `boundedElastic`) through configuration.
+- Spring WebFlux adapter exposes a `reqShieldScheduler` bean you can override for tuning thread usage.
+
+### Kotlin Coroutine Parallelism Configuration
+
+| Property | Default | Description |
+|----------|---------|-------------|
+| `reqshield.blocking.parallelism` | `availableProcessors * 2` (clamped 4-256) | Controls parallelism for blocking calls in the coroutine aspect |
+
+**Note**: This feature uses `Dispatchers.IO.limitedParallelism()` which is marked as `@ExperimentalCoroutinesApi`.
+The API may change in future Kotlin Coroutines versions.
+
## Contributing
Pull requests are welcome. For major changes, please open an issue first to discuss what you would like to
diff --git a/core-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/KeyLocalLock.kt b/core-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/KeyLocalLock.kt
index cfc8472..370befd 100644
--- a/core-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/KeyLocalLock.kt
+++ b/core-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/KeyLocalLock.kt
@@ -27,31 +27,93 @@ import kotlinx.coroutines.launch
import org.slf4j.LoggerFactory
import java.util.concurrent.ConcurrentHashMap
import java.util.concurrent.Semaphore
+import java.util.concurrent.atomic.AtomicBoolean
import kotlin.coroutines.CoroutineContext
private val log = LoggerFactory.getLogger(KeyLocalLock::class.java)
class KeyLocalLock(private val lockTimeoutMillis: Long) : KeyLock, CoroutineScope {
- private data class LockInfo(val semaphore: Semaphore, val createdAt: Long)
+ /**
+ * Internal lock state holder.
+ * Using class instead of data class to allow mutable expiresAt for atomic updates.
+ */
+ private class LockInfo(
+ val semaphore: Semaphore,
+ /**
+ * Expiration timestamp in milliseconds.
+ * @Volatile ensures visibility across threads when updated inside compute() and read by monitor.
+ */
+ @Volatile var expiresAt: Long,
+ /**
+ * Tracks whether the lock is currently held.
+ * Uses AtomicBoolean with CAS operations to prevent over-release
+ * when multiple threads race to release the same lock (e.g., tryLock expiration
+ * check vs unLock, or monitor cleanup vs unLock).
+ */
+ val isHeld: AtomicBoolean = AtomicBoolean(false),
+ )
- private val lockMap = ConcurrentHashMap()
+ companion object {
+ private val lockMap = ConcurrentHashMap()
+
+ @Volatile
+ private var monitorJob: Job? = null
+
+ private fun ensureMonitorStarted() {
+ if (monitorJob?.isActive == true) return
+ synchronized(this) {
+ if (monitorJob?.isActive == true) return
+ monitorJob =
+ CoroutineScope(Dispatchers.IO).launch {
+ while (isActive) {
+ runCatching {
+ val now = System.currentTimeMillis()
+ // Remove expired locks using compute() for atomic check-and-remove.
+ // This prevents TOCTOU race condition where removeIf's lambda returns true
+ // but the actual removal happens after a new lock is acquired.
+ // compute() guarantees atomic execution per key, so cleanup and tryLock
+ // are mutually exclusive for the same key.
+ lockMap.keys.forEach { key ->
+ lockMap.compute(key) { _, lockInfo ->
+ if (lockInfo == null) return@compute null
+
+ if (now > lockInfo.expiresAt) {
+ // Expired lock: force release regardless of isHeld state.
+ // This handles the case where unlock() was missed due to exception.
+ // CAS ensures safe release (no-op if already released).
+ if (lockInfo.isHeld.compareAndSet(true, false)) {
+ lockInfo.semaphore.release()
+ }
+ null // Atomic removal
+ } else {
+ lockInfo // Keep the entry
+ }
+ }
+ }
+ delay(LOCK_MONITOR_INTERVAL_MILLIS)
+ }.onFailure { e ->
+ log.error("Error in lock lifecycle monitoring: {}", e.message, e)
+ }
+ }
+ }
+ }
+ }
+
+ // For testing and resource cleanup
+ internal fun stopMonitoring() {
+ synchronized(this) {
+ monitorJob?.cancel()
+ monitorJob = null
+ }
+ }
+ }
private val job = Job()
override val coroutineContext: CoroutineContext
get() = Dispatchers.IO + job
init {
- launch {
- while (isActive) {
- runCatching {
- val now = System.currentTimeMillis()
- lockMap.entries.removeIf { now - it.value.createdAt > lockTimeoutMillis } // 특정 시간이 지나면 lock 여부와 상관없이 map에서 삭제한다.
- delay(LOCK_MONITOR_INTERVAL_MILLIS)
- }.onFailure { e ->
- log.error("Error in lock lifecycle monitoring : {}", e.message)
- }
- }
- }
+ ensureMonitorStarted()
}
override suspend fun tryLock(
@@ -59,9 +121,38 @@ class KeyLocalLock(private val lockTimeoutMillis: Long) : KeyLock, CoroutineScop
lockType: LockType,
): Boolean {
val completeKey = "${key}_${lockType.name}"
- val lockInfo = lockMap.computeIfAbsent(completeKey) { LockInfo(Semaphore(1), nowToEpochTime()) }
+ val now = nowToEpochTime()
+ val result = AtomicBoolean(false)
+
+ // Use compute() for atomic lock acquisition.
+ // This ensures mutual exclusion with cleanup - they cannot race on the same key.
+ lockMap.compute(completeKey) { _, existing ->
+ if (existing != null) {
+ // Force-release expired locks to allow reacquisition.
+ // Use CAS to prevent race condition with concurrent unLock().
+ // Without CAS, if unLock() executes between isHeld.get() and release(),
+ // both threads would call release(), causing over-release (permits > 1).
+ if (now > existing.expiresAt && existing.isHeld.compareAndSet(true, false)) {
+ existing.semaphore.release()
+ }
- return lockInfo.semaphore.tryAcquire()
+ // Existing entry: try to acquire semaphore
+ if (existing.semaphore.tryAcquire()) {
+ existing.isHeld.set(true)
+ existing.expiresAt = now + lockTimeoutMillis
+ result.set(true)
+ }
+ existing
+ } else {
+ // New entry: create and acquire
+ val newLock = LockInfo(Semaphore(1), now + lockTimeoutMillis)
+ newLock.semaphore.tryAcquire() // Always succeeds for new semaphore
+ newLock.isHeld.set(true)
+ result.set(true)
+ newLock
+ }
+ }
+ return result.get()
}
override suspend fun unLock(
@@ -69,15 +160,20 @@ class KeyLocalLock(private val lockTimeoutMillis: Long) : KeyLock, CoroutineScop
lockType: LockType,
): Boolean {
val completeKey = "${key}_${lockType.name}"
- val lockInfo = lockMap[completeKey]
- lockInfo?.let {
- it.semaphore.release()
- lockMap.remove(completeKey)
+ val lockInfo = lockMap[completeKey] ?: return false
+
+ // Use CAS to prevent over-release: only release if we actually hold the lock
+ return if (lockInfo.isHeld.compareAndSet(true, false)) {
+ lockInfo.semaphore.release()
+ true
+ } else {
+ log.debug("Attempted to unlock key '{}' that is not held", completeKey)
+ false
}
- return true
}
fun cancel() {
job.cancel()
+ // Monitor cleanup is handled via stopMonitoring() in tests
}
}
diff --git a/core-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/KeyGlobalLockTest.kt b/core-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/KeyGlobalLockTest.kt
index e55048a..2d7ab72 100644
--- a/core-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/KeyGlobalLockTest.kt
+++ b/core-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/KeyGlobalLockTest.kt
@@ -43,11 +43,16 @@ class KeyGlobalLockTest :
@BeforeEach
fun init() {
- val redisUrl = "redis://localhost:6379" // testContainer url
+ val host = AbstractRedisTest.redisHost
+ val port = AbstractRedisTest.redisPort
+ val redisUrl = "redis://$host:$port"
val redisClient = RedisClient.create(redisUrl)
val connection = redisClient.connect()
redisCommands = connection.async()
+ // Clean up all keys from previous tests for proper test isolation
+ connection.sync().flushdb()
+
globalLockFunc = { key, timeToLiveMillis ->
redisCommands.setnx(key, key).toCompletableFuture().await()
}
diff --git a/core-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/KeyLocalLockTest.kt b/core-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/KeyLocalLockTest.kt
index afb0667..8a4af79 100644
--- a/core-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/KeyLocalLockTest.kt
+++ b/core-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/KeyLocalLockTest.kt
@@ -18,17 +18,57 @@ package com.linecorp.cse.reqshield.kotlin.coroutine
import com.linecorp.cse.reqshield.support.BaseKeyLockTest
import kotlinx.coroutines.Dispatchers
+import kotlinx.coroutines.async
+import kotlinx.coroutines.awaitAll
import kotlinx.coroutines.delay
import kotlinx.coroutines.joinAll
import kotlinx.coroutines.launch
import kotlinx.coroutines.runBlocking
import kotlinx.coroutines.withContext
+import org.junit.jupiter.api.AfterEach
import org.junit.jupiter.api.Assertions.assertEquals
+import org.junit.jupiter.api.Assertions.assertFalse
import org.junit.jupiter.api.Assertions.assertTrue
import org.junit.jupiter.api.Test
import java.util.concurrent.atomic.AtomicInteger
class KeyLocalLockTest : BaseKeyLockTest {
+ @AfterEach
+ fun cleanup() {
+ // Ensure monitor is stopped after each test to prevent memory leaks
+ KeyLocalLock.stopMonitoring()
+ }
+
+ @Test
+ fun `should share global lockMap across multiple instances`() =
+ runBlocking {
+ val instance1 = KeyLocalLock(lockTimeoutMillis)
+ val instance2 = KeyLocalLock(lockTimeoutMillis)
+ val key = "shared-key"
+ val lockType = LockType.CREATE
+
+ assertTrue(instance1.tryLock(key, lockType))
+ assertTrue(!instance2.tryLock(key, lockType))
+
+ instance1.unLock(key, lockType)
+ }
+
+ @Test
+ fun `should maintain request collapsing across multiple instances`() =
+ runBlocking {
+ val instance1 = KeyLocalLock(lockTimeoutMillis)
+ val instance2 = KeyLocalLock(lockTimeoutMillis)
+ val instance3 = KeyLocalLock(lockTimeoutMillis)
+ val key = "collapsing-key"
+ val lockType = LockType.CREATE
+
+ val acquired = listOf(instance1, instance2, instance3).map { it.tryLock(key, lockType) }.count { it }
+ assertEquals(1, acquired)
+
+ // cleanup whoever acquired
+ listOf(instance1, instance2, instance3).forEach { it.unLock(key, lockType) }
+ }
+
@Test
override fun testConcurrencyWithOneKey() =
runBlocking {
@@ -129,5 +169,177 @@ class KeyLocalLockTest : BaseKeyLockTest {
assertTrue(keyLock.unLock(key, lockType))
}
+ @Test
+ fun `should not over-release semaphore on multiple unlock calls`() =
+ runBlocking {
+ val keyLock = KeyLocalLock(lockTimeoutMillis)
+ val key = "over-release-test"
+ val lockType = LockType.CREATE
+
+ // Acquire lock
+ assertTrue(keyLock.tryLock(key, lockType))
+
+ // First unlock should succeed
+ assertTrue(keyLock.unLock(key, lockType), "First unlock should succeed")
+
+ // Second unlock should return false (over-release prevention)
+ assertFalse(keyLock.unLock(key, lockType), "Second unlock should fail (over-release prevention)")
+
+ // Verify semaphore is not over-released: can acquire once, not twice
+ assertTrue(keyLock.tryLock(key, lockType), "Should acquire lock after proper unlock")
+ assertFalse(keyLock.tryLock(key, lockType), "Should not acquire lock twice (semaphore intact)")
+
+ // Cleanup
+ keyLock.unLock(key, lockType)
+ keyLock.cancel()
+ }
+
+ @Test
+ fun `should prevent concurrent lock acquisition after over-release attempt`() =
+ runBlocking {
+ val keyLock = KeyLocalLock(lockTimeoutMillis)
+ val key = "concurrent-over-release-test"
+ val lockType = LockType.CREATE
+ val successfulAcquisitions = AtomicInteger(0)
+
+ // Simulate over-release attempt
+ assertTrue(keyLock.tryLock(key, lockType))
+ keyLock.unLock(key, lockType)
+ // Multiple unlock attempts should all return false (not over-release)
+ repeat(5) { assertFalse(keyLock.unLock(key, lockType)) }
+
+ // Try to acquire lock concurrently - only ONE should succeed
+ val attempts =
+ (1..10).map {
+ async(Dispatchers.IO) {
+ if (keyLock.tryLock(key, lockType)) {
+ successfulAcquisitions.incrementAndGet()
+ }
+ }
+ }
+
+ attempts.awaitAll()
+
+ // Only one should have acquired the lock
+ assertEquals(1, successfulAcquisitions.get(), "Only one should acquire the lock")
+
+ // Cleanup
+ keyLock.unLock(key, lockType)
+ keyLock.cancel()
+ }
+
+ @Test
+ fun `should not over-release when tryLock and unLock race on expired lock`() =
+ runBlocking {
+ // Use a very short lock timeout to trigger expiration quickly
+ val shortLockTimeout = 50L
+ val keyLock = KeyLocalLock(shortLockTimeout)
+ val key = "race-condition-test"
+ val lockType = LockType.CREATE
+
+ repeat(100) { iteration ->
+ // Step 1: Acquire lock
+ assertTrue(keyLock.tryLock(key, lockType), "Iteration $iteration: Initial lock should succeed")
+
+ // Step 2: Wait for lock to expire (but not be cleaned up by monitor)
+ delay(shortLockTimeout + 10L)
+
+ // Step 3: Simulate race condition - tryLock and unLock concurrently
+ // tryLock will detect expiration and try to force-release
+ // unLock will also try to release
+ // Without CAS fix, both would call semaphore.release() causing over-release
+ val tryLockResult =
+ async(Dispatchers.IO) {
+ keyLock.tryLock(key, lockType)
+ }
+ val unLockResult =
+ async(Dispatchers.IO) {
+ keyLock.unLock(key, lockType)
+ }
+
+ tryLockResult.await()
+ unLockResult.await()
+
+ // Step 4: Verify no over-release by checking lock behavior
+ // If over-release occurred, permits > 1, allowing multiple acquisitions
+ val acquisitions = AtomicInteger(0)
+ val attempts =
+ (1..5).map {
+ async(Dispatchers.IO) {
+ if (keyLock.tryLock(key, lockType)) {
+ acquisitions.incrementAndGet()
+ }
+ }
+ }
+ attempts.awaitAll()
+
+ // At most 1 should succeed (0 if tryLock already holds it, 1 if it released)
+ assertTrue(
+ acquisitions.get() <= 1,
+ "Iteration $iteration: Over-release detected! " +
+ "Expected at most 1 acquisition, got ${acquisitions.get()}",
+ )
+
+ // Cleanup for next iteration
+ repeat(3) { keyLock.unLock(key, lockType) }
+ }
+
+ keyLock.cancel()
+ }
+
+ @Test
+ fun `should handle high contention tryLock and unLock without over-release`() =
+ runBlocking {
+ val shortLockTimeout = 30L
+ val keyLock = KeyLocalLock(shortLockTimeout)
+ val key = "high-contention-test"
+ val lockType = LockType.CREATE
+ val overReleaseDetected = AtomicInteger(0)
+
+ repeat(50) { iteration ->
+ // Acquire lock and let it expire
+ assertTrue(keyLock.tryLock(key, lockType))
+ delay(shortLockTimeout + 5L)
+
+ // High contention: many concurrent tryLock and unLock calls
+ val jobs =
+ (1..20).map { i ->
+ if (i % 2 == 0) {
+ async(Dispatchers.IO) { keyLock.tryLock(key, lockType) }
+ } else {
+ async(Dispatchers.IO) { keyLock.unLock(key, lockType) }
+ }
+ }
+ jobs.awaitAll()
+
+ // Verify: try to acquire lock multiple times concurrently
+ val acquisitions = AtomicInteger(0)
+ val verifyJobs =
+ (1..10).map {
+ async(Dispatchers.IO) {
+ if (keyLock.tryLock(key, lockType)) {
+ acquisitions.incrementAndGet()
+ }
+ }
+ }
+ verifyJobs.awaitAll()
+
+ if (acquisitions.get() > 1) {
+ overReleaseDetected.incrementAndGet()
+ }
+
+ // Cleanup
+ repeat(15) { keyLock.unLock(key, lockType) }
+ }
+
+ assertEquals(
+ 0,
+ overReleaseDetected.get(),
+ "Over-release detected in ${overReleaseDetected.get()} iterations",
+ )
+
+ keyLock.cancel()
+ }
+
private suspend fun doWork() = delay(1000)
}
diff --git a/core-reactor/src/main/kotlin/com/linecorp/cse/reqshield/reactor/KeyLocalLock.kt b/core-reactor/src/main/kotlin/com/linecorp/cse/reqshield/reactor/KeyLocalLock.kt
index 427744e..3ea4b06 100644
--- a/core-reactor/src/main/kotlin/com/linecorp/cse/reqshield/reactor/KeyLocalLock.kt
+++ b/core-reactor/src/main/kotlin/com/linecorp/cse/reqshield/reactor/KeyLocalLock.kt
@@ -19,34 +19,122 @@ package com.linecorp.cse.reqshield.reactor
import com.linecorp.cse.reqshield.support.constant.ConfigValues.LOCK_MONITOR_INTERVAL_MILLIS
import com.linecorp.cse.reqshield.support.utils.nowToEpochTime
import org.slf4j.LoggerFactory
+import reactor.core.Disposable
import reactor.core.publisher.Flux
import reactor.core.publisher.Mono
import reactor.core.scheduler.Schedulers
import java.time.Duration
import java.util.concurrent.ConcurrentHashMap
import java.util.concurrent.Semaphore
+import java.util.concurrent.atomic.AtomicBoolean
+import java.util.concurrent.atomic.AtomicInteger
private val log = LoggerFactory.getLogger(KeyLocalLock::class.java)
class KeyLocalLock(
private val lockTimeoutMillis: Long,
) : KeyLock {
- private data class LockInfo(
+ /**
+ * Internal lock state holder.
+ * Using class instead of data class to allow mutable expiresAt for atomic updates.
+ */
+ private class LockInfo(
val semaphore: Semaphore,
- val createdAt: Long,
+ /**
+ * Expiration timestamp in milliseconds.
+ * @Volatile ensures visibility across threads when updated inside compute() and read by monitor.
+ */
+ @Volatile var expiresAt: Long,
+ /**
+ * Tracks whether the lock is currently held.
+ * Uses AtomicBoolean with CAS operations to prevent over-release
+ * when multiple threads race to release the same lock (e.g., tryLock expiration
+ * check vs unLock, or monitor cleanup vs unLock).
+ */
+ val isHeld: AtomicBoolean = AtomicBoolean(false),
)
- private val lockMap = ConcurrentHashMap()
+ companion object {
+ private val lockMap = ConcurrentHashMap()
+
+ @Volatile
+ private var monitoringStarted: Boolean = false
+
+ @Volatile
+ private var monitorDisposable: Disposable? = null
+
+ // Track consecutive failures for backoff logging
+ private val consecutiveFailures = AtomicInteger(0)
+
+ private fun startMonitoringOnce() {
+ if (monitoringStarted) return
+ synchronized(this) {
+ if (monitoringStarted) return
+ monitorDisposable =
+ Flux
+ .interval(Duration.ofMillis(LOCK_MONITOR_INTERVAL_MILLIS), Schedulers.single())
+ .flatMap {
+ Mono
+ .fromRunnable {
+ val now = System.currentTimeMillis()
+ // Remove expired locks using compute() for atomic check-and-remove.
+ // This prevents TOCTOU race condition where removeIf's lambda returns true
+ // but the actual removal happens after a new lock is acquired.
+ // compute() guarantees atomic execution per key, so cleanup and tryLock
+ // are mutually exclusive for the same key.
+ lockMap.keys.forEach { key ->
+ lockMap.compute(key) { _, lockInfo ->
+ if (lockInfo == null) return@compute null
+
+ if (now > lockInfo.expiresAt) {
+ // Expired lock: force release regardless of isHeld state.
+ // This handles the case where unlock() was missed due to exception.
+ // CAS ensures safe release (no-op if already released).
+ if (lockInfo.isHeld.compareAndSet(true, false)) {
+ lockInfo.semaphore.release()
+ }
+ null // Atomic removal
+ } else {
+ lockInfo // Keep the entry
+ }
+ }
+ }
+ consecutiveFailures.set(0) // Reset on success
+ }.onErrorResume { e ->
+ // Log only on first failure or every 10th consecutive failure
+ val failures = consecutiveFailures.incrementAndGet()
+ if (failures == 1 || failures % 10 == 0) {
+ log.warn(
+ "Error in lock lifecycle monitoring (consecutive failures: {}): {}",
+ failures,
+ e.message,
+ )
+ }
+ // Backoff: delay on failure (max 5 seconds)
+ val backoffMs = minOf(failures * LOCK_MONITOR_INTERVAL_MILLIS, 5000L)
+ Mono.delay(Duration.ofMillis(backoffMs)).then(Mono.empty())
+ }
+ }.subscribe(
+ { /* success - no action needed */ },
+ { e -> log.error("Fatal error in lock lifecycle monitoring: {}", e.message, e) },
+ )
+ monitoringStarted = true
+ }
+ }
+
+ // For testing and resource cleanup
+ internal fun stopMonitoring() {
+ synchronized(this) {
+ monitorDisposable?.dispose()
+ monitorDisposable = null
+ consecutiveFailures.set(0)
+ monitoringStarted = false
+ }
+ }
+ }
init {
- Flux
- .interval(Duration.ofMillis(LOCK_MONITOR_INTERVAL_MILLIS), Schedulers.single())
- .doOnNext {
- val now = System.currentTimeMillis()
- lockMap.entries.removeIf { now - it.value.createdAt > lockTimeoutMillis }
- }.doOnError { e ->
- log.error("Error in lock lifecycle monitoring : {}", e.message)
- }.subscribe()
+ startMonitoringOnce()
}
override fun tryLock(
@@ -55,21 +143,55 @@ class KeyLocalLock(
): Mono =
Mono.fromCallable {
val completeKey = "${key}_${lockType.name}"
- val lockInfo = lockMap.computeIfAbsent(completeKey) { LockInfo(Semaphore(1), nowToEpochTime()) }
- lockInfo.semaphore.tryAcquire()
+ val now = nowToEpochTime()
+ val result = AtomicBoolean(false)
+
+ // Use compute() for atomic lock acquisition.
+ // This ensures mutual exclusion with cleanup - they cannot race on the same key.
+ lockMap.compute(completeKey) { _, existing ->
+ if (existing != null) {
+ // Force-release expired locks to allow reacquisition.
+ // Use CAS to prevent race condition with concurrent unLock().
+ // Without CAS, if unLock() executes between isHeld.get() and release(),
+ // both threads would call release(), causing over-release (permits > 1).
+ if (now > existing.expiresAt && existing.isHeld.compareAndSet(true, false)) {
+ existing.semaphore.release()
+ }
+
+ // Existing entry: try to acquire semaphore
+ if (existing.semaphore.tryAcquire()) {
+ existing.isHeld.set(true)
+ existing.expiresAt = now + lockTimeoutMillis
+ result.set(true)
+ }
+ existing
+ } else {
+ // New entry: create and acquire
+ val newLock = LockInfo(Semaphore(1), now + lockTimeoutMillis)
+ newLock.semaphore.tryAcquire() // Always succeeds for new semaphore
+ newLock.isHeld.set(true)
+ result.set(true)
+ newLock
+ }
+ }
+ result.get()
}
override fun unLock(
key: String,
lockType: LockType,
): Mono =
- Mono
- .fromCallable {
- val completeKey = "${key}_${lockType.name}"
- val lockInfo = lockMap[completeKey]
- lockInfo?.let {
- it.semaphore.release()
- lockMap.remove(completeKey)
- }
- }.thenReturn(true)
+ Mono.fromCallable {
+ val completeKey = "${key}_${lockType.name}"
+ val lockInfo = lockMap[completeKey] ?: return@fromCallable false
+
+ // Use CAS to prevent over-release: only release if we actually hold the lock
+ if (lockInfo.isHeld.compareAndSet(true, false)) {
+ lockInfo.semaphore.release()
+ true
+ } else {
+ log.debug("Attempted to unlock key '{}' that is not held", completeKey)
+ false
+ }
+ }
}
diff --git a/core-reactor/src/main/kotlin/com/linecorp/cse/reqshield/reactor/ReqShield.kt b/core-reactor/src/main/kotlin/com/linecorp/cse/reqshield/reactor/ReqShield.kt
index 1d0b03f..7353314 100644
--- a/core-reactor/src/main/kotlin/com/linecorp/cse/reqshield/reactor/ReqShield.kt
+++ b/core-reactor/src/main/kotlin/com/linecorp/cse/reqshield/reactor/ReqShield.kt
@@ -19,19 +19,18 @@ package com.linecorp.cse.reqshield.reactor
import com.linecorp.cse.reqshield.reactor.config.ReqShieldConfiguration
import com.linecorp.cse.reqshield.reactor.config.ReqShieldWorkMode
import com.linecorp.cse.reqshield.support.constant.ConfigValues.GET_CACHE_INTERVAL_MILLIS
-import com.linecorp.cse.reqshield.support.constant.ConfigValues.MAX_ATTEMPT_SET_CACHE
-import com.linecorp.cse.reqshield.support.constant.ConfigValues.SET_CACHE_RETRY_INTERVAL_MILLIS
import com.linecorp.cse.reqshield.support.exception.ClientException
import com.linecorp.cse.reqshield.support.exception.code.ErrorCode
import com.linecorp.cse.reqshield.support.model.ReqShieldData
import com.linecorp.cse.reqshield.support.utils.decideToUpdateCache
+import org.slf4j.LoggerFactory
import reactor.core.publisher.Flux
import reactor.core.publisher.Mono
-import reactor.core.scheduler.Schedulers
-import reactor.util.retry.Retry
import java.time.Duration
import java.util.concurrent.Callable
+private val log = LoggerFactory.getLogger(ReqShield::class.java)
+
class ReqShield(
private val reqShieldConfig: ReqShieldConfiguration,
) {
@@ -73,13 +72,13 @@ class ReqShield(
fun processMono(): Mono> =
executeCallable({ callable.call() }, true, key, lockType)
.map { data -> buildReqShieldData(data, timeToLiveMillis) }
- .doOnNext { reqShieldData ->
+ .flatMap { reqShieldData ->
setReqShieldData(
reqShieldConfig.setCacheFunction,
key,
reqShieldData,
lockType,
- )
+ ).thenReturn(reqShieldData)
}.switchIfEmpty(
Mono.defer {
val reqShieldData = buildReqShieldData(null, timeToLiveMillis)
@@ -88,22 +87,27 @@ class ReqShield(
key,
reqShieldData,
lockType,
- )
- Mono.just(reqShieldData)
+ ).thenReturn(reqShieldData)
},
)
if (reqShieldConfig.reqShieldWorkMode == ReqShieldWorkMode.ONLY_CREATE_CACHE) {
processMono()
- .subscribeOn(Schedulers.boundedElastic())
- .subscribe()
+ .subscribeOn(reqShieldConfig.scheduler)
+ .subscribe(
+ { /* success - no action needed */ },
+ { e -> log.error("Failed to update cache for key '{}': {}", key, e.message, e) },
+ )
} else {
reqShieldConfig.keyLock
.tryLock(key, lockType)
.filter { it }
.flatMap { processMono() }
- .subscribeOn(Schedulers.boundedElastic())
- .subscribe()
+ .subscribeOn(reqShieldConfig.scheduler)
+ .subscribe(
+ { /* success - no action needed */ },
+ { e -> log.error("Failed to update cache for key '{}': {}", key, e.message, e) },
+ )
}
}
@@ -137,27 +141,32 @@ class ReqShield(
): Mono> =
executeCallable({ callable.call() }, true, key, lockType)
.map { data -> buildReqShieldData(data, timeToLiveMillis) }
- .flatMap { reqShieldData ->
-
+ .doOnNext { reqShieldData ->
+ // Async fire-and-forget cache storage (matches coroutine implementation)
setReqShieldData(
reqShieldConfig.setCacheFunction,
key,
reqShieldData,
lockType,
- )
-
- Mono.just(reqShieldData)
+ ).subscribeOn(reqShieldConfig.scheduler)
+ .subscribe(
+ { /* success - no action needed */ },
+ { e -> log.error("Failed to set cache for key '{}': {}", key, e.message, e) },
+ )
}.switchIfEmpty(
Mono.defer {
val reqShieldData = buildReqShieldData(null, timeToLiveMillis)
-
+ // Async fire-and-forget cache storage (matches coroutine implementation)
setReqShieldData(
reqShieldConfig.setCacheFunction,
key,
reqShieldData,
lockType,
- )
-
+ ).subscribeOn(reqShieldConfig.scheduler)
+ .subscribe(
+ { /* success - no action needed */ },
+ { e -> log.error("Failed to set cache for key '{}': {}", key, e.message, e) },
+ )
Mono.just(reqShieldData)
},
)
@@ -191,7 +200,7 @@ class ReqShield(
Mono.just(reqShieldData)
},
),
- ).subscribeOn(Schedulers.boundedElastic())
+ ).subscribeOn(reqShieldConfig.scheduler)
private fun buildReqShieldData(
value: T?,
@@ -207,18 +216,14 @@ class ReqShield(
key: String,
reqShieldData: ReqShieldData,
lockType: LockType,
- ) {
- executeSetCacheFunction(cacheSetter, key, reqShieldData, lockType).subscribe()
- }
+ ): Mono = executeSetCacheFunction(cacheSetter, key, reqShieldData, lockType)
private fun executeGetCacheFunction(
getFunction: (String) -> Mono?>,
key: String,
): Mono?> =
getFunction(key)
- .doOnError { e ->
- throw ClientException(ErrorCode.GET_CACHE_ERROR, originErrorMessage = e.message)
- }
+ .onErrorMap { e -> ClientException(ErrorCode.GET_CACHE_ERROR, originErrorMessage = e.message) }
private fun executeSetCacheFunction(
setFunction: (String, ReqShieldData, Long) -> Mono,
@@ -227,20 +232,22 @@ class ReqShield(
lockType: LockType,
): Mono =
setFunction(key, value, value.timeToLiveMillis)
- .doOnError { e ->
- throw ClientException(ErrorCode.SET_CACHE_ERROR, originErrorMessage = e.message)
- }.doFinally {
+ .onErrorMap { e -> ClientException(ErrorCode.SET_CACHE_ERROR, originErrorMessage = e.message) }
+ .doFinally {
if (shouldAttemptUnlock(lockType)) {
+ // No retry needed: false means lock already released or expired (not an error)
reqShieldConfig.keyLock
.unLock(key, lockType)
- .retryWhen(
- Retry.fixedDelay(
- MAX_ATTEMPT_SET_CACHE - 1L,
- Duration.ofMillis(SET_CACHE_RETRY_INTERVAL_MILLIS),
- ),
- ).subscribe()
+ .doOnNext { unlocked ->
+ if (!unlocked) {
+ log.debug("Lock already released or expired for key '{}'", key)
+ }
+ }.subscribe(
+ { /* success - no action needed */ },
+ { e -> log.error("Failed to unlock key '{}': {}", key, e.message, e) },
+ )
}
- }.subscribeOn(Schedulers.boundedElastic())
+ }.subscribeOn(reqShieldConfig.scheduler)
private fun executeCallable(
callable: Callable>,
@@ -250,11 +257,24 @@ class ReqShield(
): Mono =
callable
.call()
- .doOnError { e ->
+ .doOnError { _ ->
if (isUnlockWhenException && key != null && lockType != null) {
- reqShieldConfig.keyLock.unLock(key, lockType).subscribe()
+ reqShieldConfig.keyLock
+ .unLock(key, lockType)
+ .subscribe(
+ { /* success - no action needed */ },
+ { unlockError ->
+ log.error(
+ "Failed to unlock key '{}' after callable error: {}",
+ key,
+ unlockError.message,
+ unlockError,
+ )
+ },
+ )
}
- throw ClientException(ErrorCode.SUPPLIER_ERROR, originErrorMessage = e.message)
+ }.onErrorMap { e ->
+ ClientException(ErrorCode.SUPPLIER_ERROR, originErrorMessage = e.message)
}
private fun shouldAttemptUnlock(lockType: LockType): Boolean =
diff --git a/core-reactor/src/test/kotlin/com/linecorp/cse/reqshield/reactor/KeyGlobalLockTest.kt b/core-reactor/src/test/kotlin/com/linecorp/cse/reqshield/reactor/KeyGlobalLockTest.kt
index 17867ae..613bf53 100644
--- a/core-reactor/src/test/kotlin/com/linecorp/cse/reqshield/reactor/KeyGlobalLockTest.kt
+++ b/core-reactor/src/test/kotlin/com/linecorp/cse/reqshield/reactor/KeyGlobalLockTest.kt
@@ -40,11 +40,16 @@ class KeyGlobalLockTest :
@BeforeEach
fun init() {
- val redisUrl = "redis://localhost:6379" // testContainer url
+ val host = AbstractRedisTest.redisHost
+ val port = AbstractRedisTest.redisPort
+ val redisUrl = "redis://$host:$port"
val redisClient = RedisClient.create(redisUrl)
val connection = redisClient.connect()
redisCommands = connection.async()
+ // Clean up all keys from previous tests for proper test isolation
+ connection.sync().flushdb()
+
globalLockFunc = { key, timeToLiveMillis ->
Mono.fromFuture { redisCommands.setnx(key, key).toCompletableFuture() }
}
diff --git a/core-reactor/src/test/kotlin/com/linecorp/cse/reqshield/reactor/KeyLocalLockTest.kt b/core-reactor/src/test/kotlin/com/linecorp/cse/reqshield/reactor/KeyLocalLockTest.kt
index 6460cef..f245284 100644
--- a/core-reactor/src/test/kotlin/com/linecorp/cse/reqshield/reactor/KeyLocalLockTest.kt
+++ b/core-reactor/src/test/kotlin/com/linecorp/cse/reqshield/reactor/KeyLocalLockTest.kt
@@ -17,6 +17,7 @@
package com.linecorp.cse.reqshield.reactor
import com.linecorp.cse.reqshield.support.BaseKeyLockTest
+import org.junit.jupiter.api.AfterEach
import org.junit.jupiter.api.Assertions.assertEquals
import org.junit.jupiter.api.Assertions.assertTrue
import org.junit.jupiter.api.Test
@@ -27,6 +28,47 @@ import java.time.Duration
import java.util.concurrent.atomic.AtomicInteger
class KeyLocalLockTest : BaseKeyLockTest {
+ @AfterEach
+ fun cleanup() {
+ // Ensure monitor can restart after each test to prevent test isolation issues
+ KeyLocalLock.stopMonitoring()
+ }
+
+ @Test
+ fun `should share global lockMap across multiple instances`() {
+ val instance1 = KeyLocalLock(lockTimeoutMillis)
+ val instance2 = KeyLocalLock(lockTimeoutMillis)
+ val key = "shared-key"
+ val lockType = LockType.CREATE
+
+ StepVerifier.create(instance1.tryLock(key, lockType)).expectNext(true).verifyComplete()
+ StepVerifier.create(instance2.tryLock(key, lockType)).expectNext(false).verifyComplete()
+
+ StepVerifier.create(instance1.unLock(key, lockType)).expectNext(true).verifyComplete()
+ }
+
+ @Test
+ fun `should maintain request collapsing across multiple instances`() {
+ val instance1 = KeyLocalLock(lockTimeoutMillis)
+ val instance2 = KeyLocalLock(lockTimeoutMillis)
+ val instance3 = KeyLocalLock(lockTimeoutMillis)
+ val key = "collapsing-key"
+ val lockType = LockType.CREATE
+
+ val attempts =
+ listOf(instance1, instance2, instance3).map { inst ->
+ inst.tryLock(key, lockType).map { acquired -> if (acquired) 1 else 0 }
+ }
+
+ StepVerifier
+ .create(Mono.zip(attempts) { arr -> arr.sumOf { it as Int } })
+ .expectNextMatches { it == 1 }
+ .verifyComplete()
+
+ // cleanup by unlocking whoever acquired
+ listOf(instance1, instance2, instance3).forEach { inst -> inst.unLock(key, lockType).subscribe() }
+ }
+
@Test
override fun testConcurrencyWithOneKey() {
val keyLock = KeyLocalLock(lockTimeoutMillis)
@@ -148,6 +190,82 @@ class KeyLocalLockTest : BaseKeyLockTest {
.verifyComplete()
}
+ @Test
+ fun `should not over-release semaphore on multiple unlock calls`() {
+ val keyLock = KeyLocalLock(lockTimeoutMillis)
+ val key = "over-release-test"
+ val lockType = LockType.CREATE
+
+ // Acquire lock
+ StepVerifier.create(keyLock.tryLock(key, lockType))
+ .expectNext(true)
+ .verifyComplete()
+
+ // First unlock should succeed
+ StepVerifier.create(keyLock.unLock(key, lockType))
+ .expectNext(true)
+ .verifyComplete()
+
+ // Second unlock should return false (over-release prevention)
+ StepVerifier.create(keyLock.unLock(key, lockType))
+ .expectNext(false)
+ .verifyComplete()
+
+ // Verify semaphore is not over-released: can acquire once, not twice
+ StepVerifier.create(keyLock.tryLock(key, lockType))
+ .expectNext(true)
+ .verifyComplete()
+
+ StepVerifier.create(keyLock.tryLock(key, lockType))
+ .expectNext(false)
+ .verifyComplete()
+
+ // Cleanup
+ keyLock.unLock(key, lockType).subscribe()
+ }
+
+ @Test
+ fun `should prevent concurrent lock acquisition after over-release attempt`() {
+ val keyLock = KeyLocalLock(lockTimeoutMillis)
+ val key = "concurrent-over-release-test"
+ val lockType = LockType.CREATE
+ val successfulAcquisitions = AtomicInteger(0)
+
+ // Simulate over-release attempt
+ StepVerifier.create(keyLock.tryLock(key, lockType))
+ .expectNext(true)
+ .verifyComplete()
+
+ StepVerifier.create(keyLock.unLock(key, lockType))
+ .expectNext(true)
+ .verifyComplete()
+
+ // Multiple unlock attempts should all return false
+ repeat(5) {
+ StepVerifier.create(keyLock.unLock(key, lockType))
+ .expectNext(false)
+ .verifyComplete()
+ }
+
+ // Try to acquire lock concurrently - only ONE should succeed
+ val attempts =
+ (1..10).map {
+ keyLock.tryLock(key, lockType)
+ .map { acquired -> if (acquired) successfulAcquisitions.incrementAndGet() else 0 }
+ }
+
+ StepVerifier
+ .create(Mono.zip(attempts) { it.toList() })
+ .expectNextCount(1)
+ .verifyComplete()
+
+ // Only one should have acquired the lock
+ assertEquals(1, successfulAcquisitions.get(), "Only one should acquire the lock")
+
+ // Cleanup
+ keyLock.unLock(key, lockType).subscribe()
+ }
+
private fun doWork(): Mono =
Mono
.delay(Duration.ofSeconds(1))
diff --git a/core-spring-webflux-kotlin-coroutine/build.gradle.kts b/core-spring-webflux-kotlin-coroutine/build.gradle.kts
index 52595db..9bba987 100644
--- a/core-spring-webflux-kotlin-coroutine/build.gradle.kts
+++ b/core-spring-webflux-kotlin-coroutine/build.gradle.kts
@@ -34,7 +34,9 @@ dependencies {
testImplementation(rootProject.libs.kotlin.coroutine.test)
testImplementation(rootProject.libs.kotlin.coroutine.jvm)
testImplementation(rootProject.libs.spring.context)
+ testImplementation(rootProject.libs.spring.test)
testImplementation(rootProject.libs.aspectj)
+ testImplementation(rootProject.libs.lettuce)
}
tasks.withType().configureEach {
diff --git a/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/CoroutineExtension.kt b/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/CoroutineExtension.kt
index 6513ef1..5ace0a5 100644
--- a/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/CoroutineExtension.kt
+++ b/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/CoroutineExtension.kt
@@ -18,6 +18,10 @@
package com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.aspect
+import kotlinx.coroutines.CoroutineDispatcher
+import kotlinx.coroutines.Dispatchers
+import kotlinx.coroutines.ExperimentalCoroutinesApi
+import kotlinx.coroutines.withContext
import org.aspectj.lang.ProceedingJoinPoint
import kotlin.coroutines.Continuation
import kotlin.coroutines.intrinsics.startCoroutineUninterceptedOrReturn
@@ -39,3 +43,44 @@ suspend fun ProceedingJoinPoint.proceedCoroutine(args: Array = this.corout
fun ProceedingJoinPoint.runCoroutine(block: suspend () -> Any?): Any? =
block.startCoroutineUninterceptedOrReturn(this.coroutineContinuation)
+
+/**
+ * Bounded dispatcher for non-suspend join point execution to prevent IO dispatcher exhaustion.
+ * Limits concurrent blocking calls to prevent thread pool saturation under heavy load.
+ *
+ * Parallelism can be configured via system property:
+ * - `reqshield.blocking.parallelism`: explicit parallelism value (1-1024)
+ * - Default: availableProcessors * 2, clamped to [4, 256]
+ *
+ * Examples:
+ * - `-Dreqshield.blocking.parallelism=64` for high-throughput environments
+ * - `-Dreqshield.blocking.parallelism=8` for resource-constrained environments
+ */
+@OptIn(ExperimentalCoroutinesApi::class)
+private val boundedBlockingDispatcher: CoroutineDispatcher by lazy {
+ val defaultParallelism =
+ (Runtime.getRuntime().availableProcessors() * 2)
+ .coerceIn(4, 256) // Min 4, max 256
+
+ val parallelism =
+ System.getProperty("reqshield.blocking.parallelism")
+ ?.toIntOrNull()
+ ?.coerceIn(1, 1024) // Configured value also bounded
+ ?: defaultParallelism
+
+ Dispatchers.IO.limitedParallelism(parallelism)
+}
+
+/**
+ * Proceed supporting both suspend and non-suspend join points.
+ * If the last argument is a Continuation, treat as suspend; otherwise proceed normally.
+ * Uses a bounded dispatcher for non-suspend calls to prevent IO thread pool exhaustion.
+ */
+suspend fun ProceedingJoinPoint.proceedSmart(): Any? =
+ if (this.args.isNotEmpty() && this.args.last() is Continuation<*>) {
+ this.proceedCoroutine()
+ } else {
+ // Use bounded dispatcher to prevent IO thread pool exhaustion
+ // when many synchronous methods are proxied concurrently
+ withContext(boundedBlockingDispatcher) { this@proceedSmart.proceed() }
+ }
diff --git a/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/ReqShieldAspect.kt b/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/ReqShieldAspect.kt
index 0fd9777..8cf2302 100644
--- a/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/ReqShieldAspect.kt
+++ b/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/ReqShieldAspect.kt
@@ -47,7 +47,7 @@ import kotlin.coroutines.Continuation
@Aspect
@Component
-class ReqShieldAspect(
+open class ReqShieldAspect(
private val asyncCache: AsyncCache,
) : BeanFactoryAware {
private lateinit var beanFactory: BeanFactory
@@ -58,41 +58,38 @@ class ReqShieldAspect(
private val keyGeneratorMap = ConcurrentHashMap()
internal val reqShieldMap = ConcurrentHashMap>()
- @Around("execution(@com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.annotation.* * *(.., kotlin.coroutines.Continuation))")
- fun aroundTargetCacheable(joinPoint: ProceedingJoinPoint): Any? {
- return joinPoint.runCoroutine {
- getTargetMethod(joinPoint).annotations.forEach { annotation ->
- when (annotation) {
- is ReqShieldCacheable -> {
- val reqShield = getOrCreateReqShield(joinPoint)
- val cacheKey = getCacheableCacheKey(joinPoint)
-
- return@runCoroutine reqShield
- .getAndSetReqShieldData(
- cacheKey,
- {
- joinPoint.proceedCoroutine().let { rtn ->
- if (rtn is Mono<*>) {
- rtn.awaitSingleOrNull()?.let { it as T }
- } else {
- rtn?.let { it as T }
- }
- }
- },
- annotation.timeToLiveMillis,
- ).value
- }
-
- is ReqShieldCacheEvict -> {
- val cacheKey = getCacheEvictCacheKey(joinPoint)
- return@runCoroutine asyncCache.evict(cacheKey)
- }
- }
- }
+ @Around("@annotation(com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.annotation.ReqShieldCacheable)")
+ fun aroundReqShieldCacheable(joinPoint: ProceedingJoinPoint): Any? =
+ joinPoint.runCoroutine {
+ val annotation = getCacheableAnnotation(joinPoint)
+ val reqShield = getOrCreateReqShield(joinPoint)
+ val cacheKey = getCacheableCacheKey(joinPoint)
+
+ reqShield
+ .getAndSetReqShieldData(
+ cacheKey,
+ {
+ joinPoint.proceedSmart().let { rtn ->
+ if (rtn is Mono<*>) {
+ rtn.awaitSingleOrNull()?.let { it as T }
+ } else {
+ rtn?.let { it as T }
+ }
+ }
+ },
+ annotation.timeToLiveMillis,
+ ).value
+ }
+
+ @Around("@annotation(com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.annotation.ReqShieldCacheEvict)")
+ fun aroundReqShieldCacheEvict(joinPoint: ProceedingJoinPoint): Any? =
+ joinPoint.runCoroutine {
+ val cacheKey = getCacheEvictCacheKey(joinPoint)
+ asyncCache.evict(cacheKey)
+ joinPoint.proceedSmart()
}
- }
- internal fun getTargetMethod(joinPoint: ProceedingJoinPoint): Method = (joinPoint.signature as MethodSignature).method
+ internal open fun getTargetMethod(joinPoint: ProceedingJoinPoint): Method = (joinPoint.signature as MethodSignature).method
internal fun getCacheableAnnotation(joinPoint: ProceedingJoinPoint): ReqShieldCacheable =
AnnotationUtils.getAnnotation(getTargetMethod(joinPoint), ReqShieldCacheable::class.java)
@@ -141,7 +138,15 @@ class ReqShieldAspect(
keyGenerator.generate(joinPoint.target, method, args).toString()
}
- require(!key.isNullOrBlank()) { "Null key returned for cache method : $method" }
+ require(!key.isNullOrBlank()) {
+ "Null/blank key for @ReqShieldCacheable method=${method.declaringClass.name}.${method.name} " +
+ "args=${args.joinToString(prefix = "[", postfix = "]") {
+ it?.let {
+ arg ->
+ "${arg::class.simpleName}@${arg.hashCode().toString(16)}"
+ } ?: "null"
+ }}"
+ }
return key
}
@@ -183,7 +188,9 @@ class ReqShieldAspect(
cacheKeyGenerator: String,
) {
if (cacheKey.isNotBlank() && cacheKeyGenerator.isNotBlank()) {
- throw IllegalArgumentException("The key and keyGenerator attributes are mutually exclusive.")
+ throw IllegalArgumentException(
+ "The key and keyGenerator attributes are mutually exclusive: key='$cacheKey', keyGenerator='$cacheKeyGenerator'",
+ )
}
}
@@ -206,8 +213,11 @@ class ReqShieldAspect(
return major > 6 || (major == 6 && minor >= 1)
}
- private fun generateReqShieldKey(joinPoint: ProceedingJoinPoint): String =
- "${getCacheableAnnotation(joinPoint).cacheName}-${getCacheableCacheKey(joinPoint)}"
+ private fun generateReqShieldKey(joinPoint: ProceedingJoinPoint): String {
+ val method = getTargetMethod(joinPoint)
+ return "${method.declaringClass.name}.${method.name}-" +
+ "${getCacheableAnnotation(joinPoint).cacheName}-${getCacheableCacheKey(joinPoint)}"
+ }
override fun setBeanFactory(beanFactory: BeanFactory) {
this.beanFactory = beanFactory
diff --git a/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/config/LibAutoConfiguration.kt b/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/config/LibAutoConfiguration.kt
index 34fec10..14555fa 100644
--- a/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/config/LibAutoConfiguration.kt
+++ b/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/config/LibAutoConfiguration.kt
@@ -16,11 +16,12 @@
package com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.config
-import org.springframework.context.annotation.ComponentScan
+import com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.aspect.ReqShieldAspect
import org.springframework.context.annotation.Configuration
import org.springframework.context.annotation.EnableAspectJAutoProxy
+import org.springframework.context.annotation.Import
@Configuration
-@EnableAspectJAutoProxy
-@ComponentScan(basePackages = ["com.linecorp.cse"])
+@EnableAspectJAutoProxy(proxyTargetClass = true)
+@Import(ReqShieldAspect::class)
open class LibAutoConfiguration
diff --git a/core-spring-webflux-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/InMemoryAsyncCache.kt b/core-spring-webflux-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/InMemoryAsyncCache.kt
new file mode 100644
index 0000000..926c3dc
--- /dev/null
+++ b/core-spring-webflux-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/InMemoryAsyncCache.kt
@@ -0,0 +1,56 @@
+/*
+ * Copyright 2024 LY Corporation
+ *
+ * LY Corporation licenses this file to you under the Apache License,
+ * version 2.0 (the "License"); you may not use this file except in compliance
+ * with the License. You may obtain a copy of the License at:
+ *
+ * https://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS, WITHOUT
+ * WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the
+ * License for the specific language governing permissions and limitations
+ * under the License.
+ */
+
+package com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.aspect
+
+import com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.cache.AsyncCache
+import com.linecorp.cse.reqshield.support.model.ReqShieldData
+import java.util.concurrent.ConcurrentHashMap
+import java.util.concurrent.Semaphore
+
+class InMemoryAsyncCache : AsyncCache {
+ private data class Entry(val data: ReqShieldData, val expiresAt: Long)
+
+ private val store = ConcurrentHashMap>()
+ private val locks = ConcurrentHashMap()
+
+ override suspend fun get(key: String): ReqShieldData? {
+ val now = System.currentTimeMillis()
+ return store[key]?.let { e -> if (now <= e.expiresAt) e.data else null }
+ }
+
+ override suspend fun put(
+ key: String,
+ value: ReqShieldData,
+ timeToLiveMillis: Long,
+ ): Boolean {
+ val expiresAt = System.currentTimeMillis() + timeToLiveMillis
+ store[key] = Entry(value, expiresAt)
+ return true
+ }
+
+ override suspend fun evict(key: String): Boolean = store.remove(key) != null
+
+ override suspend fun globalLock(
+ key: String,
+ timeToLiveMillis: Long,
+ ): Boolean = locks.computeIfAbsent(key) { Semaphore(1) }.tryAcquire()
+
+ override suspend fun globalUnLock(key: String): Boolean {
+ locks[key]?.release()
+ return true
+ }
+}
diff --git a/core-spring-webflux-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/ReqShieldAspectIntegrationTest.kt b/core-spring-webflux-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/ReqShieldAspectIntegrationTest.kt
new file mode 100644
index 0000000..8a253d5
--- /dev/null
+++ b/core-spring-webflux-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/ReqShieldAspectIntegrationTest.kt
@@ -0,0 +1,105 @@
+/*
+ * Copyright 2024 LY Corporation
+ *
+ * LY Corporation licenses this file to you under the Apache License,
+ * version 2.0 (the "License"); you may not use this file except in compliance
+ * with the License. You may obtain a copy of the License at:
+ *
+ * https://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS, WITHOUT
+ * WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the
+ * License for the specific language governing permissions and limitations
+ * under the License.
+ */
+
+package com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.aspect
+
+import com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.annotation.ReqShieldCacheEvict
+import com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.annotation.ReqShieldCacheable
+import com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.cache.AsyncCache
+import com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.config.LibAutoConfiguration
+import kotlinx.coroutines.Dispatchers
+import kotlinx.coroutines.async
+import kotlinx.coroutines.awaitAll
+import kotlinx.coroutines.delay
+import kotlinx.coroutines.runBlocking
+import kotlinx.coroutines.withTimeoutOrNull
+import org.junit.jupiter.api.Assertions.assertEquals
+import org.junit.jupiter.api.Assertions.assertTrue
+import org.junit.jupiter.api.Test
+import org.junit.jupiter.api.extension.ExtendWith
+import org.springframework.beans.factory.annotation.Autowired
+import org.springframework.context.annotation.Bean
+import org.springframework.context.annotation.Configuration
+import org.springframework.test.context.ContextConfiguration
+import org.springframework.test.context.junit.jupiter.SpringExtension
+import java.util.concurrent.atomic.AtomicInteger
+
+@ExtendWith(SpringExtension::class)
+@ContextConfiguration(classes = [LibAutoConfiguration::class, ReqShieldAspectIntegrationTest.TestConfig::class])
+class ReqShieldAspectIntegrationTest {
+ @Autowired
+ private lateinit var service: TestService
+
+ @Autowired
+ private lateinit var asyncCache: AsyncCache
+
+ private suspend fun awaitCachePut(
+ key: String,
+ timeoutMillis: Long = 1_000,
+ ): Boolean =
+ withTimeoutOrNull(timeoutMillis) {
+ while (asyncCache.get(key) == null) {
+ delay(5)
+ }
+ true
+ } ?: false
+
+ @Test
+ fun shouldCollapseDuplicateRequests() =
+ runBlocking {
+ val key = "dup"
+ val attempts = 20
+ val results = (1..attempts).map { async(Dispatchers.IO) { service.get(key) } }.awaitAll()
+ assertEquals(attempts, results.size)
+ val first = results.firstOrNull()
+ assertTrue(results.all { it == first })
+ }
+
+ @Test
+ fun shouldEvictAndRecompute() =
+ runBlocking {
+ val key = "evict-${System.nanoTime()}" // Use unique key for test isolation
+ val v1 = service.get(key)
+ // ReqShield stores cache asynchronously; wait until the cache write is observed.
+ assertTrue(awaitCachePut(key), "Timed out waiting for cache put for key=$key")
+ val evicted = service.evict(key)
+ val v2 = service.get(key)
+
+ assertTrue(evicted)
+ assertTrue(v1.isNotEmpty())
+ assertTrue(v2.isNotEmpty())
+ assertTrue(v1 != v2)
+ }
+
+ @Configuration
+ open class TestConfig {
+ @Bean
+ open fun asyncCache(): AsyncCache = InMemoryAsyncCache()
+
+ @Bean
+ open fun service(): TestService = TestService()
+ }
+
+ open class TestService {
+ val counter = AtomicInteger(0)
+
+ @ReqShieldCacheable(cacheName = "it", key = "#key", timeToLiveMillis = 10_000)
+ open suspend fun get(key: String): String = "value-" + counter.incrementAndGet()
+
+ @ReqShieldCacheEvict(cacheName = "it", key = "#key")
+ open suspend fun evict(key: String): Boolean = true
+ }
+}
diff --git a/core-spring-webflux-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/ReqShieldAspectRedisIntegrationTest.kt b/core-spring-webflux-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/ReqShieldAspectRedisIntegrationTest.kt
new file mode 100644
index 0000000..9518e40
--- /dev/null
+++ b/core-spring-webflux-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/ReqShieldAspectRedisIntegrationTest.kt
@@ -0,0 +1,174 @@
+/*
+ * Copyright 2024 LY Corporation
+ *
+ * LY Corporation licenses this file to you under the Apache License,
+ * version 2.0 (the "License"); you may not use this file except in compliance
+ * with the License. You may obtain a copy of the License at:
+ *
+ * https://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS, WITHOUT
+ * WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the
+ * License for the specific language governing permissions and limitations
+ * under the License.
+ */
+
+package com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.aspect
+
+import com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.annotation.ReqShieldCacheEvict
+import com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.annotation.ReqShieldCacheable
+import com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.cache.AsyncCache
+import com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.config.LibAutoConfiguration
+import com.linecorp.cse.reqshield.support.model.ReqShieldData
+import com.linecorp.cse.reqshield.support.redis.AbstractRedisTest
+import io.lettuce.core.RedisClient
+import io.lettuce.core.api.StatefulRedisConnection
+import io.lettuce.core.api.sync.RedisCommands
+import kotlinx.coroutines.Dispatchers
+import kotlinx.coroutines.async
+import kotlinx.coroutines.awaitAll
+import kotlinx.coroutines.delay
+import kotlinx.coroutines.runBlocking
+import kotlinx.coroutines.withTimeoutOrNull
+import org.junit.jupiter.api.Assertions.assertTrue
+import org.junit.jupiter.api.BeforeEach
+import org.junit.jupiter.api.Test
+import org.junit.jupiter.api.extension.ExtendWith
+import org.springframework.beans.factory.annotation.Autowired
+import org.springframework.beans.factory.annotation.Value
+import org.springframework.context.annotation.Bean
+import org.springframework.context.annotation.Configuration
+import org.springframework.test.context.ContextConfiguration
+import org.springframework.test.context.junit.jupiter.SpringExtension
+import java.util.concurrent.atomic.AtomicInteger
+
+@ExtendWith(SpringExtension::class)
+@ContextConfiguration(classes = [LibAutoConfiguration::class, ReqShieldAspectRedisIntegrationTest.TestConfig::class])
+class ReqShieldAspectRedisIntegrationTest : AbstractRedisTest() {
+ @Autowired
+ private lateinit var service: TestService
+
+ @Autowired
+ private lateinit var asyncCache: AsyncCache
+
+ @BeforeEach
+ fun resetCounter() {
+ service.resetCounter()
+ }
+
+ private suspend fun awaitCachePut(
+ key: String,
+ timeoutMillis: Long = 2_000,
+ ): Boolean =
+ withTimeoutOrNull(timeoutMillis) {
+ while (asyncCache.get(key) == null) {
+ delay(10)
+ }
+ true
+ } ?: false
+
+ @Test
+ fun shouldCollapseDuplicateRequestsWithRedis() =
+ runBlocking {
+ val key = "dup-redis-${System.nanoTime()}" // Use unique key for test isolation
+ val attempts = 20
+ val results = (1..attempts).map { async(Dispatchers.IO) { service.get(key) } }.awaitAll()
+
+ // Request collapsing core: callable should be invoked only once
+ assertTrue(
+ service.getRequestCount() == 1,
+ "Callable should be invoked only once. actual=${service.getRequestCount()}",
+ )
+
+ // All results should be valid (not null)
+ assertTrue(
+ results.size == attempts && results.all { it != null },
+ "Expected all results to be valid. results=$results",
+ )
+ }
+
+ @Test
+ fun shouldEvictAndRecomputeWithRedis() =
+ runBlocking {
+ val key = "evict-redis-${System.nanoTime()}" // Use unique key for test isolation
+ val v1 = service.get(key)
+ // ReqShield stores cache asynchronously; wait until the cache write is observed.
+ assertTrue(awaitCachePut(key), "Timed out waiting for cache put for key=$key")
+ val evicted = service.evict(key)
+ val v2 = service.get(key)
+ assertTrue(evicted, "Eviction should return true")
+ assertTrue(v1 != v2, "Values should differ after eviction: v1=$v1, v2=$v2")
+ }
+
+ @Configuration
+ open class TestConfig {
+ @Value("\${spring.redis.host}")
+ private lateinit var host: String
+
+ @Value("\${spring.redis.port}")
+ private var port: Int = 0
+
+ @Bean(destroyMethod = "shutdown")
+ open fun redisClient(): RedisClient = RedisClient.create("redis://$host:$port")
+
+ @Bean(destroyMethod = "close")
+ open fun redisConnection(redisClient: RedisClient): StatefulRedisConnection = redisClient.connect()
+
+ @Bean
+ open fun asyncCache(redisConnection: StatefulRedisConnection): AsyncCache {
+ val sync: RedisCommands = redisConnection.sync()
+ // Ensure clean DB state for tests running in CI
+ runCatching { sync.flushdb() }
+
+ return object : AsyncCache {
+ override suspend fun get(key: String): ReqShieldData? =
+ sync.get(key)?.let { ReqShieldData(value = it, timeToLiveMillis = 10_000) }
+
+ override suspend fun put(
+ key: String,
+ value: ReqShieldData,
+ timeToLiveMillis: Long,
+ ): Boolean {
+ sync.psetex(key, timeToLiveMillis, value.value ?: "")
+ return true
+ }
+
+ override suspend fun evict(key: String): Boolean = sync.del(key) > 0
+
+ override suspend fun globalLock(
+ key: String,
+ timeToLiveMillis: Long,
+ ): Boolean = sync.setnx("lock:$key", "1").also { if (it) sync.pexpire("lock:$key", timeToLiveMillis) }
+
+ override suspend fun globalUnLock(key: String): Boolean = sync.del("lock:$key") >= 0
+ }
+ }
+
+ @Bean
+ open fun service(): TestService = TestService()
+ }
+
+ open class TestService {
+ val counter = AtomicInteger(0)
+
+ open fun resetCounter() {
+ counter.set(0)
+ }
+
+ open fun getRequestCount(): Int = counter.get()
+
+ @ReqShieldCacheable(
+ cacheName = "it",
+ key = "#key",
+ timeToLiveMillis = 10_000,
+ // CI environments can be slow; give enough time for async cache put to be observed by waiters.
+ maxAttemptGetCache = 200,
+ lockTimeoutMillis = 10_000,
+ )
+ open suspend fun get(key: String): String = "value-" + counter.incrementAndGet()
+
+ @ReqShieldCacheEvict(cacheName = "it", key = "#key")
+ open suspend fun evict(key: String): Boolean = true
+ }
+}
diff --git a/core-spring-webflux-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/ReqShieldAspectTest.kt b/core-spring-webflux-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/ReqShieldAspectTest.kt
index dfc657d..223007d 100644
--- a/core-spring-webflux-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/ReqShieldAspectTest.kt
+++ b/core-spring-webflux-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/ReqShieldAspectTest.kt
@@ -50,7 +50,7 @@ private val log = LoggerFactory.getLogger(ReqShieldAspectTest::class.java)
@OptIn(ExperimentalCoroutinesApi::class)
class ReqShieldAspectTest : BaseReqShieldModuleSupportTest {
- private val asyncCache: AsyncCache = mockk()
+ private val asyncCache: AsyncCache = InMemoryAsyncCache()
private val joinPoint: ProceedingJoinPoint = mockk()
private val reqShieldAspect: ReqShieldAspect = spyk(ReqShieldAspect(asyncCache))
private val targetObject = spyk(TestBean())
@@ -80,9 +80,9 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest {
runTest {
// Mock the cache data using mockk
val reqShieldData = ReqShieldData(methodReturn, 1000)
- coEvery { asyncCache.get(any()) } returns reqShieldData
+ asyncCache.put(spelEvaluatedKey, reqShieldData, 1000)
coEvery { joinPoint.proceed() } coAnswers { targetObject.cacheableWithCustomKey(argument) }
- coEvery { reqShieldAspect.getTargetMethod(joinPoint) } returns
+ every { reqShieldAspect.getTargetMethod(joinPoint) } returns
TestBean::class
.functions
.find {
@@ -90,11 +90,13 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest {
}?.javaMethod!!
// Test the aroundTargetCacheable method
- val result = reqShieldAspect.aroundTargetCacheable(joinPoint)
+ val result = reqShieldAspect.aroundReqShieldCacheable(joinPoint)
assertEquals(result, reqShieldData.value)
assertTrue(reqShieldAspect.reqShieldMap.size == 1)
- assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$spelEvaluatedKey"])
+ val method = reqShieldAspect.getTargetMethod(joinPoint)
+ val expectedKey = "${method.declaringClass.name}.${method.name}-$cacheName-$spelEvaluatedKey"
+ assertNotNull(reqShieldAspect.reqShieldMap[expectedKey])
}
@Test
@@ -102,9 +104,9 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest {
runTest {
// Mock the cache data using mockk
val reqShieldData = ReqShieldData(methodReturn, 1000)
- coEvery { asyncCache.get(any()) } returns reqShieldData
+ asyncCache.put(spelEvaluatedKey, reqShieldData, 1000)
coEvery { joinPoint.proceed() } coAnswers { targetObject.cacheableWithCustomKey(argument) }
- coEvery { reqShieldAspect.getTargetMethod(joinPoint) } returns
+ every { reqShieldAspect.getTargetMethod(joinPoint) } returns
TestBean::class
.functions
.find {
@@ -114,46 +116,51 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest {
val jobs =
List(20) {
async {
- reqShieldAspect.aroundTargetCacheable(joinPoint)
+ reqShieldAspect.aroundReqShieldCacheable(joinPoint)
}
}
jobs.awaitAll()
assertTrue(reqShieldAspect.reqShieldMap.size == 1)
- assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$spelEvaluatedKey"])
+ val method = reqShieldAspect.getTargetMethod(joinPoint)
+ val expectedKey = "${method.declaringClass.name}.${method.name}-$cacheName-$spelEvaluatedKey"
+ assertNotNull(reqShieldAspect.reqShieldMap[expectedKey])
}
@Test
override fun verifyReqShieldCacheEviction() =
runTest {
- // Mock the cache data using mockk
val reqShieldData = ReqShieldData(methodReturn, 1000)
- coEvery { asyncCache.get(any()) } returns reqShieldData
- coEvery { joinPoint.proceed() } coAnswers { targetObject.cacheableWithCustomKey(argument) }
- coEvery { reqShieldAspect.getTargetMethod(joinPoint) } returns
+ // Use SpEL-based key to align with eviction method's key
+ every { reqShieldAspect.getTargetMethod(joinPoint) } returns
TestBean::class
.functions
.find {
- it.name == TestBean::cacheableWithDefaultKeyGenerator.name && it.parameters.size == 2
+ it.name == TestBean::cacheableWithCustomKey.name && it.parameters.size == 2
}?.javaMethod!!
+ val generatedKey = reqShieldAspect.getCacheableCacheKey(joinPoint)
+ asyncCache.put(generatedKey, reqShieldData, 1000)
+ coEvery { joinPoint.proceed() } coAnswers { targetObject.cacheableWithCustomKey(argument) }
- // Test the aroundTargetCacheable method
- val result = reqShieldAspect.aroundTargetCacheable(joinPoint)
+ // Test the aroundTargetCacheable method using the same SpEL key
+ val result = reqShieldAspect.aroundReqShieldCacheable(joinPoint)
assertEquals(reqShieldData.value, result)
- // Validate cache eviction
- coEvery { asyncCache.evict(any()) } returns true
- coEvery { reqShieldAspect.getTargetMethod(joinPoint) } returns
+ // Validate cache eviction using the eviction method (same SpEL key)
+ // real eviction call
+ every { reqShieldAspect.getTargetMethod(joinPoint) } returns
TestBean::class
.functions
.find {
it.name == TestBean::evict.name && it.parameters.size == 2
}?.javaMethod!!
- coEvery { joinPoint.proceed() } coAnswers { targetObject.evict(argument) }
+ // Mock proceed for eviction - the aspect proceeds to the original method after evicting cache
+ // proceedSmart() calls proceed(args) with continuation, so we need to mock that as well
+ coEvery { joinPoint.proceed(any>()) } coAnswers { targetObject.evict(argument) }
- val removeProductMono = reqShieldAspect.aroundTargetCacheable(joinPoint)
+ val removeProductMono = reqShieldAspect.aroundReqShieldCacheEvict(joinPoint)
assertTrue(removeProductMono as Boolean)
}
diff --git a/core-spring-webflux/build.gradle.kts b/core-spring-webflux/build.gradle.kts
index 7642a4f..c84dead 100644
--- a/core-spring-webflux/build.gradle.kts
+++ b/core-spring-webflux/build.gradle.kts
@@ -31,7 +31,9 @@ dependencies {
testImplementation(rootProject.libs.reactor)
testImplementation(rootProject.libs.reactor.test)
testImplementation(rootProject.libs.spring.context)
+ testImplementation(rootProject.libs.spring.test)
testImplementation(rootProject.libs.aspectj)
+ testImplementation(rootProject.libs.lettuce)
}
tasks.withType().configureEach {
diff --git a/core-spring-webflux/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/annotation/ReqShieldCacheable.kt b/core-spring-webflux/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/annotation/ReqShieldCacheable.kt
index c093fc7..5d7ba67 100644
--- a/core-spring-webflux/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/annotation/ReqShieldCacheable.kt
+++ b/core-spring-webflux/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/annotation/ReqShieldCacheable.kt
@@ -33,4 +33,10 @@ annotation class ReqShieldCacheable(
val maxAttemptGetCache: Int = MAX_ATTEMPT_GET_CACHE,
val timeToLiveMillis: Long = 10 * 60 * 1000,
val reqShieldWorkMode: ReqShieldWorkMode = ReqShieldWorkMode.CREATE_AND_UPDATE_CACHE,
+ val nullHandling: NullHandling = NullHandling.EMIT_EMPTY,
)
+
+enum class NullHandling {
+ EMIT_EMPTY,
+ ERROR,
+}
diff --git a/core-spring-webflux/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/aspect/ReqShieldAspect.kt b/core-spring-webflux/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/aspect/ReqShieldAspect.kt
index fa65a83..429e0b5 100644
--- a/core-spring-webflux/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/aspect/ReqShieldAspect.kt
+++ b/core-spring-webflux/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/aspect/ReqShieldAspect.kt
@@ -44,7 +44,7 @@ import java.util.concurrent.ConcurrentHashMap
@Aspect
@Component
-class ReqShieldAspect(
+open class ReqShieldAspect(
private val asyncCache: AsyncCache,
) : BeanFactoryAware {
private lateinit var beanFactory: BeanFactory
@@ -60,14 +60,28 @@ class ReqShieldAspect(
val reqShield = getOrCreateReqShield(joinPoint)
val cacheKey = getCacheableCacheKey(joinPoint)
- return reqShield
- .getAndSetReqShieldData(
- cacheKey,
- {
- joinPoint.proceed() as Mono
- },
- annotation.timeToLiveMillis,
- ).mapNotNull { it.value }
+ val resultMono =
+ reqShield
+ .getAndSetReqShieldData(
+ cacheKey,
+ {
+ joinPoint.proceed() as Mono
+ },
+ annotation.timeToLiveMillis,
+ ).map { it.value }
+
+ return when (annotation.nullHandling) {
+ com.linecorp.cse.reqshield.spring.webflux.annotation.NullHandling.EMIT_EMPTY ->
+ resultMono.flatMap { Mono.justOrEmpty(it) }
+ com.linecorp.cse.reqshield.spring.webflux.annotation.NullHandling.ERROR ->
+ resultMono.flatMap { value ->
+ if (value == null) {
+ Mono.error(IllegalStateException("ReqShieldCacheable returned null for key=$cacheKey"))
+ } else {
+ Mono.just(value)
+ }
+ }
+ }
}
@Around("@annotation(com.linecorp.cse.reqshield.spring.webflux.annotation.ReqShieldCacheEvict)")
@@ -104,7 +118,7 @@ class ReqShieldAspect(
return getCacheKeyOrDefault(annotation.key, annotation.keyGenerator, joinPoint)
}
- internal fun getTargetMethod(joinPoint: ProceedingJoinPoint): Method = (joinPoint.signature as MethodSignature).method
+ internal open fun getTargetMethod(joinPoint: ProceedingJoinPoint): Method = (joinPoint.signature as MethodSignature).method
private fun getCacheKeyOrDefault(
annotationCacheKey: String,
@@ -124,7 +138,15 @@ class ReqShieldAspect(
keyGenerator.generate(joinPoint.target, method, joinPoint.args).toString()
}
- require(!key.isNullOrBlank()) { "Null key returned for cache method : $method" }
+ require(!key.isNullOrBlank()) {
+ "Null/blank key for @ReqShieldCacheable method=${method.declaringClass.name}.${method.name} " +
+ "args=${joinPoint.args.joinToString(prefix = "[", postfix = "]") {
+ it?.let {
+ arg ->
+ "${arg::class.simpleName}@${arg.hashCode().toString(16)}"
+ } ?: "null"
+ }}"
+ }
return key
}
@@ -156,6 +178,7 @@ class ReqShieldAspect(
decisionForUpdate = annotation.decisionForUpdate,
maxAttemptGetCache = annotation.maxAttemptGetCache,
reqShieldWorkMode = annotation.reqShieldWorkMode,
+ scheduler = beanFactory.getBean("reqShieldScheduler", reactor.core.scheduler.Scheduler::class.java),
)
return ReqShield(reqShieldConfiguration)
@@ -166,7 +189,9 @@ class ReqShieldAspect(
cacheKeyGenerator: String,
) {
if (cacheKey.isNotBlank() && cacheKeyGenerator.isNotBlank()) {
- throw IllegalArgumentException("The key and keyGenerator attributes are mutually exclusive.")
+ throw IllegalArgumentException(
+ "The key and keyGenerator attributes are mutually exclusive: key='$cacheKey', keyGenerator='$cacheKeyGenerator'",
+ )
}
}
@@ -180,8 +205,11 @@ class ReqShieldAspect(
}
}
- private fun generateReqShieldKey(joinPoint: ProceedingJoinPoint): String =
- "${getCacheableAnnotation(joinPoint).cacheName}-${getCacheableCacheKey(joinPoint)}"
+ private fun generateReqShieldKey(joinPoint: ProceedingJoinPoint): String {
+ val method = getTargetMethod(joinPoint)
+ return "${method.declaringClass.name}.${method.name}-" +
+ "${getCacheableAnnotation(joinPoint).cacheName}-${getCacheableCacheKey(joinPoint)}"
+ }
override fun setBeanFactory(beanFactory: BeanFactory) {
this.beanFactory = beanFactory
diff --git a/core-spring-webflux/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/config/LibAutoConfiguration.kt b/core-spring-webflux/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/config/LibAutoConfiguration.kt
index 863cb78..2a46c27 100644
--- a/core-spring-webflux/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/config/LibAutoConfiguration.kt
+++ b/core-spring-webflux/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/config/LibAutoConfiguration.kt
@@ -16,11 +16,18 @@
package com.linecorp.cse.reqshield.spring.webflux.config
-import org.springframework.context.annotation.ComponentScan
+import com.linecorp.cse.reqshield.spring.webflux.aspect.ReqShieldAspect
+import org.springframework.context.annotation.Bean
import org.springframework.context.annotation.Configuration
import org.springframework.context.annotation.EnableAspectJAutoProxy
+import org.springframework.context.annotation.Import
+import reactor.core.scheduler.Scheduler
+import reactor.core.scheduler.Schedulers
@Configuration
-@EnableAspectJAutoProxy
-@ComponentScan(basePackages = ["com.linecorp.cse"])
-open class LibAutoConfiguration
+@EnableAspectJAutoProxy(proxyTargetClass = true)
+@Import(ReqShieldAspect::class)
+open class LibAutoConfiguration {
+ @Bean
+ open fun reqShieldScheduler(): Scheduler = Schedulers.boundedElastic()
+}
diff --git a/core-spring-webflux/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/aspect/InMemoryAsyncCache.kt b/core-spring-webflux/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/aspect/InMemoryAsyncCache.kt
new file mode 100644
index 0000000..15a5631
--- /dev/null
+++ b/core-spring-webflux/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/aspect/InMemoryAsyncCache.kt
@@ -0,0 +1,63 @@
+/*
+ * Copyright 2024 LY Corporation
+ *
+ * LY Corporation licenses this file to you under the Apache License,
+ * version 2.0 (the "License"); you may not use this file except in compliance
+ * with the License. You may obtain a copy of the License at:
+ *
+ * https://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS, WITHOUT
+ * WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the
+ * License for the specific language governing permissions and limitations
+ * under the License.
+ */
+
+package com.linecorp.cse.reqshield.spring.webflux.aspect
+
+import com.linecorp.cse.reqshield.spring.webflux.cache.AsyncCache
+import com.linecorp.cse.reqshield.support.model.ReqShieldData
+import reactor.core.publisher.Mono
+import java.util.concurrent.ConcurrentHashMap
+import java.util.concurrent.Semaphore
+
+class InMemoryAsyncCache : AsyncCache {
+ private data class Entry(val data: ReqShieldData, val expiresAt: Long)
+
+ private val store = ConcurrentHashMap>()
+ private val locks = ConcurrentHashMap()
+
+ override fun get(key: String): Mono?> =
+ Mono.fromCallable {
+ val now = System.currentTimeMillis()
+ store[key]?.let { e -> if (now <= e.expiresAt) e.data else null }
+ }
+
+ override fun put(
+ key: String,
+ value: ReqShieldData,
+ timeToLiveMillis: Long,
+ ): Mono =
+ Mono.fromCallable {
+ val expiresAt = System.currentTimeMillis() + timeToLiveMillis
+ store[key] = Entry(value, expiresAt)
+ true
+ }
+
+ override fun evict(key: String): Mono = Mono.fromCallable { store.remove(key) != null }
+
+ override fun globalLock(
+ key: String,
+ timeToLiveMillis: Long,
+ ): Mono =
+ Mono.fromCallable {
+ locks.computeIfAbsent(key) { Semaphore(1) }.tryAcquire()
+ }
+
+ override fun globalUnLock(key: String): Mono =
+ Mono.fromCallable {
+ locks[key]?.release()
+ true
+ }
+}
diff --git a/core-spring-webflux/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/aspect/ReqShieldAspectIntegrationTest.kt b/core-spring-webflux/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/aspect/ReqShieldAspectIntegrationTest.kt
new file mode 100644
index 0000000..577ac0c
--- /dev/null
+++ b/core-spring-webflux/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/aspect/ReqShieldAspectIntegrationTest.kt
@@ -0,0 +1,91 @@
+/*
+ * Copyright 2024 LY Corporation
+ *
+ * LY Corporation licenses this file to you under the Apache License,
+ * version 2.0 (the "License"); you may not use this file except in compliance
+ * with the License. You may obtain a copy of the License at:
+ *
+ * https://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS, WITHOUT
+ * WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the
+ * License for the specific language governing permissions and limitations
+ * under the License.
+ */
+
+package com.linecorp.cse.reqshield.spring.webflux.aspect
+
+import com.linecorp.cse.reqshield.spring.webflux.annotation.ReqShieldCacheEvict
+import com.linecorp.cse.reqshield.spring.webflux.annotation.ReqShieldCacheable
+import com.linecorp.cse.reqshield.spring.webflux.cache.AsyncCache
+import com.linecorp.cse.reqshield.spring.webflux.config.LibAutoConfiguration
+import org.junit.jupiter.api.Assertions.assertEquals
+import org.junit.jupiter.api.Assertions.assertTrue
+import org.junit.jupiter.api.Test
+import org.junit.jupiter.api.extension.ExtendWith
+import org.springframework.beans.factory.annotation.Autowired
+import org.springframework.context.annotation.Bean
+import org.springframework.context.annotation.Configuration
+import org.springframework.test.context.ContextConfiguration
+import org.springframework.test.context.junit.jupiter.SpringExtension
+import reactor.core.publisher.Flux
+import reactor.core.publisher.Mono
+import reactor.core.scheduler.Schedulers
+import java.util.concurrent.atomic.AtomicInteger
+
+@ExtendWith(SpringExtension::class)
+@ContextConfiguration(classes = [LibAutoConfiguration::class, ReqShieldAspectIntegrationTest.TestConfig::class])
+class ReqShieldAspectIntegrationTest {
+ @Autowired
+ private lateinit var service: TestService
+
+ @Test
+ fun shouldCollapseDuplicateRequests() {
+ val key = "dup"
+ val attempts = 20
+
+ val result =
+ Flux
+ .range(1, attempts)
+ .flatMap { service.get(key).subscribeOn(Schedulers.boundedElastic()) }
+ .collectList()
+ .block()
+
+ assertEquals(attempts, result?.size)
+ val first = result?.firstOrNull()
+ assertTrue(result?.all { it == first } == true)
+ }
+
+ @Test
+ fun shouldEvictAndRecompute() {
+ val key = "evict"
+ val v1 = service.get(key).block()
+ val evicted = service.evict(key).block()
+ val v2 = service.get(key).block()
+
+ assertTrue(evicted == true)
+ assertTrue(v1 != null)
+ assertTrue(v2 != null)
+ assertTrue(v1 != v2)
+ }
+
+ @Configuration
+ open class TestConfig {
+ @Bean
+ open fun asyncCache(): AsyncCache = InMemoryAsyncCache()
+
+ @Bean
+ open fun service(): TestService = TestService()
+ }
+
+ open class TestService {
+ val counter = AtomicInteger(0)
+
+ @ReqShieldCacheable(cacheName = "it", key = "#key", timeToLiveMillis = 10_000)
+ open fun get(key: String): Mono = Mono.fromCallable { "value-" + counter.incrementAndGet() }
+
+ @ReqShieldCacheEvict(cacheName = "it", key = "#key")
+ open fun evict(key: String): Mono = Mono.just(true)
+ }
+}
diff --git a/core-spring-webflux/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/aspect/ReqShieldAspectRedisIntegrationTest.kt b/core-spring-webflux/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/aspect/ReqShieldAspectRedisIntegrationTest.kt
new file mode 100644
index 0000000..eb69443
--- /dev/null
+++ b/core-spring-webflux/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/aspect/ReqShieldAspectRedisIntegrationTest.kt
@@ -0,0 +1,173 @@
+/*
+ * Copyright 2024 LY Corporation
+ *
+ * LY Corporation licenses this file to you under the Apache License,
+ * version 2.0 (the "License"); you may not use this file except in compliance
+ * with the License. You may obtain a copy of the License at:
+ *
+ * https://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS, WITHOUT
+ * WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the
+ * License for the specific language governing permissions and limitations
+ * under the License.
+ */
+
+package com.linecorp.cse.reqshield.spring.webflux.aspect
+
+import com.linecorp.cse.reqshield.spring.webflux.annotation.ReqShieldCacheEvict
+import com.linecorp.cse.reqshield.spring.webflux.annotation.ReqShieldCacheable
+import com.linecorp.cse.reqshield.spring.webflux.cache.AsyncCache
+import com.linecorp.cse.reqshield.spring.webflux.config.LibAutoConfiguration
+import com.linecorp.cse.reqshield.support.model.ReqShieldData
+import com.linecorp.cse.reqshield.support.redis.AbstractRedisTest
+import io.lettuce.core.RedisClient
+import io.lettuce.core.api.StatefulRedisConnection
+import io.lettuce.core.api.sync.RedisCommands
+import org.junit.jupiter.api.Assertions.assertTrue
+import org.junit.jupiter.api.BeforeEach
+import org.junit.jupiter.api.Test
+import org.junit.jupiter.api.extension.ExtendWith
+import org.springframework.beans.factory.annotation.Autowired
+import org.springframework.beans.factory.annotation.Value
+import org.springframework.context.annotation.Bean
+import org.springframework.context.annotation.Configuration
+import org.springframework.test.context.ContextConfiguration
+import org.springframework.test.context.junit.jupiter.SpringExtension
+import reactor.core.publisher.Flux
+import reactor.core.publisher.Mono
+import reactor.core.scheduler.Schedulers
+import java.util.concurrent.atomic.AtomicInteger
+
+@ExtendWith(SpringExtension::class)
+@ContextConfiguration(classes = [LibAutoConfiguration::class, ReqShieldAspectRedisIntegrationTest.TestConfig::class])
+class ReqShieldAspectRedisIntegrationTest : AbstractRedisTest() {
+ @Autowired
+ private lateinit var service: TestService
+
+ @Autowired
+ private lateinit var asyncCache: AsyncCache
+
+ @BeforeEach
+ fun resetCounter() {
+ service.resetCounter()
+ }
+
+ private fun awaitCachePut(
+ key: String,
+ timeoutMillis: Long = 2_000,
+ ): Boolean {
+ val start = System.currentTimeMillis()
+ while (System.currentTimeMillis() - start < timeoutMillis) {
+ if (asyncCache.get(key).block() != null) {
+ return true
+ }
+ Thread.sleep(10)
+ }
+ return false
+ }
+
+ @Test
+ fun shouldCollapseDuplicateRequestsWithRedis() {
+ val key = "dup-redis-${System.nanoTime()}" // Use unique key for test isolation
+ val attempts = 20
+
+ val result =
+ Flux
+ .range(1, attempts)
+ .flatMap { service.get(key).subscribeOn(Schedulers.boundedElastic()) }
+ .collectList()
+ .block()!!
+
+ // Request collapsing core: callable should be invoked only once
+ assertTrue(
+ service.getRequestCount() == 1,
+ "Callable should be invoked only once. actual=${service.getRequestCount()}",
+ )
+
+ // All results should be valid (not null)
+ assertTrue(result.size == attempts && result.all { it != null }, "Expected all results to be valid")
+ }
+
+ @Test
+ fun shouldEvictAndRecomputeWithRedis() {
+ val key = "evict-redis-${System.nanoTime()}" // Use unique key for test isolation
+ val v1 = service.get(key).block()
+ // ReqShield stores cache asynchronously; wait until the cache write is observed.
+ assertTrue(awaitCachePut(key), "Timed out waiting for cache put for key=$key")
+ val evicted = service.evict(key).block()
+ val v2 = service.get(key).block()
+
+ assertTrue(evicted == true, "Eviction should return true")
+ assertTrue(v1 != null && v2 != null && v1 != v2, "Values should differ after eviction: v1=$v1, v2=$v2")
+ }
+
+ @Configuration
+ open class TestConfig {
+ @Value("\${spring.redis.host}")
+ private lateinit var host: String
+
+ @Value("\${spring.redis.port}")
+ private var port: Int = 0
+
+ @Bean(destroyMethod = "shutdown")
+ open fun redisClient(): RedisClient = RedisClient.create("redis://$host:$port")
+
+ @Bean(destroyMethod = "close")
+ open fun redisConnection(redisClient: RedisClient): StatefulRedisConnection = redisClient.connect()
+
+ @Bean
+ open fun asyncCache(redisConnection: StatefulRedisConnection): AsyncCache {
+ val sync: RedisCommands = redisConnection.sync()
+ // Ensure clean DB state for tests running in CI
+ runCatching { sync.flushdb() }
+
+ return object : AsyncCache {
+ override fun get(key: String): Mono?> =
+ Mono.fromCallable {
+ sync.get(key)?.let { ReqShieldData(value = it, timeToLiveMillis = 10_000) }
+ }
+
+ override fun put(
+ key: String,
+ value: ReqShieldData,
+ timeToLiveMillis: Long,
+ ): Mono =
+ Mono.fromCallable {
+ sync.psetex(key, timeToLiveMillis, value.value ?: "")
+ true
+ }
+
+ override fun evict(key: String): Mono = Mono.fromCallable { sync.del(key) > 0 }
+
+ override fun globalLock(
+ key: String,
+ timeToLiveMillis: Long,
+ ): Mono =
+ Mono.fromCallable { sync.setnx("lock:$key", "1").also { if (it) sync.pexpire("lock:$key", timeToLiveMillis) } }
+
+ override fun globalUnLock(key: String): Mono = Mono.fromCallable { sync.del("lock:$key") >= 0 }
+ }
+ }
+
+ @Bean
+ open fun service(): TestService = TestService()
+ }
+
+ open class TestService {
+ val counter = AtomicInteger(0)
+
+ open fun resetCounter() {
+ counter.set(0)
+ }
+
+ open fun getRequestCount(): Int = counter.get()
+
+ @ReqShieldCacheable(cacheName = "it", key = "#key", timeToLiveMillis = 10_000)
+ open fun get(key: String): Mono = Mono.fromCallable { "value-" + counter.incrementAndGet() }
+
+ @ReqShieldCacheEvict(cacheName = "it", key = "#key")
+ open fun evict(key: String): Mono = Mono.just(true)
+ }
+}
diff --git a/core-spring-webflux/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/aspect/ReqShieldAspectTest.kt b/core-spring-webflux/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/aspect/ReqShieldAspectTest.kt
index afdfab0..8591fd8 100644
--- a/core-spring-webflux/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/aspect/ReqShieldAspectTest.kt
+++ b/core-spring-webflux/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/aspect/ReqShieldAspectTest.kt
@@ -42,7 +42,7 @@ import kotlin.test.assertEquals
import kotlin.test.assertTrue
class ReqShieldAspectTest : BaseReqShieldModuleSupportTest {
- private val asyncCache: AsyncCache = mockk()
+ private val asyncCache: AsyncCache = InMemoryAsyncCache()
private val joinPoint = mockk()
private val reqShieldAspect = spyk(ReqShieldAspect(asyncCache))
private val targetObject = spyk(TestBean())
@@ -63,13 +63,17 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest {
every { joinPoint.target } returns targetObject
reqShieldAspect.setBeanFactory(beanFactory)
+ // Provide scheduler bean expected by aspect configuration
+ every {
+ beanFactory.getBean("reqShieldScheduler", reactor.core.scheduler.Scheduler::class.java)
+ } returns Schedulers.boundedElastic()
}
@Test
override fun verifyReqShieldCacheCreation() {
- // Mock the cache data using mockk
val reqShieldData = ReqShieldData(methodReturn, 1000)
- every { asyncCache.get(any()) } returns Mono.just(reqShieldData)
+ // pre-populate cache
+ asyncCache.put(spelEvaluatedKey, reqShieldData, 1000).block()
every { joinPoint.proceed() } answers { targetObject.cacheableWithCustomKey(argument) }
every { reqShieldAspect.getTargetMethod(joinPoint) } returns
ReflectionUtils.findMethod(
@@ -87,15 +91,16 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest {
.assertNext { value ->
assertEquals(reqShieldData.value, value)
Assertions.assertTrue(reqShieldAspect.reqShieldMap.size == 1)
- Assertions.assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$spelEvaluatedKey"])
+ val method = reqShieldAspect.getTargetMethod(joinPoint)
+ val expectedKey = "${method.declaringClass.name}.${method.name}-$cacheName-$spelEvaluatedKey"
+ Assertions.assertNotNull(reqShieldAspect.reqShieldMap[expectedKey])
}.verifyComplete()
}
@Test
override fun reqShieldObjectShouldBeCreatedOnce() {
- // Mock the cache data using mockk
val reqShieldData = ReqShieldData(methodReturn, 1000)
- every { asyncCache.get(any()) } returns Mono.just(reqShieldData)
+ asyncCache.put(spelEvaluatedKey, reqShieldData, 1000).block()
every { joinPoint.proceed() } answers { targetObject.cacheableWithCustomKey(argument) }
every { reqShieldAspect.getTargetMethod(joinPoint) } returns
ReflectionUtils.findMethod(
@@ -117,16 +122,16 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest {
.create(flux)
.assertNext { productList ->
Assertions.assertTrue(reqShieldAspect.reqShieldMap.size == 1)
- println(reqShieldAspect.reqShieldMap.keys().toList())
- Assertions.assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$spelEvaluatedKey"])
+ val method = reqShieldAspect.getTargetMethod(joinPoint)
+ val expectedKey = "${method.declaringClass.name}.${method.name}-$cacheName-$spelEvaluatedKey"
+ Assertions.assertNotNull(reqShieldAspect.reqShieldMap[expectedKey])
}.verifyComplete()
}
@Test
override fun verifyReqShieldCacheEviction() {
- // Mock the cache data using mockk
val reqShieldData = ReqShieldData(methodReturn, 1000)
- every { asyncCache.get(any()) } returns Mono.just(reqShieldData)
+ asyncCache.put("${SimpleKeyGenerator.generateKey(arrayOf(argument))}", reqShieldData, 1000).block()
every { reqShieldAspect.getTargetMethod(joinPoint) } returns
ReflectionUtils.findMethod(
TestBean::class.java,
@@ -145,7 +150,7 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest {
}.verifyComplete()
// Validate cache eviction
- every { asyncCache.evict(any()) } returns Mono.just(true)
+ // real eviction call
every { reqShieldAspect.getTargetMethod(joinPoint) } returns
ReflectionUtils.findMethod(
TestBean::class.java,
diff --git a/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/aspect/ReqShieldAspect.kt b/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/aspect/ReqShieldAspect.kt
index bf2cae7..9136337 100644
--- a/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/aspect/ReqShieldAspect.kt
+++ b/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/aspect/ReqShieldAspect.kt
@@ -151,7 +151,15 @@ class ReqShieldAspect(
keyGenerator.generate(joinPoint.target, method, joinPoint.args).toString()
}
- require(!key.isNullOrBlank()) { "Null key returned for cache method : $method" }
+ require(!key.isNullOrBlank()) {
+ "Null/blank key for @ReqShieldCacheable method=${method.declaringClass.name}.${method.name} " +
+ "args=${joinPoint.args.joinToString(prefix = "[", postfix = "]") {
+ it?.let {
+ arg ->
+ "${arg::class.simpleName}@${arg.hashCode().toString(16)}"
+ } ?: "null"
+ }}"
+ }
return key
}
@@ -161,7 +169,9 @@ class ReqShieldAspect(
cacheKeyGenerator: String,
) {
if (cacheKey.isNotBlank() && cacheKeyGenerator.isNotBlank()) {
- throw IllegalArgumentException("The key and keyGenerator attributes are mutually exclusive.")
+ throw IllegalArgumentException(
+ "The key and keyGenerator attributes are mutually exclusive: key='$cacheKey', keyGenerator='$cacheKeyGenerator'",
+ )
}
}
@@ -175,8 +185,11 @@ class ReqShieldAspect(
}
}
- private fun generateReqShieldKey(joinPoint: ProceedingJoinPoint): String =
- "${getCacheableAnnotation(joinPoint).cacheName}-${getCacheableCacheKey(joinPoint)}"
+ private fun generateReqShieldKey(joinPoint: ProceedingJoinPoint): String {
+ val method = getTargetMethod(joinPoint)
+ return "${method.declaringClass.name}.${method.name}-" +
+ "${getCacheableAnnotation(joinPoint).cacheName}-${getCacheableCacheKey(joinPoint)}"
+ }
override fun setBeanFactory(beanFactory: BeanFactory) {
this.beanFactory = beanFactory
diff --git a/core-spring/src/test/kotlin/aspect/ReqShieldAspectTest.kt b/core-spring/src/test/kotlin/aspect/ReqShieldAspectTest.kt
index 619bd08..38b246c 100644
--- a/core-spring/src/test/kotlin/aspect/ReqShieldAspectTest.kt
+++ b/core-spring/src/test/kotlin/aspect/ReqShieldAspectTest.kt
@@ -90,7 +90,9 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest {
// then
assertEquals(reqShieldData.value, result)
assertTrue(reqShieldAspect.reqShieldMap.size == 1)
- assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$spelEvaluatedKey"])
+ val method = reqShieldAspect.getTargetMethod(joinPoint)
+ val expectedKey = "${method.declaringClass.name}.${method.name}-$cacheName-$spelEvaluatedKey"
+ assertNotNull(reqShieldAspect.reqShieldMap[expectedKey])
}
}
@@ -117,7 +119,9 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest {
Awaitility.await().atMost(Duration.ofMillis(BaseReqShieldTest.AWAIT_TIMEOUT)).untilAsserted {
// then
assertTrue(reqShieldAspect.reqShieldMap.size == 1)
- assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$spelEvaluatedKey"])
+ val method = reqShieldAspect.getTargetMethod(joinPoint)
+ val expectedKey = "${method.declaringClass.name}.${method.name}-$cacheName-$spelEvaluatedKey"
+ assertNotNull(reqShieldAspect.reqShieldMap[expectedKey])
}
}
diff --git a/core/src/main/kotlin/com/linecorp/cse/reqshield/KeyLocalLock.kt b/core/src/main/kotlin/com/linecorp/cse/reqshield/KeyLocalLock.kt
index 59153dd..7abf230 100644
--- a/core/src/main/kotlin/com/linecorp/cse/reqshield/KeyLocalLock.kt
+++ b/core/src/main/kotlin/com/linecorp/cse/reqshield/KeyLocalLock.kt
@@ -24,11 +24,30 @@ import java.util.concurrent.Executors
import java.util.concurrent.ScheduledExecutorService
import java.util.concurrent.Semaphore
import java.util.concurrent.TimeUnit
+import java.util.concurrent.atomic.AtomicBoolean
private val log = LoggerFactory.getLogger(KeyLocalLock::class.java)
class KeyLocalLock(private val lockTimeoutMillis: Long) : KeyLock {
- private data class LockInfo(val semaphore: Semaphore, val createdAt: Long)
+ /**
+ * Internal lock state holder.
+ * Using class instead of data class to allow mutable expiresAt for atomic updates.
+ */
+ private class LockInfo(
+ val semaphore: Semaphore,
+ /**
+ * Expiration timestamp in milliseconds.
+ * @Volatile ensures visibility across threads when updated inside compute() and read by monitor.
+ */
+ @Volatile var expiresAt: Long,
+ /**
+ * Tracks whether the lock is currently held.
+ * Uses AtomicBoolean with CAS operations to prevent over-release
+ * when multiple threads race to release the same lock (e.g., tryLock expiration
+ * check vs unLock, or monitor cleanup vs unLock).
+ */
+ val isHeld: AtomicBoolean = AtomicBoolean(false),
+ )
companion object {
// Global lockMap shared by all instances - CRITICAL FIX for request collapsing
@@ -38,9 +57,6 @@ class KeyLocalLock(private val lockTimeoutMillis: Long) : KeyLock {
@Volatile
private var sharedScheduler: ScheduledExecutorService? = null
- // Track active instances
- private val instances = ConcurrentHashMap.newKeySet()
-
// Thread-safe lazy initialization
private fun getOrCreateScheduler(): ScheduledExecutorService {
return sharedScheduler ?: synchronized(this) {
@@ -61,14 +77,39 @@ class KeyLocalLock(private val lockTimeoutMillis: Long) : KeyLock {
}
private fun startMonitoring(scheduler: ScheduledExecutorService) {
- // Batch cleanup for all instances (10ms → 1000ms)
+ // Single cleanup task operating on the global lockMap
scheduler.scheduleWithFixedDelay({
try {
- instances.forEach { instance ->
- instance.cleanupExpiredLocks()
+ val now = System.currentTimeMillis()
+ val before = lockMap.size
+ // Remove expired locks using compute() for atomic check-and-remove.
+ // This prevents TOCTOU race condition where removeIf's lambda returns true
+ // but the actual removal happens after a new lock is acquired.
+ // compute() guarantees atomic execution per key, so cleanup and tryLock
+ // are mutually exclusive for the same key.
+ lockMap.keys.forEach { key ->
+ lockMap.compute(key) { _, lockInfo ->
+ if (lockInfo == null) return@compute null
+
+ if (now > lockInfo.expiresAt) {
+ // Expired lock: force release regardless of isHeld state.
+ // This handles the case where unlock() was missed due to exception.
+ // CAS ensures safe release (no-op if already released).
+ if (lockInfo.isHeld.compareAndSet(true, false)) {
+ lockInfo.semaphore.release()
+ }
+ null // Atomic removal
+ } else {
+ lockInfo // Keep the entry
+ }
+ }
+ }
+ val after = lockMap.size
+ if (log.isTraceEnabled && before > after) {
+ log.trace("Cleaned up {} expired locks, {} remaining", before - after, after)
}
} catch (e: Exception) {
- log.error("Error in shared lock lifecycle monitoring: {}", e.message)
+ log.error("Error in shared lock lifecycle monitoring: {}", e.message, e)
}
}, 0, LOCK_MONITOR_INTERVAL_MILLIS, TimeUnit.MILLISECONDS)
}
@@ -94,35 +135,49 @@ class KeyLocalLock(private val lockTimeoutMillis: Long) : KeyLock {
}
init {
- // Register instance and initialize scheduler
- instances.add(this)
+ // Initialize scheduler on first instance creation
getOrCreateScheduler()
}
- // Internal cleanup method (called by shared scheduler)
- internal fun cleanupExpiredLocks() {
- val now = System.currentTimeMillis()
- val expiredCount = lockMap.size
- lockMap.entries.removeIf { now - it.value.createdAt > lockTimeoutMillis }
- val remainingCount = lockMap.size
-
- if (log.isTraceEnabled && expiredCount > remainingCount) {
- log.trace(
- "Cleaned up {} expired locks, {} remaining",
- expiredCount - remainingCount,
- remainingCount,
- )
- }
- }
+ // Internal cleanup method no longer needed per-instance with single shared cleanup
override fun tryLock(
key: String,
lockType: LockType,
): Boolean {
val completeKey = "${key}_${lockType.name}"
- val lockInfo = lockMap.computeIfAbsent(completeKey) { LockInfo(Semaphore(1), nowToEpochTime()) }
+ val now = nowToEpochTime()
+ val result = AtomicBoolean(false)
+
+ // Use compute() for atomic lock acquisition.
+ // This ensures mutual exclusion with cleanup - they cannot race on the same key.
+ lockMap.compute(completeKey) { _, existing ->
+ if (existing != null) {
+ // Force-release expired locks to allow reacquisition.
+ // Use CAS to prevent race condition with concurrent unLock().
+ // Without CAS, if unLock() executes between isHeld.get() and release(),
+ // both threads would call release(), causing over-release (permits > 1).
+ if (now > existing.expiresAt && existing.isHeld.compareAndSet(true, false)) {
+ existing.semaphore.release()
+ }
- return lockInfo.semaphore.tryAcquire()
+ // Existing entry: try to acquire semaphore
+ if (existing.semaphore.tryAcquire()) {
+ existing.isHeld.set(true)
+ existing.expiresAt = now + lockTimeoutMillis
+ result.set(true)
+ }
+ existing
+ } else {
+ // New entry: create and acquire
+ val newLock = LockInfo(Semaphore(1), now + lockTimeoutMillis)
+ newLock.semaphore.tryAcquire() // Always succeeds for new semaphore
+ newLock.isHeld.set(true)
+ result.set(true)
+ newLock
+ }
+ }
+ return result.get()
}
override fun unLock(
@@ -130,19 +185,20 @@ class KeyLocalLock(private val lockTimeoutMillis: Long) : KeyLock {
lockType: LockType,
): Boolean {
val completeKey = "${key}_${lockType.name}"
- val lockInfo = lockMap[completeKey]
- lockInfo?.let {
- it.semaphore.release()
- lockMap.remove(completeKey)
+ val lockInfo = lockMap[completeKey] ?: return false
+
+ // Use CAS to prevent over-release: only release if we actually hold the lock
+ return if (lockInfo.isHeld.compareAndSet(true, false)) {
+ lockInfo.semaphore.release()
+ true
+ } else {
+ log.debug("Attempted to unlock key '{}' that is not held", completeKey)
+ false
}
- return true
}
fun shutdown() {
- // Deregister instance
- instances.remove(this)
-
// Shared scheduler is managed globally, no individual shutdown needed
- log.debug("KeyLocalLock instance deregistered from shared monitoring")
+ log.debug("KeyLocalLock instance shutdown (scheduler managed globally)")
}
}
diff --git a/core/src/main/kotlin/com/linecorp/cse/reqshield/ReqShield.kt b/core/src/main/kotlin/com/linecorp/cse/reqshield/ReqShield.kt
index 20a9f2a..7ae3136 100644
--- a/core/src/main/kotlin/com/linecorp/cse/reqshield/ReqShield.kt
+++ b/core/src/main/kotlin/com/linecorp/cse/reqshield/ReqShield.kt
@@ -19,12 +19,11 @@ package com.linecorp.cse.reqshield
import com.linecorp.cse.reqshield.config.ReqShieldConfiguration
import com.linecorp.cse.reqshield.config.ReqShieldWorkMode
import com.linecorp.cse.reqshield.support.constant.ConfigValues.GET_CACHE_INTERVAL_MILLIS
-import com.linecorp.cse.reqshield.support.constant.ConfigValues.MAX_ATTEMPT_SET_CACHE
-import com.linecorp.cse.reqshield.support.constant.ConfigValues.SET_CACHE_RETRY_INTERVAL_MILLIS
import com.linecorp.cse.reqshield.support.exception.ClientException
import com.linecorp.cse.reqshield.support.exception.code.ErrorCode
import com.linecorp.cse.reqshield.support.model.ReqShieldData
import com.linecorp.cse.reqshield.support.utils.decideToUpdateCache
+import org.slf4j.LoggerFactory
import java.util.concurrent.Callable
import java.util.concurrent.CompletableFuture
import java.util.concurrent.ScheduledExecutorService
@@ -32,6 +31,8 @@ import java.util.concurrent.ScheduledFuture
import java.util.concurrent.TimeUnit
import java.util.concurrent.atomic.AtomicInteger
+private val log = LoggerFactory.getLogger(ReqShield::class.java)
+
class ReqShield(
private val reqShieldConfig: ReqShieldConfiguration,
) {
@@ -163,26 +164,45 @@ class ReqShield(
callable: Callable,
key: String,
) {
- fun schedule(): ScheduledFuture<*> =
- executor.schedule({
- if (!future.isDone) {
+ val scheduled: ScheduledFuture<*> =
+ executor.scheduleAtFixedRate({
+ try {
+ // Early exit if future is already completed to avoid unnecessary work
+ if (future.isDone) {
+ return@scheduleAtFixedRate
+ }
+
val funcResult = executeGetCacheFunction(cacheGetter, key)
if (funcResult != null) {
+ // Use CAS-like complete to handle race condition safely
+ // If another thread already completed, this is a no-op
future.complete(funcResult.value)
- } else if (counter.incrementAndGet() >= reqShieldConfig.maxAttemptGetCache) {
- future.complete(
- executeCallable({ callable.call() }, false),
- )
+ return@scheduleAtFixedRate
}
+
+ // Increment first, then check - ensures atomic decision making
+ val attempts = counter.incrementAndGet()
+ if (attempts >= reqShieldConfig.maxAttemptGetCache && !future.isDone) {
+ // Use complete() which handles concurrent completion safely
+ // If another thread completed between our check and this call, it's ignored
+ future.complete(executeCallable({ callable.call() }, false))
+ }
+ } catch (e: Exception) {
+ // Handle exception to prevent scheduleAtFixedRate from stopping
+ // Fallback to callable to ensure service availability
+ log.error("Error in scheduled cache getter for key '{}', falling back to callable", key, e)
if (!future.isDone) {
- schedule() // Schedule the next execution
+ try {
+ future.complete(executeCallable({ callable.call() }, false))
+ } catch (fallbackException: Exception) {
+ log.error("Fallback callable also failed for key '{}'", key, fallbackException)
+ future.completeExceptionally(fallbackException)
+ }
}
}
- }, GET_CACHE_INTERVAL_MILLIS, TimeUnit.MILLISECONDS)
-
- val scheduleFuture = schedule()
+ }, GET_CACHE_INTERVAL_MILLIS, GET_CACHE_INTERVAL_MILLIS, TimeUnit.MILLISECONDS)
- future.whenComplete { _, _ -> scheduleFuture.cancel(false) }
+ future.whenComplete { _, _ -> scheduled.cancel(false) }
}
private fun executeGetCacheFunction(
@@ -207,15 +227,10 @@ class ReqShield(
throw ClientException(ErrorCode.SET_CACHE_ERROR, originErrorMessage = e.message)
} finally {
if (shouldAttemptUnlock(lockType)) {
- var unlockSuccess = false
- var retryCount = 0
- while (!unlockSuccess && retryCount < MAX_ATTEMPT_SET_CACHE) {
- if (reqShieldConfig.keyLock.unLock(key, lockType)) {
- unlockSuccess = true
- } else {
- retryCount++
- Thread.sleep(SET_CACHE_RETRY_INTERVAL_MILLIS)
- }
+ // No retry needed: false means lock already released or expired (not an error)
+ val unlocked = reqShieldConfig.keyLock.unLock(key, lockType)
+ if (!unlocked) {
+ log.debug("Lock already released or expired for key '{}'", key)
}
}
}
diff --git a/core/src/test/kotlin/com/linecorp/cse/reqshield/KeyGlobalLockTest.kt b/core/src/test/kotlin/com/linecorp/cse/reqshield/KeyGlobalLockTest.kt
index 9e9d443..7c70f48 100644
--- a/core/src/test/kotlin/com/linecorp/cse/reqshield/KeyGlobalLockTest.kt
+++ b/core/src/test/kotlin/com/linecorp/cse/reqshield/KeyGlobalLockTest.kt
@@ -40,11 +40,16 @@ class KeyGlobalLockTest :
@BeforeEach
fun init() {
- val redisUrl = "redis://localhost:6379" // testContainer url
+ val host = AbstractRedisTest.redisHost
+ val port = AbstractRedisTest.redisPort
+ val redisUrl = "redis://$host:$port"
val redisClient = RedisClient.create(redisUrl)
val connection = redisClient.connect()
redisCommands = connection.sync()
+ // Clean up all keys from previous tests for proper test isolation
+ redisCommands.flushdb()
+
globalLockFunc = { key, timeToLiveMillis ->
redisCommands.setnx(key, key)
}
diff --git a/core/src/test/kotlin/com/linecorp/cse/reqshield/KeyLocalLockTest.kt b/core/src/test/kotlin/com/linecorp/cse/reqshield/KeyLocalLockTest.kt
index 237e405..9f835c9 100644
--- a/core/src/test/kotlin/com/linecorp/cse/reqshield/KeyLocalLockTest.kt
+++ b/core/src/test/kotlin/com/linecorp/cse/reqshield/KeyLocalLockTest.kt
@@ -20,6 +20,7 @@ import com.linecorp.cse.reqshield.support.BaseKeyLockTest
import com.linecorp.cse.reqshield.support.BaseReqShieldTest.Companion.AWAIT_TIMEOUT
import org.awaitility.Awaitility.await
import org.junit.jupiter.api.Assertions.assertEquals
+import org.junit.jupiter.api.Assertions.assertFalse
import org.junit.jupiter.api.Assertions.assertTrue
import org.junit.jupiter.api.Test
import java.lang.management.ManagementFactory
@@ -160,20 +161,22 @@ class KeyLocalLockTest : BaseKeyLockTest {
@Test
fun testLockCleanupEfficiency() {
- // Given: KeyLocalLock instance
- val keyLock = KeyLocalLock(1000L) // Expires after 1 second
+ // Given: KeyLocalLock instance with timeout that allows cleanup to run
+ // LOCK_MONITOR_INTERVAL_MILLIS is 1000ms, so we need timeout > interval
+ val lockTimeout = 1500L
+ val keyLock = KeyLocalLock(lockTimeout)
val key = "testKey"
val lockType = LockType.CREATE
- // When: Acquire lock and wait for expiration
+ // When: Acquire lock and wait for expiration + cleanup interval
assertTrue(keyLock.tryLock(key, lockType))
// Then: Cleanup should work efficiently
- // Previously executed excessively at 10ms intervals
- Thread.sleep(1200L) // Expiration time + buffer
+ // Wait for: lockTimeout + cleanup interval (1000ms) + buffer
+ Thread.sleep(lockTimeout + 1500L)
// Expired locks should be cleaned up, allowing new lock acquisition
- await().atMost(Duration.ofSeconds(2)).untilAsserted {
+ await().atMost(Duration.ofSeconds(3)).untilAsserted {
assertTrue(keyLock.tryLock(key, lockType))
keyLock.unLock(key, lockType)
}
@@ -194,10 +197,10 @@ class KeyLocalLockTest : BaseKeyLockTest {
// Then - Instance2 should not be able to acquire the same lock
val lock2Result = instance2.tryLock(key, lockType)
-
+
assertTrue(lock1Result)
assertTrue(!lock2Result, "Instance2 should not acquire lock held by Instance1")
-
+
// Cleanup
instance1.unLock(key, lockType)
instance1.shutdown()
@@ -208,7 +211,7 @@ class KeyLocalLockTest : BaseKeyLockTest {
fun `should maintain request collapsing across multiple instances`() {
// Given
val instance1 = KeyLocalLock(lockTimeoutMillis)
- val instance2 = KeyLocalLock(lockTimeoutMillis)
+ val instance2 = KeyLocalLock(lockTimeoutMillis)
val instance3 = KeyLocalLock(lockTimeoutMillis)
val key = "collapsing-key"
val lockType = LockType.CREATE
@@ -220,11 +223,12 @@ class KeyLocalLockTest : BaseKeyLockTest {
// When - Multiple instances try to acquire the same lock concurrently
repeat(3) { index ->
executor.submit {
- val instance = when (index) {
- 0 -> instance1
- 1 -> instance2
- else -> instance3
- }
+ val instance =
+ when (index) {
+ 0 -> instance1
+ 1 -> instance2
+ else -> instance3
+ }
attemptCount.incrementAndGet()
if (instance.tryLock(key, lockType)) {
successCount.incrementAndGet()
@@ -241,7 +245,7 @@ class KeyLocalLockTest : BaseKeyLockTest {
// Then - Only one should succeed in acquiring the lock
assertEquals(3, attemptCount.get())
assertEquals(1, successCount.get(), "Only one instance should acquire the lock")
-
+
// Cleanup
instance1.shutdown()
instance2.shutdown()
@@ -249,27 +253,90 @@ class KeyLocalLockTest : BaseKeyLockTest {
}
@Test
- fun `should allow different instances to unlock the same key`() {
+ fun `should allow different instances to unlock the same key via global lockMap`() {
// Given
val instance1 = KeyLocalLock(lockTimeoutMillis)
val instance2 = KeyLocalLock(lockTimeoutMillis)
val key = "unlock-shared-key"
val lockType = LockType.CREATE
- // When - Instance1 acquires lock, Instance2 unlocks
+ // When - Instance1 acquires lock, Instance2 can also unlock (global lockMap shared)
assertTrue(instance1.tryLock(key, lockType))
- instance2.unLock(key, lockType) // Should work even from different instance
-
+ // Instance2 can unlock because isHeld state is global
+ assertTrue(instance2.unLock(key, lockType), "Global unlock should succeed from any instance")
+
// Then - New lock acquisition should succeed
val newLockResult = instance2.tryLock(key, lockType)
assertTrue(newLockResult, "Should be able to acquire lock after global unlock")
-
+
// Cleanup
instance2.unLock(key, lockType)
instance1.shutdown()
instance2.shutdown()
}
+ @Test
+ fun `should not over-release semaphore on multiple unlock calls`() {
+ // Given
+ val keyLock = KeyLocalLock(lockTimeoutMillis)
+ val key = "over-release-test"
+ val lockType = LockType.CREATE
+
+ // When - Acquire lock
+ assertTrue(keyLock.tryLock(key, lockType))
+
+ // Then - First unlock should succeed
+ assertTrue(keyLock.unLock(key, lockType), "First unlock should succeed")
+
+ // Second unlock should return false (lock not held)
+ assertFalse(keyLock.unLock(key, lockType), "Second unlock should fail (over-release prevention)")
+
+ // Verify semaphore is not over-released: can acquire once, not twice
+ assertTrue(keyLock.tryLock(key, lockType), "Should acquire lock after proper unlock")
+ assertFalse(keyLock.tryLock(key, lockType), "Should not acquire lock twice (semaphore intact)")
+
+ // Cleanup
+ keyLock.unLock(key, lockType)
+ keyLock.shutdown()
+ }
+
+ @Test
+ fun `should prevent concurrent lock acquisition after over-release attempt`() {
+ // Given
+ val keyLock = KeyLocalLock(lockTimeoutMillis)
+ val key = "concurrent-over-release-test"
+ val lockType = LockType.CREATE
+ val executor = Executors.newFixedThreadPool(10)
+ val successfulAcquisitions = AtomicInteger(0)
+ val latch = CountDownLatch(10)
+
+ // Simulate over-release attempt
+ assertTrue(keyLock.tryLock(key, lockType))
+ keyLock.unLock(key, lockType)
+ // Multiple unlock attempts should all return false (not over-release)
+ repeat(5) { assertFalse(keyLock.unLock(key, lockType)) }
+
+ // When - Try to acquire lock concurrently
+ repeat(10) {
+ executor.submit {
+ if (keyLock.tryLock(key, lockType)) {
+ successfulAcquisitions.incrementAndGet()
+ }
+ latch.countDown()
+ }
+ }
+
+ latch.await(5, TimeUnit.SECONDS)
+ executor.shutdown()
+
+ // Then - Only ONE thread should succeed (semaphore not corrupted by over-release)
+ assertEquals(1, successfulAcquisitions.get(), "Only one thread should acquire the lock")
+
+ // Cleanup
+ keyLock.unLock(key, lockType)
+ keyLock.shutdown()
+ }
+
@Test
fun `should handle concurrent operations from multiple instances safely`() {
// Given - 5 instances operating on 10 different keys concurrently
@@ -305,19 +372,197 @@ class KeyLocalLockTest : BaseKeyLockTest {
// Then - Verify thread safety and concurrent operations handling
assertEquals(0, errors.get(), "No errors should occur during concurrent operations")
-
+
// Due to sequential nature of ThreadPool(10) and brief work duration (10ms),
// multiple operations can succeed on the same key at different times
- assertTrue(operations.get() >= 10,
- "At least one operation per key should succeed (minimum 10)")
- assertTrue(operations.get() <= 50,
- "No more operations than total attempts should succeed (maximum 50)")
-
+ assertTrue(
+ operations.get() >= 10,
+ "At least one operation per key should succeed (minimum 10)",
+ )
+ assertTrue(
+ operations.get() <= 50,
+ "No more operations than total attempts should succeed (maximum 50)",
+ )
+
println("Successful operations: ${operations.get()}/50 total attempts")
-
+
// Cleanup
instances.forEach { it.shutdown() }
}
+ @Test
+ fun `should not remove lock that was just acquired during cleanup window`() {
+ // This test verifies that compute() based cleanup and tryLock are mutually exclusive.
+ // With compute(), cleanup and acquisition cannot race on the same key because
+ // compute() provides per-key atomic execution.
+
+ // Given: Lock with timeout matching cleanup interval to maximize cleanup opportunities
+ val lockTimeout = 500L
+ val keyLock = KeyLocalLock(lockTimeout)
+ val key = "race-condition-test"
+ val lockType = LockType.CREATE
+ val errors = AtomicInteger(0)
+ val successfulCycles = AtomicInteger(0)
+
+ // When: Sequentially acquire, let expire, release, and reacquire
+ // This validates that compute() atomicity prevents race conditions
+ repeat(10) {
+ // Acquire lock
+ assertTrue(keyLock.tryLock(key, lockType), "Should acquire lock")
+
+ // Hold until expiration
+ Thread.sleep(lockTimeout + 200)
+
+ // Release
+ keyLock.unLock(key, lockType)
+
+ // Immediately reacquire - compute() ensures this doesn't race with cleanup
+ val reacquired = keyLock.tryLock(key, lockType)
+ if (reacquired) {
+ // Verify lock exclusivity - second acquire must fail
+ if (keyLock.tryLock(key, lockType)) {
+ // This indicates lock was incorrectly removed during acquisition
+ errors.incrementAndGet()
+ keyLock.unLock(key, lockType)
+ }
+ successfulCycles.incrementAndGet()
+ keyLock.unLock(key, lockType)
+ }
+ }
+
+ // Then: No errors should occur due to compute() atomicity
+ assertEquals(0, errors.get(), "No race condition errors should occur")
+ assertTrue(successfulCycles.get() >= 5, "Most reacquisitions should succeed")
+
+ keyLock.shutdown()
+ }
+
+ @Test
+ fun `should verify compute atomicity prevents TOCTOU race condition during concurrent cleanup and acquisition`() {
+ // Given: Lock with timeout to trigger cleanup
+ // Using compute() for both cleanup and tryLock ensures mutual exclusion per key.
+ val lockTimeout = 500L
+ val keyLock = KeyLocalLock(lockTimeout)
+ val lockType = LockType.CREATE
+ val executor = Executors.newFixedThreadPool(5)
+ val successfulCycles = AtomicInteger(0)
+ val lockRemovedWhileHeld = AtomicInteger(0)
+ val iterations = 10
+ val latch = CountDownLatch(iterations)
+
+ // When: Concurrently acquire, let expire, release, and re-acquire on different keys
+ // compute() guarantees each operation is atomic per key
+ repeat(iterations) { i ->
+ val key = "toctou-key-$i"
+ executor.submit {
+ try {
+ // Acquire lock
+ if (keyLock.tryLock(key, lockType)) {
+ // Hold past expiration to trigger cleanup consideration
+ Thread.sleep(lockTimeout + 200)
+
+ // Release and immediately re-acquire
+ keyLock.unLock(key, lockType)
+
+ // With compute(), this operation is atomic with respect to cleanup
+ val reacquired = keyLock.tryLock(key, lockType)
+ if (reacquired) {
+ // Verify lock exclusivity
+ if (keyLock.tryLock(key, lockType)) {
+ // This should never happen - compute() ensures atomicity
+ lockRemovedWhileHeld.incrementAndGet()
+ keyLock.unLock(key, lockType)
+ }
+ successfulCycles.incrementAndGet()
+ keyLock.unLock(key, lockType)
+ }
+ }
+ } finally {
+ latch.countDown()
+ }
+ }
+ }
+
+ latch.await(30, TimeUnit.SECONDS)
+ executor.shutdown()
+ executor.awaitTermination(5, TimeUnit.SECONDS)
+
+ // Then: No lock corruption due to compute() atomicity
+ println("Successful cycles: ${successfulCycles.get()}, Lock removed while held: ${lockRemovedWhileHeld.get()}")
+ assertEquals(0, lockRemovedWhileHeld.get(), "No lock should be removed while still held")
+
+ keyLock.shutdown()
+ }
+
+ @Test
+ fun `should cleanup expired lock even when unlock is never called`() {
+ // This test verifies that cleanup properly handles the scenario where unlock() is missed
+ // (e.g., due to exception). Previously, isHeld=true locks were never cleaned up,
+ // causing memory leaks.
+
+ // Given: Short-lived lock
+ val lockTimeout = 500L
+ val keyLock = KeyLocalLock(lockTimeout)
+ val key = "missed-unlock-test"
+ val lockType = LockType.CREATE
+
+ // When: Acquire lock but never unlock (simulating exception scenario)
+ assertTrue(keyLock.tryLock(key, lockType), "Should acquire lock")
+ // DO NOT call unlock - simulating exception scenario
+
+ // Wait for expiration + cleanup interval + buffer
+ // Cleanup runs every 1000ms (LOCK_MONITOR_INTERVAL_MILLIS)
+ Thread.sleep(lockTimeout + 1500L)
+
+ // Then: Cleanup should have force-released and removed the expired lock
+ // A new lock acquisition should succeed
+ await().atMost(Duration.ofSeconds(3)).untilAsserted {
+ assertTrue(
+ keyLock.tryLock(key, lockType),
+ "Should acquire lock after cleanup removed expired held lock",
+ )
+ }
+
+ // Verify lock is working normally
+ assertFalse(keyLock.tryLock(key, lockType), "Second acquire should fail (lock is held)")
+ assertTrue(keyLock.unLock(key, lockType), "Unlock should succeed")
+
+ keyLock.shutdown()
+ }
+
+ @Test
+ fun `should cleanup multiple expired held locks without memory leak`() {
+ // This test verifies that cleanup prevents memory leaks when many locks expire
+ // without being unlocked.
+
+ // Given: Short-lived locks
+ val lockTimeout = 300L
+ val keyLock = KeyLocalLock(lockTimeout)
+ val lockType = LockType.CREATE
+ val keyCount = 20
+
+ // When: Acquire many locks but never unlock them
+ repeat(keyCount) { i ->
+ assertTrue(keyLock.tryLock("leak-test-$i", lockType), "Should acquire lock $i")
+ }
+
+ // Wait for all locks to expire and be cleaned up
+ // Cleanup interval is 1000ms, so we need to wait for expiration + cleanup cycle
+ Thread.sleep(lockTimeout + 1500L)
+
+ // Then: All expired locks should be cleaned up, allowing reacquisition
+ await().atMost(Duration.ofSeconds(5)).untilAsserted {
+ repeat(keyCount) { i ->
+ assertTrue(
+ keyLock.tryLock("leak-test-$i", lockType),
+ "Should acquire lock $i after cleanup",
+ )
+ keyLock.unLock("leak-test-$i", lockType)
+ }
+ }
+
+ keyLock.shutdown()
+ }
+
private fun doWork() = Thread.sleep(1000)
}
diff --git a/core/src/test/kotlin/com/linecorp/cse/reqshield/ReqShieldTest.kt b/core/src/test/kotlin/com/linecorp/cse/reqshield/ReqShieldTest.kt
index e6536ca..54d994b 100644
--- a/core/src/test/kotlin/com/linecorp/cse/reqshield/ReqShieldTest.kt
+++ b/core/src/test/kotlin/com/linecorp/cse/reqshield/ReqShieldTest.kt
@@ -482,4 +482,88 @@ class ReqShieldTest : BaseReqShieldTest {
verify { cacheSetter.invoke(key, reqShieldData, 1000L) }
verify { keyLock.unLock(any(), any()) }
}
+
+ @Test
+ fun `should complete future with callable result when cache getter throws exception in scheduled task`() {
+ // Given: First cache check returns null (triggers scheduleTask path),
+ // then subsequent calls in scheduleTask throw exception
+ var callCount = 0
+ every { cacheGetter.invoke(key) } answers {
+ callCount++
+ if (callCount == 1) {
+ null // First call: cache miss, triggers handleLockForCacheCreation
+ } else {
+ throw Exception("cache connection error") // Subsequent calls: exception in scheduleTask
+ }
+ }
+ every { keyLock.tryLock(key, LockType.CREATE) } returns false
+
+ // When: getAndSetReqShieldData is called
+ // The scheduled task will hit exception, should fallback to callable
+ val result = reqShield.getAndSetReqShieldData(key, callable, timeToLiveMillis)
+
+ // Then: Should return callable result (fallback), not hang
+ await().atMost(Duration.ofMillis(AWAIT_TIMEOUT)).untilAsserted {
+ assertNotNull(result)
+ assertEquals(value, result.value)
+ verify { callable.call() }
+ }
+ }
+
+ @Test
+ fun `should not hang when scheduled task encounters repeated cache getter exceptions`() {
+ // Given: First cache check returns null (triggers scheduleTask path),
+ // then subsequent calls always fail
+ var callCount = 0
+ every { cacheGetter.invoke(key) } answers {
+ callCount++
+ if (callCount == 1) {
+ null // First call: cache miss, triggers handleLockForCacheCreation
+ } else {
+ throw Exception("persistent cache error $callCount") // Subsequent calls: exception in scheduleTask
+ }
+ }
+ every { keyLock.tryLock(key, LockType.CREATE) } returns false
+
+ // When: getAndSetReqShieldData is called
+ val startTime = System.currentTimeMillis()
+ val result = reqShield.getAndSetReqShieldData(key, callable, timeToLiveMillis)
+ val elapsed = System.currentTimeMillis() - startTime
+
+ // Then: Should complete within reasonable time (not hang), using callable fallback
+ assertNotNull(result)
+ assertEquals(value, result.value)
+ // Should complete within 5 seconds (way less than infinite hang)
+ assertTrue(elapsed < 5000, "Should not hang - completed in ${elapsed}ms")
+ verify { callable.call() }
+ }
+
+ @Test
+ fun `should propagate exception when both cache getter and callable fail`() {
+ // Given: First cache check returns null (triggers scheduleTask path),
+ // then cache getter fails and callable also fails
+ var callCount = 0
+ every { cacheGetter.invoke(key) } answers {
+ callCount++
+ if (callCount == 1) {
+ null // First call: cache miss
+ } else {
+ throw Exception("cache error") // Subsequent calls: exception in scheduleTask
+ }
+ }
+ every { keyLock.tryLock(key, LockType.CREATE) } returns false
+ every { callable.call() } throws Exception("callable also failed")
+
+ // When/Then: Should propagate the callable exception (via completeExceptionally)
+ val exception =
+ assertThrows {
+ reqShield.getAndSetReqShieldData(key, callable, timeToLiveMillis)
+ }
+
+ // The exception should be from the fallback callable failure
+ assertTrue(
+ exception is ClientException || exception.cause is ClientException,
+ "Should propagate ClientException from failed callable",
+ )
+ }
}
diff --git a/libs.versions.toml b/libs.versions.toml
index 80d2491..48a0376 100644
--- a/libs.versions.toml
+++ b/libs.versions.toml
@@ -1,6 +1,7 @@
[versions]
kotlin = "1.8.20"
kotlinCoroutine = "1.7.3"
+kotlinCoroutineSpring = "1.6.4"
reactor = "3.4.23"
spring = "5.3.30"
springBoot3 = "3.3.1"
@@ -22,6 +23,7 @@ kotlin-coroutine = { module = "org.jetbrains.kotlinx:kotlinx-coroutines-core", v
kotlin-coroutine-jvm = { module = "org.jetbrains.kotlinx:kotlinx-coroutines-core-jvm", version.ref = "kotlinCoroutine" }
kotlin-coroutine-jdk8 = { module = "org.jetbrains.kotlinx:kotlinx-coroutines-jdk8", version.ref = "kotlinCoroutine" }
kotlin-coroutine-reactor = { module = "org.jetbrains.kotlinx:kotlinx-coroutines-reactor", version.ref = "kotlinCoroutine" }
+kotlin-coroutine-spring = { module = "org.jetbrains.kotlinx:kotlinx-coroutines-spring", version.ref = "kotlinCoroutineSpring" }
kotlin-coroutine-reactive = { module = "org.jetbrains.kotlinx:kotlinx-coroutines-reactive", version.ref = "kotlinCoroutine" }
# log
diff --git a/req-shield-spring-boot3-example/src/test/kotlin/com/linecorp/cse/reqshield/spring3/mvc/example/service/CacheAnnotationTest.kt b/req-shield-spring-boot3-example/src/test/kotlin/com/linecorp/cse/reqshield/spring3/mvc/example/service/CacheAnnotationTest.kt
index 5f90aa7..d5404c3 100644
--- a/req-shield-spring-boot3-example/src/test/kotlin/com/linecorp/cse/reqshield/spring3/mvc/example/service/CacheAnnotationTest.kt
+++ b/req-shield-spring-boot3-example/src/test/kotlin/com/linecorp/cse/reqshield/spring3/mvc/example/service/CacheAnnotationTest.kt
@@ -1,7 +1,6 @@
package com.linecorp.cse.reqshield.spring3.mvc.example.service
import com.linecorp.cse.reqshield.spring.cache.ReqShieldCache
-import com.linecorp.cse.reqshield.support.BaseReqShieldTest
import com.linecorp.cse.reqshield.support.model.Product
import com.linecorp.cse.reqshield.support.redis.AbstractRedisTest
import org.awaitility.Awaitility.await
@@ -14,7 +13,6 @@ import org.junit.jupiter.api.extension.ExtendWith
import org.springframework.beans.factory.annotation.Autowired
import org.springframework.boot.test.context.SpringBootTest
import org.springframework.test.context.junit.jupiter.SpringExtension
-import java.time.Duration
import java.util.UUID
import java.util.concurrent.Executors
import java.util.concurrent.TimeUnit
@@ -48,7 +46,7 @@ class CacheAnnotationTest : AbstractRedisTest() {
executorService.shutdown()
executorService.awaitTermination(3000, TimeUnit.SECONDS)
- await().atMost(Duration.ofMillis(BaseReqShieldTest.AWAIT_TIMEOUT)).untilAsserted {
+ await().atMost(5, TimeUnit.SECONDS).untilAsserted {
assertEquals(1, sampleService.getRequestCount())
assertNotNull(reqShieldCache.get("product-$testProductId"))
}
@@ -69,7 +67,7 @@ class CacheAnnotationTest : AbstractRedisTest() {
executorService.shutdown()
executorService.awaitTermination(3000, TimeUnit.SECONDS)
- await().atMost(Duration.ofMillis(BaseReqShieldTest.AWAIT_TIMEOUT)).untilAsserted {
+ await().atMost(5, TimeUnit.SECONDS).untilAsserted {
assertEquals(100, sampleService.getRequestCount())
assertNotNull(reqShieldCache.get("product-$testProductId"))
}
@@ -92,7 +90,7 @@ class CacheAnnotationTest : AbstractRedisTest() {
Thread.sleep(1000)
- await().atMost(Duration.ofMillis(BaseReqShieldTest.AWAIT_TIMEOUT)).untilAsserted {
+ await().atMost(5, TimeUnit.SECONDS).untilAsserted {
assertEquals(1, sampleService.getRequestCount())
assertNotNull(reqShieldCache.get("product-$testProductId"))
}
@@ -105,7 +103,7 @@ class CacheAnnotationTest : AbstractRedisTest() {
sampleService.getProduct(testProductId)
await().atMost(5, TimeUnit.SECONDS).until {
- reqShieldCache.get("product-$testProductId") != null
+ runCatching { reqShieldCache.get("product-$testProductId") != null }.getOrDefault(false)
}
assertNotNull(reqShieldCache.get("product-$testProductId"))
@@ -115,7 +113,7 @@ class CacheAnnotationTest : AbstractRedisTest() {
// then
await().atMost(5, TimeUnit.SECONDS).until {
- reqShieldCache.get("product-$testProductId") == null
+ runCatching { reqShieldCache.get("product-$testProductId") == null }.getOrDefault(false)
}
assertNull(reqShieldCache.get("product-$testProductId"))
diff --git a/req-shield-spring-boot3-webflux-example/src/test/kotlin/com/linecorp/cse/reqshield/spring3/webflux/example/IntegrationSmokeTest.kt b/req-shield-spring-boot3-webflux-example/src/test/kotlin/com/linecorp/cse/reqshield/spring3/webflux/example/IntegrationSmokeTest.kt
new file mode 100644
index 0000000..70f5c2c
--- /dev/null
+++ b/req-shield-spring-boot3-webflux-example/src/test/kotlin/com/linecorp/cse/reqshield/spring3/webflux/example/IntegrationSmokeTest.kt
@@ -0,0 +1,15 @@
+package com.linecorp.cse.reqshield.spring3.webflux.example
+
+import org.junit.jupiter.api.Test
+import org.junit.jupiter.api.extension.ExtendWith
+import org.springframework.boot.test.context.SpringBootTest
+import org.springframework.test.context.junit.jupiter.SpringExtension
+
+@SpringBootTest(webEnvironment = SpringBootTest.WebEnvironment.RANDOM_PORT)
+@ExtendWith(SpringExtension::class)
+class IntegrationSmokeTest {
+ @Test
+ fun contextLoads() {
+ // just ensure context starts with Testcontainers
+ }
+}
diff --git a/req-shield-spring-boot3-webflux-example/src/test/kotlin/com/linecorp/cse/reqshield/spring3/webflux/example/service/CacheAnnotationTest.kt b/req-shield-spring-boot3-webflux-example/src/test/kotlin/com/linecorp/cse/reqshield/spring3/webflux/example/service/CacheAnnotationTest.kt
index 6421006..75ce088 100644
--- a/req-shield-spring-boot3-webflux-example/src/test/kotlin/com/linecorp/cse/reqshield/spring3/webflux/example/service/CacheAnnotationTest.kt
+++ b/req-shield-spring-boot3-webflux-example/src/test/kotlin/com/linecorp/cse/reqshield/spring3/webflux/example/service/CacheAnnotationTest.kt
@@ -73,9 +73,15 @@ class CacheAnnotationTest : AbstractRedisTest() {
StepVerifier
.create(flux)
.assertNext { productList ->
- assertEquals(19, sampleService.getRequestCount(), "Request count should be 19")
+ assertEquals(20, productList.size, "Result count should be 20")
}.verifyComplete()
+ // Wait for all doFinally callbacks to complete (async cache storage may cause timing issues)
+ await().atMost(5, TimeUnit.SECONDS).until {
+ sampleService.getRequestCount() == 20
+ }
+ assertEquals(20, sampleService.getRequestCount(), "Request count should be 20")
+
await().atMost(5, TimeUnit.SECONDS).until {
asyncCache.get("product-$testProductId").block() != null
}
diff --git a/req-shield-spring-webflux-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/example/service/CacheAnnotationTest.kt b/req-shield-spring-webflux-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/example/service/CacheAnnotationTest.kt
index d825672..0af2d08 100644
--- a/req-shield-spring-webflux-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/example/service/CacheAnnotationTest.kt
+++ b/req-shield-spring-webflux-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/example/service/CacheAnnotationTest.kt
@@ -90,9 +90,15 @@ class CacheAnnotationTest : AbstractRedisTest() {
StepVerifier
.create(flux)
.assertNext { productList ->
- assertEquals(19, sampleService.getRequestCount(), "Request count should be 19")
+ assertEquals(20, productList.size, "Result count should be 20")
}.verifyComplete()
+ // Wait for all doFinally callbacks to complete (async cache storage may cause timing issues)
+ await().atMost(5, TimeUnit.SECONDS).until {
+ sampleService.getRequestCount() == 20
+ }
+ assertEquals(20, sampleService.getRequestCount(), "Request count should be 20")
+
await().atMost(5, TimeUnit.SECONDS).until {
asyncCache.get("product-$testProductId").block() != null
}
diff --git a/support/build.gradle.kts b/support/build.gradle.kts
index 936dab9..110db86 100644
--- a/support/build.gradle.kts
+++ b/support/build.gradle.kts
@@ -23,8 +23,11 @@ plugins {
}
dependencies {
- testFixturesImplementation(rootProject.libs.testcontainers)
- testFixturesImplementation(rootProject.libs.junit.jupiter.testcontainers)
+ testFixturesImplementation(rootProject.libs.junit)
+ // Expose Testcontainers to consumers of test fixtures because RedisContainer
+ // leaks GenericContainer type in its API (instance property)
+ testFixturesApi(rootProject.libs.testcontainers)
+ testFixturesApi(rootProject.libs.junit.jupiter.testcontainers)
testFixturesImplementation(rootProject.libs.spring.context)
testFixturesImplementation(rootProject.libs.spring.test)
testFixturesImplementation(rootProject.libs.spring.boot.test)
diff --git a/support/src/main/kotlin/com/linecorp/cse/reqshield/support/constant/ConfigValues.kt b/support/src/main/kotlin/com/linecorp/cse/reqshield/support/constant/ConfigValues.kt
index feea90e..bb713ec 100644
--- a/support/src/main/kotlin/com/linecorp/cse/reqshield/support/constant/ConfigValues.kt
+++ b/support/src/main/kotlin/com/linecorp/cse/reqshield/support/constant/ConfigValues.kt
@@ -22,7 +22,7 @@ object ConfigValues {
const val LOCK_MONITOR_INTERVAL_MILLIS = 1000L
- const val MAX_ATTEMPT_GET_CACHE = 50
+ const val MAX_ATTEMPT_GET_CACHE = 60
const val GET_CACHE_INTERVAL_MILLIS = 50L
const val MAX_ATTEMPT_SET_CACHE = 3
diff --git a/support/src/testFixtures/kotlin/com/linecorp/cse/reqshield/support/redis/AbstractRedisTest.kt b/support/src/testFixtures/kotlin/com/linecorp/cse/reqshield/support/redis/AbstractRedisTest.kt
index 50ca6e8..fa6548b 100644
--- a/support/src/testFixtures/kotlin/com/linecorp/cse/reqshield/support/redis/AbstractRedisTest.kt
+++ b/support/src/testFixtures/kotlin/com/linecorp/cse/reqshield/support/redis/AbstractRedisTest.kt
@@ -15,28 +15,61 @@
*/
package com.linecorp.cse.reqshield.support.redis
-
import org.springframework.context.ApplicationContextInitializer
import org.springframework.context.ConfigurableApplicationContext
import org.springframework.core.env.MapPropertySource
import org.springframework.test.context.ContextConfiguration
-import org.testcontainers.junit.jupiter.Container
-import org.testcontainers.junit.jupiter.Testcontainers
-@Testcontainers
@ContextConfiguration(initializers = [AbstractRedisTest.Companion.Initializer::class])
abstract class AbstractRedisTest {
companion object {
- @Container
- private val redisContainer = RedisContainer.instance
+ // Lazy initialization to avoid starting Testcontainers when external Redis is available
+ private val redisContainer by lazy { RedisContainer.instance }
+
+ // Lazy-initialized Redis connection info - computed once on first access
+ private val connectionInfo: Pair by lazy {
+ val externalHost =
+ System.getProperty("test.redis.host")
+ ?: System.getenv("TEST_REDIS_HOST")
+ val externalPortStr =
+ System.getProperty("test.redis.port")
+ ?: System.getenv("TEST_REDIS_PORT")
+
+ if (!externalHost.isNullOrBlank() && !externalPortStr.isNullOrBlank()) {
+ val parsedPort =
+ externalPortStr.toIntOrNull()
+ ?: throw IllegalArgumentException(
+ "Invalid TEST_REDIS_PORT value: '$externalPortStr'. Expected a valid integer.",
+ )
+ externalHost to parsedPort
+ } else {
+ // Ensure the Testcontainers Redis is started before reading host/port
+ if (!redisContainer.isRunning) {
+ redisContainer.start()
+ }
+ redisContainer.host to redisContainer.getMappedPort(6379)
+ }
+ }
+
+ // Redis connection info accessible to subclasses
+ val redisHost: String get() = connectionInfo.first
+ val redisPort: Int get() = connectionInfo.second
internal class Initializer : ApplicationContextInitializer {
override fun initialize(context: ConfigurableApplicationContext) {
val env = context.environment
-
val properties: HashMap = hashMapOf()
- properties["spring.redis.host"] = redisContainer.host
- properties["spring.redis.port"] = redisContainer.getMappedPort(6379)
+
+ // Trigger lazy initialization and get host/port
+ val host = redisHost
+ val port = redisPort
+
+ // Spring Boot 2.x style
+ properties["spring.redis.host"] = host
+ properties["spring.redis.port"] = port
+ // Spring Boot 3.x (Spring Data Redis) style
+ properties["spring.data.redis.host"] = host
+ properties["spring.data.redis.port"] = port
val propertySource = MapPropertySource("testProperties", properties)
env.propertySources.addFirst(propertySource)
diff --git a/support/src/testFixtures/kotlin/com/linecorp/cse/reqshield/support/redis/RedisContainer.kt b/support/src/testFixtures/kotlin/com/linecorp/cse/reqshield/support/redis/RedisContainer.kt
index 0453baa..ce15ada 100644
--- a/support/src/testFixtures/kotlin/com/linecorp/cse/reqshield/support/redis/RedisContainer.kt
+++ b/support/src/testFixtures/kotlin/com/linecorp/cse/reqshield/support/redis/RedisContainer.kt
@@ -21,10 +21,7 @@ import org.testcontainers.utility.DockerImageName
object RedisContainer {
val instance =
- GenericContainer(
- DockerImageName.parse("redis:6.2.7-alpine"),
- ).apply {
- portBindings = listOf("6379:6379")
- withReuse(true)
- }
+ GenericContainer(DockerImageName.parse("redis:6.2.7-alpine"))
+ .withExposedPorts(6379)
+ .withReuse(true)
}