diff --git a/src/main/kotlin/com/didit/application/retrospect/RetrospectiveCompletionCoordinator.kt b/src/main/kotlin/com/didit/application/retrospect/RetrospectiveCompletionCoordinator.kt index ffa19151..116aae6e 100644 --- a/src/main/kotlin/com/didit/application/retrospect/RetrospectiveCompletionCoordinator.kt +++ b/src/main/kotlin/com/didit/application/retrospect/RetrospectiveCompletionCoordinator.kt @@ -26,6 +26,7 @@ class RetrospectiveCompletionCoordinator( private val aiClient: AIClient, private val transactionTemplate: TransactionTemplate, private val metrics: RetrospectiveAiMetrics, + private val summarySaveConcurrencyLimiter: SummarySaveConcurrencyLimiter, ) { companion object { private val logger = LoggerFactory.getLogger(RetrospectiveCompletionCoordinator::class.java) @@ -43,7 +44,11 @@ class RetrospectiveCompletionCoordinator( metrics.recordStage("summary", "openai") { aiClient.generateSummaryWithTitle(snapshot.job, snapshot.answers, snapshot.deepQuestion) } - metrics.recordStage("summary", "save") { save(retrospectiveId, userId, summary) } + metrics.recordStage("summary", "save") { + summarySaveConcurrencyLimiter.execute { + save(retrospectiveId, userId, summary) + } + } summary } catch (exception: Exception) { reset(retrospectiveId, userId) diff --git a/src/main/kotlin/com/didit/application/retrospect/SummarySaveConcurrencyLimiter.kt b/src/main/kotlin/com/didit/application/retrospect/SummarySaveConcurrencyLimiter.kt new file mode 100644 index 00000000..a81921b0 --- /dev/null +++ b/src/main/kotlin/com/didit/application/retrospect/SummarySaveConcurrencyLimiter.kt @@ -0,0 +1,27 @@ +package com.didit.application.retrospect + +import org.springframework.beans.factory.annotation.Value +import org.springframework.stereotype.Component +import java.util.concurrent.Semaphore + +@Component +class SummarySaveConcurrencyLimiter( + @Value("\${retrospective.summary-save-max-concurrency:4}") maxConcurrency: Int, +) { + private val semaphore = + Semaphore( + requireNotNull(maxConcurrency.takeIf { it > 0 }) { + "retrospective.summary-save-max-concurrency must be greater than 0" + }, + true, + ) + + fun execute(action: () -> T): T { + semaphore.acquire() + try { + return action() + } finally { + semaphore.release() + } + } +} diff --git a/src/test/kotlin/com/didit/application/retrospect/RetrospectiveCompletionCoordinatorTest.kt b/src/test/kotlin/com/didit/application/retrospect/RetrospectiveCompletionCoordinatorTest.kt index 97354e39..c0757e5b 100644 --- a/src/test/kotlin/com/didit/application/retrospect/RetrospectiveCompletionCoordinatorTest.kt +++ b/src/test/kotlin/com/didit/application/retrospect/RetrospectiveCompletionCoordinatorTest.kt @@ -54,6 +54,7 @@ class RetrospectiveCompletionCoordinatorTest { aiClient, TransactionTemplate(transactionManager), metrics, + SummarySaveConcurrencyLimiter(4), ) } diff --git a/src/test/kotlin/com/didit/application/retrospect/SummarySaveConcurrencyLimiterTest.kt b/src/test/kotlin/com/didit/application/retrospect/SummarySaveConcurrencyLimiterTest.kt new file mode 100644 index 00000000..f0b4489d --- /dev/null +++ b/src/test/kotlin/com/didit/application/retrospect/SummarySaveConcurrencyLimiterTest.kt @@ -0,0 +1,75 @@ +package com.didit.application.retrospect + +import org.assertj.core.api.Assertions.assertThat +import org.junit.jupiter.api.Test +import org.junit.jupiter.api.assertThrows +import java.util.concurrent.CountDownLatch +import java.util.concurrent.Executors +import java.util.concurrent.TimeUnit +import java.util.concurrent.atomic.AtomicInteger + +class SummarySaveConcurrencyLimiterTest { + @Test + fun `execute - simultaneously runs no more than configured number of actions`() { + val limiter = SummarySaveConcurrencyLimiter(4) + val executor = Executors.newFixedThreadPool(10) + val ready = CountDownLatch(10) + val start = CountDownLatch(1) + val firstWaveEntered = CountDownLatch(4) + val release = CountDownLatch(1) + val active = AtomicInteger() + val maxActive = AtomicInteger() + val completed = AtomicInteger() + + try { + val futures = + (1..10).map { + executor.submit { + ready.countDown() + start.await() + limiter.execute { + val currentActive = active.incrementAndGet() + maxActive.accumulateAndGet(currentActive, ::maxOf) + firstWaveEntered.countDown() + try { + release.await() + completed.incrementAndGet() + } finally { + active.decrementAndGet() + } + } + } + } + + assertThat(ready.await(1, TimeUnit.SECONDS)).isTrue() + start.countDown() + assertThat(firstWaveEntered.await(1, TimeUnit.SECONDS)).isTrue() + assertThat(active.get()).isEqualTo(4) + + release.countDown() + futures.forEach { it.get(1, TimeUnit.SECONDS) } + + assertThat(maxActive.get()).isEqualTo(4) + assertThat(completed.get()).isEqualTo(10) + } finally { + release.countDown() + executor.shutdownNow() + } + } + + @Test + fun `execute - releases permit when action fails`() { + val limiter = SummarySaveConcurrencyLimiter(1) + + assertThrows { + limiter.execute { throw IllegalStateException("failed") } + } + + assertThat(limiter.execute { "completed" }).isEqualTo("completed") + } + + @Test + fun `constructor - rejects non-positive max concurrency`() { + assertThrows { SummarySaveConcurrencyLimiter(0) } + } +}