From 9995930873574295b5f235eeae934eedeb5d64e6 Mon Sep 17 00:00:00 2001 From: "kanghyun.yang" Date: Mon, 6 Oct 2025 22:46:07 +0900 Subject: [PATCH] [ISSUE #42] Fix critical issues in the code --- .github/workflows/pull_request_event.yml | 13 + AGENTS.md | 28 ++ CLAUDE.md | 2 +- README.md | 38 +++ .../kotlin/coroutine/KeyLocalLock.kt | 136 ++++++-- .../kotlin/coroutine/KeyGlobalLockTest.kt | 7 +- .../kotlin/coroutine/KeyLocalLockTest.kt | 212 +++++++++++++ .../cse/reqshield/reactor/KeyLocalLock.kt | 166 ++++++++-- .../cse/reqshield/reactor/ReqShield.kt | 100 +++--- .../reqshield/reactor/KeyGlobalLockTest.kt | 7 +- .../cse/reqshield/reactor/KeyLocalLockTest.kt | 118 +++++++ .../build.gradle.kts | 2 + .../coroutine/aspect/CoroutineExtension.kt | 45 +++ .../coroutine/aspect/ReqShieldAspect.kt | 86 ++--- .../coroutine/config/LibAutoConfiguration.kt | 7 +- .../coroutine/aspect/InMemoryAsyncCache.kt | 56 ++++ .../aspect/ReqShieldAspectIntegrationTest.kt | 105 ++++++ .../ReqShieldAspectRedisIntegrationTest.kt | 174 ++++++++++ .../coroutine/aspect/ReqShieldAspectTest.kt | 49 +-- core-spring-webflux/build.gradle.kts | 2 + .../webflux/annotation/ReqShieldCacheable.kt | 6 + .../spring/webflux/aspect/ReqShieldAspect.kt | 56 +++- .../webflux/config/LibAutoConfiguration.kt | 15 +- .../webflux/aspect/InMemoryAsyncCache.kt | 63 ++++ .../aspect/ReqShieldAspectIntegrationTest.kt | 91 ++++++ .../ReqShieldAspectRedisIntegrationTest.kt | 173 ++++++++++ .../webflux/aspect/ReqShieldAspectTest.kt | 27 +- .../spring/aspect/ReqShieldAspect.kt | 21 +- .../test/kotlin/aspect/ReqShieldAspectTest.kt | 8 +- .../linecorp/cse/reqshield/KeyLocalLock.kt | 128 +++++--- .../com/linecorp/cse/reqshield/ReqShield.kt | 61 ++-- .../cse/reqshield/KeyGlobalLockTest.kt | 7 +- .../cse/reqshield/KeyLocalLockTest.kt | 299 ++++++++++++++++-- .../linecorp/cse/reqshield/ReqShieldTest.kt | 84 +++++ libs.versions.toml | 2 + .../example/service/CacheAnnotationTest.kt | 12 +- .../webflux/example/IntegrationSmokeTest.kt | 15 + .../example/service/CacheAnnotationTest.kt | 8 +- .../example/service/CacheAnnotationTest.kt | 8 +- support/build.gradle.kts | 7 +- .../support/constant/ConfigValues.kt | 2 +- .../support/redis/AbstractRedisTest.kt | 51 ++- .../reqshield/support/redis/RedisContainer.kt | 9 +- 43 files changed, 2210 insertions(+), 296 deletions(-) create mode 100644 AGENTS.md create mode 100644 core-spring-webflux-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/InMemoryAsyncCache.kt create mode 100644 core-spring-webflux-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/ReqShieldAspectIntegrationTest.kt create mode 100644 core-spring-webflux-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/aspect/ReqShieldAspectRedisIntegrationTest.kt create mode 100644 core-spring-webflux/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/aspect/InMemoryAsyncCache.kt create mode 100644 core-spring-webflux/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/aspect/ReqShieldAspectIntegrationTest.kt create mode 100644 core-spring-webflux/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/aspect/ReqShieldAspectRedisIntegrationTest.kt create mode 100644 req-shield-spring-boot3-webflux-example/src/test/kotlin/com/linecorp/cse/reqshield/spring3/webflux/example/IntegrationSmokeTest.kt 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) }