From 8be905cddea2e53f932a30e6c5e30b212380b425 Mon Sep 17 00:00:00 2001 From: "kanghyun.yang" Date: Fri, 3 Jan 2025 19:57:56 +0900 Subject: [PATCH 01/11] Add ready_for_review type to pull-request-event --- .github/workflows/pull_request_event.yml | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/.github/workflows/pull_request_event.yml b/.github/workflows/pull_request_event.yml index c31f458..6b2db40 100644 --- a/.github/workflows/pull_request_event.yml +++ b/.github/workflows/pull_request_event.yml @@ -15,7 +15,9 @@ jobs: build: name: Test runs-on: ubuntu-latest - if: ${{ github.event_name == 'pull_request' && (github.event.action == 'opened' || github.event.action == 'synchronize' || github.event.action == 'reopened') }} + if: ${{ github.event_name == 'pull_request' + && (github.event.action == 'opened' || github.event.action == 'synchronize' || + github.event.action == 'reopened' || github.event.action == 'ready_for_review') }} steps: - uses: actions/checkout@v4 with: From 2ea999a025123142a0f771d5a8879f1801d4030a Mon Sep 17 00:00:00 2001 From: "kanghyun.yang" Date: Thu, 13 Mar 2025 15:38:24 +0900 Subject: [PATCH 02/11] add annotation target (ANNOTATION_CLASS) --- .../webflux/kotlin/coroutine/annotation/ReqShieldCacheable.kt | 2 +- .../reqshield/spring/webflux/annotation/ReqShieldCacheable.kt | 2 +- .../cse/reqshield/spring/annotation/ReqShieldCacheable.kt | 2 +- 3 files changed, 3 insertions(+), 3 deletions(-) diff --git a/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheable.kt b/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheable.kt index 9591bfe..7febcb9 100644 --- a/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheable.kt +++ b/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheable.kt @@ -19,7 +19,7 @@ package com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.annotation import com.linecorp.cse.reqshield.support.constant.ConfigValues.MAX_ATTEMPT_GET_CACHE import java.lang.annotation.Inherited -@Target(AnnotationTarget.FUNCTION) +@Target(AnnotationTarget.FUNCTION, AnnotationTarget.ANNOTATION_CLASS) @Retention(AnnotationRetention.RUNTIME) @Inherited annotation class ReqShieldCacheable( 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 5386176..960305c 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 @@ -19,7 +19,7 @@ package com.linecorp.cse.reqshield.spring.webflux.annotation import com.linecorp.cse.reqshield.support.constant.ConfigValues.MAX_ATTEMPT_GET_CACHE import java.lang.annotation.Inherited -@Target(AnnotationTarget.FUNCTION) +@Target(AnnotationTarget.FUNCTION, AnnotationTarget.ANNOTATION_CLASS) @Retention(AnnotationRetention.RUNTIME) @Inherited annotation class ReqShieldCacheable( diff --git a/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheable.kt b/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheable.kt index 5b2e062..c9dcc13 100644 --- a/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheable.kt +++ b/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheable.kt @@ -19,7 +19,7 @@ package com.linecorp.cse.reqshield.spring.annotation import com.linecorp.cse.reqshield.support.constant.ConfigValues.MAX_ATTEMPT_GET_CACHE import java.lang.annotation.Inherited -@Target(AnnotationTarget.FUNCTION) +@Target(AnnotationTarget.FUNCTION, AnnotationTarget.ANNOTATION_CLASS) @Retention(AnnotationRetention.RUNTIME) @Inherited annotation class ReqShieldCacheable( From c7ef485ac08c6ae9e1935bbf19c0bad1be269566 Mon Sep 17 00:00:00 2001 From: "kanghyun.yang" Date: Tue, 25 Mar 2025 16:22:16 +0900 Subject: [PATCH 03/11] [ISSUE #27] Allow us to choose when the req-shield is triggered (creation or modification of the cache) --- .../reqshield/kotlin/coroutine/ReqShield.kt | 23 ++++- .../config/ReqShieldConfiguration.kt | 7 ++ .../kotlin/coroutine/ReqShieldTest.kt | 64 ++++++++++++++ .../cse/reqshield/reactor/ReqShield.kt | 79 +++++++++++------- .../reactor/config/ReqShieldConfiguration.kt | 7 ++ .../cse/reqshield/reactor/ReqShieldTest.kt | 83 +++++++++++++++++++ .../annotation/ReqShieldCacheable.kt | 2 + .../coroutine/aspect/ReqShieldAspect.kt | 1 + .../webflux/annotation/ReqShieldCacheable.kt | 2 + .../spring/webflux/aspect/ReqShieldAspect.kt | 1 + .../spring/annotation/ReqShieldCacheable.kt | 2 + .../spring/aspect/ReqShieldAspect.kt | 1 + .../com/linecorp/cse/reqshield/ReqShield.kt | 55 +++++++----- .../config/ReqShieldConfiguration.kt | 7 ++ .../linecorp/cse/reqshield/ReqShieldTest.kt | 61 ++++++++++++++ .../cse/reqshield/service/SampleService.kt | 15 ++++ .../spring/service/CacheAnnotationTest.kt | 20 +++++ .../webflux/example/service/SampleService.kt | 18 ++++ .../example/service/CacheAnnotationTest.kt | 20 +++++ .../example/service/SampleService.kt | 16 ++++ .../example/service/CacheAnnotationTest.kt | 16 ++++ .../reqshield/support/BaseReqShieldTest.kt | 4 + 22 files changed, 452 insertions(+), 52 deletions(-) diff --git a/core-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/ReqShield.kt b/core-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/ReqShield.kt index 601ab83..d1568c0 100644 --- a/core-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/ReqShield.kt +++ b/core-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/ReqShield.kt @@ -17,6 +17,7 @@ package com.linecorp.cse.reqshield.kotlin.coroutine import com.linecorp.cse.reqshield.kotlin.coroutine.config.ReqShieldConfiguration +import com.linecorp.cse.reqshield.kotlin.coroutine.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 @@ -60,7 +61,8 @@ class ReqShield( timeToLiveMillis: Long, ) { val lockType = LockType.UPDATE - if (reqShieldConfig.keyLock.tryLock(key, lockType)) { + + fun executeAsyncTask() { CoroutineScope(Dispatchers.IO).launch { val reqShieldData = buildReqShieldData( @@ -75,6 +77,12 @@ class ReqShield( ) } } + + if (reqShieldConfig.reqShieldWorkMode == ReqShieldWorkMode.ONLY_CREATE_CACHE || + reqShieldConfig.keyLock.tryLock(key, lockType) + ) { + return executeAsyncTask() + } } private suspend fun handleLockForCacheCreation( @@ -83,7 +91,10 @@ class ReqShield( timeToLiveMillis: Long, ): ReqShieldData { val lockType = LockType.CREATE - return if (reqShieldConfig.keyLock.tryLock(key, lockType)) { + + return if (reqShieldConfig.reqShieldWorkMode == ReqShieldWorkMode.ONLY_UPDATE_CACHE || + reqShieldConfig.keyLock.tryLock(key, lockType) + ) { createReqShieldData(key, callable, timeToLiveMillis, lockType) } else { handleLockFailure(key, callable, timeToLiveMillis) @@ -177,7 +188,9 @@ class ReqShield( } catch (e: Exception) { throw ClientException(ErrorCode.SET_CACHE_ERROR, originErrorMessage = e.message) } finally { - unlockWithRetry(key, lockType) + if (shouldAttemptUnlock(lockType)) { + unlockWithRetry(key, lockType) + } } } @@ -208,4 +221,8 @@ class ReqShield( } throw ClientException(ErrorCode.SUPPLIER_ERROR, originErrorMessage = it.message) } + + private fun shouldAttemptUnlock(lockType: LockType): Boolean = + (lockType == LockType.UPDATE && reqShieldConfig.reqShieldWorkMode != ReqShieldWorkMode.ONLY_CREATE_CACHE) || + (lockType == LockType.CREATE && reqShieldConfig.reqShieldWorkMode != ReqShieldWorkMode.ONLY_UPDATE_CACHE) } diff --git a/core-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/config/ReqShieldConfiguration.kt b/core-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/config/ReqShieldConfiguration.kt index 86fce1d..9d900ae 100644 --- a/core-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/config/ReqShieldConfiguration.kt +++ b/core-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/config/ReqShieldConfiguration.kt @@ -40,6 +40,7 @@ data class ReqShieldConfiguration( KeyGlobalLock(globalLockFunction!!, globalUnLockFunction!!, lockTimeoutMillis) }, val maxAttemptGetCache: Int = MAX_ATTEMPT_GET_CACHE, + val reqShieldWorkMode: ReqShieldWorkMode = ReqShieldWorkMode.CREATE_AND_UPDATE_CACHE, ) { init { if (!isLocalLock) { @@ -52,3 +53,9 @@ data class ReqShieldConfiguration( } } } + +enum class ReqShieldWorkMode { + CREATE_AND_UPDATE_CACHE, + ONLY_CREATE_CACHE, + ONLY_UPDATE_CACHE, +} diff --git a/core-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/ReqShieldTest.kt b/core-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/ReqShieldTest.kt index 1a55e76..e2ada9b 100644 --- a/core-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/ReqShieldTest.kt +++ b/core-kotlin-coroutine/src/test/kotlin/com/linecorp/cse/reqshield/kotlin/coroutine/ReqShieldTest.kt @@ -17,6 +17,7 @@ package com.linecorp.cse.reqshield.kotlin.coroutine import com.linecorp.cse.reqshield.kotlin.coroutine.config.ReqShieldConfiguration +import com.linecorp.cse.reqshield.kotlin.coroutine.config.ReqShieldWorkMode import com.linecorp.cse.reqshield.support.BaseReqShieldTest import com.linecorp.cse.reqshield.support.exception.ClientException import com.linecorp.cse.reqshield.support.exception.code.ErrorCode @@ -51,6 +52,8 @@ import kotlin.test.assertTrue @OptIn(ExperimentalCoroutinesApi::class) class ReqShieldTest : BaseReqShieldTest { private lateinit var reqShield: ReqShield + private lateinit var reqShieldOnlyUpdateCache: ReqShield + private lateinit var reqShieldOnlyCreateCache: ReqShield private lateinit var reqShieldForGlobalLock: ReqShield private lateinit var reqShieldForGlobalLockForError: ReqShield private lateinit var cacheSetter: suspend (String, ReqShieldData, Long) -> Boolean @@ -88,6 +91,26 @@ class ReqShieldTest : BaseReqShieldTest { ), ) + reqShieldOnlyUpdateCache = + ReqShield( + ReqShieldConfiguration( + cacheSetter, + cacheGetter, + keyLock = keyLock, + reqShieldWorkMode = ReqShieldWorkMode.ONLY_UPDATE_CACHE, + ), + ) + + reqShieldOnlyCreateCache = + ReqShield( + ReqShieldConfiguration( + cacheSetter, + cacheGetter, + keyLock = keyLock, + reqShieldWorkMode = ReqShieldWorkMode.ONLY_CREATE_CACHE, + ), + ) + reqShieldForGlobalLock = ReqShield( ReqShieldConfiguration( @@ -128,6 +151,24 @@ class ReqShieldTest : BaseReqShieldTest { coVerify { callable() } } + @Test + override fun testSetMethodCacheNotExistsAndOnlyUpdateCache() { + runBlocking { + coEvery { cacheGetter.invoke(key) } returns null + coEvery { cacheSetter.invoke(key, any(), any()) } returns true + + val result = reqShieldOnlyUpdateCache.getAndSetReqShieldData(key, callable, timeToLiveMillis) + delay(100) + + assertNotNull(result) + coVerify { cacheGetter.invoke(key) } + coVerify { cacheSetter.invoke(key, result, timeToLiveMillis) } + coVerify(inverse = true) { keyLock.tryLock(key, LockType.CREATE) } + coVerify(inverse = true) { keyLock.unLock(key, LockType.CREATE) } + coVerify { callable() } + } + } + @Test override fun testSetMethodCacheNotExistsAndGlobalLockAcquired() = runBlocking { @@ -394,6 +435,29 @@ class ReqShieldTest : BaseReqShieldTest { coVerify { callable() } } + @Test + override fun testSetMethodCacheExistsAndTheUpdateTargetOnlyCreateCache() { + runBlocking { + val timeToLiveMillis: Long = 1000 + val reqShieldData = ReqShieldData(oldValue, timeToLiveMillis) + val newReqShieldData = ReqShieldData(value, timeToLiveMillis) + + coEvery { cacheGetter.invoke(key) } returns reqShieldData + coEvery { cacheSetter.invoke(key, any(), any()) } coAnswers { true } + + val result = reqShieldOnlyCreateCache.getAndSetReqShieldData(key, callable, timeToLiveMillis) + + delay(100) + + assertEquals(reqShieldData, result) + coVerify { cacheGetter.invoke(key) } + coVerify { cacheSetter.invoke(key, newReqShieldData, timeToLiveMillis) } + coVerify(inverse = true) { keyLock.tryLock(key, LockType.UPDATE) } + coVerify(inverse = true) { keyLock.unLock(key, LockType.UPDATE) } + coVerify { callable() } + } + } + @Test override fun testSetMethodCacheExistsAndTheUpdateTargetAndCallableReturnNull() = runBlocking { 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 e038534..1d0b03f 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 @@ -17,6 +17,7 @@ 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 @@ -68,35 +69,42 @@ class ReqShield( timeToLiveMillis: Long, ) { val lockType = LockType.UPDATE - reqShieldConfig.keyLock - .tryLock(key, lockType) - .filter { it } - .flatMap { - executeCallable({ callable.call() }, true, key, lockType) - .map { data -> buildReqShieldData(data, timeToLiveMillis) } - .doOnNext { reqShieldData -> + + fun processMono(): Mono> = + executeCallable({ callable.call() }, true, key, lockType) + .map { data -> buildReqShieldData(data, timeToLiveMillis) } + .doOnNext { reqShieldData -> + setReqShieldData( + reqShieldConfig.setCacheFunction, + key, + reqShieldData, + lockType, + ) + }.switchIfEmpty( + Mono.defer { + val reqShieldData = buildReqShieldData(null, timeToLiveMillis) setReqShieldData( reqShieldConfig.setCacheFunction, key, reqShieldData, lockType, ) - }.switchIfEmpty( - Mono.defer { - val reqShieldData = buildReqShieldData(null, timeToLiveMillis) - - setReqShieldData( - reqShieldConfig.setCacheFunction, - key, - reqShieldData, - lockType, - ) + Mono.just(reqShieldData) + }, + ) - Mono.just(reqShieldData) - }, - ) - }.subscribeOn(Schedulers.boundedElastic()) - .subscribe() + if (reqShieldConfig.reqShieldWorkMode == ReqShieldWorkMode.ONLY_CREATE_CACHE) { + processMono() + .subscribeOn(Schedulers.boundedElastic()) + .subscribe() + } else { + reqShieldConfig.keyLock + .tryLock(key, lockType) + .filter { it } + .flatMap { processMono() } + .subscribeOn(Schedulers.boundedElastic()) + .subscribe() + } } private fun handleLockForCacheCreation( @@ -105,6 +113,11 @@ class ReqShield( timeToLiveMillis: Long, ): Mono> { val lockType = LockType.CREATE + + if (reqShieldConfig.reqShieldWorkMode == ReqShieldWorkMode.ONLY_UPDATE_CACHE) { + return createReqShieldData(key, callable, timeToLiveMillis, lockType) + } + return reqShieldConfig.keyLock .tryLock(key, lockType) .flatMap { acquired -> @@ -217,14 +230,16 @@ class ReqShield( .doOnError { e -> throw ClientException(ErrorCode.SET_CACHE_ERROR, originErrorMessage = e.message) }.doFinally { - reqShieldConfig.keyLock - .unLock(key, lockType) - .retryWhen( - Retry.fixedDelay( - MAX_ATTEMPT_SET_CACHE - 1L, - Duration.ofMillis(SET_CACHE_RETRY_INTERVAL_MILLIS), - ), - ).subscribe() + if (shouldAttemptUnlock(lockType)) { + reqShieldConfig.keyLock + .unLock(key, lockType) + .retryWhen( + Retry.fixedDelay( + MAX_ATTEMPT_SET_CACHE - 1L, + Duration.ofMillis(SET_CACHE_RETRY_INTERVAL_MILLIS), + ), + ).subscribe() + } }.subscribeOn(Schedulers.boundedElastic()) private fun executeCallable( @@ -241,4 +256,8 @@ class ReqShield( } throw ClientException(ErrorCode.SUPPLIER_ERROR, originErrorMessage = e.message) } + + private fun shouldAttemptUnlock(lockType: LockType): Boolean = + (lockType == LockType.UPDATE && reqShieldConfig.reqShieldWorkMode != ReqShieldWorkMode.ONLY_CREATE_CACHE) || + (lockType == LockType.CREATE && reqShieldConfig.reqShieldWorkMode != ReqShieldWorkMode.ONLY_UPDATE_CACHE) } diff --git a/core-reactor/src/main/kotlin/com/linecorp/cse/reqshield/reactor/config/ReqShieldConfiguration.kt b/core-reactor/src/main/kotlin/com/linecorp/cse/reqshield/reactor/config/ReqShieldConfiguration.kt index e5c9340..abfe4ab 100644 --- a/core-reactor/src/main/kotlin/com/linecorp/cse/reqshield/reactor/config/ReqShieldConfiguration.kt +++ b/core-reactor/src/main/kotlin/com/linecorp/cse/reqshield/reactor/config/ReqShieldConfiguration.kt @@ -44,6 +44,7 @@ data class ReqShieldConfiguration( KeyGlobalLock(globalLockFunction!!, globalUnLockFunction!!, lockTimeoutMillis) }, val maxAttemptGetCache: Int = MAX_ATTEMPT_GET_CACHE, + val reqShieldWorkMode: ReqShieldWorkMode = ReqShieldWorkMode.CREATE_AND_UPDATE_CACHE, ) { init { if (!isLocalLock) { @@ -56,3 +57,9 @@ data class ReqShieldConfiguration( } } } + +enum class ReqShieldWorkMode { + CREATE_AND_UPDATE_CACHE, + ONLY_CREATE_CACHE, + ONLY_UPDATE_CACHE, +} diff --git a/core-reactor/src/test/kotlin/com/linecorp/cse/reqshield/reactor/ReqShieldTest.kt b/core-reactor/src/test/kotlin/com/linecorp/cse/reqshield/reactor/ReqShieldTest.kt index 1234709..be82f40 100644 --- a/core-reactor/src/test/kotlin/com/linecorp/cse/reqshield/reactor/ReqShieldTest.kt +++ b/core-reactor/src/test/kotlin/com/linecorp/cse/reqshield/reactor/ReqShieldTest.kt @@ -17,6 +17,7 @@ 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.BaseReqShieldTest import com.linecorp.cse.reqshield.support.exception.ClientException import com.linecorp.cse.reqshield.support.exception.code.ErrorCode @@ -45,6 +46,8 @@ import kotlin.test.assertNull class ReqShieldTest : BaseReqShieldTest { private lateinit var reqShield: ReqShield + private lateinit var reqShieldOnlyUpdateCache: ReqShield + private lateinit var reqShieldOnlyCreateCache: ReqShield private lateinit var reqShieldForGlobalLock: ReqShield private lateinit var reqShieldForGlobalLockForError: ReqShield private lateinit var cacheSetter: (String, ReqShieldData, Long) -> Mono @@ -82,6 +85,26 @@ class ReqShieldTest : BaseReqShieldTest { ), ) + reqShieldOnlyUpdateCache = + ReqShield( + ReqShieldConfiguration( + cacheSetter, + cacheGetter, + keyLock = keyLock, + reqShieldWorkMode = ReqShieldWorkMode.ONLY_UPDATE_CACHE, + ), + ) + + reqShieldOnlyCreateCache = + ReqShield( + ReqShieldConfiguration( + cacheSetter, + cacheGetter, + keyLock = keyLock, + reqShieldWorkMode = ReqShieldWorkMode.ONLY_CREATE_CACHE, + ), + ) + reqShieldForGlobalLock = ReqShield( ReqShieldConfiguration( @@ -132,6 +155,33 @@ class ReqShieldTest : BaseReqShieldTest { verify { callable.call() } } + @Test + override fun testSetMethodCacheNotExistsAndOnlyUpdateCache() { + every { cacheGetter.invoke(key) } returns Mono.empty() + every { cacheSetter.invoke(key, any(), any()) } returns Mono.just(true) + + val result = reqShieldOnlyUpdateCache.getAndSetReqShieldData(key, callable, timeToLiveMillis) + + StepVerifier + .create(result) + .assertNext { + assertNotNull(it) + }.verifyComplete() + + StepVerifier + .create(Mono.delay(Duration.ofMillis(100))) + .expectSubscription() + .thenAwait(Duration.ofMillis(100)) + .expectNextCount(1) + .verifyComplete() + + verify { cacheGetter.invoke(key) } + verify { cacheSetter.invoke(key, any(), any()) } + verify(inverse = true) { keyLock.tryLock(key, LockType.CREATE) } + verify(inverse = true) { keyLock.unLock(key, LockType.CREATE) } + verify { callable.call() } + } + @Test override fun testSetMethodCacheNotExistsAndGlobalLockAcquired() { every { cacheGetter.invoke(key) } returns Mono.empty() @@ -457,6 +507,39 @@ class ReqShieldTest : BaseReqShieldTest { verify { callable.call() } } + @Test + override fun testSetMethodCacheExistsAndTheUpdateTargetOnlyCreateCache() { + timeToLiveMillis = 1000 + val reqShieldData = ReqShieldData(oldValue, timeToLiveMillis) + val newReqShieldData = ReqShieldData(value, timeToLiveMillis) + + every { cacheGetter.invoke(key) } returns Mono.just(reqShieldData) + every { cacheSetter.invoke(key, any(), any()) } answers { Mono.just(true) } + + val result = reqShieldOnlyCreateCache.getAndSetReqShieldData(key, callable, timeToLiveMillis) + + StepVerifier + .create(result) + .expectNextMatches { + assertEquals(reqShieldData, it) + true + }.expectComplete() + .verify() + + StepVerifier + .create(Mono.delay(Duration.ofMillis(100))) + .expectSubscription() + .thenAwait(Duration.ofMillis(100)) + .expectNextCount(1) + .verifyComplete() + + verify(inverse = true) { keyLock.tryLock(key, LockType.UPDATE) } + verify { cacheGetter.invoke(key) } + verify { cacheSetter.invoke(key, newReqShieldData, timeToLiveMillis) } + verify(inverse = true) { keyLock.unLock(key, LockType.UPDATE) } + verify { callable.call() } + } + @Test override fun testSetMethodCacheExistsAndTheUpdateTargetAndCallableReturnNull() { timeToLiveMillis = 1000 diff --git a/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheable.kt b/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheable.kt index 7febcb9..ce606c5 100644 --- a/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheable.kt +++ b/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheable.kt @@ -16,6 +16,7 @@ package com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.annotation +import com.linecorp.cse.reqshield.kotlin.coroutine.config.ReqShieldWorkMode import com.linecorp.cse.reqshield.support.constant.ConfigValues.MAX_ATTEMPT_GET_CACHE import java.lang.annotation.Inherited @@ -30,4 +31,5 @@ annotation class ReqShieldCacheable( val decisionForUpdate: Int = 90, val maxAttemptGetCache: Int = MAX_ATTEMPT_GET_CACHE, val timeToLiveMillis: Long = 10 * 60 * 1000, + val reqShieldWorkMode: ReqShieldWorkMode = ReqShieldWorkMode.CREATE_AND_UPDATE_CACHE, ) 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 e0fefa5..d0f35bb 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 @@ -142,6 +142,7 @@ class ReqShieldAspect( lockTimeoutMillis = annotation.lockTimeoutMillis, decisionForUpdate = annotation.decisionForUpdate, maxAttemptGetCache = annotation.maxAttemptGetCache, + reqShieldWorkMode = annotation.reqShieldWorkMode, ) return ReqShield(reqShieldConfiguration) 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 960305c..70653aa 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 @@ -16,6 +16,7 @@ package com.linecorp.cse.reqshield.spring.webflux.annotation +import com.linecorp.cse.reqshield.reactor.config.ReqShieldWorkMode import com.linecorp.cse.reqshield.support.constant.ConfigValues.MAX_ATTEMPT_GET_CACHE import java.lang.annotation.Inherited @@ -30,4 +31,5 @@ annotation class ReqShieldCacheable( val decisionForUpdate: Int = 90, val maxAttemptGetCache: Int = MAX_ATTEMPT_GET_CACHE, val timeToLiveMillis: Long = 10 * 60 * 1000, + val reqShieldWorkMode: ReqShieldWorkMode = ReqShieldWorkMode.CREATE_AND_UPDATE_CACHE, ) 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 c8d3203..87b5f27 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 @@ -128,6 +128,7 @@ class ReqShieldAspect( lockTimeoutMillis = annotation.lockTimeoutMillis, decisionForUpdate = annotation.decisionForUpdate, maxAttemptGetCache = annotation.maxAttemptGetCache, + reqShieldWorkMode = annotation.reqShieldWorkMode, ) return ReqShield(reqShieldConfiguration) diff --git a/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheable.kt b/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheable.kt index c9dcc13..f0d01d8 100644 --- a/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheable.kt +++ b/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheable.kt @@ -16,6 +16,7 @@ package com.linecorp.cse.reqshield.spring.annotation +import com.linecorp.cse.reqshield.config.ReqShieldWorkMode import com.linecorp.cse.reqshield.support.constant.ConfigValues.MAX_ATTEMPT_GET_CACHE import java.lang.annotation.Inherited @@ -30,4 +31,5 @@ annotation class ReqShieldCacheable( val decisionForUpdate: Int = 90, val maxAttemptGetCache: Int = MAX_ATTEMPT_GET_CACHE, val timeToLiveMillis: Long = 10 * 60 * 1000, + val reqShieldWorkMode: ReqShieldWorkMode = ReqShieldWorkMode.CREATE_AND_UPDATE_CACHE, ) 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 8b8f24c..be4f91c 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 @@ -88,6 +88,7 @@ class ReqShieldAspect( lockTimeoutMillis = annotation.lockTimeoutMillis, decisionForUpdate = annotation.decisionForUpdate, maxAttemptGetCache = annotation.maxAttemptGetCache, + reqShieldWorkMode = annotation.reqShieldWorkMode, ) return ReqShield(reqShieldConfiguration) 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 b432187..166bfca 100644 --- a/core/src/main/kotlin/com/linecorp/cse/reqshield/ReqShield.kt +++ b/core/src/main/kotlin/com/linecorp/cse/reqshield/ReqShield.kt @@ -17,6 +17,7 @@ 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 @@ -32,14 +33,14 @@ import java.util.concurrent.TimeUnit import java.util.concurrent.atomic.AtomicInteger class ReqShield( - private val reqShieldConfiguration: ReqShieldConfiguration, + private val reqShieldConfig: ReqShieldConfiguration, ) { fun getAndSetReqShieldData( key: String, callable: Callable, timeToLiveMillis: Long, ): ReqShieldData { - val currentReqShieldData = executeGetCacheFunction(reqShieldConfiguration.getCacheFunction, key) + val currentReqShieldData = executeGetCacheFunction(reqShieldConfig.getCacheFunction, key) currentReqShieldData?.let { if (shouldUpdateCache(it)) { updateReqShieldData(key, callable, timeToLiveMillis) @@ -51,7 +52,7 @@ class ReqShield( } private fun shouldUpdateCache(reqShieldData: ReqShieldData): Boolean = - decideToUpdateCache(reqShieldData.createdAt, reqShieldData.timeToLiveMillis, reqShieldConfiguration.decisionForUpdate) + decideToUpdateCache(reqShieldData.createdAt, reqShieldData.timeToLiveMillis, reqShieldConfig.decisionForUpdate) private fun updateReqShieldData( key: String, @@ -59,7 +60,8 @@ class ReqShield( timeToLiveMillis: Long, ) { val lockType = LockType.UPDATE - if (reqShieldConfiguration.keyLock.tryLock(key, lockType)) { + + fun executeAsyncTask() { CompletableFuture.runAsync({ val reqShieldData = buildReqShieldData( @@ -67,12 +69,18 @@ class ReqShield( timeToLiveMillis, ) setReqShieldData( - reqShieldConfiguration.setCacheFunction, + reqShieldConfig.setCacheFunction, key, reqShieldData, - LockType.UPDATE, + lockType, ) - }, reqShieldConfiguration.executor) + }, reqShieldConfig.executor) + } + + if (reqShieldConfig.reqShieldWorkMode == ReqShieldWorkMode.ONLY_CREATE_CACHE || + reqShieldConfig.keyLock.tryLock(key, lockType) + ) { + return executeAsyncTask() } } @@ -82,7 +90,10 @@ class ReqShield( timeToLiveMillis: Long, ): ReqShieldData { val lockType = LockType.CREATE - return if (reqShieldConfiguration.keyLock.tryLock(key, lockType)) { + + return if (reqShieldConfig.reqShieldWorkMode == ReqShieldWorkMode.ONLY_UPDATE_CACHE || + reqShieldConfig.keyLock.tryLock(key, lockType) + ) { createReqShieldData(key, callable, timeToLiveMillis, lockType) } else { handleLockFailure(key, callable, timeToLiveMillis) @@ -101,8 +112,8 @@ class ReqShield( timeToLiveMillis, ) CompletableFuture.runAsync({ - setReqShieldData(reqShieldConfiguration.setCacheFunction, key, reqShieldData, lockType) - }, reqShieldConfiguration.executor) + setReqShieldData(reqShieldConfig.setCacheFunction, key, reqShieldData, lockType) + }, reqShieldConfig.executor) return reqShieldData } @@ -115,7 +126,7 @@ class ReqShield( val future = createFuture() val counter = createCounter() - scheduleTask(reqShieldConfiguration.executor, future, counter, reqShieldConfiguration.getCacheFunction, callable, key) + scheduleTask(reqShieldConfig.executor, future, counter, reqShieldConfig.getCacheFunction, callable, key) val result = future.get() @@ -158,7 +169,7 @@ class ReqShield( val funcResult = executeGetCacheFunction(cacheGetter, key) if (funcResult != null) { future.complete(funcResult.value) - } else if (counter.incrementAndGet() >= reqShieldConfiguration.maxAttemptGetCache) { + } else if (counter.incrementAndGet() >= reqShieldConfig.maxAttemptGetCache) { future.complete( executeCallable({ callable.call() }, false), ) @@ -198,12 +209,14 @@ class ReqShield( } catch (e: Exception) { throw ClientException(ErrorCode.SET_CACHE_ERROR, originErrorMessage = e.message) } finally { - while (!unlockSuccess && retryCount < MAX_ATTEMPT_SET_CACHE) { - if (reqShieldConfiguration.keyLock.unLock(key, lockType)) { - unlockSuccess = true - } else { - retryCount++ - Thread.sleep(SET_CACHE_RETRY_INTERVAL_MILLIS) + if (shouldAttemptUnlock(lockType)) { + while (!unlockSuccess && retryCount < MAX_ATTEMPT_SET_CACHE) { + if (reqShieldConfig.keyLock.unLock(key, lockType)) { + unlockSuccess = true + } else { + retryCount++ + Thread.sleep(SET_CACHE_RETRY_INTERVAL_MILLIS) + } } } } @@ -219,8 +232,12 @@ class ReqShield( callable.call() }.getOrElse { if (isUnlockWhenException && key != null && lockType != null) { - reqShieldConfiguration.keyLock.unLock(key, lockType) + reqShieldConfig.keyLock.unLock(key, lockType) } throw ClientException(ErrorCode.SUPPLIER_ERROR, originErrorMessage = it.message) } + + private fun shouldAttemptUnlock(lockType: LockType): Boolean = + (lockType == LockType.UPDATE && reqShieldConfig.reqShieldWorkMode != ReqShieldWorkMode.ONLY_CREATE_CACHE) || + (lockType == LockType.CREATE && reqShieldConfig.reqShieldWorkMode != ReqShieldWorkMode.ONLY_UPDATE_CACHE) } diff --git a/core/src/main/kotlin/com/linecorp/cse/reqshield/config/ReqShieldConfiguration.kt b/core/src/main/kotlin/com/linecorp/cse/reqshield/config/ReqShieldConfiguration.kt index dce7d33..68994de 100644 --- a/core/src/main/kotlin/com/linecorp/cse/reqshield/config/ReqShieldConfiguration.kt +++ b/core/src/main/kotlin/com/linecorp/cse/reqshield/config/ReqShieldConfiguration.kt @@ -46,6 +46,7 @@ data class ReqShieldConfiguration( KeyGlobalLock(globalLockFunction!!, globalUnLockFunction!!, lockTimeoutMillis) }, val maxAttemptGetCache: Int = MAX_ATTEMPT_GET_CACHE, + val reqShieldWorkMode: ReqShieldWorkMode = ReqShieldWorkMode.CREATE_AND_UPDATE_CACHE, ) { init { if (!isLocalLock) { @@ -58,3 +59,9 @@ data class ReqShieldConfiguration( } } } + +enum class ReqShieldWorkMode { + CREATE_AND_UPDATE_CACHE, + ONLY_CREATE_CACHE, + ONLY_UPDATE_CACHE, +} 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 4dae75c..e6536ca 100644 --- a/core/src/test/kotlin/com/linecorp/cse/reqshield/ReqShieldTest.kt +++ b/core/src/test/kotlin/com/linecorp/cse/reqshield/ReqShieldTest.kt @@ -17,6 +17,7 @@ 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.BaseReqShieldTest import com.linecorp.cse.reqshield.support.BaseReqShieldTest.Companion.AWAIT_TIMEOUT import com.linecorp.cse.reqshield.support.exception.ClientException @@ -46,6 +47,8 @@ import kotlin.test.assertNull class ReqShieldTest : BaseReqShieldTest { private lateinit var reqShield: ReqShield + private lateinit var reqShieldOnlyUpdateCache: ReqShield + private lateinit var reqShieldOnlyCreateCache: ReqShield private lateinit var reqShieldForGlobalLock: ReqShield private lateinit var reqShieldForGlobalLockForError: ReqShield private lateinit var cacheSetter: (String, ReqShieldData, Long) -> Boolean @@ -83,6 +86,26 @@ class ReqShieldTest : BaseReqShieldTest { ), ) + reqShieldOnlyUpdateCache = + ReqShield( + ReqShieldConfiguration( + cacheSetter, + cacheGetter, + keyLock = keyLock, + reqShieldWorkMode = ReqShieldWorkMode.ONLY_UPDATE_CACHE, + ), + ) + + reqShieldOnlyCreateCache = + ReqShield( + ReqShieldConfiguration( + cacheSetter, + cacheGetter, + keyLock = keyLock, + reqShieldWorkMode = ReqShieldWorkMode.ONLY_CREATE_CACHE, + ), + ) + reqShieldForGlobalLock = ReqShield( ReqShieldConfiguration( @@ -123,6 +146,23 @@ class ReqShieldTest : BaseReqShieldTest { } } + @Test + override fun testSetMethodCacheNotExistsAndOnlyUpdateCache() { + every { cacheGetter.invoke(key) } returns null + every { cacheSetter.invoke(key, any(), any()) } returns true + + val result = reqShieldOnlyUpdateCache.getAndSetReqShieldData(key, callable, timeToLiveMillis) + + await().atMost(Duration.ofMillis(AWAIT_TIMEOUT)).untilAsserted { + assertNotNull(result) + verify { cacheGetter.invoke(key) } + verify { cacheSetter.invoke(key, result, timeToLiveMillis) } + verify(inverse = true) { keyLock.tryLock(key, LockType.CREATE) } + verify(inverse = true) { keyLock.unLock(key, LockType.CREATE) } + verify { callable.call() } + } + } + @Test override fun testSetMethodCacheNotExistsAndGlobalLockAcquired() { every { cacheGetter.invoke(key) } returns null @@ -368,6 +408,27 @@ class ReqShieldTest : BaseReqShieldTest { } } + @Test + override fun testSetMethodCacheExistsAndTheUpdateTargetOnlyCreateCache() { + timeToLiveMillis = 1000 + val reqShieldData = ReqShieldData(oldValue, timeToLiveMillis) + val newReqShieldData = ReqShieldData(value, timeToLiveMillis) + + every { cacheGetter.invoke(key) } returns reqShieldData + every { cacheSetter.invoke(key, any(), any()) } answers { true } + + val result = reqShieldOnlyCreateCache.getAndSetReqShieldData(key, callable, timeToLiveMillis) + + await().atMost(Duration.ofMillis(AWAIT_TIMEOUT)).untilAsserted { + assertEquals(reqShieldData, result) + verify { cacheGetter.invoke(key) } + verify { cacheSetter.invoke(key, newReqShieldData, timeToLiveMillis) } + verify(inverse = true) { keyLock.tryLock(key, LockType.UPDATE) } + verify(inverse = true) { keyLock.unLock(key, LockType.UPDATE) } + verify { callable.call() } + } + } + @Test override fun testSetMethodCacheExistsAndTheUpdateTargetAndCallableReturnNull() { timeToLiveMillis = 1000 diff --git a/req-shield-spring-example/src/main/kotlin/com/linecorp/cse/reqshield/service/SampleService.kt b/req-shield-spring-example/src/main/kotlin/com/linecorp/cse/reqshield/service/SampleService.kt index 3bef7c6..8449659 100644 --- a/req-shield-spring-example/src/main/kotlin/com/linecorp/cse/reqshield/service/SampleService.kt +++ b/req-shield-spring-example/src/main/kotlin/com/linecorp/cse/reqshield/service/SampleService.kt @@ -16,6 +16,7 @@ package com.linecorp.cse.reqshield.service +import com.linecorp.cse.reqshield.config.ReqShieldWorkMode import com.linecorp.cse.reqshield.dto.Product import com.linecorp.cse.reqshield.spring.annotation.ReqShieldCacheEvict import com.linecorp.cse.reqshield.spring.annotation.ReqShieldCacheable @@ -38,6 +39,20 @@ class SampleService { return Product(productId, "product_$productId") } + @ReqShieldCacheable( + cacheName = "productOnlyUpdateCache", + decisionForUpdate = 90, + timeToLiveMillis = 60 * 1000, + reqShieldWorkMode = ReqShieldWorkMode.ONLY_UPDATE_CACHE, + ) + fun getProductOnlyUpdateCache(productId: String): Product { + log.info("find product with db request - req-shield local lock only update cache (will take 1 second)") + Thread.sleep(500) + + atomicInteger.incrementAndGet() + return Product(productId, "product_$productId") + } + @ReqShieldCacheable(cacheName = "product", isLocalLock = false, decisionForUpdate = 90) fun getProductForGlobalLock(productId: String): Product { log.info("find product with db request - req-shield global lock (will take 1 second)") diff --git a/req-shield-spring-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/service/CacheAnnotationTest.kt b/req-shield-spring-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/service/CacheAnnotationTest.kt index cb2eedc..1337691 100644 --- a/req-shield-spring-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/service/CacheAnnotationTest.kt +++ b/req-shield-spring-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/service/CacheAnnotationTest.kt @@ -70,6 +70,26 @@ class CacheAnnotationTest : AbstractRedisTest() { } } + @Test + fun `ReqShieldCacheable test - request to 'sampleService' should be request count times (only update cache mode)`() { + val executorService = Executors.newFixedThreadPool(100) + + val testProductId: String = UUID.randomUUID().toString() + + for (i in 1..100) { + executorService.submit { + sampleService.getProductOnlyUpdateCache(testProductId) + } + } + + executorService.shutdown() + executorService.awaitTermination(3000, TimeUnit.SECONDS) + + await().atMost(Duration.ofMillis(BaseReqShieldTest.AWAIT_TIMEOUT)).untilAsserted { + assertEquals(100, sampleService.getRequestCount()) + } + } + @Test fun `ReqShieldCacheable test - request to 'sampleService' should be only one times for global lock`() { val executorService = Executors.newFixedThreadPool(100) diff --git a/req-shield-spring-webflux-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/example/service/SampleService.kt b/req-shield-spring-webflux-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/example/service/SampleService.kt index 181cc35..d15f73d 100644 --- a/req-shield-spring-webflux-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/example/service/SampleService.kt +++ b/req-shield-spring-webflux-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/example/service/SampleService.kt @@ -17,6 +17,7 @@ package com.linecorp.cse.reqshield.spring.webflux.example.service import com.linecorp.cse.reqshield.reactor.ReqShield +import com.linecorp.cse.reqshield.reactor.config.ReqShieldWorkMode 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.example.dto.Product @@ -46,6 +47,23 @@ class SampleService( }.doFinally { atomicInteger.incrementAndGet() }, ) + @ReqShieldCacheable( + cacheName = "productOnlyUpdataCache", + decisionForUpdate = 80, + timeToLiveMillis = 60 * 1000, + reqShieldWorkMode = ReqShieldWorkMode.ONLY_UPDATE_CACHE, + ) + fun getProductOnlyUpdateCache(productId: String): Mono = + Mono + .delay(Duration.ofMillis(500)) + .then( + Mono + .just(Product(productId, "product_$productId")) + .doOnNext { + log.info("find product with db request - req-shield local lock (will take 1 second)") + }.doFinally { atomicInteger.incrementAndGet() }, + ) + fun getProductNoAnno(productId: String): Mono = reqShield .getAndSetReqShieldData( 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 ff7c0b9..b84a973 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 @@ -70,6 +70,26 @@ class CacheAnnotationTest : AbstractRedisTest() { }.verifyComplete() } + @Test + fun `ReqShieldCacheable test - request to 'sampleService' should be request count times(only update cache mode)`() { + val testProductId: String = UUID.randomUUID().toString() + + val flux = + Flux + .range(1, 20) + .flatMap { + sampleService + .getProductOnlyUpdateCache(testProductId) + .subscribeOn(Schedulers.boundedElastic()) + }.collectList() + + StepVerifier + .create(flux) + .assertNext { productList -> + assertEquals(19, sampleService.getRequestCount(), "Request count should be 19") + }.verifyComplete() + } + @Test fun `ReqShieldCacheable test - request to 'sampleService' should be only one times For Global Lock`() { val testProductId: String = UUID.randomUUID().toString() diff --git a/req-shield-spring-webflux-kotlin-coroutine-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/SampleService.kt b/req-shield-spring-webflux-kotlin-coroutine-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/SampleService.kt index a82aec8..de120e1 100644 --- a/req-shield-spring-webflux-kotlin-coroutine-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/SampleService.kt +++ b/req-shield-spring-webflux-kotlin-coroutine-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/SampleService.kt @@ -17,6 +17,7 @@ package com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.example.service import com.linecorp.cse.reqshield.kotlin.coroutine.ReqShield +import com.linecorp.cse.reqshield.kotlin.coroutine.config.ReqShieldWorkMode 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.example.dto.Product @@ -43,6 +44,21 @@ class SampleService( return Product(productId, "product_$productId") } + @ReqShieldCacheable( + cacheName = "productOnlyUpdateCache", + decisionForUpdate = 80, + timeToLiveMillis = 60 * 1000, + reqShieldWorkMode = ReqShieldWorkMode.ONLY_UPDATE_CACHE, + ) + suspend fun getProductOnlyUpdateCache(productId: String): Product { + log.info("find product with db request with req-shield local lock (will take 1 second)") + + delay(500) + atomicInteger.incrementAndGet() + + return Product(productId, "product_$productId") + } + suspend fun getProductNoAnno(productId: String): Product? { val result = reqShield diff --git a/req-shield-spring-webflux-kotlin-coroutine-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/CacheAnnotationTest.kt b/req-shield-spring-webflux-kotlin-coroutine-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/CacheAnnotationTest.kt index 207ee90..0dafade 100644 --- a/req-shield-spring-webflux-kotlin-coroutine-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/CacheAnnotationTest.kt +++ b/req-shield-spring-webflux-kotlin-coroutine-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/CacheAnnotationTest.kt @@ -66,6 +66,22 @@ class CacheAnnotationTest : AbstractRedisTest() { Assertions.assertEquals(1, sampleService.getRequestCount()) } + @Test + fun `ReqShieldCacheable test - request to 'sampleService' should be request count times(only update cache mode)`() = + runBlocking { + val testProductId: String = UUID.randomUUID().toString() + + List(20) { + async { + sampleService.getProductOnlyUpdateCache(testProductId) + } + }.awaitAll() + + delay(500) + + Assertions.assertEquals(20, sampleService.getRequestCount()) + } + @Test fun `ReqShieldCacheable test - request to 'sampleService' should be only one times For Global Lock`() = runBlocking { diff --git a/support/src/testFixtures/kotlin/com/linecorp/cse/reqshield/support/BaseReqShieldTest.kt b/support/src/testFixtures/kotlin/com/linecorp/cse/reqshield/support/BaseReqShieldTest.kt index ae555db..31f6e49 100644 --- a/support/src/testFixtures/kotlin/com/linecorp/cse/reqshield/support/BaseReqShieldTest.kt +++ b/support/src/testFixtures/kotlin/com/linecorp/cse/reqshield/support/BaseReqShieldTest.kt @@ -19,6 +19,8 @@ package com.linecorp.cse.reqshield.support interface BaseReqShieldTest { fun testSetMethodCacheNotExistsAndLocalLockAcquired() + fun testSetMethodCacheNotExistsAndOnlyUpdateCache() + fun testSetMethodCacheNotExistsAndGlobalLockAcquired() fun testSetMethodCacheNotExistsAndGlobalLockAcquiredAndDoesNotExistGlobalLockFunction() @@ -43,6 +45,8 @@ interface BaseReqShieldTest { fun testSetMethodCacheExistsAndTheUpdateTarget() + fun testSetMethodCacheExistsAndTheUpdateTargetOnlyCreateCache() + fun testSetMethodCacheExistsAndTheUpdateTargetAndCallableReturnNull() fun executeSetCacheFunctionShouldHandleExceptionFromCacheSetter() From 744276483683b47d1acff2e46ffe9150d205a4fc Mon Sep 17 00:00:00 2001 From: byungchan-lee Date: Thu, 27 Mar 2025 13:44:51 +0900 Subject: [PATCH 04/11] SpEL Support --- .../coroutine/aspect/ReqShieldAspect.kt | 40 ++++++++++------ .../coroutine/aspect/ReqShieldAspectTest.kt | 21 +++++---- .../spring/webflux/aspect/ReqShieldAspect.kt | 33 +++++++++---- .../webflux/aspect/ReqShieldAspectTest.kt | 22 +++++---- .../spring/aspect/ReqShieldAspect.kt | 34 ++++++++++---- .../test/kotlin/aspect/ReqShieldAspectTest.kt | 46 +++++++++---------- .../mvc/example/service/SampleService.kt | 6 +-- .../example/service/CacheAnnotationTest.kt | 8 ++-- .../webflux/example/service/SampleService.kt | 6 +-- .../example/service/CacheAnnotationTest.kt | 8 ++-- .../example/service/SampleService.kt | 11 +++-- .../coroutine/example/CacheAnnotationTest.kt | 8 ++-- .../cse/reqshield/service/SampleService.kt | 6 +-- .../spring/service/CacheAnnotationTest.kt | 8 ++-- .../webflux/example/service/SampleService.kt | 6 +-- .../example/service/CacheAnnotationTest.kt | 8 ++-- .../example/service/SampleService.kt | 6 +-- .../example/service/CacheAnnotationTest.kt | 8 ++-- 18 files changed, 166 insertions(+), 119 deletions(-) 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 d0f35bb..c40c548 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 @@ -27,7 +27,12 @@ import org.aspectj.lang.annotation.Around import org.aspectj.lang.annotation.Aspect import org.aspectj.lang.reflect.MethodSignature import org.springframework.cache.interceptor.SimpleKeyGenerator +import org.springframework.context.expression.MethodBasedEvaluationContext +import org.springframework.core.DefaultParameterNameDiscoverer import org.springframework.core.annotation.AnnotationUtils +import org.springframework.expression.EvaluationContext +import org.springframework.expression.Expression +import org.springframework.expression.spel.standard.SpelExpressionParser import org.springframework.stereotype.Component import org.springframework.util.StringUtils import reactor.core.publisher.Mono @@ -40,7 +45,7 @@ import kotlin.coroutines.Continuation class ReqShieldAspect( private val asyncCache: AsyncCache, ) { - private val keyGenerator = SimpleKeyGenerator() + private val spelParser = SpelExpressionParser() internal val reqShieldMap = ConcurrentHashMap>() @Around("execution(@com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.annotation.* * *(.., kotlin.coroutines.Continuation))") @@ -89,31 +94,36 @@ class ReqShieldAspect( internal fun getCacheableCacheKey(joinPoint: ProceedingJoinPoint): String { val annotation = getCacheableAnnotation(joinPoint) - return getCacheKeyOrDefault(annotation.cacheName, annotation.key, joinPoint) + return getCacheKeyOrDefault(annotation.key, joinPoint) } internal fun getCacheEvictCacheKey(joinPoint: ProceedingJoinPoint): String { val annotation = getCacheEvictAnnotation(joinPoint) - return getCacheKeyOrDefault(annotation.cacheName, annotation.key, joinPoint) + return getCacheKeyOrDefault(annotation.key, joinPoint) } private fun getCacheKeyOrDefault( - annotationCacheName: String, annotationCacheKey: String, joinPoint: ProceedingJoinPoint, ): String { - return annotationCacheKey.ifBlank { - run { - val params = - StringUtils.arrayToDelimitedString( - joinPoint.args - .filter { it !is Continuation<*> } - .toTypedArray(), - "_", - ) - return "$annotationCacheName-[$annotationCacheKey$params]" + val method = getTargetMethod(joinPoint) + val args = joinPoint.args.filter { it !is Continuation<*> }.toTypedArray() + val context: EvaluationContext = + MethodBasedEvaluationContext(joinPoint.target, method, args, DefaultParameterNameDiscoverer()) + + val key: String? = + if (StringUtils.hasText(annotationCacheKey)) { + val expression: Expression = spelParser.parseExpression(annotationCacheKey) + expression.getValue(context, String::class.java) + } else { + SimpleKeyGenerator.generateKey(joinPoint.target, method, args).toString() } + + if (key.isNullOrBlank()) { + throw IllegalArgumentException("Null key returned for cache method : $method") } + + return key } private fun getOrCreateReqShield(joinPoint: ProceedingJoinPoint): ReqShield = @@ -149,5 +159,5 @@ class ReqShieldAspect( } private fun generateReqShieldKey(joinPoint: ProceedingJoinPoint): String = - "${getCacheableAnnotation(joinPoint).cacheName}-${getCacheableAnnotation(joinPoint).key}" + "${getCacheableAnnotation(joinPoint).cacheName}-${getCacheableCacheKey(joinPoint)}" } 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 f4ab912..c2c5421 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 @@ -60,8 +60,9 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { private val method = kotlinMethod?.javaMethod private val cacheName = "testCacheName" - private val cacheKey = "testCacheKey" - private val argument = "testArgument" + private val cacheKey = "#paramMap['x'] + #paramMap['y']" + private val argument = mapOf("x" to "paramX", "y" to "paramY") + private val evaluatedKey = "paramXparamY" private val methodReturn = Product("testProduct", "testCategory") @BeforeEach @@ -109,7 +110,7 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { assertEquals(result, reqShieldData.value) assertTrue(reqShieldAspect.reqShieldMap.size == 1) - assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$cacheKey"]) + assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$evaluatedKey"]) } @Test @@ -130,7 +131,7 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { jobs.awaitAll() assertTrue(reqShieldAspect.reqShieldMap.size == 1) - assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$cacheKey"]) + assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$evaluatedKey"]) } @Test @@ -169,7 +170,7 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { reqShieldAspect.aroundTargetCacheable(joinPoint) - assertEquals("testCacheKey", reqShieldAspect.getCacheableCacheKey(joinPoint)) + assertEquals(evaluatedKey, reqShieldAspect.getCacheableCacheKey(joinPoint)) } @Test @@ -180,15 +181,15 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { key = cacheKey, ) - assertEquals(cacheKey, reqShieldAspect.getCacheableCacheKey(joinPoint)) + assertEquals(evaluatedKey, reqShieldAspect.getCacheableCacheKey(joinPoint)) } class TestBean { - @ReqShieldCacheable(cacheName = "TestCacheName") - suspend fun cacheableWithSingleArgument(testArgument: String): String = "ReturnValue: $testArgument" + @ReqShieldCacheable(cacheName = "TestCacheName", key = "#paramMap['x'] + #paramMap['y']") + suspend fun cacheableWithSingleArgument(paramMap: Map): String = "ReturnValue: $paramMap" - @ReqShieldCacheEvict(cacheName = "TestCacheName") - suspend fun evictWithSingleArgument(testArgument: String): Boolean { + @ReqShieldCacheEvict(cacheName = "TestCacheName", key = "#paramMap['x'] + #paramMap['y']") + suspend fun evictWithSingleArgument(paramMap: Map): Boolean { log.debug("cache eviction") return true } 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 87b5f27..1d0558b 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 @@ -26,7 +26,12 @@ import org.aspectj.lang.annotation.Around import org.aspectj.lang.annotation.Aspect import org.aspectj.lang.reflect.MethodSignature import org.springframework.cache.interceptor.SimpleKeyGenerator +import org.springframework.context.expression.MethodBasedEvaluationContext +import org.springframework.core.DefaultParameterNameDiscoverer import org.springframework.core.annotation.AnnotationUtils +import org.springframework.expression.EvaluationContext +import org.springframework.expression.Expression +import org.springframework.expression.spel.standard.SpelExpressionParser import org.springframework.stereotype.Component import org.springframework.util.StringUtils import reactor.core.publisher.Mono @@ -38,7 +43,7 @@ import java.util.concurrent.ConcurrentHashMap class ReqShieldAspect( private val asyncCache: AsyncCache, ) { - private val keyGenerator = SimpleKeyGenerator() + private val spelParser = SpelExpressionParser() internal val reqShieldMap = ConcurrentHashMap>() @Around("@annotation(com.linecorp.cse.reqshield.spring.webflux.annotation.ReqShieldCacheable)") @@ -75,7 +80,7 @@ class ReqShieldAspect( internal fun getCacheEvictCacheKey(joinPoint: ProceedingJoinPoint): String { val annotation = getCacheEvictAnnotation(joinPoint) - return getCacheKeyOrDefault(annotation.cacheName, annotation.key, joinPoint) + return getCacheKeyOrDefault(annotation.key, joinPoint) } internal fun getTargetMethod(joinPoint: ProceedingJoinPoint): Method = (joinPoint.signature as MethodSignature).method @@ -86,20 +91,30 @@ class ReqShieldAspect( internal fun getCacheableCacheKey(joinPoint: ProceedingJoinPoint): String { val annotation = getCacheableAnnotation(joinPoint) - return getCacheKeyOrDefault(annotation.cacheName, annotation.key, joinPoint) + return getCacheKeyOrDefault(annotation.key, joinPoint) } private fun getCacheKeyOrDefault( - annotationCacheName: String, annotationCacheKey: String, joinPoint: ProceedingJoinPoint, ): String { - return annotationCacheKey.ifBlank { - run { - val params = StringUtils.arrayToDelimitedString(joinPoint.args, "_") - return "$annotationCacheName-[$annotationCacheKey$params]" + val method = getTargetMethod(joinPoint) + val context: EvaluationContext = + MethodBasedEvaluationContext(joinPoint.target, method, joinPoint.args, DefaultParameterNameDiscoverer()) + + val key = + if (StringUtils.hasText(annotationCacheKey)) { + val expression: Expression = spelParser.parseExpression(annotationCacheKey) + expression.getValue(context, String::class.java) + } else { + SimpleKeyGenerator.generateKey(joinPoint.target, method, joinPoint.args).toString() } + + if (key.isNullOrBlank()) { + throw IllegalArgumentException("Null key returned for cache method : $method") } + + return key } private fun getOrCreateReqShield(joinPoint: ProceedingJoinPoint): ReqShield = @@ -135,5 +150,5 @@ class ReqShieldAspect( } private fun generateReqShieldKey(joinPoint: ProceedingJoinPoint): String = - "${getCacheableAnnotation(joinPoint).cacheName}-${getCacheableAnnotation(joinPoint).key}" + "${getCacheableAnnotation(joinPoint).cacheName}-${getCacheableCacheKey(joinPoint)}" } 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 62aee5b..cb80292 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 @@ -48,12 +48,13 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { ReflectionUtils.findMethod( TestBean::class.java, TestBean::cacheableWithSingleArgument.name, - String::class.java, + Map::class.java, ) private val cacheName = "testCacheName" - private val cacheKey = "testCacheKey" - private val argument = "testArgument" + private val cacheKey = "#paramMap['x'] + #paramMap['y']" + private val argument = mapOf("x" to "paramX", "y" to "paramY") + private val evaulatedKey = "paramXparamY" private val methodReturn = Product("testProduct", "testCategory") @BeforeEach @@ -103,7 +104,7 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { .assertNext { value -> assertEquals(reqShieldData.value, value) Assertions.assertTrue(reqShieldAspect.reqShieldMap.size == 1) - Assertions.assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$cacheKey"]) + Assertions.assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$evaulatedKey"]) }.verifyComplete() } @@ -128,7 +129,7 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { .assertNext { productList -> Assertions.assertTrue(reqShieldAspect.reqShieldMap.size == 1) println(reqShieldAspect.reqShieldMap.keys().toList()) - Assertions.assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$cacheKey"]) + Assertions.assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$evaulatedKey"]) }.verifyComplete() } @@ -174,7 +175,7 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { .create(result) .assertNext { value -> Assertions.assertEquals( - "testCacheKey", + evaulatedKey, reqShieldAspect.getCacheableCacheKey(joinPoint), ) }.verifyComplete() @@ -188,14 +189,15 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { key = cacheKey, ) - Assertions.assertEquals(cacheKey, reqShieldAspect.getCacheableCacheKey(joinPoint)) + Assertions.assertEquals(evaulatedKey, reqShieldAspect.getCacheableCacheKey(joinPoint)) } class TestBean { - @ReqShieldCacheable(cacheName = "TestCacheName") - fun cacheableWithSingleArgument(testArgument: String): Mono = Mono.justOrEmpty(Product("testProduct", "testCategory")) + @ReqShieldCacheable(cacheName = "TestCacheName", key = "#paramMap['x'] + #paramMap['y']") + fun cacheableWithSingleArgument(paramMap: Map): Mono = + Mono.justOrEmpty(Product("testProduct", "testCategory")) @ReqShieldCacheEvict(cacheName = "TestCacheName") - fun evictWithSingleArgument(testArgument: String): Mono = Mono.just(true) + fun evictWithSingleArgument(paramMap: Map): Mono = Mono.just(true) } } 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 be4f91c..2c24632 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 @@ -25,7 +25,13 @@ import org.aspectj.lang.ProceedingJoinPoint import org.aspectj.lang.annotation.Around import org.aspectj.lang.annotation.Aspect import org.aspectj.lang.reflect.MethodSignature +import org.springframework.cache.interceptor.SimpleKeyGenerator +import org.springframework.context.expression.MethodBasedEvaluationContext +import org.springframework.core.DefaultParameterNameDiscoverer import org.springframework.core.annotation.AnnotationUtils +import org.springframework.expression.EvaluationContext +import org.springframework.expression.Expression +import org.springframework.expression.spel.standard.SpelExpressionParser import org.springframework.stereotype.Component import org.springframework.util.StringUtils import java.lang.reflect.Method @@ -36,6 +42,7 @@ import java.util.concurrent.ConcurrentHashMap class ReqShieldAspect( private val reqShieldCache: ReqShieldCache, ) { + private val spelParser = SpelExpressionParser() internal val reqShieldMap = ConcurrentHashMap>() @Around("@annotation(com.linecorp.cse.reqshield.spring.annotation.ReqShieldCacheable)") @@ -106,28 +113,37 @@ class ReqShieldAspect( internal fun getCacheableCacheKey(joinPoint: ProceedingJoinPoint): String { val annotation = getCacheableAnnotation(joinPoint) - return getCacheKeyOrDefault(annotation.cacheName, annotation.key, joinPoint) + return getCacheKeyOrDefault(annotation.key, joinPoint) } internal fun getCacheEvictCacheKey(joinPoint: ProceedingJoinPoint): String { val annotation = getCacheEvictAnnotation(joinPoint) - return getCacheKeyOrDefault(annotation.cacheName, annotation.key, joinPoint) + return getCacheKeyOrDefault(annotation.key, joinPoint) } private fun getCacheKeyOrDefault( - annotationCacheName: String, annotationCacheKey: String, joinPoint: ProceedingJoinPoint, ): String { - return annotationCacheKey.ifBlank { - run { - val params = StringUtils.arrayToDelimitedString(joinPoint.args, "_") - - return "$annotationCacheName-[$annotationCacheKey$params]" + val method = getTargetMethod(joinPoint) + val context: EvaluationContext = + MethodBasedEvaluationContext(joinPoint.target, method, joinPoint.args, DefaultParameterNameDiscoverer()) + + val key = + if (StringUtils.hasText(annotationCacheKey)) { + val expression: Expression = spelParser.parseExpression(annotationCacheKey) + expression.getValue(context, String::class.java) + } else { + SimpleKeyGenerator.generateKey(joinPoint.target, method, joinPoint.args).toString() } + + if (key.isNullOrBlank()) { + throw IllegalArgumentException("Null key returned for cache method : $method") } + + return key } private fun generateReqShieldKey(joinPoint: ProceedingJoinPoint): String = - "${getCacheableAnnotation(joinPoint).cacheName}-${getCacheableAnnotation(joinPoint).key}" + "${getCacheableAnnotation(joinPoint).cacheName}-${getCacheableCacheKey(joinPoint)}" } diff --git a/core-spring/src/test/kotlin/aspect/ReqShieldAspectTest.kt b/core-spring/src/test/kotlin/aspect/ReqShieldAspectTest.kt index 72b1d87..d227ae8 100644 --- a/core-spring/src/test/kotlin/aspect/ReqShieldAspectTest.kt +++ b/core-spring/src/test/kotlin/aspect/ReqShieldAspectTest.kt @@ -44,13 +44,20 @@ private val log = LoggerFactory.getLogger(ReqShieldAspectTest::class.java) class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { private val reqShieldCache: ReqShieldCache = mockk() - private val targetObject = spyk(TestBean()) private val joinPoint = mockk() private val reqShieldAspect = spyk(ReqShieldAspect(reqShieldCache)) + private val targetObject = spyk(TestBean()) + private val method = + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithSingleArgument.name, + Map::class.java, + ) - private val cacheKey = "testCacheKey" private val cacheName = "testCacheName" - private val argument = "testArgument" + private val cacheKey = "#paramMap['x'] + #paramMap['y']" + private val argument = mapOf("x" to "paramX", "y" to "paramY") + private val evaluatedKey = "paramXparamY" private val methodReturn = Product("testProduct", "testCategory") @BeforeEach @@ -59,12 +66,7 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { every { joinPoint.args } returns arrayOf(argument) every { joinPoint.proceed() } answers { targetObject.cacheableWithSingleArgument(argument) } - every { reqShieldAspect.getTargetMethod(joinPoint) } returns - ReflectionUtils.findMethod( - TestBean::class.java, - TestBean::cacheableWithSingleArgument.name, - String::class.java, - )!! + every { reqShieldAspect.getTargetMethod(joinPoint) } returns method!! every { reqShieldAspect.getCacheableAnnotation(joinPoint) } returns ReqShieldCacheable( @@ -93,7 +95,7 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { Awaitility.await().atMost(Duration.ofMillis(BaseReqShieldTest.AWAIT_TIMEOUT)).untilAsserted { // then assertTrue(reqShieldAspect.reqShieldMap.size == 1) - assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$cacheKey"]) + assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$evaluatedKey"]) } } @@ -114,7 +116,7 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { Awaitility.await().atMost(Duration.ofMillis(BaseReqShieldTest.AWAIT_TIMEOUT)).untilAsserted { // then assertTrue(reqShieldAspect.reqShieldMap.size == 1) - assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$cacheKey"]) + assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$evaluatedKey"]) } } @@ -141,20 +143,16 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { override fun testCacheKeyGenerationUseGeneratedKey() { every { joinPoint.target } returns targetObject every { joinPoint.args } returns arrayOf(argument) - every { reqShieldAspect.getTargetMethod(joinPoint) } returns - ReflectionUtils.findMethod( - TestBean::class.java, - TestBean::cacheableWithSingleArgument.name, - String::class.java, - )!! + every { reqShieldAspect.getTargetMethod(joinPoint) } returns method!! every { reqShieldAspect.getCacheableAnnotation(joinPoint) } returns ReqShieldCacheable( cacheName = cacheName, + key = cacheKey, ) assertEquals( - "$cacheName-[testArgument]", + evaluatedKey, reqShieldAspect.getCacheableCacheKey(joinPoint), ) } @@ -167,18 +165,18 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { key = cacheKey, ) - assertEquals(cacheKey, reqShieldAspect.getCacheableCacheKey(joinPoint)) + assertEquals(evaluatedKey, reqShieldAspect.getCacheableCacheKey(joinPoint)) } class TestBean { - @ReqShieldCacheable(cacheName = "TestCacheName") - fun cacheableWithSingleArgument(testArgument: String): String { + @ReqShieldCacheable(cacheName = "TestCacheName", key = "#paramMap['x'] + #paramMap['y']") + fun cacheableWithSingleArgument(paramMap: Map): String { log.debug("method invoked") - return "ReturnValue: $testArgument" + return "ReturnValue: $paramMap" } - @ReqShieldCacheEvict(cacheName = "TestCacheName") - fun evictWithSingleArgument(testArgument: String) { + @ReqShieldCacheEvict(cacheName = "TestCacheName", key = "#paramMap['x'] + #paramMap['y']") + fun evictWithSingleArgument(paramMap: Map) { log.debug("cache eviction") } } diff --git a/req-shield-spring-boot3-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/mvc/example/service/SampleService.kt b/req-shield-spring-boot3-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/mvc/example/service/SampleService.kt index 3945e61..0585267 100644 --- a/req-shield-spring-boot3-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/mvc/example/service/SampleService.kt +++ b/req-shield-spring-boot3-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/mvc/example/service/SampleService.kt @@ -13,7 +13,7 @@ private val log = LoggerFactory.getLogger(SampleService::class.java) class SampleService { private val atomicInteger: AtomicInteger = AtomicInteger(0) - @ReqShieldCacheable(cacheName = "product", decisionForUpdate = 90, timeToLiveMillis = 60 * 1000) + @ReqShieldCacheable(cacheName = "product", key = "'product-' + #productId", decisionForUpdate = 90, timeToLiveMillis = 60 * 1000) fun getProduct(productId: String): Product { log.info("find product with db request - req-shield local lock (will take 1 second)") Thread.sleep(500) @@ -22,7 +22,7 @@ class SampleService { return Product(productId, "product_$productId") } - @ReqShieldCacheable(cacheName = "product", isLocalLock = false, decisionForUpdate = 90) + @ReqShieldCacheable(cacheName = "product", key = "'product-' + #productId", isLocalLock = false, decisionForUpdate = 90) fun getProductForGlobalLock(productId: String): Product { log.info("find product with db request - req-shield global lock (will take 1 second)") Thread.sleep(500) @@ -31,7 +31,7 @@ class SampleService { return Product(productId, "product_$productId") } - @ReqShieldCacheEvict(cacheName = "product") + @ReqShieldCacheEvict(cacheName = "product", key = "'product-' + #productId") fun removeProduct(productId: String) { log.info("remove product ($productId)") } 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 601a196..9e6ab2e 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 @@ -82,19 +82,19 @@ class CacheAnnotationTest : AbstractRedisTest() { sampleService.getProduct(testProductId) await().atMost(5, TimeUnit.SECONDS).until { - reqShieldCache.get("product-[$testProductId]") != null + reqShieldCache.get("product-$testProductId") != null } - assertNotNull(reqShieldCache.get("product-[$testProductId]")) + assertNotNull(reqShieldCache.get("product-$testProductId")) // when sampleService.removeProduct(testProductId) // then await().atMost(5, TimeUnit.SECONDS).until { - reqShieldCache.get("product-[$testProductId]") == null + reqShieldCache.get("product-$testProductId") == null } - assertNull(reqShieldCache.get("product-[$testProductId]")) + assertNull(reqShieldCache.get("product-$testProductId")) } } diff --git a/req-shield-spring-boot3-webflux-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/webflux/example/service/SampleService.kt b/req-shield-spring-boot3-webflux-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/webflux/example/service/SampleService.kt index 728e848..3011b8d 100644 --- a/req-shield-spring-boot3-webflux-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/webflux/example/service/SampleService.kt +++ b/req-shield-spring-boot3-webflux-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/webflux/example/service/SampleService.kt @@ -18,7 +18,7 @@ class SampleService( ) { private val atomicInteger: AtomicInteger = AtomicInteger(0) - @ReqShieldCacheable(cacheName = "product", decisionForUpdate = 80, timeToLiveMillis = 60 * 1000) + @ReqShieldCacheable(cacheName = "product", decisionForUpdate = 80, timeToLiveMillis = 60 * 1000, key = "'product-' + #productId") fun getProduct(productId: String): Mono = Mono .delay(Duration.ofMillis(500)) @@ -48,7 +48,7 @@ class SampleService( 60 * 1000, ).mapNotNull { it.value } - @ReqShieldCacheable(cacheName = "product", isLocalLock = false, decisionForUpdate = 80) + @ReqShieldCacheable(cacheName = "product", isLocalLock = false, decisionForUpdate = 80, key = "'product-' + #productId") fun getProductForGlobalLock(productId: String): Mono = Mono .delay(Duration.ofMillis(500)) @@ -60,7 +60,7 @@ class SampleService( }.doFinally { atomicInteger.incrementAndGet() }, ) - @ReqShieldCacheEvict(cacheName = "product") + @ReqShieldCacheEvict(cacheName = "product", key = "'product-' + #productId") fun removeProduct(productId: String): Mono { log.info("remove product ($productId)") return Mono.fromCallable { true } 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 f7396c7..c424d99 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 @@ -88,9 +88,9 @@ class CacheAnnotationTest : AbstractRedisTest() { // then await().atMost(5, TimeUnit.SECONDS).until { - asyncCache.get("product-[$testProductId]").block() != null + asyncCache.get("product-$testProductId").block() != null } - val cacheMono = asyncCache.get("product-[$testProductId]").block() + val cacheMono = asyncCache.get("product-$testProductId").block() assertNotNull(cacheMono) // when @@ -106,9 +106,9 @@ class CacheAnnotationTest : AbstractRedisTest() { // then await().atMost(5, TimeUnit.SECONDS).until { - asyncCache.get("product-[$testProductId]").block() == null + asyncCache.get("product-$testProductId").block() == null } - val cacheMonoNull = asyncCache.get("product-[$testProductId]").block() + val cacheMonoNull = asyncCache.get("product-$testProductId").block() assertNull(cacheMonoNull) } } diff --git a/req-shield-spring-boot3-webflux-kotlin-coroutine-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/webflux/kotlin/coroutine/example/service/SampleService.kt b/req-shield-spring-boot3-webflux-kotlin-coroutine-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/webflux/kotlin/coroutine/example/service/SampleService.kt index 2071153..812b0d0 100644 --- a/req-shield-spring-boot3-webflux-kotlin-coroutine-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/webflux/kotlin/coroutine/example/service/SampleService.kt +++ b/req-shield-spring-boot3-webflux-kotlin-coroutine-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/webflux/kotlin/coroutine/example/service/SampleService.kt @@ -17,7 +17,12 @@ class SampleService( ) { private val atomicInteger: AtomicInteger = AtomicInteger(0) - @ReqShieldCacheable(cacheName = "product", decisionForUpdate = 80, timeToLiveMillis = 60 * 1000) + @ReqShieldCacheable( + cacheName = "product", + key = "'product-' + #productId", + decisionForUpdate = 80, + timeToLiveMillis = 60 * 1000, + ) suspend fun getProduct(productId: String): Product { log.info("find product with db request with req-shield local lock (will take 1 second)") @@ -45,7 +50,7 @@ class SampleService( @ReqShieldCacheable( cacheName = "product", - key = "global_lock", + key = "'product-' + #productId", isLocalLock = false, decisionForUpdate = 70, timeToLiveMillis = 60 * 1000, @@ -59,7 +64,7 @@ class SampleService( return Product(productId, "product_$productId") } - @ReqShieldCacheEvict(cacheName = "product") + @ReqShieldCacheEvict(cacheName = "product", key = "'product-' + #productId") suspend fun removeProduct(productId: String) { log.info("remove product ($productId)") } diff --git a/req-shield-spring-boot3-webflux-kotlin-coroutine-example/src/test/kotlin/com/linecorp/cse/reqshield/spring3/webflux/kotlin/coroutine/example/CacheAnnotationTest.kt b/req-shield-spring-boot3-webflux-kotlin-coroutine-example/src/test/kotlin/com/linecorp/cse/reqshield/spring3/webflux/kotlin/coroutine/example/CacheAnnotationTest.kt index 53c2eba..78be8f6 100644 --- a/req-shield-spring-boot3-webflux-kotlin-coroutine-example/src/test/kotlin/com/linecorp/cse/reqshield/spring3/webflux/kotlin/coroutine/example/CacheAnnotationTest.kt +++ b/req-shield-spring-boot3-webflux-kotlin-coroutine-example/src/test/kotlin/com/linecorp/cse/reqshield/spring3/webflux/kotlin/coroutine/example/CacheAnnotationTest.kt @@ -77,7 +77,7 @@ class CacheAnnotationTest : AbstractRedisTest() { val maxAttempts = 30 var attempts = 0 - while (asyncCache.get("product-[$testProductId]") == null) { + while (asyncCache.get("product-$testProductId") == null) { if (attempts >= maxAttempts) { break } @@ -85,13 +85,13 @@ class CacheAnnotationTest : AbstractRedisTest() { delay(100) } - assertNotNull(asyncCache.get("product-[$testProductId]")) + assertNotNull(asyncCache.get("product-$testProductId")) // when sampleService.removeProduct(testProductId) var attemptsSecond = 0 - while (asyncCache.get("product-[$testProductId]") != null) { + while (asyncCache.get("product-$testProductId") != null) { if (attemptsSecond >= maxAttempts) { break } @@ -99,6 +99,6 @@ class CacheAnnotationTest : AbstractRedisTest() { delay(100) } - assertNull(asyncCache.get("product-[$testProductId]")) + assertNull(asyncCache.get("product-$testProductId")) } } diff --git a/req-shield-spring-example/src/main/kotlin/com/linecorp/cse/reqshield/service/SampleService.kt b/req-shield-spring-example/src/main/kotlin/com/linecorp/cse/reqshield/service/SampleService.kt index 8449659..3e7e64f 100644 --- a/req-shield-spring-example/src/main/kotlin/com/linecorp/cse/reqshield/service/SampleService.kt +++ b/req-shield-spring-example/src/main/kotlin/com/linecorp/cse/reqshield/service/SampleService.kt @@ -30,7 +30,7 @@ private val log = LoggerFactory.getLogger(SampleService::class.java) class SampleService { private val atomicInteger: AtomicInteger = AtomicInteger(0) - @ReqShieldCacheable(cacheName = "product", decisionForUpdate = 90, timeToLiveMillis = 60 * 1000) + @ReqShieldCacheable(cacheName = "product", key = "'product-' + #productId", decisionForUpdate = 90, timeToLiveMillis = 60 * 1000) fun getProduct(productId: String): Product { log.info("find product with db request - req-shield local lock (will take 1 second)") Thread.sleep(500) @@ -53,7 +53,7 @@ class SampleService { return Product(productId, "product_$productId") } - @ReqShieldCacheable(cacheName = "product", isLocalLock = false, decisionForUpdate = 90) + @ReqShieldCacheable(cacheName = "product", key = "'product-' + #productId", isLocalLock = false, decisionForUpdate = 90) fun getProductForGlobalLock(productId: String): Product { log.info("find product with db request - req-shield global lock (will take 1 second)") Thread.sleep(500) @@ -62,7 +62,7 @@ class SampleService { return Product(productId, "product_$productId") } - @ReqShieldCacheEvict(cacheName = "product") + @ReqShieldCacheEvict(cacheName = "product", key = "'product-' + #productId") fun removeProduct(productId: String) { log.info("remove product ($productId)") } diff --git a/req-shield-spring-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/service/CacheAnnotationTest.kt b/req-shield-spring-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/service/CacheAnnotationTest.kt index 1337691..b1718a5 100644 --- a/req-shield-spring-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/service/CacheAnnotationTest.kt +++ b/req-shield-spring-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/service/CacheAnnotationTest.kt @@ -119,19 +119,19 @@ class CacheAnnotationTest : AbstractRedisTest() { sampleService.getProduct(testProductId) await().atMost(5, TimeUnit.SECONDS).until { - reqShieldCache.get("product-[$testProductId]") != null + reqShieldCache.get("product-$testProductId") != null } - assertNotNull(reqShieldCache.get("product-[$testProductId]")) + assertNotNull(reqShieldCache.get("product-$testProductId")) // when sampleService.removeProduct(testProductId) // then await().atMost(5, TimeUnit.SECONDS).until { - reqShieldCache.get("product-[$testProductId]") == null + reqShieldCache.get("product-$testProductId") == null } - assertNull(reqShieldCache.get("product-[$testProductId]")) + assertNull(reqShieldCache.get("product-$testProductId")) } } diff --git a/req-shield-spring-webflux-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/example/service/SampleService.kt b/req-shield-spring-webflux-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/example/service/SampleService.kt index d15f73d..798ec24 100644 --- a/req-shield-spring-webflux-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/example/service/SampleService.kt +++ b/req-shield-spring-webflux-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/example/service/SampleService.kt @@ -35,7 +35,7 @@ class SampleService( ) { private val atomicInteger: AtomicInteger = AtomicInteger(0) - @ReqShieldCacheable(cacheName = "product", decisionForUpdate = 80, timeToLiveMillis = 60 * 1000) + @ReqShieldCacheable(cacheName = "product", key = "'product-' + #productId", decisionForUpdate = 80, timeToLiveMillis = 60 * 1000) fun getProduct(productId: String): Mono = Mono .delay(Duration.ofMillis(500)) @@ -82,7 +82,7 @@ class SampleService( 60 * 1000, ).mapNotNull { it.value } - @ReqShieldCacheable(cacheName = "product", isLocalLock = false, decisionForUpdate = 80) + @ReqShieldCacheable(cacheName = "product", key = "'product-' + #productId", isLocalLock = false, decisionForUpdate = 80) fun getProductForGlobalLock(productId: String): Mono = Mono .delay(Duration.ofMillis(500)) @@ -94,7 +94,7 @@ class SampleService( }.doFinally { atomicInteger.incrementAndGet() }, ) - @ReqShieldCacheEvict(cacheName = "product") + @ReqShieldCacheEvict(cacheName = "product", key = "'product-' + #productId") fun removeProduct(productId: String): Mono { log.info("remove product ($productId)") return Mono.fromCallable { true } 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 b84a973..b5c3175 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 @@ -126,9 +126,9 @@ class CacheAnnotationTest : AbstractRedisTest() { // then await().atMost(5, TimeUnit.SECONDS).until { - asyncCache.get("product-[$testProductId]").block() != null + asyncCache.get("product-$testProductId").block() != null } - val cacheMono = asyncCache.get("product-[$testProductId]").block() + val cacheMono = asyncCache.get("product-$testProductId").block() assertNotNull(cacheMono) // when @@ -144,9 +144,9 @@ class CacheAnnotationTest : AbstractRedisTest() { // then await().atMost(5, TimeUnit.SECONDS).until { - asyncCache.get("product-[$testProductId]").block() == null + asyncCache.get("product-$testProductId").block() == null } - val cacheMonoNull = asyncCache.get("product-[$testProductId]").block() + val cacheMonoNull = asyncCache.get("product-$testProductId").block() assertNull(cacheMonoNull) } } diff --git a/req-shield-spring-webflux-kotlin-coroutine-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/SampleService.kt b/req-shield-spring-webflux-kotlin-coroutine-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/SampleService.kt index de120e1..0efd2ae 100644 --- a/req-shield-spring-webflux-kotlin-coroutine-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/SampleService.kt +++ b/req-shield-spring-webflux-kotlin-coroutine-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/SampleService.kt @@ -34,7 +34,7 @@ class SampleService( ) { private val atomicInteger: AtomicInteger = AtomicInteger(0) - @ReqShieldCacheable(cacheName = "product", decisionForUpdate = 80, timeToLiveMillis = 60 * 1000) + @ReqShieldCacheable(cacheName = "product", key = "'product-' + #productId", decisionForUpdate = 80, timeToLiveMillis = 60 * 1000) suspend fun getProduct(productId: String): Product { log.info("find product with db request with req-shield local lock (will take 1 second)") @@ -77,7 +77,7 @@ class SampleService( @ReqShieldCacheable( cacheName = "product", - key = "global_lock", + key = "'product-' + #productId", isLocalLock = false, decisionForUpdate = 70, timeToLiveMillis = 60 * 1000, @@ -91,7 +91,7 @@ class SampleService( return Product(productId, "product_$productId") } - @ReqShieldCacheEvict(cacheName = "product") + @ReqShieldCacheEvict(cacheName = "product", key = "'product-' + #productId") suspend fun removeProduct(productId: String) { log.info("remove product ($productId)") } diff --git a/req-shield-spring-webflux-kotlin-coroutine-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/CacheAnnotationTest.kt b/req-shield-spring-webflux-kotlin-coroutine-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/CacheAnnotationTest.kt index 0dafade..49a9763 100644 --- a/req-shield-spring-webflux-kotlin-coroutine-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/CacheAnnotationTest.kt +++ b/req-shield-spring-webflux-kotlin-coroutine-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/CacheAnnotationTest.kt @@ -108,7 +108,7 @@ class CacheAnnotationTest : AbstractRedisTest() { val maxAttempts = 30 var attempts = 0 - while (asyncCache.get("product-[$testProductId]") == null) { + while (asyncCache.get("product-$testProductId") == null) { if (attempts >= maxAttempts) { break } @@ -116,13 +116,13 @@ class CacheAnnotationTest : AbstractRedisTest() { delay(100) } - assertNotNull(asyncCache.get("product-[$testProductId]")) + assertNotNull(asyncCache.get("product-$testProductId")) // when sampleService.removeProduct(testProductId) var attemptsSecond = 0 - while (asyncCache.get("product-[$testProductId]") != null) { + while (asyncCache.get("product-$testProductId") != null) { if (attemptsSecond >= maxAttempts) { break } @@ -130,6 +130,6 @@ class CacheAnnotationTest : AbstractRedisTest() { delay(100) } - assertNull(asyncCache.get("product-[$testProductId]")) + assertNull(asyncCache.get("product-$testProductId")) } } From b068e1eb808ca09437ae5f4195de4196b0f5f9f7 Mon Sep 17 00:00:00 2001 From: byungchan-lee Date: Thu, 27 Mar 2025 13:45:28 +0900 Subject: [PATCH 05/11] SpEL Support (by spring branch for kotlin coroutine) --- .../coroutine/aspect/ReqShieldAspect.kt | 19 ++++++++++++++++++- 1 file changed, 18 insertions(+), 1 deletion(-) 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 c40c548..5f14ab0 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 @@ -29,6 +29,7 @@ import org.aspectj.lang.reflect.MethodSignature import org.springframework.cache.interceptor.SimpleKeyGenerator import org.springframework.context.expression.MethodBasedEvaluationContext import org.springframework.core.DefaultParameterNameDiscoverer +import org.springframework.core.SpringVersion import org.springframework.core.annotation.AnnotationUtils import org.springframework.expression.EvaluationContext import org.springframework.expression.Expression @@ -45,6 +46,7 @@ import kotlin.coroutines.Continuation class ReqShieldAspect( private val asyncCache: AsyncCache, ) { + private val springVersion = SpringVersion.getVersion() private val spelParser = SpelExpressionParser() internal val reqShieldMap = ConcurrentHashMap>() @@ -107,7 +109,13 @@ class ReqShieldAspect( joinPoint: ProceedingJoinPoint, ): String { val method = getTargetMethod(joinPoint) - val args = joinPoint.args.filter { it !is Continuation<*> }.toTypedArray() + val args = + if (isCoroutineSupportedSpringVersion()) { + joinPoint.args + } else { + joinPoint.args.filter { it !is Continuation<*> }.toTypedArray() + } + val context: EvaluationContext = MethodBasedEvaluationContext(joinPoint.target, method, args, DefaultParameterNameDiscoverer()) @@ -158,6 +166,15 @@ class ReqShieldAspect( return ReqShield(reqShieldConfiguration) } + private fun isCoroutineSupportedSpringVersion(): Boolean { + val version = springVersion ?: return false + val parts = version.split(".") + val major = parts.getOrNull(0)?.toIntOrNull() ?: return false + val minor = parts.getOrNull(1)?.toIntOrNull() ?: return false + + return major > 6 || (major == 6 && minor >= 1) + } + private fun generateReqShieldKey(joinPoint: ProceedingJoinPoint): String = "${getCacheableAnnotation(joinPoint).cacheName}-${getCacheableCacheKey(joinPoint)}" } From 2797e0841be5df86816e27db2b88929698c1a7d8 Mon Sep 17 00:00:00 2001 From: byungchan-lee Date: Fri, 28 Mar 2025 18:25:16 +0900 Subject: [PATCH 06/11] keyGenerator support --- .../annotation/ReqShieldCacheEvict.kt | 1 + .../annotation/ReqShieldCacheable.kt | 1 + .../coroutine/aspect/ReqShieldAspect.kt | 53 ++++- .../coroutine/aspect/ReqShieldAspectTest.kt | 173 ++++++++++------- .../webflux/annotation/ReqShieldCacheEvict.kt | 1 + .../webflux/annotation/ReqShieldCacheable.kt | 1 + .../spring/webflux/aspect/ReqShieldAspect.kt | 69 +++++-- .../webflux/aspect/ReqShieldAspectTest.kt | 183 +++++++++++------- .../spring/annotation/ReqShieldCacheEvict.kt | 1 + .../spring/annotation/ReqShieldCacheable.kt | 1 + .../spring/aspect/ReqShieldAspect.kt | 51 ++++- .../test/kotlin/aspect/ReqShieldAspectTest.kt | 179 +++++++++++------ .../support/BaseReqShieldModuleSupportTest.kt | 12 +- 13 files changed, 486 insertions(+), 240 deletions(-) diff --git a/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheEvict.kt b/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheEvict.kt index 925902f..b6e09af 100644 --- a/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheEvict.kt +++ b/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheEvict.kt @@ -24,6 +24,7 @@ import java.lang.annotation.Inherited annotation class ReqShieldCacheEvict( val cacheName: String, val key: String = "", + val keyGenerator: String = "", val isLocalLock: Boolean = true, val lockTimeoutMillis: Long = 3000, val condition: String = "", diff --git a/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheable.kt b/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheable.kt index ce606c5..ff101b7 100644 --- a/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheable.kt +++ b/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheable.kt @@ -26,6 +26,7 @@ import java.lang.annotation.Inherited annotation class ReqShieldCacheable( val cacheName: String, val key: String = "", + val keyGenerator: String = "", val isLocalLock: Boolean = true, val lockTimeoutMillis: Long = 3000, val decisionForUpdate: Int = 90, 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 5f14ab0..0fd9777 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 @@ -26,6 +26,9 @@ import org.aspectj.lang.ProceedingJoinPoint import org.aspectj.lang.annotation.Around import org.aspectj.lang.annotation.Aspect import org.aspectj.lang.reflect.MethodSignature +import org.springframework.beans.factory.BeanFactory +import org.springframework.beans.factory.BeanFactoryAware +import org.springframework.cache.interceptor.KeyGenerator import org.springframework.cache.interceptor.SimpleKeyGenerator import org.springframework.context.expression.MethodBasedEvaluationContext import org.springframework.core.DefaultParameterNameDiscoverer @@ -36,6 +39,7 @@ import org.springframework.expression.Expression import org.springframework.expression.spel.standard.SpelExpressionParser import org.springframework.stereotype.Component import org.springframework.util.StringUtils +import org.springframework.util.function.SingletonSupplier import reactor.core.publisher.Mono import java.lang.reflect.Method import java.util.concurrent.ConcurrentHashMap @@ -45,9 +49,13 @@ import kotlin.coroutines.Continuation @Component class ReqShieldAspect( private val asyncCache: AsyncCache, -) { +) : BeanFactoryAware { + private lateinit var beanFactory: BeanFactory private val springVersion = SpringVersion.getVersion() private val spelParser = SpelExpressionParser() + private var defaultKeyGenerator = SingletonSupplier.of { SimpleKeyGenerator() } + + private val keyGeneratorMap = ConcurrentHashMap() internal val reqShieldMap = ConcurrentHashMap>() @Around("execution(@com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.annotation.* * *(.., kotlin.coroutines.Continuation))") @@ -92,20 +100,25 @@ class ReqShieldAspect( fun getCacheEvictAnnotation(joinPoint: ProceedingJoinPoint): ReqShieldCacheEvict = AnnotationUtils.getAnnotation(getTargetMethod(joinPoint), ReqShieldCacheEvict::class.java) - ?: throw IllegalArgumentException("ReqShieldCacheable annotation is required") + ?: throw IllegalArgumentException("ReqShieldCacheEvict annotation is required") internal fun getCacheableCacheKey(joinPoint: ProceedingJoinPoint): String { val annotation = getCacheableAnnotation(joinPoint) - return getCacheKeyOrDefault(annotation.key, joinPoint) + validateCacheKey(annotation.key, annotation.keyGenerator) + + return getCacheKeyOrDefault(annotation.key, annotation.keyGenerator, joinPoint) } internal fun getCacheEvictCacheKey(joinPoint: ProceedingJoinPoint): String { val annotation = getCacheEvictAnnotation(joinPoint) - return getCacheKeyOrDefault(annotation.key, joinPoint) + validateCacheKey(annotation.key, annotation.keyGenerator) + + return getCacheKeyOrDefault(annotation.key, annotation.keyGenerator, joinPoint) } private fun getCacheKeyOrDefault( annotationCacheKey: String, + annotationCacheKeyGenerator: String, joinPoint: ProceedingJoinPoint, ): String { val method = getTargetMethod(joinPoint) @@ -119,17 +132,16 @@ class ReqShieldAspect( val context: EvaluationContext = MethodBasedEvaluationContext(joinPoint.target, method, args, DefaultParameterNameDiscoverer()) - val key: String? = + val key = if (StringUtils.hasText(annotationCacheKey)) { val expression: Expression = spelParser.parseExpression(annotationCacheKey) expression.getValue(context, String::class.java) } else { - SimpleKeyGenerator.generateKey(joinPoint.target, method, args).toString() + val keyGenerator = getOrCreateKeyGenerator(annotationCacheKeyGenerator) + keyGenerator.generate(joinPoint.target, method, args).toString() } - if (key.isNullOrBlank()) { - throw IllegalArgumentException("Null key returned for cache method : $method") - } + require(!key.isNullOrBlank()) { "Null key returned for cache method : $method" } return key } @@ -166,6 +178,25 @@ class ReqShieldAspect( return ReqShield(reqShieldConfiguration) } + private fun validateCacheKey( + cacheKey: String, + cacheKeyGenerator: String, + ) { + if (cacheKey.isNotBlank() && cacheKeyGenerator.isNotBlank()) { + throw IllegalArgumentException("The key and keyGenerator attributes are mutually exclusive.") + } + } + + private fun getOrCreateKeyGenerator(keyGeneratorBeanName: String?): KeyGenerator { + if (keyGeneratorBeanName.isNullOrBlank()) { + return defaultKeyGenerator.obtain() + } + + return keyGeneratorMap.computeIfAbsent(keyGeneratorBeanName) { + beanFactory.getBean(it, KeyGenerator::class.java) + } + } + private fun isCoroutineSupportedSpringVersion(): Boolean { val version = springVersion ?: return false val parts = version.split(".") @@ -177,4 +208,8 @@ class ReqShieldAspect( private fun generateReqShieldKey(joinPoint: ProceedingJoinPoint): String = "${getCacheableAnnotation(joinPoint).cacheName}-${getCacheableCacheKey(joinPoint)}" + + override fun setBeanFactory(beanFactory: BeanFactory) { + this.beanFactory = beanFactory + } } 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 c2c5421..85d2398 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 @@ -31,12 +31,15 @@ import kotlinx.coroutines.async import kotlinx.coroutines.awaitAll import kotlinx.coroutines.test.runTest import org.aspectj.lang.ProceedingJoinPoint -import org.aspectj.lang.reflect.MethodSignature import org.junit.jupiter.api.Assertions.assertNotNull import org.junit.jupiter.api.Assertions.assertTrue import org.junit.jupiter.api.BeforeEach import org.junit.jupiter.api.Test import org.slf4j.LoggerFactory +import org.springframework.beans.factory.BeanFactory +import org.springframework.cache.interceptor.KeyGenerator +import org.springframework.cache.interceptor.SimpleKeyGenerator +import java.lang.reflect.Method import kotlin.coroutines.Continuation import kotlin.coroutines.EmptyCoroutineContext import kotlin.reflect.full.functions @@ -49,77 +52,69 @@ private val log = LoggerFactory.getLogger(ReqShieldAspectTest::class.java) class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { private val asyncCache: AsyncCache = mockk() private val joinPoint: ProceedingJoinPoint = mockk() - private val methodSignature: MethodSignature = mockk() private val reqShieldAspect: ReqShieldAspect = spyk(ReqShieldAspect(asyncCache)) private val targetObject = spyk(TestBean()) + private val argument = mapOf("x" to "paramX", "y" to "paramY") private val mockContinuation = mockk>() - private val kotlinMethod = - TestBean::class.functions.find { - it.name == TestBean::cacheableWithSingleArgument.name && it.parameters.size == 2 - } - private val method = kotlinMethod?.javaMethod +// private val kotlinMethod = +// TestBean::class.functions.find { +// it.name == TestBean::cacheableWithSingleArgument.name && it.parameters.size == 2 +// } +// private val method = kotlinMethod?.javaMethod + + private val cacheName = "TestCacheName" + private val cacheKeyGenerator = "customGenerator" + private val spelEvaluatedKey = "paramXparamY" + private val keyGeneratorKey = "KeyGeneratedByGenerator" + + private val beanFactory = mockk() - private val cacheName = "testCacheName" - private val cacheKey = "#paramMap['x'] + #paramMap['y']" - private val argument = mapOf("x" to "paramX", "y" to "paramY") - private val evaluatedKey = "paramXparamY" private val methodReturn = Product("testProduct", "testCategory") @BeforeEach fun setUp() { - every { joinPoint.signature } returns methodSignature - - every { methodSignature.method } returns method every { mockContinuation.context } returns EmptyCoroutineContext every { joinPoint.args } returns arrayOf(argument, mockContinuation) every { joinPoint.target } returns targetObject - val reqShieldCacheable: ReqShieldCacheable = mockk() - every { reqShieldCacheable.key } returns cacheKey - every { reqShieldCacheable.timeToLiveMillis } returns 60 - every { reqShieldCacheable.isLocalLock } returns false - every { reqShieldCacheable.lockTimeoutMillis } returns 0 - every { reqShieldCacheable.decisionForUpdate } returns 70 - every { reqShieldCacheable.cacheName } returns "product" - - every { reqShieldAspect.getCacheableAnnotation(joinPoint) } returns - ReqShieldCacheable( - cacheName = cacheName, - key = cacheKey, - lockTimeoutMillis = 1000, - timeToLiveMillis = 1000, - ) - - every { reqShieldAspect.getCacheEvictAnnotation(joinPoint) } returns - ReqShieldCacheEvict( - cacheName = cacheName, - lockTimeoutMillis = 1000, - ) + reqShieldAspect.setBeanFactory(beanFactory) } @Test - override fun testAspectOperationVerifyReqShieldAndCacheCreation() = + override fun verifyReqShieldCacheCreation() = runTest { // Mock the cache data using mockk val reqShieldData = ReqShieldData(methodReturn, 1000) coEvery { asyncCache.get(any()) } returns reqShieldData - coEvery { joinPoint.proceed() } coAnswers { targetObject.cacheableWithSingleArgument(argument) } + coEvery { joinPoint.proceed() } coAnswers { targetObject.cacheableWithCustomKey(argument) } + coEvery { reqShieldAspect.getTargetMethod(joinPoint) } returns + TestBean::class + .functions + .find { + it.name == TestBean::cacheableWithCustomKey.name && it.parameters.size == 2 + }?.javaMethod!! // Test the aroundTargetCacheable method val result = reqShieldAspect.aroundTargetCacheable(joinPoint) assertEquals(result, reqShieldData.value) assertTrue(reqShieldAspect.reqShieldMap.size == 1) - assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$evaluatedKey"]) + assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$spelEvaluatedKey"]) } @Test - override fun testAspectOperationReqShieldObjectShouldBeCreatedOnce() = + override fun reqShieldObjectShouldBeCreatedOnce() = runTest { // Mock the cache data using mockk val reqShieldData = ReqShieldData(methodReturn, 1000) coEvery { asyncCache.get(any()) } returns reqShieldData - coEvery { joinPoint.proceed() } coAnswers { targetObject.cacheableWithSingleArgument(argument) } + coEvery { joinPoint.proceed() } coAnswers { targetObject.cacheableWithCustomKey(argument) } + coEvery { reqShieldAspect.getTargetMethod(joinPoint) } returns + TestBean::class + .functions + .find { + it.name == TestBean::cacheableWithCustomKey.name && it.parameters.size == 2 + }?.javaMethod!! val jobs = List(20) { @@ -131,67 +126,109 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { jobs.awaitAll() assertTrue(reqShieldAspect.reqShieldMap.size == 1) - assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$evaluatedKey"]) + assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$spelEvaluatedKey"]) } @Test - override fun testAspectOperationCacheEviction() = + override fun verifyReqShieldCacheEviction() = runTest { // Mock the cache data using mockk val reqShieldData = ReqShieldData(methodReturn, 1000) coEvery { asyncCache.get(any()) } returns reqShieldData - coEvery { asyncCache.evict(any()) } returns true - coEvery { joinPoint.proceed() } coAnswers { targetObject.cacheableWithSingleArgument(argument) } + coEvery { joinPoint.proceed() } coAnswers { targetObject.cacheableWithCustomKey(argument) } + coEvery { reqShieldAspect.getTargetMethod(joinPoint) } returns + TestBean::class + .functions + .find { + it.name == TestBean::cacheableWithDefaultKeyGenerator.name && it.parameters.size == 2 + }?.javaMethod!! // Test the aroundTargetCacheable method val result = reqShieldAspect.aroundTargetCacheable(joinPoint) assertEquals(reqShieldData.value, result) - val kotlinMethod = - TestBean::class.functions.find { - it.name == TestBean::evictWithSingleArgument.name && it.parameters.size == 2 - } - val method = kotlinMethod?.javaMethod - every { methodSignature.method } returns method + // Validate cache eviction + coEvery { asyncCache.evict(any()) } returns true + coEvery { 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) } - coEvery { joinPoint.proceed() } coAnswers { targetObject.evictWithSingleArgument(argument) } val removeProductMono = reqShieldAspect.aroundTargetCacheable(joinPoint) assertTrue(removeProductMono as Boolean) } @Test - override fun testCacheKeyGenerationUseGeneratedKey() = + override fun verifyCacheKeyGenerationWithSpEL() = runTest { - val reqShieldData = ReqShieldData(methodReturn, 1000) - coEvery { asyncCache.get(any()) } returns reqShieldData - coEvery { joinPoint.proceed() } coAnswers { targetObject.cacheableWithSingleArgument(argument) } - - reqShieldAspect.aroundTargetCacheable(joinPoint) + coEvery { reqShieldAspect.getTargetMethod(joinPoint) } returns + TestBean::class + .functions + .find { + it.name == TestBean::cacheableWithCustomKey.name && it.parameters.size == 2 + }?.javaMethod!! + + assertEquals(spelEvaluatedKey, reqShieldAspect.getCacheableCacheKey(joinPoint)) + } - assertEquals(evaluatedKey, reqShieldAspect.getCacheableCacheKey(joinPoint)) + @Test + override fun verifyCacheKeyGenerationWithKeyGenerator() = + runTest { + coEvery { beanFactory.getBean(cacheKeyGenerator, KeyGenerator::class.java) } returns + CustomGenerator() + coEvery { reqShieldAspect.getTargetMethod(joinPoint) } returns + TestBean::class + .functions + .find { + it.name == TestBean::cacheableWithKeyGenerator.name && it.parameters.size == 2 + }?.javaMethod!! + + assertEquals(keyGeneratorKey, reqShieldAspect.getCacheableCacheKey(joinPoint)) } @Test - override fun testCacheKeyGenerationCacheKeyShouldBeSuppliedKey() { - every { reqShieldAspect.getCacheableAnnotation(joinPoint) } returns - ReqShieldCacheable( - cacheName = cacheName, - key = cacheKey, + override fun verifyCacheKeyGenerationWithDefaultGenerator() = + runTest { + coEvery { reqShieldAspect.getTargetMethod(joinPoint) } returns + TestBean::class + .functions + .find { + it.name == TestBean::cacheableWithDefaultKeyGenerator.name && it.parameters.size == 2 + }?.javaMethod!! + + assertEquals( + SimpleKeyGenerator.generateKey(arrayOf(argument)).toString(), + reqShieldAspect.getCacheableCacheKey(joinPoint), ) - - assertEquals(evaluatedKey, reqShieldAspect.getCacheableCacheKey(joinPoint)) - } + } class TestBean { @ReqShieldCacheable(cacheName = "TestCacheName", key = "#paramMap['x'] + #paramMap['y']") - suspend fun cacheableWithSingleArgument(paramMap: Map): String = "ReturnValue: $paramMap" + suspend fun cacheableWithCustomKey(paramMap: Map): String = "ReturnValue: $paramMap" + + @ReqShieldCacheable(cacheName = "TestCacheName") + suspend fun cacheableWithDefaultKeyGenerator(paramMap: Map): String = "ReturnValue: $paramMap" + + @ReqShieldCacheable(cacheName = "TestCacheName", keyGenerator = "customGenerator") + fun cacheableWithKeyGenerator(paramMap: Map): String = "ReturnValue: $paramMap" @ReqShieldCacheEvict(cacheName = "TestCacheName", key = "#paramMap['x'] + #paramMap['y']") - suspend fun evictWithSingleArgument(paramMap: Map): Boolean { + suspend fun evict(paramMap: Map): Boolean { log.debug("cache eviction") return true } } + + class CustomGenerator : KeyGenerator { + override fun generate( + target: Any, + method: Method, + vararg params: Any?, + ): Any = "KeyGeneratedByGenerator" + } } diff --git a/core-spring-webflux/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/annotation/ReqShieldCacheEvict.kt b/core-spring-webflux/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/annotation/ReqShieldCacheEvict.kt index 25a9ff1..7296434 100644 --- a/core-spring-webflux/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/annotation/ReqShieldCacheEvict.kt +++ b/core-spring-webflux/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/annotation/ReqShieldCacheEvict.kt @@ -21,6 +21,7 @@ package com.linecorp.cse.reqshield.spring.webflux.annotation annotation class ReqShieldCacheEvict( val cacheName: String, val key: String = "", + val keyGenerator: String = "", val isLocalLock: Boolean = true, val lockTimeoutMillis: Long = 3000, val condition: String = "", 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 70653aa..c093fc7 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 @@ -26,6 +26,7 @@ import java.lang.annotation.Inherited annotation class ReqShieldCacheable( val cacheName: String, val key: String = "", + val keyGenerator: String = "", val isLocalLock: Boolean = true, val lockTimeoutMillis: Long = 3000, val decisionForUpdate: Int = 90, 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 1d0558b..fa65a83 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 @@ -25,6 +25,9 @@ import org.aspectj.lang.ProceedingJoinPoint import org.aspectj.lang.annotation.Around import org.aspectj.lang.annotation.Aspect import org.aspectj.lang.reflect.MethodSignature +import org.springframework.beans.factory.BeanFactory +import org.springframework.beans.factory.BeanFactoryAware +import org.springframework.cache.interceptor.KeyGenerator import org.springframework.cache.interceptor.SimpleKeyGenerator import org.springframework.context.expression.MethodBasedEvaluationContext import org.springframework.core.DefaultParameterNameDiscoverer @@ -34,6 +37,7 @@ import org.springframework.expression.Expression import org.springframework.expression.spel.standard.SpelExpressionParser import org.springframework.stereotype.Component import org.springframework.util.StringUtils +import org.springframework.util.function.SingletonSupplier import reactor.core.publisher.Mono import java.lang.reflect.Method import java.util.concurrent.ConcurrentHashMap @@ -42,8 +46,12 @@ import java.util.concurrent.ConcurrentHashMap @Component class ReqShieldAspect( private val asyncCache: AsyncCache, -) { +) : BeanFactoryAware { + private lateinit var beanFactory: BeanFactory private val spelParser = SpelExpressionParser() + private var defaultKeyGenerator = SingletonSupplier.of { SimpleKeyGenerator() } + + private val keyGeneratorMap = ConcurrentHashMap() internal val reqShieldMap = ConcurrentHashMap>() @Around("@annotation(com.linecorp.cse.reqshield.spring.webflux.annotation.ReqShieldCacheable)") @@ -74,28 +82,33 @@ class ReqShieldAspect( } } - fun getCacheEvictAnnotation(joinPoint: ProceedingJoinPoint): ReqShieldCacheEvict = - AnnotationUtils.getAnnotation(getTargetMethod(joinPoint), ReqShieldCacheEvict::class.java) - ?: throw IllegalArgumentException("ReqShieldCacheable annotation is required") - - internal fun getCacheEvictCacheKey(joinPoint: ProceedingJoinPoint): String { - val annotation = getCacheEvictAnnotation(joinPoint) - return getCacheKeyOrDefault(annotation.key, joinPoint) - } - - internal fun getTargetMethod(joinPoint: ProceedingJoinPoint): Method = (joinPoint.signature as MethodSignature).method - internal fun getCacheableAnnotation(joinPoint: ProceedingJoinPoint): ReqShieldCacheable = AnnotationUtils.getAnnotation(getTargetMethod(joinPoint), ReqShieldCacheable::class.java) ?: throw IllegalArgumentException("ReqShieldCacheable annotation is required") internal fun getCacheableCacheKey(joinPoint: ProceedingJoinPoint): String { val annotation = getCacheableAnnotation(joinPoint) - return getCacheKeyOrDefault(annotation.key, joinPoint) + validateCacheKey(annotation.key, annotation.keyGenerator) + + return getCacheKeyOrDefault(annotation.key, annotation.keyGenerator, joinPoint) } + internal fun getCacheEvictAnnotation(joinPoint: ProceedingJoinPoint): ReqShieldCacheEvict = + AnnotationUtils.getAnnotation(getTargetMethod(joinPoint), ReqShieldCacheEvict::class.java) + ?: throw IllegalArgumentException("ReqShieldCacheEvict annotation is required") + + internal fun getCacheEvictCacheKey(joinPoint: ProceedingJoinPoint): String { + val annotation = getCacheEvictAnnotation(joinPoint) + validateCacheKey(annotation.key, annotation.keyGenerator) + + return getCacheKeyOrDefault(annotation.key, annotation.keyGenerator, joinPoint) + } + + internal fun getTargetMethod(joinPoint: ProceedingJoinPoint): Method = (joinPoint.signature as MethodSignature).method + private fun getCacheKeyOrDefault( annotationCacheKey: String, + annotationCacheKeyGenerator: String, joinPoint: ProceedingJoinPoint, ): String { val method = getTargetMethod(joinPoint) @@ -107,12 +120,11 @@ class ReqShieldAspect( val expression: Expression = spelParser.parseExpression(annotationCacheKey) expression.getValue(context, String::class.java) } else { - SimpleKeyGenerator.generateKey(joinPoint.target, method, joinPoint.args).toString() + val keyGenerator = getOrCreateKeyGenerator(annotationCacheKeyGenerator) + keyGenerator.generate(joinPoint.target, method, joinPoint.args).toString() } - if (key.isNullOrBlank()) { - throw IllegalArgumentException("Null key returned for cache method : $method") - } + require(!key.isNullOrBlank()) { "Null key returned for cache method : $method" } return key } @@ -149,6 +161,29 @@ class ReqShieldAspect( return ReqShield(reqShieldConfiguration) } + private fun validateCacheKey( + cacheKey: String, + cacheKeyGenerator: String, + ) { + if (cacheKey.isNotBlank() && cacheKeyGenerator.isNotBlank()) { + throw IllegalArgumentException("The key and keyGenerator attributes are mutually exclusive.") + } + } + + private fun getOrCreateKeyGenerator(keyGeneratorBeanName: String?): KeyGenerator { + if (keyGeneratorBeanName.isNullOrBlank()) { + return defaultKeyGenerator.obtain() + } + + return keyGeneratorMap.computeIfAbsent(keyGeneratorBeanName) { + beanFactory.getBean(it, KeyGenerator::class.java) + } + } + private fun generateReqShieldKey(joinPoint: ProceedingJoinPoint): String = "${getCacheableAnnotation(joinPoint).cacheName}-${getCacheableCacheKey(joinPoint)}" + + override fun setBeanFactory(beanFactory: BeanFactory) { + this.beanFactory = beanFactory + } } 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 cb80292..afdfab0 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 @@ -26,74 +26,57 @@ import io.mockk.every import io.mockk.mockk import io.mockk.spyk import org.aspectj.lang.ProceedingJoinPoint -import org.aspectj.lang.reflect.MethodSignature import org.junit.jupiter.api.Assertions import org.junit.jupiter.api.BeforeEach import org.junit.jupiter.api.Test +import org.springframework.beans.factory.BeanFactory +import org.springframework.cache.interceptor.KeyGenerator +import org.springframework.cache.interceptor.SimpleKeyGenerator import org.springframework.util.ReflectionUtils import reactor.core.publisher.Flux import reactor.core.publisher.Mono import reactor.core.scheduler.Schedulers import reactor.test.StepVerifier +import java.lang.reflect.Method import kotlin.test.assertEquals import kotlin.test.assertTrue class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { private val asyncCache: AsyncCache = mockk() private val joinPoint = mockk() - private val methodSignature: MethodSignature = mockk() private val reqShieldAspect = spyk(ReqShieldAspect(asyncCache)) private val targetObject = spyk(TestBean()) - private val method = - ReflectionUtils.findMethod( - TestBean::class.java, - TestBean::cacheableWithSingleArgument.name, - Map::class.java, - ) - - private val cacheName = "testCacheName" - private val cacheKey = "#paramMap['x'] + #paramMap['y']" private val argument = mapOf("x" to "paramX", "y" to "paramY") - private val evaulatedKey = "paramXparamY" + + private val cacheName = "TestCacheName" + private val cacheKeyGenerator = "customGenerator" + private val spelEvaluatedKey = "paramXparamY" + private val keyGeneratorKey = "KeyGeneratedByGenerator" + + private val beanFactory = mockk() + private val methodReturn = Product("testProduct", "testCategory") @BeforeEach fun setUp() { - every { joinPoint.signature } returns methodSignature - - every { methodSignature.method } returns method every { joinPoint.args } returns arrayOf(argument) every { joinPoint.target } returns targetObject - val reqShieldCacheable: ReqShieldCacheable = mockk() - every { reqShieldCacheable.key } returns cacheKey - every { reqShieldCacheable.timeToLiveMillis } returns 60 - every { reqShieldCacheable.isLocalLock } returns false - every { reqShieldCacheable.lockTimeoutMillis } returns 0 - every { reqShieldCacheable.decisionForUpdate } returns 70 - every { reqShieldCacheable.cacheName } returns "product" - - every { reqShieldAspect.getCacheableAnnotation(joinPoint) } returns - ReqShieldCacheable( - cacheName = cacheName, - key = cacheKey, - lockTimeoutMillis = 1000, - timeToLiveMillis = 1000, - ) - - every { reqShieldAspect.getCacheEvictAnnotation(joinPoint) } returns - ReqShieldCacheEvict( - cacheName = cacheName, - lockTimeoutMillis = 1000, - ) + reqShieldAspect.setBeanFactory(beanFactory) } @Test - override fun testAspectOperationVerifyReqShieldAndCacheCreation() { + override fun verifyReqShieldCacheCreation() { // Mock the cache data using mockk val reqShieldData = ReqShieldData(methodReturn, 1000) every { asyncCache.get(any()) } returns Mono.just(reqShieldData) - every { joinPoint.proceed() } answers { targetObject.cacheableWithSingleArgument(argument) } + every { joinPoint.proceed() } answers { targetObject.cacheableWithCustomKey(argument) } + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithCustomKey.name, + Map::class.java, + )!! // Test the aroundTargetCacheable method val result = reqShieldAspect.aroundTargetCacheable(joinPoint) @@ -104,16 +87,22 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { .assertNext { value -> assertEquals(reqShieldData.value, value) Assertions.assertTrue(reqShieldAspect.reqShieldMap.size == 1) - Assertions.assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$evaulatedKey"]) + Assertions.assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$spelEvaluatedKey"]) }.verifyComplete() } @Test - override fun testAspectOperationReqShieldObjectShouldBeCreatedOnce() { + override fun reqShieldObjectShouldBeCreatedOnce() { // Mock the cache data using mockk val reqShieldData = ReqShieldData(methodReturn, 1000) every { asyncCache.get(any()) } returns Mono.just(reqShieldData) - every { joinPoint.proceed() } answers { targetObject.cacheableWithSingleArgument(argument) } + every { joinPoint.proceed() } answers { targetObject.cacheableWithCustomKey(argument) } + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithCustomKey.name, + Map::class.java, + )!! val flux = Flux @@ -129,17 +118,22 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { .assertNext { productList -> Assertions.assertTrue(reqShieldAspect.reqShieldMap.size == 1) println(reqShieldAspect.reqShieldMap.keys().toList()) - Assertions.assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$evaulatedKey"]) + Assertions.assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$spelEvaluatedKey"]) }.verifyComplete() } @Test - override fun testAspectOperationCacheEviction() { + override fun verifyReqShieldCacheEviction() { // Mock the cache data using mockk val reqShieldData = ReqShieldData(methodReturn, 1000) every { asyncCache.get(any()) } returns Mono.just(reqShieldData) - every { asyncCache.evict(any()) } returns Mono.just(true) - every { joinPoint.proceed() } answers { targetObject.cacheableWithSingleArgument(argument) } + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithDefaultKeyGenerator.name, + Map::class.java, + )!! + every { joinPoint.proceed() } answers { targetObject.cacheableWithCustomKey(argument) } // Test the aroundTargetCacheable method val result = reqShieldAspect.aroundTargetCacheable(joinPoint) @@ -150,7 +144,16 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { assertEquals(reqShieldData.value, value) }.verifyComplete() - every { joinPoint.proceed() } answers { targetObject.evictWithSingleArgument(argument) } + // Validate cache eviction + every { asyncCache.evict(any()) } returns Mono.just(true) + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::evict.name, + Map::class.java, + )!! + every { joinPoint.proceed() } answers { targetObject.evict(argument) } + val removeProductMono = reqShieldAspect.aroundReqShieldCacheEvict(joinPoint) StepVerifier @@ -161,43 +164,77 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { } @Test - override fun testCacheKeyGenerationUseGeneratedKey() { - // Mock the cache data using mockk - val reqShieldData = ReqShieldData(methodReturn, 1000) - every { asyncCache.get(any()) } returns Mono.just(reqShieldData) - every { joinPoint.proceed() } answers { targetObject.cacheableWithSingleArgument(argument) } - - // Test the aroundTargetCacheable method - val result = reqShieldAspect.aroundTargetCacheable(joinPoint) + override fun verifyCacheKeyGenerationWithSpEL() { + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithCustomKey.name, + Map::class.java, + )!! + + // when, then + Assertions.assertEquals( + spelEvaluatedKey, + reqShieldAspect.getCacheableCacheKey(joinPoint), + ) + } - // Verify the behavior using StepVerifier - StepVerifier - .create(result) - .assertNext { value -> - Assertions.assertEquals( - evaulatedKey, - reqShieldAspect.getCacheableCacheKey(joinPoint), - ) - }.verifyComplete() + override fun verifyCacheKeyGenerationWithKeyGenerator() { + // given + every { beanFactory.getBean(cacheKeyGenerator, KeyGenerator::class.java) } returns + CustomGenerator() + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithKeyGenerator.name, + Map::class.java, + )!! + + // when, then + Assertions.assertEquals( + keyGeneratorKey, + reqShieldAspect.getCacheableCacheKey(joinPoint), + ) } @Test - override fun testCacheKeyGenerationCacheKeyShouldBeSuppliedKey() { - every { reqShieldAspect.getCacheableAnnotation(joinPoint) } returns - ReqShieldCacheable( - cacheName = cacheName, - key = cacheKey, - ) - - Assertions.assertEquals(evaulatedKey, reqShieldAspect.getCacheableCacheKey(joinPoint)) + override fun verifyCacheKeyGenerationWithDefaultGenerator() { + // given + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithDefaultKeyGenerator.name, + Map::class.java, + )!! + + // when, then + Assertions.assertEquals( + SimpleKeyGenerator.generateKey(arrayOf(argument)).toString(), + reqShieldAspect.getCacheableCacheKey(joinPoint), + ) } class TestBean { @ReqShieldCacheable(cacheName = "TestCacheName", key = "#paramMap['x'] + #paramMap['y']") - fun cacheableWithSingleArgument(paramMap: Map): Mono = + fun cacheableWithCustomKey(paramMap: Map): Mono = Mono.justOrEmpty(Product("testProduct", "testCategory")) + + @ReqShieldCacheable(cacheName = "TestCacheName") + fun cacheableWithDefaultKeyGenerator(paramMap: Map): Mono = + Mono.justOrEmpty(Product("testProduct", "testCategory")) + + @ReqShieldCacheable(cacheName = "TestCacheName", keyGenerator = "customGenerator") + fun cacheableWithKeyGenerator(paramMap: Map): Mono = Mono.justOrEmpty(Product("testProduct", "testCategory")) @ReqShieldCacheEvict(cacheName = "TestCacheName") - fun evictWithSingleArgument(paramMap: Map): Mono = Mono.just(true) + fun evict(paramMap: Map): Mono = Mono.just(true) + } + + class CustomGenerator : KeyGenerator { + override fun generate( + target: Any, + method: Method, + vararg params: Any?, + ): Any = "KeyGeneratedByGenerator" } } diff --git a/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheEvict.kt b/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheEvict.kt index 2b9b130..44b5635 100644 --- a/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheEvict.kt +++ b/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheEvict.kt @@ -24,6 +24,7 @@ import java.lang.annotation.Inherited annotation class ReqShieldCacheEvict( val cacheName: String, val key: String = "", + val keyGenerator: String = "", val isLocalLock: Boolean = true, val lockTimeoutMillis: Long = 3000, val condition: String = "", diff --git a/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheable.kt b/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheable.kt index f0d01d8..edaf2c2 100644 --- a/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheable.kt +++ b/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheable.kt @@ -26,6 +26,7 @@ import java.lang.annotation.Inherited annotation class ReqShieldCacheable( val cacheName: String, val key: String = "", + val keyGenerator: String = "", val isLocalLock: Boolean = true, val lockTimeoutMillis: Long = 30000, val decisionForUpdate: Int = 90, 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 2c24632..bf2cae7 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 @@ -25,6 +25,9 @@ import org.aspectj.lang.ProceedingJoinPoint import org.aspectj.lang.annotation.Around import org.aspectj.lang.annotation.Aspect import org.aspectj.lang.reflect.MethodSignature +import org.springframework.beans.factory.BeanFactory +import org.springframework.beans.factory.BeanFactoryAware +import org.springframework.cache.interceptor.KeyGenerator import org.springframework.cache.interceptor.SimpleKeyGenerator import org.springframework.context.expression.MethodBasedEvaluationContext import org.springframework.core.DefaultParameterNameDiscoverer @@ -34,6 +37,7 @@ import org.springframework.expression.Expression import org.springframework.expression.spel.standard.SpelExpressionParser import org.springframework.stereotype.Component import org.springframework.util.StringUtils +import org.springframework.util.function.SingletonSupplier import java.lang.reflect.Method import java.util.concurrent.ConcurrentHashMap @@ -41,8 +45,12 @@ import java.util.concurrent.ConcurrentHashMap @Component class ReqShieldAspect( private val reqShieldCache: ReqShieldCache, -) { +) : BeanFactoryAware { + private lateinit var beanFactory: BeanFactory private val spelParser = SpelExpressionParser() + private val defaultKeyGenerator = SingletonSupplier.of { SimpleKeyGenerator() } + + private val keyGeneratorMap = ConcurrentHashMap() internal val reqShieldMap = ConcurrentHashMap>() @Around("@annotation(com.linecorp.cse.reqshield.spring.annotation.ReqShieldCacheable)") @@ -109,20 +117,25 @@ class ReqShieldAspect( internal fun getCacheEvictAnnotation(joinPoint: ProceedingJoinPoint): ReqShieldCacheEvict = AnnotationUtils.getAnnotation(getTargetMethod(joinPoint), ReqShieldCacheEvict::class.java) - ?: throw IllegalArgumentException("ReqShieldCacheable annotation is required") + ?: throw IllegalArgumentException("ReqShieldCacheEvict annotation is required") internal fun getCacheableCacheKey(joinPoint: ProceedingJoinPoint): String { val annotation = getCacheableAnnotation(joinPoint) - return getCacheKeyOrDefault(annotation.key, joinPoint) + validateCacheKey(annotation.key, annotation.keyGenerator) + + return getCacheKeyOrDefault(annotation.key, annotation.keyGenerator, joinPoint) } internal fun getCacheEvictCacheKey(joinPoint: ProceedingJoinPoint): String { val annotation = getCacheEvictAnnotation(joinPoint) - return getCacheKeyOrDefault(annotation.key, joinPoint) + validateCacheKey(annotation.key, annotation.keyGenerator) + + return getCacheKeyOrDefault(annotation.key, annotation.keyGenerator, joinPoint) } private fun getCacheKeyOrDefault( annotationCacheKey: String, + annotationCacheKeyGenerator: String, joinPoint: ProceedingJoinPoint, ): String { val method = getTargetMethod(joinPoint) @@ -134,16 +147,38 @@ class ReqShieldAspect( val expression: Expression = spelParser.parseExpression(annotationCacheKey) expression.getValue(context, String::class.java) } else { - SimpleKeyGenerator.generateKey(joinPoint.target, method, joinPoint.args).toString() + val keyGenerator = getOrCreateKeyGenerator(annotationCacheKeyGenerator) + keyGenerator.generate(joinPoint.target, method, joinPoint.args).toString() } - if (key.isNullOrBlank()) { - throw IllegalArgumentException("Null key returned for cache method : $method") - } + require(!key.isNullOrBlank()) { "Null key returned for cache method : $method" } return key } + private fun validateCacheKey( + cacheKey: String, + cacheKeyGenerator: String, + ) { + if (cacheKey.isNotBlank() && cacheKeyGenerator.isNotBlank()) { + throw IllegalArgumentException("The key and keyGenerator attributes are mutually exclusive.") + } + } + + private fun getOrCreateKeyGenerator(keyGeneratorBeanName: String?): KeyGenerator { + if (keyGeneratorBeanName.isNullOrBlank()) { + return defaultKeyGenerator.obtain() + } + + return keyGeneratorMap.computeIfAbsent(keyGeneratorBeanName) { + beanFactory.getBean(it, KeyGenerator::class.java) + } + } + private fun generateReqShieldKey(joinPoint: ProceedingJoinPoint): String = "${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 d227ae8..619bd08 100644 --- a/core-spring/src/test/kotlin/aspect/ReqShieldAspectTest.kt +++ b/core-spring/src/test/kotlin/aspect/ReqShieldAspectTest.kt @@ -36,7 +36,11 @@ import org.junit.jupiter.api.Assertions.assertTrue import org.junit.jupiter.api.BeforeEach import org.junit.jupiter.api.Test import org.slf4j.LoggerFactory +import org.springframework.beans.factory.BeanFactory +import org.springframework.cache.interceptor.KeyGenerator +import org.springframework.cache.interceptor.SimpleKeyGenerator import org.springframework.util.ReflectionUtils +import java.lang.reflect.Method import java.time.Duration import java.util.concurrent.Executors @@ -47,63 +51,60 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { private val joinPoint = mockk() private val reqShieldAspect = spyk(ReqShieldAspect(reqShieldCache)) private val targetObject = spyk(TestBean()) - private val method = - ReflectionUtils.findMethod( - TestBean::class.java, - TestBean::cacheableWithSingleArgument.name, - Map::class.java, - ) - - private val cacheName = "testCacheName" - private val cacheKey = "#paramMap['x'] + #paramMap['y']" private val argument = mapOf("x" to "paramX", "y" to "paramY") - private val evaluatedKey = "paramXparamY" + + private val cacheName = "TestCacheName" + private val cacheKeyGenerator = "customGenerator" + private val spelEvaluatedKey = "paramXparamY" + private val keyGeneratorKey = "KeyGeneratedByGenerator" + + private val beanFactory = mockk() + private val methodReturn = Product("testProduct", "testCategory") @BeforeEach fun setUp() { every { joinPoint.target } returns targetObject every { joinPoint.args } returns arrayOf(argument) - every { joinPoint.proceed() } answers { targetObject.cacheableWithSingleArgument(argument) } - - every { reqShieldAspect.getTargetMethod(joinPoint) } returns method!! - - every { reqShieldAspect.getCacheableAnnotation(joinPoint) } returns - ReqShieldCacheable( - cacheName = cacheName, - key = cacheKey, - lockTimeoutMillis = 1000, - timeToLiveMillis = 1000, - ) - - every { reqShieldAspect.getCacheEvictAnnotation(joinPoint) } returns - ReqShieldCacheEvict( - cacheName = cacheName, - key = cacheKey, - ) + + reqShieldAspect.setBeanFactory(beanFactory) } @Test - override fun testAspectOperationVerifyReqShieldAndCacheCreation() { + override fun verifyReqShieldCacheCreation() { // given - every { reqShieldCache.get(any()) } returns ReqShieldData(methodReturn, 1000) - every { joinPoint.proceed() } answers { targetObject.cacheableWithSingleArgument(argument) } + val reqShieldData = ReqShieldData(methodReturn, 1000) + every { reqShieldCache.get(any()) } returns reqShieldData + every { joinPoint.proceed() } answers { targetObject.cacheableWithCustomKey(argument) } + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithCustomKey.name, + Map::class.java, + )!! // when - reqShieldAspect.aroundReqShieldCacheable(joinPoint) + val result = reqShieldAspect.aroundReqShieldCacheable(joinPoint) Awaitility.await().atMost(Duration.ofMillis(BaseReqShieldTest.AWAIT_TIMEOUT)).untilAsserted { // then + assertEquals(reqShieldData.value, result) assertTrue(reqShieldAspect.reqShieldMap.size == 1) - assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$evaluatedKey"]) + assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$spelEvaluatedKey"]) } } @Test - override fun testAspectOperationReqShieldObjectShouldBeCreatedOnce() { + override fun reqShieldObjectShouldBeCreatedOnce() { // given every { reqShieldCache.get(any()) } returns ReqShieldData(methodReturn, 1000) - every { joinPoint.proceed() } answers { targetObject.cacheableWithSingleArgument(argument) } + every { joinPoint.proceed() } answers { targetObject.cacheableWithCustomKey(argument) } + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithCustomKey.name, + Map::class.java, + )!! // when val executorService = Executors.newFixedThreadPool(10) @@ -116,22 +117,36 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { Awaitility.await().atMost(Duration.ofMillis(BaseReqShieldTest.AWAIT_TIMEOUT)).untilAsserted { // then assertTrue(reqShieldAspect.reqShieldMap.size == 1) - assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$evaluatedKey"]) + assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$spelEvaluatedKey"]) } } @Test - override fun testAspectOperationCacheEviction() { + override fun verifyReqShieldCacheEviction() { // given val reqShieldData = ReqShieldData(methodReturn, 10000) every { reqShieldCache.get(any()) } returns reqShieldData - every { reqShieldCache.evict(any()) } returns true - every { joinPoint.proceed() } answers { targetObject.cacheableWithSingleArgument(argument) } + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithCustomKey.name, + Map::class.java, + )!! + every { joinPoint.proceed() } answers { targetObject.cacheableWithCustomKey(argument) } val cachedResult = reqShieldAspect.aroundReqShieldCacheable(joinPoint) assertEquals(reqShieldData.value, cachedResult) + // Validate cache eviction + every { reqShieldCache.evict(any()) } returns true + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::evict.name, + Map::class.java, + )!! + // when reqShieldAspect.aroundReqShieldCacheEvict(joinPoint) @@ -140,44 +155,88 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { } @Test - override fun testCacheKeyGenerationUseGeneratedKey() { - every { joinPoint.target } returns targetObject - every { joinPoint.args } returns arrayOf(argument) - every { reqShieldAspect.getTargetMethod(joinPoint) } returns method!! - - every { reqShieldAspect.getCacheableAnnotation(joinPoint) } returns - ReqShieldCacheable( - cacheName = cacheName, - key = cacheKey, - ) + override fun verifyCacheKeyGenerationWithSpEL() { + // given + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithCustomKey.name, + Map::class.java, + )!! + + // when, then + assertEquals( + spelEvaluatedKey, + reqShieldAspect.getCacheableCacheKey(joinPoint), + ) + } + @Test + override fun verifyCacheKeyGenerationWithKeyGenerator() { + // given + every { beanFactory.getBean(cacheKeyGenerator, KeyGenerator::class.java) } returns + CustomGenerator() + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithKeyGenerator.name, + Map::class.java, + )!! + + // when, then assertEquals( - evaluatedKey, + keyGeneratorKey, reqShieldAspect.getCacheableCacheKey(joinPoint), ) } @Test - override fun testCacheKeyGenerationCacheKeyShouldBeSuppliedKey() { - every { reqShieldAspect.getCacheableAnnotation(joinPoint) } returns - ReqShieldCacheable( - cacheName = cacheName, - key = cacheKey, - ) - - assertEquals(evaluatedKey, reqShieldAspect.getCacheableCacheKey(joinPoint)) + override fun verifyCacheKeyGenerationWithDefaultGenerator() { + // given + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithDefaultKeyGenerator.name, + Map::class.java, + )!! + + // when, then + assertEquals( + SimpleKeyGenerator.generateKey(arrayOf(argument)).toString(), + reqShieldAspect.getCacheableCacheKey(joinPoint), + ) } class TestBean { @ReqShieldCacheable(cacheName = "TestCacheName", key = "#paramMap['x'] + #paramMap['y']") - fun cacheableWithSingleArgument(paramMap: Map): String { - log.debug("method invoked") + fun cacheableWithCustomKey(paramMap: Map): String { + log.debug("cacheableWithCustomKey method invoked") + return "ReturnValue: $paramMap" + } + + @ReqShieldCacheable(cacheName = "TestCacheName") + fun cacheableWithDefaultKeyGenerator(paramMap: Map): String { + log.debug("cacheableWithDefaultKeyGenerator method invoked") + return "ReturnValue: $paramMap" + } + + @ReqShieldCacheable(cacheName = "TestCacheName", keyGenerator = "customGenerator") + fun cacheableWithKeyGenerator(paramMap: Map): String { + log.debug("cacheableWithCustomGenerator method invoked") return "ReturnValue: $paramMap" } @ReqShieldCacheEvict(cacheName = "TestCacheName", key = "#paramMap['x'] + #paramMap['y']") - fun evictWithSingleArgument(paramMap: Map) { + fun evict(paramMap: Map) { log.debug("cache eviction") } } + + class CustomGenerator : KeyGenerator { + override fun generate( + target: Any, + method: Method, + vararg params: Any?, + ): Any = "KeyGeneratedByGenerator" + } } diff --git a/support/src/testFixtures/kotlin/com/linecorp/cse/reqshield/support/BaseReqShieldModuleSupportTest.kt b/support/src/testFixtures/kotlin/com/linecorp/cse/reqshield/support/BaseReqShieldModuleSupportTest.kt index b37c98d..35dac2c 100644 --- a/support/src/testFixtures/kotlin/com/linecorp/cse/reqshield/support/BaseReqShieldModuleSupportTest.kt +++ b/support/src/testFixtures/kotlin/com/linecorp/cse/reqshield/support/BaseReqShieldModuleSupportTest.kt @@ -17,13 +17,15 @@ package com.linecorp.cse.reqshield.support interface BaseReqShieldModuleSupportTest { - fun testAspectOperationVerifyReqShieldAndCacheCreation() + fun verifyReqShieldCacheCreation() - fun testAspectOperationReqShieldObjectShouldBeCreatedOnce() + fun reqShieldObjectShouldBeCreatedOnce() - fun testAspectOperationCacheEviction() + fun verifyReqShieldCacheEviction() - fun testCacheKeyGenerationUseGeneratedKey() + fun verifyCacheKeyGenerationWithSpEL() - fun testCacheKeyGenerationCacheKeyShouldBeSuppliedKey() + fun verifyCacheKeyGenerationWithKeyGenerator() + + fun verifyCacheKeyGenerationWithDefaultGenerator() } From 275ed6002a958e449fdc465992b26fcb496a61a3 Mon Sep 17 00:00:00 2001 From: byungchan-lee Date: Tue, 1 Apr 2025 15:59:37 +0900 Subject: [PATCH 07/11] SpEL Support - set SpEL key in new test code --- .../mvc/example/service/SampleService.kt | 16 +++++++++ .../example/service/CacheAnnotationTest.kt | 23 ++++++++++++ .../webflux/example/service/SampleService.kt | 19 ++++++++++ .../example/service/CacheAnnotationTest.kt | 36 ++++++++++++++++++- .../example/service/SampleService.kt | 17 +++++++++ .../coroutine/example/CacheAnnotationTest.kt | 25 ++++++++++--- .../cse/reqshield/service/SampleService.kt | 1 + .../spring/service/CacheAnnotationTest.kt | 7 ++-- .../webflux/example/service/SampleService.kt | 1 + .../example/service/CacheAnnotationTest.kt | 12 +++++++ .../example/service/SampleService.kt | 1 + .../example/service/CacheAnnotationTest.kt | 11 +++--- 12 files changed, 155 insertions(+), 14 deletions(-) diff --git a/req-shield-spring-boot3-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/mvc/example/service/SampleService.kt b/req-shield-spring-boot3-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/mvc/example/service/SampleService.kt index 0585267..2dc12dc 100644 --- a/req-shield-spring-boot3-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/mvc/example/service/SampleService.kt +++ b/req-shield-spring-boot3-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/mvc/example/service/SampleService.kt @@ -1,5 +1,6 @@ package com.linecorp.cse.reqshield.spring3.mvc.example.service +import com.linecorp.cse.reqshield.config.ReqShieldWorkMode import com.linecorp.cse.reqshield.spring.annotation.ReqShieldCacheEvict import com.linecorp.cse.reqshield.spring.annotation.ReqShieldCacheable import com.linecorp.cse.reqshield.spring3.mvc.example.dto.Product @@ -22,6 +23,21 @@ class SampleService { return Product(productId, "product_$productId") } + @ReqShieldCacheable( + cacheName = "productOnlyUpdateCache", + key = "'product-' + #productId", + decisionForUpdate = 90, + timeToLiveMillis = 60 * 1000, + reqShieldWorkMode = ReqShieldWorkMode.ONLY_UPDATE_CACHE, + ) + fun getProductOnlyUpdateCache(productId: String): Product { + log.info("find product with db request - req-shield local lock only update cache (will take 1 second)") + Thread.sleep(500) + + atomicInteger.incrementAndGet() + return Product(productId, "product_$productId") + } + @ReqShieldCacheable(cacheName = "product", key = "'product-' + #productId", isLocalLock = false, decisionForUpdate = 90) fun getProductForGlobalLock(productId: String): Product { log.info("find product with db request - req-shield global lock (will take 1 second)") 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 9e6ab2e..5f90aa7 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 @@ -50,6 +50,28 @@ class CacheAnnotationTest : AbstractRedisTest() { await().atMost(Duration.ofMillis(BaseReqShieldTest.AWAIT_TIMEOUT)).untilAsserted { assertEquals(1, sampleService.getRequestCount()) + assertNotNull(reqShieldCache.get("product-$testProductId")) + } + } + + @Test + fun `ReqShieldCacheable test - request to 'sampleService' should be request count times (only update cache mode)`() { + val executorService = Executors.newFixedThreadPool(100) + + val testProductId: String = UUID.randomUUID().toString() + + for (i in 1..100) { + executorService.submit { + sampleService.getProductOnlyUpdateCache(testProductId) + } + } + + executorService.shutdown() + executorService.awaitTermination(3000, TimeUnit.SECONDS) + + await().atMost(Duration.ofMillis(BaseReqShieldTest.AWAIT_TIMEOUT)).untilAsserted { + assertEquals(100, sampleService.getRequestCount()) + assertNotNull(reqShieldCache.get("product-$testProductId")) } } @@ -72,6 +94,7 @@ class CacheAnnotationTest : AbstractRedisTest() { await().atMost(Duration.ofMillis(BaseReqShieldTest.AWAIT_TIMEOUT)).untilAsserted { assertEquals(1, sampleService.getRequestCount()) + assertNotNull(reqShieldCache.get("product-$testProductId")) } } diff --git a/req-shield-spring-boot3-webflux-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/webflux/example/service/SampleService.kt b/req-shield-spring-boot3-webflux-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/webflux/example/service/SampleService.kt index 3011b8d..f3d99c6 100644 --- a/req-shield-spring-boot3-webflux-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/webflux/example/service/SampleService.kt +++ b/req-shield-spring-boot3-webflux-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/webflux/example/service/SampleService.kt @@ -1,6 +1,7 @@ package com.linecorp.cse.reqshield.spring3.webflux.example.service import com.linecorp.cse.reqshield.reactor.ReqShield +import com.linecorp.cse.reqshield.reactor.config.ReqShieldWorkMode import com.linecorp.cse.reqshield.spring.webflux.annotation.ReqShieldCacheEvict import com.linecorp.cse.reqshield.spring.webflux.annotation.ReqShieldCacheable import com.linecorp.cse.reqshield.spring3.webflux.example.dto.Product @@ -30,6 +31,24 @@ class SampleService( }.doFinally { atomicInteger.incrementAndGet() }, ) + @ReqShieldCacheable( + cacheName = "productOnlyUpdataCache", + key = "'product-' + #productId", + decisionForUpdate = 80, + timeToLiveMillis = 60 * 1000, + reqShieldWorkMode = ReqShieldWorkMode.ONLY_UPDATE_CACHE, + ) + fun getProductOnlyUpdateCache(productId: String): Mono = + Mono + .delay(Duration.ofMillis(500)) + .then( + Mono + .just(Product(productId, "product_$productId")) + .doOnNext { + log.info("find product with db request - req-shield local lock (will take 1 second)") + }.doFinally { atomicInteger.incrementAndGet() }, + ) + fun getProductNoAnno(productId: String): Mono = reqShield .getAndSetReqShieldData( 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 c424d99..6421006 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 @@ -49,7 +49,36 @@ class CacheAnnotationTest : AbstractRedisTest() { .create(flux) .assertNext { productList -> assertEquals(1, sampleService.getRequestCount(), "Request count should be 1") + }.expectComplete() + .verify() + + await().atMost(5, TimeUnit.SECONDS).until { + asyncCache.get("product-$testProductId").block() != null + } + } + + @Test + fun `ReqShieldCacheable test - request to 'sampleService' should be request count times(only update cache mode)`() { + val testProductId: String = UUID.randomUUID().toString() + + val flux = + Flux + .range(1, 20) + .flatMap { + sampleService + .getProductOnlyUpdateCache(testProductId) + .subscribeOn(Schedulers.boundedElastic()) + }.collectList() + + StepVerifier + .create(flux) + .assertNext { productList -> + assertEquals(19, sampleService.getRequestCount(), "Request count should be 19") }.verifyComplete() + + await().atMost(5, TimeUnit.SECONDS).until { + asyncCache.get("product-$testProductId").block() != null + } } @Test @@ -69,7 +98,12 @@ class CacheAnnotationTest : AbstractRedisTest() { .create(flux) .assertNext { productList -> assertEquals(1, sampleService.getRequestCount(), "Request count should be 1") - }.verifyComplete() + }.expectComplete() + .verify() + + await().atMost(5, TimeUnit.SECONDS).until { + asyncCache.get("product-$testProductId").block() != null + } } @Test diff --git a/req-shield-spring-boot3-webflux-kotlin-coroutine-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/webflux/kotlin/coroutine/example/service/SampleService.kt b/req-shield-spring-boot3-webflux-kotlin-coroutine-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/webflux/kotlin/coroutine/example/service/SampleService.kt index 812b0d0..66752ea 100644 --- a/req-shield-spring-boot3-webflux-kotlin-coroutine-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/webflux/kotlin/coroutine/example/service/SampleService.kt +++ b/req-shield-spring-boot3-webflux-kotlin-coroutine-example/src/main/kotlin/com/linecorp/cse/reqshield/spring3/webflux/kotlin/coroutine/example/service/SampleService.kt @@ -1,6 +1,7 @@ package com.linecorp.cse.reqshield.spring3.webflux.kotlin.coroutine.example.service import com.linecorp.cse.reqshield.kotlin.coroutine.ReqShield +import com.linecorp.cse.reqshield.kotlin.coroutine.config.ReqShieldWorkMode 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.spring3.webflux.kotlin.coroutine.example.dto.Product @@ -32,6 +33,22 @@ class SampleService( return Product(productId, "product_$productId") } + @ReqShieldCacheable( + cacheName = "productOnlyUpdateCache", + key = "'product-' + #productId", + decisionForUpdate = 80, + timeToLiveMillis = 60 * 1000, + reqShieldWorkMode = ReqShieldWorkMode.ONLY_UPDATE_CACHE, + ) + suspend fun getProductOnlyUpdateCache(productId: String): Product { + log.info("find product with db request with req-shield local lock (will take 1 second)") + + delay(500) + atomicInteger.incrementAndGet() + + return Product(productId, "product_$productId") + } + suspend fun getProductNoAnno(productId: String): Product? { val result = reqShield diff --git a/req-shield-spring-boot3-webflux-kotlin-coroutine-example/src/test/kotlin/com/linecorp/cse/reqshield/spring3/webflux/kotlin/coroutine/example/CacheAnnotationTest.kt b/req-shield-spring-boot3-webflux-kotlin-coroutine-example/src/test/kotlin/com/linecorp/cse/reqshield/spring3/webflux/kotlin/coroutine/example/CacheAnnotationTest.kt index 78be8f6..b0a29fe 100644 --- a/req-shield-spring-boot3-webflux-kotlin-coroutine-example/src/test/kotlin/com/linecorp/cse/reqshield/spring3/webflux/kotlin/coroutine/example/CacheAnnotationTest.kt +++ b/req-shield-spring-boot3-webflux-kotlin-coroutine-example/src/test/kotlin/com/linecorp/cse/reqshield/spring3/webflux/kotlin/coroutine/example/CacheAnnotationTest.kt @@ -10,8 +10,7 @@ import kotlinx.coroutines.delay import kotlinx.coroutines.runBlocking import kotlinx.coroutines.test.runTest import org.junit.jupiter.api.Assertions -import org.junit.jupiter.api.Assertions.assertNotNull -import org.junit.jupiter.api.Assertions.assertNull +import org.junit.jupiter.api.Assertions.* import org.junit.jupiter.api.BeforeEach import org.junit.jupiter.api.Test import org.junit.jupiter.api.extension.ExtendWith @@ -48,7 +47,24 @@ class CacheAnnotationTest : AbstractRedisTest() { delay(500) - Assertions.assertEquals(1, sampleService.getRequestCount()) + assertEquals(1, sampleService.getRequestCount()) + assertNotNull(asyncCache.get("product-$testProductId")) + } + + @Test + fun `ReqShieldCacheable test - request to 'sampleService' should be request count times(only update cache mode)`() = + runBlocking { + val testProductId: String = UUID.randomUUID().toString() + + List(20) { + async { + sampleService.getProductOnlyUpdateCache(testProductId) + } + }.awaitAll() + + delay(500) + + Assertions.assertEquals(20, sampleService.getRequestCount()) } @Test @@ -64,7 +80,8 @@ class CacheAnnotationTest : AbstractRedisTest() { delay(500) - Assertions.assertEquals(1, sampleService.getRequestCount()) + assertEquals(1, sampleService.getRequestCount()) + assertNotNull(asyncCache.get("product-$testProductId")) } @Test diff --git a/req-shield-spring-example/src/main/kotlin/com/linecorp/cse/reqshield/service/SampleService.kt b/req-shield-spring-example/src/main/kotlin/com/linecorp/cse/reqshield/service/SampleService.kt index 3e7e64f..0e70271 100644 --- a/req-shield-spring-example/src/main/kotlin/com/linecorp/cse/reqshield/service/SampleService.kt +++ b/req-shield-spring-example/src/main/kotlin/com/linecorp/cse/reqshield/service/SampleService.kt @@ -41,6 +41,7 @@ class SampleService { @ReqShieldCacheable( cacheName = "productOnlyUpdateCache", + key = "'product-' + #productId", decisionForUpdate = 90, timeToLiveMillis = 60 * 1000, reqShieldWorkMode = ReqShieldWorkMode.ONLY_UPDATE_CACHE, diff --git a/req-shield-spring-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/service/CacheAnnotationTest.kt b/req-shield-spring-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/service/CacheAnnotationTest.kt index b1718a5..4509c35 100644 --- a/req-shield-spring-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/service/CacheAnnotationTest.kt +++ b/req-shield-spring-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/service/CacheAnnotationTest.kt @@ -22,9 +22,7 @@ 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 -import org.junit.jupiter.api.Assertions.assertEquals -import org.junit.jupiter.api.Assertions.assertNotNull -import org.junit.jupiter.api.Assertions.assertNull +import org.junit.jupiter.api.Assertions.* import org.junit.jupiter.api.BeforeEach import org.junit.jupiter.api.Test import org.junit.jupiter.api.extension.ExtendWith @@ -67,6 +65,7 @@ class CacheAnnotationTest : AbstractRedisTest() { await().atMost(Duration.ofMillis(BaseReqShieldTest.AWAIT_TIMEOUT)).untilAsserted { assertEquals(1, sampleService.getRequestCount()) + assertNotNull(reqShieldCache.get("product-$testProductId")) } } @@ -87,6 +86,7 @@ class CacheAnnotationTest : AbstractRedisTest() { await().atMost(Duration.ofMillis(BaseReqShieldTest.AWAIT_TIMEOUT)).untilAsserted { assertEquals(100, sampleService.getRequestCount()) + assertNotNull(reqShieldCache.get("product-$testProductId")) } } @@ -109,6 +109,7 @@ class CacheAnnotationTest : AbstractRedisTest() { await().atMost(Duration.ofMillis(BaseReqShieldTest.AWAIT_TIMEOUT)).untilAsserted { assertEquals(1, sampleService.getRequestCount()) + assertNotNull(reqShieldCache.get("product-$testProductId")) } } diff --git a/req-shield-spring-webflux-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/example/service/SampleService.kt b/req-shield-spring-webflux-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/example/service/SampleService.kt index 798ec24..f98208b 100644 --- a/req-shield-spring-webflux-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/example/service/SampleService.kt +++ b/req-shield-spring-webflux-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/example/service/SampleService.kt @@ -49,6 +49,7 @@ class SampleService( @ReqShieldCacheable( cacheName = "productOnlyUpdataCache", + key = "'product-' + #productId", decisionForUpdate = 80, timeToLiveMillis = 60 * 1000, reqShieldWorkMode = ReqShieldWorkMode.ONLY_UPDATE_CACHE, 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 b5c3175..d825672 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 @@ -68,6 +68,10 @@ class CacheAnnotationTest : AbstractRedisTest() { .assertNext { productList -> assertEquals(1, sampleService.getRequestCount(), "Request count should be 1") }.verifyComplete() + + await().atMost(5, TimeUnit.SECONDS).until { + asyncCache.get("product-$testProductId").block() != null + } } @Test @@ -88,6 +92,10 @@ class CacheAnnotationTest : AbstractRedisTest() { .assertNext { productList -> assertEquals(19, sampleService.getRequestCount(), "Request count should be 19") }.verifyComplete() + + await().atMost(5, TimeUnit.SECONDS).until { + asyncCache.get("product-$testProductId").block() != null + } } @Test @@ -108,6 +116,10 @@ class CacheAnnotationTest : AbstractRedisTest() { .assertNext { productList -> assertEquals(1, sampleService.getRequestCount(), "Request count should be 1") }.verifyComplete() + + await().atMost(5, TimeUnit.SECONDS).until { + asyncCache.get("product-$testProductId").block() != null + } } @Test diff --git a/req-shield-spring-webflux-kotlin-coroutine-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/SampleService.kt b/req-shield-spring-webflux-kotlin-coroutine-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/SampleService.kt index 0efd2ae..f930c28 100644 --- a/req-shield-spring-webflux-kotlin-coroutine-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/SampleService.kt +++ b/req-shield-spring-webflux-kotlin-coroutine-example/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/SampleService.kt @@ -46,6 +46,7 @@ class SampleService( @ReqShieldCacheable( cacheName = "productOnlyUpdateCache", + key = "'product-' + #productId", decisionForUpdate = 80, timeToLiveMillis = 60 * 1000, reqShieldWorkMode = ReqShieldWorkMode.ONLY_UPDATE_CACHE, diff --git a/req-shield-spring-webflux-kotlin-coroutine-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/CacheAnnotationTest.kt b/req-shield-spring-webflux-kotlin-coroutine-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/CacheAnnotationTest.kt index 49a9763..1425d28 100644 --- a/req-shield-spring-webflux-kotlin-coroutine-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/CacheAnnotationTest.kt +++ b/req-shield-spring-webflux-kotlin-coroutine-example/src/test/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/example/service/CacheAnnotationTest.kt @@ -24,9 +24,7 @@ import kotlinx.coroutines.awaitAll import kotlinx.coroutines.delay import kotlinx.coroutines.runBlocking import kotlinx.coroutines.test.runTest -import org.junit.jupiter.api.Assertions -import org.junit.jupiter.api.Assertions.assertNotNull -import org.junit.jupiter.api.Assertions.assertNull +import org.junit.jupiter.api.Assertions.* import org.junit.jupiter.api.BeforeEach import org.junit.jupiter.api.Test import org.junit.jupiter.api.extension.ExtendWith @@ -63,7 +61,8 @@ class CacheAnnotationTest : AbstractRedisTest() { delay(500) - Assertions.assertEquals(1, sampleService.getRequestCount()) + assertEquals(1, sampleService.getRequestCount()) + assertNotNull(asyncCache.get("product-$testProductId")) } @Test @@ -79,7 +78,7 @@ class CacheAnnotationTest : AbstractRedisTest() { delay(500) - Assertions.assertEquals(20, sampleService.getRequestCount()) + assertEquals(20, sampleService.getRequestCount()) } @Test @@ -95,7 +94,7 @@ class CacheAnnotationTest : AbstractRedisTest() { delay(500) - Assertions.assertEquals(1, sampleService.getRequestCount()) + assertEquals(1, sampleService.getRequestCount()) } @Test From 2efc9e367a19f07c05c76bfb0dce2ba49a4657f6 Mon Sep 17 00:00:00 2001 From: byungchan-lee Date: Fri, 28 Mar 2025 18:25:16 +0900 Subject: [PATCH 08/11] keyGenerator support --- .../annotation/ReqShieldCacheEvict.kt | 1 + .../annotation/ReqShieldCacheable.kt | 1 + .../coroutine/aspect/ReqShieldAspect.kt | 53 ++++- .../coroutine/aspect/ReqShieldAspectTest.kt | 173 ++++++++++------- .../webflux/annotation/ReqShieldCacheEvict.kt | 1 + .../webflux/annotation/ReqShieldCacheable.kt | 1 + .../spring/webflux/aspect/ReqShieldAspect.kt | 69 +++++-- .../webflux/aspect/ReqShieldAspectTest.kt | 183 +++++++++++------- .../spring/annotation/ReqShieldCacheEvict.kt | 1 + .../spring/annotation/ReqShieldCacheable.kt | 1 + .../spring/aspect/ReqShieldAspect.kt | 51 ++++- .../test/kotlin/aspect/ReqShieldAspectTest.kt | 179 +++++++++++------ .../support/BaseReqShieldModuleSupportTest.kt | 12 +- 13 files changed, 486 insertions(+), 240 deletions(-) diff --git a/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheEvict.kt b/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheEvict.kt index 925902f..b6e09af 100644 --- a/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheEvict.kt +++ b/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheEvict.kt @@ -24,6 +24,7 @@ import java.lang.annotation.Inherited annotation class ReqShieldCacheEvict( val cacheName: String, val key: String = "", + val keyGenerator: String = "", val isLocalLock: Boolean = true, val lockTimeoutMillis: Long = 3000, val condition: String = "", diff --git a/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheable.kt b/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheable.kt index ce606c5..ff101b7 100644 --- a/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheable.kt +++ b/core-spring-webflux-kotlin-coroutine/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/kotlin/coroutine/annotation/ReqShieldCacheable.kt @@ -26,6 +26,7 @@ import java.lang.annotation.Inherited annotation class ReqShieldCacheable( val cacheName: String, val key: String = "", + val keyGenerator: String = "", val isLocalLock: Boolean = true, val lockTimeoutMillis: Long = 3000, val decisionForUpdate: Int = 90, 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 5f14ab0..0fd9777 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 @@ -26,6 +26,9 @@ import org.aspectj.lang.ProceedingJoinPoint import org.aspectj.lang.annotation.Around import org.aspectj.lang.annotation.Aspect import org.aspectj.lang.reflect.MethodSignature +import org.springframework.beans.factory.BeanFactory +import org.springframework.beans.factory.BeanFactoryAware +import org.springframework.cache.interceptor.KeyGenerator import org.springframework.cache.interceptor.SimpleKeyGenerator import org.springframework.context.expression.MethodBasedEvaluationContext import org.springframework.core.DefaultParameterNameDiscoverer @@ -36,6 +39,7 @@ import org.springframework.expression.Expression import org.springframework.expression.spel.standard.SpelExpressionParser import org.springframework.stereotype.Component import org.springframework.util.StringUtils +import org.springframework.util.function.SingletonSupplier import reactor.core.publisher.Mono import java.lang.reflect.Method import java.util.concurrent.ConcurrentHashMap @@ -45,9 +49,13 @@ import kotlin.coroutines.Continuation @Component class ReqShieldAspect( private val asyncCache: AsyncCache, -) { +) : BeanFactoryAware { + private lateinit var beanFactory: BeanFactory private val springVersion = SpringVersion.getVersion() private val spelParser = SpelExpressionParser() + private var defaultKeyGenerator = SingletonSupplier.of { SimpleKeyGenerator() } + + private val keyGeneratorMap = ConcurrentHashMap() internal val reqShieldMap = ConcurrentHashMap>() @Around("execution(@com.linecorp.cse.reqshield.spring.webflux.kotlin.coroutine.annotation.* * *(.., kotlin.coroutines.Continuation))") @@ -92,20 +100,25 @@ class ReqShieldAspect( fun getCacheEvictAnnotation(joinPoint: ProceedingJoinPoint): ReqShieldCacheEvict = AnnotationUtils.getAnnotation(getTargetMethod(joinPoint), ReqShieldCacheEvict::class.java) - ?: throw IllegalArgumentException("ReqShieldCacheable annotation is required") + ?: throw IllegalArgumentException("ReqShieldCacheEvict annotation is required") internal fun getCacheableCacheKey(joinPoint: ProceedingJoinPoint): String { val annotation = getCacheableAnnotation(joinPoint) - return getCacheKeyOrDefault(annotation.key, joinPoint) + validateCacheKey(annotation.key, annotation.keyGenerator) + + return getCacheKeyOrDefault(annotation.key, annotation.keyGenerator, joinPoint) } internal fun getCacheEvictCacheKey(joinPoint: ProceedingJoinPoint): String { val annotation = getCacheEvictAnnotation(joinPoint) - return getCacheKeyOrDefault(annotation.key, joinPoint) + validateCacheKey(annotation.key, annotation.keyGenerator) + + return getCacheKeyOrDefault(annotation.key, annotation.keyGenerator, joinPoint) } private fun getCacheKeyOrDefault( annotationCacheKey: String, + annotationCacheKeyGenerator: String, joinPoint: ProceedingJoinPoint, ): String { val method = getTargetMethod(joinPoint) @@ -119,17 +132,16 @@ class ReqShieldAspect( val context: EvaluationContext = MethodBasedEvaluationContext(joinPoint.target, method, args, DefaultParameterNameDiscoverer()) - val key: String? = + val key = if (StringUtils.hasText(annotationCacheKey)) { val expression: Expression = spelParser.parseExpression(annotationCacheKey) expression.getValue(context, String::class.java) } else { - SimpleKeyGenerator.generateKey(joinPoint.target, method, args).toString() + val keyGenerator = getOrCreateKeyGenerator(annotationCacheKeyGenerator) + keyGenerator.generate(joinPoint.target, method, args).toString() } - if (key.isNullOrBlank()) { - throw IllegalArgumentException("Null key returned for cache method : $method") - } + require(!key.isNullOrBlank()) { "Null key returned for cache method : $method" } return key } @@ -166,6 +178,25 @@ class ReqShieldAspect( return ReqShield(reqShieldConfiguration) } + private fun validateCacheKey( + cacheKey: String, + cacheKeyGenerator: String, + ) { + if (cacheKey.isNotBlank() && cacheKeyGenerator.isNotBlank()) { + throw IllegalArgumentException("The key and keyGenerator attributes are mutually exclusive.") + } + } + + private fun getOrCreateKeyGenerator(keyGeneratorBeanName: String?): KeyGenerator { + if (keyGeneratorBeanName.isNullOrBlank()) { + return defaultKeyGenerator.obtain() + } + + return keyGeneratorMap.computeIfAbsent(keyGeneratorBeanName) { + beanFactory.getBean(it, KeyGenerator::class.java) + } + } + private fun isCoroutineSupportedSpringVersion(): Boolean { val version = springVersion ?: return false val parts = version.split(".") @@ -177,4 +208,8 @@ class ReqShieldAspect( private fun generateReqShieldKey(joinPoint: ProceedingJoinPoint): String = "${getCacheableAnnotation(joinPoint).cacheName}-${getCacheableCacheKey(joinPoint)}" + + override fun setBeanFactory(beanFactory: BeanFactory) { + this.beanFactory = beanFactory + } } 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 c2c5421..85d2398 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 @@ -31,12 +31,15 @@ import kotlinx.coroutines.async import kotlinx.coroutines.awaitAll import kotlinx.coroutines.test.runTest import org.aspectj.lang.ProceedingJoinPoint -import org.aspectj.lang.reflect.MethodSignature import org.junit.jupiter.api.Assertions.assertNotNull import org.junit.jupiter.api.Assertions.assertTrue import org.junit.jupiter.api.BeforeEach import org.junit.jupiter.api.Test import org.slf4j.LoggerFactory +import org.springframework.beans.factory.BeanFactory +import org.springframework.cache.interceptor.KeyGenerator +import org.springframework.cache.interceptor.SimpleKeyGenerator +import java.lang.reflect.Method import kotlin.coroutines.Continuation import kotlin.coroutines.EmptyCoroutineContext import kotlin.reflect.full.functions @@ -49,77 +52,69 @@ private val log = LoggerFactory.getLogger(ReqShieldAspectTest::class.java) class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { private val asyncCache: AsyncCache = mockk() private val joinPoint: ProceedingJoinPoint = mockk() - private val methodSignature: MethodSignature = mockk() private val reqShieldAspect: ReqShieldAspect = spyk(ReqShieldAspect(asyncCache)) private val targetObject = spyk(TestBean()) + private val argument = mapOf("x" to "paramX", "y" to "paramY") private val mockContinuation = mockk>() - private val kotlinMethod = - TestBean::class.functions.find { - it.name == TestBean::cacheableWithSingleArgument.name && it.parameters.size == 2 - } - private val method = kotlinMethod?.javaMethod +// private val kotlinMethod = +// TestBean::class.functions.find { +// it.name == TestBean::cacheableWithSingleArgument.name && it.parameters.size == 2 +// } +// private val method = kotlinMethod?.javaMethod + + private val cacheName = "TestCacheName" + private val cacheKeyGenerator = "customGenerator" + private val spelEvaluatedKey = "paramXparamY" + private val keyGeneratorKey = "KeyGeneratedByGenerator" + + private val beanFactory = mockk() - private val cacheName = "testCacheName" - private val cacheKey = "#paramMap['x'] + #paramMap['y']" - private val argument = mapOf("x" to "paramX", "y" to "paramY") - private val evaluatedKey = "paramXparamY" private val methodReturn = Product("testProduct", "testCategory") @BeforeEach fun setUp() { - every { joinPoint.signature } returns methodSignature - - every { methodSignature.method } returns method every { mockContinuation.context } returns EmptyCoroutineContext every { joinPoint.args } returns arrayOf(argument, mockContinuation) every { joinPoint.target } returns targetObject - val reqShieldCacheable: ReqShieldCacheable = mockk() - every { reqShieldCacheable.key } returns cacheKey - every { reqShieldCacheable.timeToLiveMillis } returns 60 - every { reqShieldCacheable.isLocalLock } returns false - every { reqShieldCacheable.lockTimeoutMillis } returns 0 - every { reqShieldCacheable.decisionForUpdate } returns 70 - every { reqShieldCacheable.cacheName } returns "product" - - every { reqShieldAspect.getCacheableAnnotation(joinPoint) } returns - ReqShieldCacheable( - cacheName = cacheName, - key = cacheKey, - lockTimeoutMillis = 1000, - timeToLiveMillis = 1000, - ) - - every { reqShieldAspect.getCacheEvictAnnotation(joinPoint) } returns - ReqShieldCacheEvict( - cacheName = cacheName, - lockTimeoutMillis = 1000, - ) + reqShieldAspect.setBeanFactory(beanFactory) } @Test - override fun testAspectOperationVerifyReqShieldAndCacheCreation() = + override fun verifyReqShieldCacheCreation() = runTest { // Mock the cache data using mockk val reqShieldData = ReqShieldData(methodReturn, 1000) coEvery { asyncCache.get(any()) } returns reqShieldData - coEvery { joinPoint.proceed() } coAnswers { targetObject.cacheableWithSingleArgument(argument) } + coEvery { joinPoint.proceed() } coAnswers { targetObject.cacheableWithCustomKey(argument) } + coEvery { reqShieldAspect.getTargetMethod(joinPoint) } returns + TestBean::class + .functions + .find { + it.name == TestBean::cacheableWithCustomKey.name && it.parameters.size == 2 + }?.javaMethod!! // Test the aroundTargetCacheable method val result = reqShieldAspect.aroundTargetCacheable(joinPoint) assertEquals(result, reqShieldData.value) assertTrue(reqShieldAspect.reqShieldMap.size == 1) - assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$evaluatedKey"]) + assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$spelEvaluatedKey"]) } @Test - override fun testAspectOperationReqShieldObjectShouldBeCreatedOnce() = + override fun reqShieldObjectShouldBeCreatedOnce() = runTest { // Mock the cache data using mockk val reqShieldData = ReqShieldData(methodReturn, 1000) coEvery { asyncCache.get(any()) } returns reqShieldData - coEvery { joinPoint.proceed() } coAnswers { targetObject.cacheableWithSingleArgument(argument) } + coEvery { joinPoint.proceed() } coAnswers { targetObject.cacheableWithCustomKey(argument) } + coEvery { reqShieldAspect.getTargetMethod(joinPoint) } returns + TestBean::class + .functions + .find { + it.name == TestBean::cacheableWithCustomKey.name && it.parameters.size == 2 + }?.javaMethod!! val jobs = List(20) { @@ -131,67 +126,109 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { jobs.awaitAll() assertTrue(reqShieldAspect.reqShieldMap.size == 1) - assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$evaluatedKey"]) + assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$spelEvaluatedKey"]) } @Test - override fun testAspectOperationCacheEviction() = + override fun verifyReqShieldCacheEviction() = runTest { // Mock the cache data using mockk val reqShieldData = ReqShieldData(methodReturn, 1000) coEvery { asyncCache.get(any()) } returns reqShieldData - coEvery { asyncCache.evict(any()) } returns true - coEvery { joinPoint.proceed() } coAnswers { targetObject.cacheableWithSingleArgument(argument) } + coEvery { joinPoint.proceed() } coAnswers { targetObject.cacheableWithCustomKey(argument) } + coEvery { reqShieldAspect.getTargetMethod(joinPoint) } returns + TestBean::class + .functions + .find { + it.name == TestBean::cacheableWithDefaultKeyGenerator.name && it.parameters.size == 2 + }?.javaMethod!! // Test the aroundTargetCacheable method val result = reqShieldAspect.aroundTargetCacheable(joinPoint) assertEquals(reqShieldData.value, result) - val kotlinMethod = - TestBean::class.functions.find { - it.name == TestBean::evictWithSingleArgument.name && it.parameters.size == 2 - } - val method = kotlinMethod?.javaMethod - every { methodSignature.method } returns method + // Validate cache eviction + coEvery { asyncCache.evict(any()) } returns true + coEvery { 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) } - coEvery { joinPoint.proceed() } coAnswers { targetObject.evictWithSingleArgument(argument) } val removeProductMono = reqShieldAspect.aroundTargetCacheable(joinPoint) assertTrue(removeProductMono as Boolean) } @Test - override fun testCacheKeyGenerationUseGeneratedKey() = + override fun verifyCacheKeyGenerationWithSpEL() = runTest { - val reqShieldData = ReqShieldData(methodReturn, 1000) - coEvery { asyncCache.get(any()) } returns reqShieldData - coEvery { joinPoint.proceed() } coAnswers { targetObject.cacheableWithSingleArgument(argument) } - - reqShieldAspect.aroundTargetCacheable(joinPoint) + coEvery { reqShieldAspect.getTargetMethod(joinPoint) } returns + TestBean::class + .functions + .find { + it.name == TestBean::cacheableWithCustomKey.name && it.parameters.size == 2 + }?.javaMethod!! + + assertEquals(spelEvaluatedKey, reqShieldAspect.getCacheableCacheKey(joinPoint)) + } - assertEquals(evaluatedKey, reqShieldAspect.getCacheableCacheKey(joinPoint)) + @Test + override fun verifyCacheKeyGenerationWithKeyGenerator() = + runTest { + coEvery { beanFactory.getBean(cacheKeyGenerator, KeyGenerator::class.java) } returns + CustomGenerator() + coEvery { reqShieldAspect.getTargetMethod(joinPoint) } returns + TestBean::class + .functions + .find { + it.name == TestBean::cacheableWithKeyGenerator.name && it.parameters.size == 2 + }?.javaMethod!! + + assertEquals(keyGeneratorKey, reqShieldAspect.getCacheableCacheKey(joinPoint)) } @Test - override fun testCacheKeyGenerationCacheKeyShouldBeSuppliedKey() { - every { reqShieldAspect.getCacheableAnnotation(joinPoint) } returns - ReqShieldCacheable( - cacheName = cacheName, - key = cacheKey, + override fun verifyCacheKeyGenerationWithDefaultGenerator() = + runTest { + coEvery { reqShieldAspect.getTargetMethod(joinPoint) } returns + TestBean::class + .functions + .find { + it.name == TestBean::cacheableWithDefaultKeyGenerator.name && it.parameters.size == 2 + }?.javaMethod!! + + assertEquals( + SimpleKeyGenerator.generateKey(arrayOf(argument)).toString(), + reqShieldAspect.getCacheableCacheKey(joinPoint), ) - - assertEquals(evaluatedKey, reqShieldAspect.getCacheableCacheKey(joinPoint)) - } + } class TestBean { @ReqShieldCacheable(cacheName = "TestCacheName", key = "#paramMap['x'] + #paramMap['y']") - suspend fun cacheableWithSingleArgument(paramMap: Map): String = "ReturnValue: $paramMap" + suspend fun cacheableWithCustomKey(paramMap: Map): String = "ReturnValue: $paramMap" + + @ReqShieldCacheable(cacheName = "TestCacheName") + suspend fun cacheableWithDefaultKeyGenerator(paramMap: Map): String = "ReturnValue: $paramMap" + + @ReqShieldCacheable(cacheName = "TestCacheName", keyGenerator = "customGenerator") + fun cacheableWithKeyGenerator(paramMap: Map): String = "ReturnValue: $paramMap" @ReqShieldCacheEvict(cacheName = "TestCacheName", key = "#paramMap['x'] + #paramMap['y']") - suspend fun evictWithSingleArgument(paramMap: Map): Boolean { + suspend fun evict(paramMap: Map): Boolean { log.debug("cache eviction") return true } } + + class CustomGenerator : KeyGenerator { + override fun generate( + target: Any, + method: Method, + vararg params: Any?, + ): Any = "KeyGeneratedByGenerator" + } } diff --git a/core-spring-webflux/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/annotation/ReqShieldCacheEvict.kt b/core-spring-webflux/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/annotation/ReqShieldCacheEvict.kt index 25a9ff1..7296434 100644 --- a/core-spring-webflux/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/annotation/ReqShieldCacheEvict.kt +++ b/core-spring-webflux/src/main/kotlin/com/linecorp/cse/reqshield/spring/webflux/annotation/ReqShieldCacheEvict.kt @@ -21,6 +21,7 @@ package com.linecorp.cse.reqshield.spring.webflux.annotation annotation class ReqShieldCacheEvict( val cacheName: String, val key: String = "", + val keyGenerator: String = "", val isLocalLock: Boolean = true, val lockTimeoutMillis: Long = 3000, val condition: String = "", 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 70653aa..c093fc7 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 @@ -26,6 +26,7 @@ import java.lang.annotation.Inherited annotation class ReqShieldCacheable( val cacheName: String, val key: String = "", + val keyGenerator: String = "", val isLocalLock: Boolean = true, val lockTimeoutMillis: Long = 3000, val decisionForUpdate: Int = 90, 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 1d0558b..fa65a83 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 @@ -25,6 +25,9 @@ import org.aspectj.lang.ProceedingJoinPoint import org.aspectj.lang.annotation.Around import org.aspectj.lang.annotation.Aspect import org.aspectj.lang.reflect.MethodSignature +import org.springframework.beans.factory.BeanFactory +import org.springframework.beans.factory.BeanFactoryAware +import org.springframework.cache.interceptor.KeyGenerator import org.springframework.cache.interceptor.SimpleKeyGenerator import org.springframework.context.expression.MethodBasedEvaluationContext import org.springframework.core.DefaultParameterNameDiscoverer @@ -34,6 +37,7 @@ import org.springframework.expression.Expression import org.springframework.expression.spel.standard.SpelExpressionParser import org.springframework.stereotype.Component import org.springframework.util.StringUtils +import org.springframework.util.function.SingletonSupplier import reactor.core.publisher.Mono import java.lang.reflect.Method import java.util.concurrent.ConcurrentHashMap @@ -42,8 +46,12 @@ import java.util.concurrent.ConcurrentHashMap @Component class ReqShieldAspect( private val asyncCache: AsyncCache, -) { +) : BeanFactoryAware { + private lateinit var beanFactory: BeanFactory private val spelParser = SpelExpressionParser() + private var defaultKeyGenerator = SingletonSupplier.of { SimpleKeyGenerator() } + + private val keyGeneratorMap = ConcurrentHashMap() internal val reqShieldMap = ConcurrentHashMap>() @Around("@annotation(com.linecorp.cse.reqshield.spring.webflux.annotation.ReqShieldCacheable)") @@ -74,28 +82,33 @@ class ReqShieldAspect( } } - fun getCacheEvictAnnotation(joinPoint: ProceedingJoinPoint): ReqShieldCacheEvict = - AnnotationUtils.getAnnotation(getTargetMethod(joinPoint), ReqShieldCacheEvict::class.java) - ?: throw IllegalArgumentException("ReqShieldCacheable annotation is required") - - internal fun getCacheEvictCacheKey(joinPoint: ProceedingJoinPoint): String { - val annotation = getCacheEvictAnnotation(joinPoint) - return getCacheKeyOrDefault(annotation.key, joinPoint) - } - - internal fun getTargetMethod(joinPoint: ProceedingJoinPoint): Method = (joinPoint.signature as MethodSignature).method - internal fun getCacheableAnnotation(joinPoint: ProceedingJoinPoint): ReqShieldCacheable = AnnotationUtils.getAnnotation(getTargetMethod(joinPoint), ReqShieldCacheable::class.java) ?: throw IllegalArgumentException("ReqShieldCacheable annotation is required") internal fun getCacheableCacheKey(joinPoint: ProceedingJoinPoint): String { val annotation = getCacheableAnnotation(joinPoint) - return getCacheKeyOrDefault(annotation.key, joinPoint) + validateCacheKey(annotation.key, annotation.keyGenerator) + + return getCacheKeyOrDefault(annotation.key, annotation.keyGenerator, joinPoint) } + internal fun getCacheEvictAnnotation(joinPoint: ProceedingJoinPoint): ReqShieldCacheEvict = + AnnotationUtils.getAnnotation(getTargetMethod(joinPoint), ReqShieldCacheEvict::class.java) + ?: throw IllegalArgumentException("ReqShieldCacheEvict annotation is required") + + internal fun getCacheEvictCacheKey(joinPoint: ProceedingJoinPoint): String { + val annotation = getCacheEvictAnnotation(joinPoint) + validateCacheKey(annotation.key, annotation.keyGenerator) + + return getCacheKeyOrDefault(annotation.key, annotation.keyGenerator, joinPoint) + } + + internal fun getTargetMethod(joinPoint: ProceedingJoinPoint): Method = (joinPoint.signature as MethodSignature).method + private fun getCacheKeyOrDefault( annotationCacheKey: String, + annotationCacheKeyGenerator: String, joinPoint: ProceedingJoinPoint, ): String { val method = getTargetMethod(joinPoint) @@ -107,12 +120,11 @@ class ReqShieldAspect( val expression: Expression = spelParser.parseExpression(annotationCacheKey) expression.getValue(context, String::class.java) } else { - SimpleKeyGenerator.generateKey(joinPoint.target, method, joinPoint.args).toString() + val keyGenerator = getOrCreateKeyGenerator(annotationCacheKeyGenerator) + keyGenerator.generate(joinPoint.target, method, joinPoint.args).toString() } - if (key.isNullOrBlank()) { - throw IllegalArgumentException("Null key returned for cache method : $method") - } + require(!key.isNullOrBlank()) { "Null key returned for cache method : $method" } return key } @@ -149,6 +161,29 @@ class ReqShieldAspect( return ReqShield(reqShieldConfiguration) } + private fun validateCacheKey( + cacheKey: String, + cacheKeyGenerator: String, + ) { + if (cacheKey.isNotBlank() && cacheKeyGenerator.isNotBlank()) { + throw IllegalArgumentException("The key and keyGenerator attributes are mutually exclusive.") + } + } + + private fun getOrCreateKeyGenerator(keyGeneratorBeanName: String?): KeyGenerator { + if (keyGeneratorBeanName.isNullOrBlank()) { + return defaultKeyGenerator.obtain() + } + + return keyGeneratorMap.computeIfAbsent(keyGeneratorBeanName) { + beanFactory.getBean(it, KeyGenerator::class.java) + } + } + private fun generateReqShieldKey(joinPoint: ProceedingJoinPoint): String = "${getCacheableAnnotation(joinPoint).cacheName}-${getCacheableCacheKey(joinPoint)}" + + override fun setBeanFactory(beanFactory: BeanFactory) { + this.beanFactory = beanFactory + } } 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 cb80292..afdfab0 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 @@ -26,74 +26,57 @@ import io.mockk.every import io.mockk.mockk import io.mockk.spyk import org.aspectj.lang.ProceedingJoinPoint -import org.aspectj.lang.reflect.MethodSignature import org.junit.jupiter.api.Assertions import org.junit.jupiter.api.BeforeEach import org.junit.jupiter.api.Test +import org.springframework.beans.factory.BeanFactory +import org.springframework.cache.interceptor.KeyGenerator +import org.springframework.cache.interceptor.SimpleKeyGenerator import org.springframework.util.ReflectionUtils import reactor.core.publisher.Flux import reactor.core.publisher.Mono import reactor.core.scheduler.Schedulers import reactor.test.StepVerifier +import java.lang.reflect.Method import kotlin.test.assertEquals import kotlin.test.assertTrue class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { private val asyncCache: AsyncCache = mockk() private val joinPoint = mockk() - private val methodSignature: MethodSignature = mockk() private val reqShieldAspect = spyk(ReqShieldAspect(asyncCache)) private val targetObject = spyk(TestBean()) - private val method = - ReflectionUtils.findMethod( - TestBean::class.java, - TestBean::cacheableWithSingleArgument.name, - Map::class.java, - ) - - private val cacheName = "testCacheName" - private val cacheKey = "#paramMap['x'] + #paramMap['y']" private val argument = mapOf("x" to "paramX", "y" to "paramY") - private val evaulatedKey = "paramXparamY" + + private val cacheName = "TestCacheName" + private val cacheKeyGenerator = "customGenerator" + private val spelEvaluatedKey = "paramXparamY" + private val keyGeneratorKey = "KeyGeneratedByGenerator" + + private val beanFactory = mockk() + private val methodReturn = Product("testProduct", "testCategory") @BeforeEach fun setUp() { - every { joinPoint.signature } returns methodSignature - - every { methodSignature.method } returns method every { joinPoint.args } returns arrayOf(argument) every { joinPoint.target } returns targetObject - val reqShieldCacheable: ReqShieldCacheable = mockk() - every { reqShieldCacheable.key } returns cacheKey - every { reqShieldCacheable.timeToLiveMillis } returns 60 - every { reqShieldCacheable.isLocalLock } returns false - every { reqShieldCacheable.lockTimeoutMillis } returns 0 - every { reqShieldCacheable.decisionForUpdate } returns 70 - every { reqShieldCacheable.cacheName } returns "product" - - every { reqShieldAspect.getCacheableAnnotation(joinPoint) } returns - ReqShieldCacheable( - cacheName = cacheName, - key = cacheKey, - lockTimeoutMillis = 1000, - timeToLiveMillis = 1000, - ) - - every { reqShieldAspect.getCacheEvictAnnotation(joinPoint) } returns - ReqShieldCacheEvict( - cacheName = cacheName, - lockTimeoutMillis = 1000, - ) + reqShieldAspect.setBeanFactory(beanFactory) } @Test - override fun testAspectOperationVerifyReqShieldAndCacheCreation() { + override fun verifyReqShieldCacheCreation() { // Mock the cache data using mockk val reqShieldData = ReqShieldData(methodReturn, 1000) every { asyncCache.get(any()) } returns Mono.just(reqShieldData) - every { joinPoint.proceed() } answers { targetObject.cacheableWithSingleArgument(argument) } + every { joinPoint.proceed() } answers { targetObject.cacheableWithCustomKey(argument) } + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithCustomKey.name, + Map::class.java, + )!! // Test the aroundTargetCacheable method val result = reqShieldAspect.aroundTargetCacheable(joinPoint) @@ -104,16 +87,22 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { .assertNext { value -> assertEquals(reqShieldData.value, value) Assertions.assertTrue(reqShieldAspect.reqShieldMap.size == 1) - Assertions.assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$evaulatedKey"]) + Assertions.assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$spelEvaluatedKey"]) }.verifyComplete() } @Test - override fun testAspectOperationReqShieldObjectShouldBeCreatedOnce() { + override fun reqShieldObjectShouldBeCreatedOnce() { // Mock the cache data using mockk val reqShieldData = ReqShieldData(methodReturn, 1000) every { asyncCache.get(any()) } returns Mono.just(reqShieldData) - every { joinPoint.proceed() } answers { targetObject.cacheableWithSingleArgument(argument) } + every { joinPoint.proceed() } answers { targetObject.cacheableWithCustomKey(argument) } + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithCustomKey.name, + Map::class.java, + )!! val flux = Flux @@ -129,17 +118,22 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { .assertNext { productList -> Assertions.assertTrue(reqShieldAspect.reqShieldMap.size == 1) println(reqShieldAspect.reqShieldMap.keys().toList()) - Assertions.assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$evaulatedKey"]) + Assertions.assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$spelEvaluatedKey"]) }.verifyComplete() } @Test - override fun testAspectOperationCacheEviction() { + override fun verifyReqShieldCacheEviction() { // Mock the cache data using mockk val reqShieldData = ReqShieldData(methodReturn, 1000) every { asyncCache.get(any()) } returns Mono.just(reqShieldData) - every { asyncCache.evict(any()) } returns Mono.just(true) - every { joinPoint.proceed() } answers { targetObject.cacheableWithSingleArgument(argument) } + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithDefaultKeyGenerator.name, + Map::class.java, + )!! + every { joinPoint.proceed() } answers { targetObject.cacheableWithCustomKey(argument) } // Test the aroundTargetCacheable method val result = reqShieldAspect.aroundTargetCacheable(joinPoint) @@ -150,7 +144,16 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { assertEquals(reqShieldData.value, value) }.verifyComplete() - every { joinPoint.proceed() } answers { targetObject.evictWithSingleArgument(argument) } + // Validate cache eviction + every { asyncCache.evict(any()) } returns Mono.just(true) + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::evict.name, + Map::class.java, + )!! + every { joinPoint.proceed() } answers { targetObject.evict(argument) } + val removeProductMono = reqShieldAspect.aroundReqShieldCacheEvict(joinPoint) StepVerifier @@ -161,43 +164,77 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { } @Test - override fun testCacheKeyGenerationUseGeneratedKey() { - // Mock the cache data using mockk - val reqShieldData = ReqShieldData(methodReturn, 1000) - every { asyncCache.get(any()) } returns Mono.just(reqShieldData) - every { joinPoint.proceed() } answers { targetObject.cacheableWithSingleArgument(argument) } - - // Test the aroundTargetCacheable method - val result = reqShieldAspect.aroundTargetCacheable(joinPoint) + override fun verifyCacheKeyGenerationWithSpEL() { + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithCustomKey.name, + Map::class.java, + )!! + + // when, then + Assertions.assertEquals( + spelEvaluatedKey, + reqShieldAspect.getCacheableCacheKey(joinPoint), + ) + } - // Verify the behavior using StepVerifier - StepVerifier - .create(result) - .assertNext { value -> - Assertions.assertEquals( - evaulatedKey, - reqShieldAspect.getCacheableCacheKey(joinPoint), - ) - }.verifyComplete() + override fun verifyCacheKeyGenerationWithKeyGenerator() { + // given + every { beanFactory.getBean(cacheKeyGenerator, KeyGenerator::class.java) } returns + CustomGenerator() + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithKeyGenerator.name, + Map::class.java, + )!! + + // when, then + Assertions.assertEquals( + keyGeneratorKey, + reqShieldAspect.getCacheableCacheKey(joinPoint), + ) } @Test - override fun testCacheKeyGenerationCacheKeyShouldBeSuppliedKey() { - every { reqShieldAspect.getCacheableAnnotation(joinPoint) } returns - ReqShieldCacheable( - cacheName = cacheName, - key = cacheKey, - ) - - Assertions.assertEquals(evaulatedKey, reqShieldAspect.getCacheableCacheKey(joinPoint)) + override fun verifyCacheKeyGenerationWithDefaultGenerator() { + // given + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithDefaultKeyGenerator.name, + Map::class.java, + )!! + + // when, then + Assertions.assertEquals( + SimpleKeyGenerator.generateKey(arrayOf(argument)).toString(), + reqShieldAspect.getCacheableCacheKey(joinPoint), + ) } class TestBean { @ReqShieldCacheable(cacheName = "TestCacheName", key = "#paramMap['x'] + #paramMap['y']") - fun cacheableWithSingleArgument(paramMap: Map): Mono = + fun cacheableWithCustomKey(paramMap: Map): Mono = Mono.justOrEmpty(Product("testProduct", "testCategory")) + + @ReqShieldCacheable(cacheName = "TestCacheName") + fun cacheableWithDefaultKeyGenerator(paramMap: Map): Mono = + Mono.justOrEmpty(Product("testProduct", "testCategory")) + + @ReqShieldCacheable(cacheName = "TestCacheName", keyGenerator = "customGenerator") + fun cacheableWithKeyGenerator(paramMap: Map): Mono = Mono.justOrEmpty(Product("testProduct", "testCategory")) @ReqShieldCacheEvict(cacheName = "TestCacheName") - fun evictWithSingleArgument(paramMap: Map): Mono = Mono.just(true) + fun evict(paramMap: Map): Mono = Mono.just(true) + } + + class CustomGenerator : KeyGenerator { + override fun generate( + target: Any, + method: Method, + vararg params: Any?, + ): Any = "KeyGeneratedByGenerator" } } diff --git a/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheEvict.kt b/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheEvict.kt index 2b9b130..44b5635 100644 --- a/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheEvict.kt +++ b/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheEvict.kt @@ -24,6 +24,7 @@ import java.lang.annotation.Inherited annotation class ReqShieldCacheEvict( val cacheName: String, val key: String = "", + val keyGenerator: String = "", val isLocalLock: Boolean = true, val lockTimeoutMillis: Long = 3000, val condition: String = "", diff --git a/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheable.kt b/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheable.kt index f0d01d8..edaf2c2 100644 --- a/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheable.kt +++ b/core-spring/src/main/kotlin/com/linecorp/cse/reqshield/spring/annotation/ReqShieldCacheable.kt @@ -26,6 +26,7 @@ import java.lang.annotation.Inherited annotation class ReqShieldCacheable( val cacheName: String, val key: String = "", + val keyGenerator: String = "", val isLocalLock: Boolean = true, val lockTimeoutMillis: Long = 30000, val decisionForUpdate: Int = 90, 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 2c24632..bf2cae7 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 @@ -25,6 +25,9 @@ import org.aspectj.lang.ProceedingJoinPoint import org.aspectj.lang.annotation.Around import org.aspectj.lang.annotation.Aspect import org.aspectj.lang.reflect.MethodSignature +import org.springframework.beans.factory.BeanFactory +import org.springframework.beans.factory.BeanFactoryAware +import org.springframework.cache.interceptor.KeyGenerator import org.springframework.cache.interceptor.SimpleKeyGenerator import org.springframework.context.expression.MethodBasedEvaluationContext import org.springframework.core.DefaultParameterNameDiscoverer @@ -34,6 +37,7 @@ import org.springframework.expression.Expression import org.springframework.expression.spel.standard.SpelExpressionParser import org.springframework.stereotype.Component import org.springframework.util.StringUtils +import org.springframework.util.function.SingletonSupplier import java.lang.reflect.Method import java.util.concurrent.ConcurrentHashMap @@ -41,8 +45,12 @@ import java.util.concurrent.ConcurrentHashMap @Component class ReqShieldAspect( private val reqShieldCache: ReqShieldCache, -) { +) : BeanFactoryAware { + private lateinit var beanFactory: BeanFactory private val spelParser = SpelExpressionParser() + private val defaultKeyGenerator = SingletonSupplier.of { SimpleKeyGenerator() } + + private val keyGeneratorMap = ConcurrentHashMap() internal val reqShieldMap = ConcurrentHashMap>() @Around("@annotation(com.linecorp.cse.reqshield.spring.annotation.ReqShieldCacheable)") @@ -109,20 +117,25 @@ class ReqShieldAspect( internal fun getCacheEvictAnnotation(joinPoint: ProceedingJoinPoint): ReqShieldCacheEvict = AnnotationUtils.getAnnotation(getTargetMethod(joinPoint), ReqShieldCacheEvict::class.java) - ?: throw IllegalArgumentException("ReqShieldCacheable annotation is required") + ?: throw IllegalArgumentException("ReqShieldCacheEvict annotation is required") internal fun getCacheableCacheKey(joinPoint: ProceedingJoinPoint): String { val annotation = getCacheableAnnotation(joinPoint) - return getCacheKeyOrDefault(annotation.key, joinPoint) + validateCacheKey(annotation.key, annotation.keyGenerator) + + return getCacheKeyOrDefault(annotation.key, annotation.keyGenerator, joinPoint) } internal fun getCacheEvictCacheKey(joinPoint: ProceedingJoinPoint): String { val annotation = getCacheEvictAnnotation(joinPoint) - return getCacheKeyOrDefault(annotation.key, joinPoint) + validateCacheKey(annotation.key, annotation.keyGenerator) + + return getCacheKeyOrDefault(annotation.key, annotation.keyGenerator, joinPoint) } private fun getCacheKeyOrDefault( annotationCacheKey: String, + annotationCacheKeyGenerator: String, joinPoint: ProceedingJoinPoint, ): String { val method = getTargetMethod(joinPoint) @@ -134,16 +147,38 @@ class ReqShieldAspect( val expression: Expression = spelParser.parseExpression(annotationCacheKey) expression.getValue(context, String::class.java) } else { - SimpleKeyGenerator.generateKey(joinPoint.target, method, joinPoint.args).toString() + val keyGenerator = getOrCreateKeyGenerator(annotationCacheKeyGenerator) + keyGenerator.generate(joinPoint.target, method, joinPoint.args).toString() } - if (key.isNullOrBlank()) { - throw IllegalArgumentException("Null key returned for cache method : $method") - } + require(!key.isNullOrBlank()) { "Null key returned for cache method : $method" } return key } + private fun validateCacheKey( + cacheKey: String, + cacheKeyGenerator: String, + ) { + if (cacheKey.isNotBlank() && cacheKeyGenerator.isNotBlank()) { + throw IllegalArgumentException("The key and keyGenerator attributes are mutually exclusive.") + } + } + + private fun getOrCreateKeyGenerator(keyGeneratorBeanName: String?): KeyGenerator { + if (keyGeneratorBeanName.isNullOrBlank()) { + return defaultKeyGenerator.obtain() + } + + return keyGeneratorMap.computeIfAbsent(keyGeneratorBeanName) { + beanFactory.getBean(it, KeyGenerator::class.java) + } + } + private fun generateReqShieldKey(joinPoint: ProceedingJoinPoint): String = "${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 d227ae8..619bd08 100644 --- a/core-spring/src/test/kotlin/aspect/ReqShieldAspectTest.kt +++ b/core-spring/src/test/kotlin/aspect/ReqShieldAspectTest.kt @@ -36,7 +36,11 @@ import org.junit.jupiter.api.Assertions.assertTrue import org.junit.jupiter.api.BeforeEach import org.junit.jupiter.api.Test import org.slf4j.LoggerFactory +import org.springframework.beans.factory.BeanFactory +import org.springframework.cache.interceptor.KeyGenerator +import org.springframework.cache.interceptor.SimpleKeyGenerator import org.springframework.util.ReflectionUtils +import java.lang.reflect.Method import java.time.Duration import java.util.concurrent.Executors @@ -47,63 +51,60 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { private val joinPoint = mockk() private val reqShieldAspect = spyk(ReqShieldAspect(reqShieldCache)) private val targetObject = spyk(TestBean()) - private val method = - ReflectionUtils.findMethod( - TestBean::class.java, - TestBean::cacheableWithSingleArgument.name, - Map::class.java, - ) - - private val cacheName = "testCacheName" - private val cacheKey = "#paramMap['x'] + #paramMap['y']" private val argument = mapOf("x" to "paramX", "y" to "paramY") - private val evaluatedKey = "paramXparamY" + + private val cacheName = "TestCacheName" + private val cacheKeyGenerator = "customGenerator" + private val spelEvaluatedKey = "paramXparamY" + private val keyGeneratorKey = "KeyGeneratedByGenerator" + + private val beanFactory = mockk() + private val methodReturn = Product("testProduct", "testCategory") @BeforeEach fun setUp() { every { joinPoint.target } returns targetObject every { joinPoint.args } returns arrayOf(argument) - every { joinPoint.proceed() } answers { targetObject.cacheableWithSingleArgument(argument) } - - every { reqShieldAspect.getTargetMethod(joinPoint) } returns method!! - - every { reqShieldAspect.getCacheableAnnotation(joinPoint) } returns - ReqShieldCacheable( - cacheName = cacheName, - key = cacheKey, - lockTimeoutMillis = 1000, - timeToLiveMillis = 1000, - ) - - every { reqShieldAspect.getCacheEvictAnnotation(joinPoint) } returns - ReqShieldCacheEvict( - cacheName = cacheName, - key = cacheKey, - ) + + reqShieldAspect.setBeanFactory(beanFactory) } @Test - override fun testAspectOperationVerifyReqShieldAndCacheCreation() { + override fun verifyReqShieldCacheCreation() { // given - every { reqShieldCache.get(any()) } returns ReqShieldData(methodReturn, 1000) - every { joinPoint.proceed() } answers { targetObject.cacheableWithSingleArgument(argument) } + val reqShieldData = ReqShieldData(methodReturn, 1000) + every { reqShieldCache.get(any()) } returns reqShieldData + every { joinPoint.proceed() } answers { targetObject.cacheableWithCustomKey(argument) } + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithCustomKey.name, + Map::class.java, + )!! // when - reqShieldAspect.aroundReqShieldCacheable(joinPoint) + val result = reqShieldAspect.aroundReqShieldCacheable(joinPoint) Awaitility.await().atMost(Duration.ofMillis(BaseReqShieldTest.AWAIT_TIMEOUT)).untilAsserted { // then + assertEquals(reqShieldData.value, result) assertTrue(reqShieldAspect.reqShieldMap.size == 1) - assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$evaluatedKey"]) + assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$spelEvaluatedKey"]) } } @Test - override fun testAspectOperationReqShieldObjectShouldBeCreatedOnce() { + override fun reqShieldObjectShouldBeCreatedOnce() { // given every { reqShieldCache.get(any()) } returns ReqShieldData(methodReturn, 1000) - every { joinPoint.proceed() } answers { targetObject.cacheableWithSingleArgument(argument) } + every { joinPoint.proceed() } answers { targetObject.cacheableWithCustomKey(argument) } + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithCustomKey.name, + Map::class.java, + )!! // when val executorService = Executors.newFixedThreadPool(10) @@ -116,22 +117,36 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { Awaitility.await().atMost(Duration.ofMillis(BaseReqShieldTest.AWAIT_TIMEOUT)).untilAsserted { // then assertTrue(reqShieldAspect.reqShieldMap.size == 1) - assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$evaluatedKey"]) + assertNotNull(reqShieldAspect.reqShieldMap["$cacheName-$spelEvaluatedKey"]) } } @Test - override fun testAspectOperationCacheEviction() { + override fun verifyReqShieldCacheEviction() { // given val reqShieldData = ReqShieldData(methodReturn, 10000) every { reqShieldCache.get(any()) } returns reqShieldData - every { reqShieldCache.evict(any()) } returns true - every { joinPoint.proceed() } answers { targetObject.cacheableWithSingleArgument(argument) } + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithCustomKey.name, + Map::class.java, + )!! + every { joinPoint.proceed() } answers { targetObject.cacheableWithCustomKey(argument) } val cachedResult = reqShieldAspect.aroundReqShieldCacheable(joinPoint) assertEquals(reqShieldData.value, cachedResult) + // Validate cache eviction + every { reqShieldCache.evict(any()) } returns true + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::evict.name, + Map::class.java, + )!! + // when reqShieldAspect.aroundReqShieldCacheEvict(joinPoint) @@ -140,44 +155,88 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { } @Test - override fun testCacheKeyGenerationUseGeneratedKey() { - every { joinPoint.target } returns targetObject - every { joinPoint.args } returns arrayOf(argument) - every { reqShieldAspect.getTargetMethod(joinPoint) } returns method!! - - every { reqShieldAspect.getCacheableAnnotation(joinPoint) } returns - ReqShieldCacheable( - cacheName = cacheName, - key = cacheKey, - ) + override fun verifyCacheKeyGenerationWithSpEL() { + // given + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithCustomKey.name, + Map::class.java, + )!! + + // when, then + assertEquals( + spelEvaluatedKey, + reqShieldAspect.getCacheableCacheKey(joinPoint), + ) + } + @Test + override fun verifyCacheKeyGenerationWithKeyGenerator() { + // given + every { beanFactory.getBean(cacheKeyGenerator, KeyGenerator::class.java) } returns + CustomGenerator() + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithKeyGenerator.name, + Map::class.java, + )!! + + // when, then assertEquals( - evaluatedKey, + keyGeneratorKey, reqShieldAspect.getCacheableCacheKey(joinPoint), ) } @Test - override fun testCacheKeyGenerationCacheKeyShouldBeSuppliedKey() { - every { reqShieldAspect.getCacheableAnnotation(joinPoint) } returns - ReqShieldCacheable( - cacheName = cacheName, - key = cacheKey, - ) - - assertEquals(evaluatedKey, reqShieldAspect.getCacheableCacheKey(joinPoint)) + override fun verifyCacheKeyGenerationWithDefaultGenerator() { + // given + every { reqShieldAspect.getTargetMethod(joinPoint) } returns + ReflectionUtils.findMethod( + TestBean::class.java, + TestBean::cacheableWithDefaultKeyGenerator.name, + Map::class.java, + )!! + + // when, then + assertEquals( + SimpleKeyGenerator.generateKey(arrayOf(argument)).toString(), + reqShieldAspect.getCacheableCacheKey(joinPoint), + ) } class TestBean { @ReqShieldCacheable(cacheName = "TestCacheName", key = "#paramMap['x'] + #paramMap['y']") - fun cacheableWithSingleArgument(paramMap: Map): String { - log.debug("method invoked") + fun cacheableWithCustomKey(paramMap: Map): String { + log.debug("cacheableWithCustomKey method invoked") + return "ReturnValue: $paramMap" + } + + @ReqShieldCacheable(cacheName = "TestCacheName") + fun cacheableWithDefaultKeyGenerator(paramMap: Map): String { + log.debug("cacheableWithDefaultKeyGenerator method invoked") + return "ReturnValue: $paramMap" + } + + @ReqShieldCacheable(cacheName = "TestCacheName", keyGenerator = "customGenerator") + fun cacheableWithKeyGenerator(paramMap: Map): String { + log.debug("cacheableWithCustomGenerator method invoked") return "ReturnValue: $paramMap" } @ReqShieldCacheEvict(cacheName = "TestCacheName", key = "#paramMap['x'] + #paramMap['y']") - fun evictWithSingleArgument(paramMap: Map) { + fun evict(paramMap: Map) { log.debug("cache eviction") } } + + class CustomGenerator : KeyGenerator { + override fun generate( + target: Any, + method: Method, + vararg params: Any?, + ): Any = "KeyGeneratedByGenerator" + } } diff --git a/support/src/testFixtures/kotlin/com/linecorp/cse/reqshield/support/BaseReqShieldModuleSupportTest.kt b/support/src/testFixtures/kotlin/com/linecorp/cse/reqshield/support/BaseReqShieldModuleSupportTest.kt index b37c98d..35dac2c 100644 --- a/support/src/testFixtures/kotlin/com/linecorp/cse/reqshield/support/BaseReqShieldModuleSupportTest.kt +++ b/support/src/testFixtures/kotlin/com/linecorp/cse/reqshield/support/BaseReqShieldModuleSupportTest.kt @@ -17,13 +17,15 @@ package com.linecorp.cse.reqshield.support interface BaseReqShieldModuleSupportTest { - fun testAspectOperationVerifyReqShieldAndCacheCreation() + fun verifyReqShieldCacheCreation() - fun testAspectOperationReqShieldObjectShouldBeCreatedOnce() + fun reqShieldObjectShouldBeCreatedOnce() - fun testAspectOperationCacheEviction() + fun verifyReqShieldCacheEviction() - fun testCacheKeyGenerationUseGeneratedKey() + fun verifyCacheKeyGenerationWithSpEL() - fun testCacheKeyGenerationCacheKeyShouldBeSuppliedKey() + fun verifyCacheKeyGenerationWithKeyGenerator() + + fun verifyCacheKeyGenerationWithDefaultGenerator() } From 651102051c9f1d10aae97431088e85944bca2ba6 Mon Sep 17 00:00:00 2001 From: byungchan-lee Date: Fri, 4 Apr 2025 15:23:13 +0900 Subject: [PATCH 09/11] SpEL Support - set SpEL key in new test code * remove comments --- .../webflux/kotlin/coroutine/aspect/ReqShieldAspectTest.kt | 5 ----- 1 file changed, 5 deletions(-) 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 85d2398..dfc657d 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 @@ -56,11 +56,6 @@ class ReqShieldAspectTest : BaseReqShieldModuleSupportTest { private val targetObject = spyk(TestBean()) private val argument = mapOf("x" to "paramX", "y" to "paramY") private val mockContinuation = mockk>() -// private val kotlinMethod = -// TestBean::class.functions.find { -// it.name == TestBean::cacheableWithSingleArgument.name && it.parameters.size == 2 -// } -// private val method = kotlinMethod?.javaMethod private val cacheName = "TestCacheName" private val cacheKeyGenerator = "customGenerator" From 40e7c303acc7d22317d8575ca60e2759754fec6a Mon Sep 17 00:00:00 2001 From: "kanghyun.yang" Date: Sun, 10 Aug 2025 19:23:52 +0900 Subject: [PATCH 10/11] [ISSUE #37] Add CLAUDE.md --- CLAUDE.md | 131 ++++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 131 insertions(+) create mode 100644 CLAUDE.md diff --git a/CLAUDE.md b/CLAUDE.md new file mode 100644 index 0000000..58d5131 --- /dev/null +++ b/CLAUDE.md @@ -0,0 +1,131 @@ +# CLAUDE.md + +This file provides guidance to Claude Code (claude.ai/code) when working with code in this repository. + +## Project Overview + +Req-Shield is a Kotlin library that provides request-collapsing functionality for cache-based applications. It prevents the thundering herd problem by ensuring only one request for the same cache key is processed at a time, while subsequent concurrent requests wait for the result. + +## Core Architecture + +The library is organized into several core modules: + +- **core**: Base implementation with synchronous operations +- **core-reactor**: Reactive implementation using Project Reactor +- **core-kotlin-coroutine**: Coroutine-based implementation +- **core-spring**: Spring integration with traditional caching +- **core-spring-webflux**: Spring WebFlux integration +- **core-spring-webflux-kotlin-coroutine**: Spring WebFlux + Kotlin Coroutines integration +- **support**: Shared utilities, models, and constants + +### Key Components + +1. **ReqShield**: Main orchestrator that manages cache operations and request collapsing +2. **KeyLock**: Locking mechanism (local or global) to prevent concurrent cache operations +3. **ReqShieldConfiguration**: Configuration object that defines cache functions, locking behavior, and timeouts +4. **ReqShieldData**: Wrapper for cached data with metadata (creation time, TTL) +5. **Spring Aspects**: AOP-based implementations that provide annotation-driven caching + +### Design Patterns + +- **Template Method**: Core ReqShield logic is template-based with pluggable cache and lock functions +- **Strategy Pattern**: Different locking strategies (local vs global) and work modes +- **Aspect-Oriented Programming**: Spring modules use AOP for transparent caching +- **Factory Pattern**: Configuration objects create appropriate lock implementations + +## Development Commands + +### Building the Project +```bash +./gradlew build # Build all modules +./gradlew :core:build # Build specific module +./gradlew clean build # Clean and build +``` + +### Running Tests +```bash +./gradlew test # Run all tests +./gradlew :core:test # Run tests for specific module +./gradlew test --tests "*ReqShieldTest*" # Run specific test pattern +``` + +### Code Quality +```bash +./gradlew ktlintCheck # Check Kotlin code style +./gradlew ktlintFormat # Format Kotlin code +./gradlew jacocoTestReport # Generate test coverage report +``` + +### Running Examples +```bash +./gradlew :req-shield-spring-boot3-example:bootRun +./gradlew :req-shield-spring-webflux-example:bootRun +./gradlew :req-shield-spring-webflux-kotlin-coroutine-example:bootRun +``` + +## Module Structure + +### Core Modules +Each core module follows the same package structure: +- `com.linecorp.cse.reqshield.{variant}/` - Main classes (ReqShield, KeyLock implementations) +- `com.linecorp.cse.reqshield.{variant}/config/` - Configuration classes + +### Spring Integration Modules +Spring modules add: +- `annotation/` - Cache annotations (@ReqShieldCacheable, @ReqShieldCacheEvict) +- `aspect/` - AOP implementation for intercepting annotated methods +- `cache/` - Cache interface implementations +- `config/` - Auto-configuration for Spring Boot + +### Support Module +Contains shared: +- `constant/` - Configuration constants and defaults +- `exception/` - Custom exceptions and error codes +- `model/` - Data models (ReqShieldData) +- `utils/` - Utility functions for cache decisions + +## Key Configuration Options + +### ReqShieldConfiguration Parameters +- `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) +- `reqShieldWorkMode`: CREATE_AND_UPDATE_CACHE | ONLY_CREATE_CACHE | ONLY_UPDATE_CACHE + +### Work Modes +- **CREATE_AND_UPDATE_CACHE**: Full functionality (default) +- **ONLY_CREATE_CACHE**: Never updates existing cache entries +- **ONLY_UPDATE_CACHE**: Never creates new cache entries + +## Testing Guidelines + +### Test Infrastructure +- Uses JUnit 5 platform +- MockK for Kotlin mocking +- Testcontainers for integration tests (Redis) +- Awaitility for asynchronous testing +- Separate test fixtures in `support` module + +### Test Coverage Requirements +- **Minimum test coverage**: 80% must be maintained across all modules +- Coverage reports generated via `./gradlew jacocoTestReport` +- Coverage enforced through Jacoco plugin configuration + +### Test Categories +- **Unit Tests**: Test individual components in isolation +- **Integration Tests**: Test Spring integration with real Redis containers +- **Base Test Classes**: Located in `support/src/testFixtures/` for reuse across modules + +## Version Compatibility + +### Java/Kotlin Compatibility +- **Core modules**: Java 8+, Kotlin 1.8+ +- **Spring Boot 3 examples**: Java 17+ +- **Spring Boot 2 examples**: Java 8+ + +### Framework Support +- Spring Boot 2.7+ (Spring Framework 5.3+) +- Spring Boot 3.3+ +- Project Reactor 3.4+ +- Kotlin Coroutines 1.7+ From a735c14b8c8141f5eecf3bdfa92f784334e7667b Mon Sep 17 00:00:00 2001 From: "kanghyun.yang" Date: Sun, 10 Aug 2025 20:18:34 +0900 Subject: [PATCH 11/11] [ISSUE #39] Performance Optimization: Core Module Thread Pool and Lock Management --- .../linecorp/cse/reqshield/KeyLocalLock.kt | 33 +++-- .../com/linecorp/cse/reqshield/ReqShield.kt | 5 +- .../config/ReqShieldConfiguration.kt | 2 +- .../cse/reqshield/KeyLocalLockShutdownTest.kt | 138 ++++++++++++++++++ .../config/ReqShieldConfigurationTest.kt | 66 +++++++++ 5 files changed, 228 insertions(+), 16 deletions(-) create mode 100644 core/src/test/kotlin/com/linecorp/cse/reqshield/KeyLocalLockShutdownTest.kt create mode 100644 core/src/test/kotlin/com/linecorp/cse/reqshield/config/ReqShieldConfigurationTest.kt 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 ed563d2..8855306 100644 --- a/core/src/main/kotlin/com/linecorp/cse/reqshield/KeyLocalLock.kt +++ b/core/src/main/kotlin/com/linecorp/cse/reqshield/KeyLocalLock.kt @@ -20,9 +20,10 @@ import com.linecorp.cse.reqshield.support.constant.ConfigValues.LOCK_MONITOR_INT import com.linecorp.cse.reqshield.support.utils.nowToEpochTime import org.slf4j.LoggerFactory import java.util.concurrent.ConcurrentHashMap -import java.util.concurrent.ExecutorService import java.util.concurrent.Executors +import java.util.concurrent.ScheduledExecutorService import java.util.concurrent.Semaphore +import java.util.concurrent.TimeUnit private val log = LoggerFactory.getLogger(KeyLocalLock::class.java) @@ -31,20 +32,17 @@ class KeyLocalLock(private val lockTimeoutMillis: Long) : KeyLock { private val lockMap = ConcurrentHashMap() - private val executorService: ExecutorService = Executors.newSingleThreadExecutor() + private val scheduledExecutor: ScheduledExecutorService = Executors.newSingleThreadScheduledExecutor() init { - executorService.execute { - while (true) { - try { - val now = System.currentTimeMillis() - lockMap.entries.removeIf { now - it.value.createdAt > lockTimeoutMillis } - Thread.sleep(LOCK_MONITOR_INTERVAL_MILLIS) - } catch (e: InterruptedException) { - log.error("Error in lock lifecycle monitoring : {}", e.message) - } + scheduledExecutor.scheduleWithFixedDelay({ + try { + val now = System.currentTimeMillis() + lockMap.entries.removeIf { now - it.value.createdAt > lockTimeoutMillis } + } catch (e: Exception) { + log.error("Error in lock lifecycle monitoring : {}", e.message) } - } + }, 0, LOCK_MONITOR_INTERVAL_MILLIS, TimeUnit.MILLISECONDS) } override fun tryLock( @@ -69,4 +67,15 @@ class KeyLocalLock(private val lockTimeoutMillis: Long) : KeyLock { } return true } + + fun shutdown() { + scheduledExecutor.shutdown() + try { + if (!scheduledExecutor.awaitTermination(5, TimeUnit.SECONDS)) { + scheduledExecutor.shutdownNow() + } + } catch (_: InterruptedException) { + scheduledExecutor.shutdownNow() + } + } } 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 166bfca..20a9f2a 100644 --- a/core/src/main/kotlin/com/linecorp/cse/reqshield/ReqShield.kt +++ b/core/src/main/kotlin/com/linecorp/cse/reqshield/ReqShield.kt @@ -201,15 +201,14 @@ class ReqShield( value: ReqShieldData, lockType: LockType, ) { - var unlockSuccess = false - var retryCount = 0 - try { setFunction.invoke(key, value, value.timeToLiveMillis) } catch (e: Exception) { 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 diff --git a/core/src/main/kotlin/com/linecorp/cse/reqshield/config/ReqShieldConfiguration.kt b/core/src/main/kotlin/com/linecorp/cse/reqshield/config/ReqShieldConfiguration.kt index 68994de..58f5a58 100644 --- a/core/src/main/kotlin/com/linecorp/cse/reqshield/config/ReqShieldConfiguration.kt +++ b/core/src/main/kotlin/com/linecorp/cse/reqshield/config/ReqShieldConfiguration.kt @@ -36,7 +36,7 @@ data class ReqShieldConfiguration( val lockTimeoutMillis: Long = DEFAULT_LOCK_TIMEOUT_MILLIS, val executor: ScheduledExecutorService = Executors.newScheduledThreadPool( - Runtime.getRuntime().availableProcessors() * 10, + maxOf(2, Runtime.getRuntime().availableProcessors() * 2), ), val decisionForUpdate: Int = DEFAULT_DECISION_FOR_UPDATE, val keyLock: KeyLock = diff --git a/core/src/test/kotlin/com/linecorp/cse/reqshield/KeyLocalLockShutdownTest.kt b/core/src/test/kotlin/com/linecorp/cse/reqshield/KeyLocalLockShutdownTest.kt new file mode 100644 index 0000000..fa336b8 --- /dev/null +++ b/core/src/test/kotlin/com/linecorp/cse/reqshield/KeyLocalLockShutdownTest.kt @@ -0,0 +1,138 @@ +/* + * 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 + +import org.awaitility.Awaitility.await +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Assertions.assertTrue +import org.junit.jupiter.api.Test +import java.time.Duration + +class KeyLocalLockShutdownTest { + @Test + fun testShutdownPreventsMemoryLeak() { + val keyLock = KeyLocalLock(5000L) // 5 second timeout + + // Create multiple locks to verify monitoring works + val key1 = "testKey1" + val key2 = "testKey2" + + keyLock.tryLock(key1, LockType.CREATE) + keyLock.tryLock(key2, LockType.CREATE) + + // Call shutdown + keyLock.shutdown() + + // Executor should be shutdown after shutdown() call + assertTrue( + keyLock.javaClass.getDeclaredField("scheduledExecutor").let { field -> + field.isAccessible = true + val executor = field.get(keyLock) as java.util.concurrent.ScheduledExecutorService + executor.isShutdown + }, + "ScheduledExecutor should be shutdown", + ) + } + + @Test + fun testScheduledCleanupRemovesExpiredLocks() { + val shortTimeout = 100L // Fast test with 100ms timeout + val keyLock = KeyLocalLock(shortTimeout) + + val key = "expireTestKey" + val lockType = LockType.CREATE + + // Acquire lock + assertTrue(keyLock.tryLock(key, lockType), "Should acquire lock initially") + + // Should fail to acquire same lock again (already acquired) + assertTrue(!keyLock.tryLock(key, lockType), "Should not acquire same lock again") + + // Should be automatically cleaned up after timeout + await().atMost(Duration.ofMillis(shortTimeout + 200L)).untilAsserted { + // Should be able to acquire new lock after timeout + assertTrue(keyLock.tryLock(key, lockType), "Should be able to acquire lock after timeout") + keyLock.unLock(key, lockType) // Cleanup + } + + keyLock.shutdown() + } + + @Test + fun testScheduledExecutorIntervalWorksCorrectly() { + val lockTimeoutMillis = 500L + val keyLock = KeyLocalLock(lockTimeoutMillis) + + // Create multiple expired locks + for (i in 1..5) { + keyLock.tryLock("key$i", LockType.CREATE) + } + + // Verify cleanup after certain time + await().atMost(Duration.ofMillis(lockTimeoutMillis + 500L)).untilAsserted { + // Should be able to acquire locks again after all keys are cleaned up + var availableCount = 0 + for (i in 1..5) { + if (keyLock.tryLock("key$i", LockType.CREATE)) { + availableCount++ + keyLock.unLock("key$i", LockType.CREATE) + } + } + assertEquals(5, availableCount, "All expired locks should be cleaned up") + } + + keyLock.shutdown() + } + + @Test + fun testShutdownWithTimeout() { + val keyLock = KeyLocalLock(1000L) + + val startTime = System.currentTimeMillis() + keyLock.shutdown() + val endTime = System.currentTimeMillis() + + // Shutdown should complete within 5 seconds + assertTrue( + endTime - startTime < 6000, + "Shutdown should complete within timeout period", + ) + } + + @Test + fun testInterruptedShutdown() { + val keyLock = KeyLocalLock(1000L) + + // Call shutdown with current thread in interrupt state + Thread.currentThread().interrupt() + + keyLock.shutdown() + + // Shutdown should complete normally (regardless of interrupt state) + assertTrue( + keyLock.javaClass.getDeclaredField("scheduledExecutor").let { field -> + field.isAccessible = true + val executor = field.get(keyLock) as java.util.concurrent.ScheduledExecutorService + executor.isShutdown + }, + "ScheduledExecutor should be shutdown even when interrupted", + ) + + // Clear interrupt state + Thread.interrupted() + } +} diff --git a/core/src/test/kotlin/com/linecorp/cse/reqshield/config/ReqShieldConfigurationTest.kt b/core/src/test/kotlin/com/linecorp/cse/reqshield/config/ReqShieldConfigurationTest.kt new file mode 100644 index 0000000..74a566e --- /dev/null +++ b/core/src/test/kotlin/com/linecorp/cse/reqshield/config/ReqShieldConfigurationTest.kt @@ -0,0 +1,66 @@ +/* + * 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.config + +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Assertions.assertTrue +import org.junit.jupiter.api.Test +import java.util.concurrent.ThreadPoolExecutor + +class ReqShieldConfigurationTest { + @Test + fun testDefaultThreadPoolSizeIsOptimal() { + val config = + ReqShieldConfiguration( + setCacheFunction = { _, _, _ -> true }, + getCacheFunction = { null }, + ) + + val executor = config.executor as? ThreadPoolExecutor + assertTrue(executor != null, "Executor should be ThreadPoolExecutor") + + val expectedSize = maxOf(2, Runtime.getRuntime().availableProcessors() * 2) + assertEquals(expectedSize, executor!!.corePoolSize, "Thread pool size should be optimal") + assertTrue( + executor.corePoolSize <= Runtime.getRuntime().availableProcessors() * 2, + "Thread pool should not be excessive", + ) + } + + @Test + fun testMinimumThreadPoolSize() { + // Simulate single core environment + val expectedMinimum = 2 + val calculatedSize = maxOf(2, 1 * 2) // Case when availableProcessors() = 1 + + assertEquals( + expectedMinimum, + calculatedSize, + "Minimum thread pool size should be 2 even on single core systems", + ) + } + + @Test + fun testMultiCoreThreadPoolSize() { + // Calculate appropriate size for multi-core environment + val coreCount = Runtime.getRuntime().availableProcessors() + val expectedSize = maxOf(2, coreCount * 2) + + assertTrue(expectedSize >= 2, "Thread pool size should be at least 2") + assertTrue(expectedSize <= coreCount * 2, "Thread pool should not exceed cores * 2") + } +}