diff --git a/core/retries/pom.xml b/core/retries/pom.xml
index aba7d96cd5d8..623bd4a00e6a 100644
--- a/core/retries/pom.xml
+++ b/core/retries/pom.xml
@@ -70,5 +70,10 @@
assertj-core
test
+
+ org.mockito
+ mockito-core
+ test
+
diff --git a/core/retries/src/main/java/software/amazon/awssdk/retries/internal/DefaultAdaptiveRetryStrategy.java b/core/retries/src/main/java/software/amazon/awssdk/retries/internal/DefaultAdaptiveRetryStrategy.java
index c52c7859526f..e63e683a18ab 100644
--- a/core/retries/src/main/java/software/amazon/awssdk/retries/internal/DefaultAdaptiveRetryStrategy.java
+++ b/core/retries/src/main/java/software/amazon/awssdk/retries/internal/DefaultAdaptiveRetryStrategy.java
@@ -43,15 +43,12 @@ public final class DefaultAdaptiveRetryStrategy
@Override
protected Duration computeInitialBackoff(AcquireInitialTokenRequest request) {
- RateLimiterTokenBucket bucket = rateLimiterTokenBucketStore.tokenBucketForScope(request.scope());
- return bucket.tryAcquire().delay();
+ throw new UnsupportedOperationException("TODO");
}
@Override
protected Duration computeBackoff(RefreshRetryTokenRequest request, DefaultRetryToken token) {
- Duration backoff = super.computeBackoff(request, token);
- RateLimiterTokenBucket bucket = rateLimiterTokenBucketStore.tokenBucketForScope(token.scope());
- return backoff.plus(bucket.tryAcquire().delay());
+ throw new UnsupportedOperationException("TODO");
}
@Override
diff --git a/core/retries/src/main/java/software/amazon/awssdk/retries/internal/ratelimiter/RateLimiterTokenBucket.java b/core/retries/src/main/java/software/amazon/awssdk/retries/internal/ratelimiter/RateLimiterTokenBucket.java
index 2d8743af9924..86a718f18906 100644
--- a/core/retries/src/main/java/software/amazon/awssdk/retries/internal/ratelimiter/RateLimiterTokenBucket.java
+++ b/core/retries/src/main/java/software/amazon/awssdk/retries/internal/ratelimiter/RateLimiterTokenBucket.java
@@ -16,42 +16,150 @@
package software.amazon.awssdk.retries.internal.ratelimiter;
import java.time.Duration;
-import java.util.concurrent.atomic.AtomicReference;
+import java.util.ArrayDeque;
+import java.util.Deque;
+import java.util.concurrent.CompletableFuture;
+import java.util.concurrent.ScheduledExecutorService;
+import java.util.concurrent.TimeUnit;
import java.util.function.Consumer;
import java.util.function.Function;
import software.amazon.awssdk.annotations.SdkInternalApi;
+import software.amazon.awssdk.annotations.SdkTestInternalApi;
+import software.amazon.awssdk.annotations.ThreadSafe;
+import software.amazon.awssdk.utils.CompletableFutureUtils;
+import software.amazon.awssdk.utils.SdkAutoCloseable;
/**
* The {@link RateLimiterTokenBucket} keeps track of past throttling responses and adapts to slow down the send rate to adapt to
- * the service. It does this by suggesting a delay amount as result of a {@link #tryAcquire()} call. Callers must update its
- * internal state by calling {@link #updateRateAfterThrottling()} when getting a throttling response or
- * {@link #updateRateAfterSuccess()} when getting successful response.
+ * the service. It does this by delaying the completion of the future returned by {@link #acquireAsync()} until the requested
+ * capacity is available. Callers must update its internal state by calling {@link #updateRateAfterThrottling()} when getting a
+ * throttling response or {@link #updateRateAfterSuccess()} when getting successful response.
*
- *
This class is thread-safe, its internal current state is kept in the inner class {@link PersistentState} which is stored
- * using an {@link AtomicReference}. This class is converted to {@link TransientState} when the state needs to be mutated and
- * converted back to a {@link PersistentState} and stored using {@link AtomicReference#compareAndSet(Object, Object)}.
+ *
This class is thread-safe.
*
*
The algorithm used is adapted from the network congestion avoidance algorithm
* CUBIC.
*/
@SdkInternalApi
-public class RateLimiterTokenBucket {
- private final AtomicReference stateReference;
+@ThreadSafe
+public class RateLimiterTokenBucket implements SdkAutoCloseable {
+ // Thread used for capacity waiting and notifying.
+ private final ScheduledExecutorService scheduler;
+
+ // Protect access to other members below
+ private final Object lock = new Object();
+
private final RateLimiterClock clock;
+ // the collection of futures returned to threads currently waiting for capacity.
+ // Futures are completed in FIFO order.
+ // The size of this equal to the number of threads concurrently accessing this bucket.
+ private final Deque> waiting = new ArrayDeque<>();
+ private PersistentState state;
+ private boolean open = true;
+ private boolean notifierRunning = false;
- RateLimiterTokenBucket(RateLimiterClock clock) {
+ RateLimiterTokenBucket(RateLimiterClock clock, ScheduledExecutorService scheduler) {
this.clock = clock;
- this.stateReference = new AtomicReference<>(new PersistentState());
+ this.scheduler = scheduler;
+ this.state = new PersistentState();
+ }
+
+ @Override
+ public void close() {
+ doClose(null);
+ }
+
+ private void doClose(Throwable cause) {
+ synchronized (lock) {
+ open = false;
+ IllegalStateException closedException = new IllegalStateException("Rate limiter bucket is closed", cause);
+ while (true) {
+ CompletableFuture w = waiting.poll();
+ if (w == null) {
+ break;
+ }
+ w.completeExceptionally(closedException);
+ }
+ }
}
/**
- * Acquire tokens from the bucket. If the bucket contains enough capacity to satisfy the request, this method will return in
- * {@link RateLimiterAcquireResponse#delay()} a {@link Duration#ZERO} value, otherwise it will return the amount of time the
- * callers need to wait until enough tokens are refilled.
+ * Acquire a token from the bucket.
+ *
+ * @return A future that is completed when the requested amount is acquired from this bucket.
*/
- public RateLimiterAcquireResponse tryAcquire() {
- StateUpdate update = updateState(ts -> ts.tokenBucketAcquire(clock, 1.0));
- return RateLimiterAcquireResponse.create(update.result);
+ public CompletableFuture acquireAsync() {
+ synchronized (lock) {
+ if (!open) {
+ return CompletableFutureUtils.failedFuture(new IllegalStateException("Rate limiter bucket is closed"));
+ }
+
+ // fast path and avoid scheduling in the executor if throttling isn't enabled.
+ if (!state.enabled) {
+ return CompletableFuture.completedFuture(null);
+ }
+
+ CompletableFuture future = new CompletableFuture<>();
+ waiting.add(future);
+ if (!notifierRunning) {
+ notifierRunning = scheduleOrClose(this::doNotify, Duration.ZERO);
+ }
+ return future;
+ }
+ }
+
+
+ @SdkTestInternalApi
+ Deque> waiting() {
+ return waiting;
+ }
+
+ @SdkTestInternalApi
+ boolean isClosed() {
+ return !open;
+ }
+
+ private void doNotify() {
+ while (true) {
+ CompletableFuture w;
+ synchronized (lock) {
+ w = waiting.poll();
+ if (w == null) {
+ notifierRunning = false;
+ return;
+ }
+
+ TransientState.AcquireResult acquireResult = updateState(t -> t.tokenBucketAcquire(clock, 1.0)).result;
+
+ // Not enough capacity. Try again later when enough time has
+ // passed to refill the bucket at the current rate.
+ if (!acquireResult.isSuccessful()) {
+ waiting.push(w);
+ notifierRunning = scheduleOrClose(this::doNotify, acquireResult.refillWait());
+ return;
+ }
+
+ }
+ // Acquire was successful, signal the waiting thread.
+ w.complete(null);
+ }
+ }
+
+ private void schedule(Runnable command, Duration d) {
+ scheduler.schedule(command, d.toMillis(), TimeUnit.MILLISECONDS);
+ }
+
+ /**
+ * @return true if schedule was successful, false otherwise. If the schedule failed, this bucket will be closed.
+ */
+ private boolean scheduleOrClose(Runnable command, Duration d) {
+ try {
+ schedule(command, d);
+ return true;
+ } catch (Throwable t) {
+ doClose(t);
+ }
+ return false;
}
/**
@@ -88,24 +196,15 @@ private StateUpdate consumeState(Consumer mutator) {
}
/**
- * Converts the stored persistent state into a transient one and transforms it using the provided function. The provided
- * function is expected to update the transient state in-place and return a value that will be returned to the caller in the
- * {@link StateUpdate#result} field. The mutated transient value is converted back to a persistent one and stored in the
- * atomic reference if no changes were made in-between. If another thread changes the value in-between, the operation is
- * retried until succeeded.
+ * Converts the stored persistent state into a transient one and transforms it using the provided function.
*/
private StateUpdate updateState(Function mutator) {
- PersistentState current;
- PersistentState updated;
- T result;
- do {
- current = stateReference.get();
- TransientState transientState = current.toTransient();
- result = mutator.apply(transientState);
- updated = transientState.toPersistent();
- } while (!stateReference.compareAndSet(current, updated));
-
- return new StateUpdate<>(updated, result);
+ synchronized (lock) {
+ TransientState transientState = state.toTransient();
+ T result = mutator.apply(transientState);
+ state = transientState.toPersistent();
+ return new StateUpdate<>(state, result);
+ }
}
static class StateUpdate {
@@ -163,17 +262,38 @@ PersistentState toPersistent() {
* a {@link Duration#ZERO} value, otherwise it will return the amount of time the callers need to wait until enough tokens
* are refilled.
*/
- Duration tokenBucketAcquire(RateLimiterClock clock, double amount) {
+ AcquireResult tokenBucketAcquire(RateLimiterClock clock, double amount) {
if (!this.enabled) {
- return Duration.ZERO;
+ return new AcquireResult(true, Duration.ZERO);
}
refill(clock);
- double waitTime = 0.0;
if (this.currentCapacity < amount) {
- waitTime = (amount - this.currentCapacity) / this.fillRate;
+ double diff = amount - currentCapacity;
+ double waitTime = diff / this.fillRate;
+ double waitTimeMs = waitTime * 1_000.0;
+ Duration duration = Duration.ofMillis((long) Math.ceil(waitTimeMs));
+ return new AcquireResult(false, duration);
}
this.currentCapacity -= amount;
- return Duration.ofNanos((long) (waitTime * 1_000_000_000.0));
+ return new AcquireResult(true, Duration.ZERO);
+ }
+
+ private static class AcquireResult {
+ final boolean successful;
+ final Duration refillWait;
+
+ AcquireResult(boolean successful, Duration refillWait) {
+ this.successful = successful;
+ this.refillWait = refillWait;
+ }
+
+ boolean isSuccessful() {
+ return successful;
+ }
+
+ Duration refillWait() {
+ return refillWait;
+ }
}
/**
diff --git a/core/retries/src/main/java/software/amazon/awssdk/retries/internal/ratelimiter/RateLimiterTokenBucketStore.java b/core/retries/src/main/java/software/amazon/awssdk/retries/internal/ratelimiter/RateLimiterTokenBucketStore.java
index 303e6d61d7f5..804649e5ee23 100644
--- a/core/retries/src/main/java/software/amazon/awssdk/retries/internal/ratelimiter/RateLimiterTokenBucketStore.java
+++ b/core/retries/src/main/java/software/amazon/awssdk/retries/internal/ratelimiter/RateLimiterTokenBucketStore.java
@@ -15,8 +15,10 @@
package software.amazon.awssdk.retries.internal.ratelimiter;
+import java.util.concurrent.ScheduledExecutorService;
import software.amazon.awssdk.annotations.SdkInternalApi;
import software.amazon.awssdk.annotations.ToBuilderIgnoreField;
+import software.amazon.awssdk.utils.SdkAutoCloseable;
import software.amazon.awssdk.utils.Validate;
import software.amazon.awssdk.utils.builder.CopyableBuilder;
import software.amazon.awssdk.utils.builder.ToCopyableBuilder;
@@ -27,19 +29,28 @@
*/
@SdkInternalApi
public final class RateLimiterTokenBucketStore
- implements ToCopyableBuilder {
+ implements ToCopyableBuilder, SdkAutoCloseable {
private static final int MAX_ENTRIES = 128;
private static final RateLimiterClock DEFAULT_CLOCK = new SystemClock();
private final LruCache scopeToTokenBucket;
private final RateLimiterClock clock;
+ private final ScheduledExecutorService scheduler;
private RateLimiterTokenBucketStore(Builder builder) {
this.clock = Validate.paramNotNull(builder.clock, "clock");
- this.scopeToTokenBucket = LruCache.builder(x -> new RateLimiterTokenBucket(clock))
+ this.scheduler = Validate.paramNotNull(builder.scheduler, "scheduler");
+ this.scopeToTokenBucket = LruCache.builder(
+ x -> new RateLimiterTokenBucket(clock, scheduler))
.maxSize(MAX_ENTRIES)
.build();
}
+ @Override
+ public void close() {
+ scopeToTokenBucket.evictAll();
+ scheduler.shutdownNow();
+ }
+
public RateLimiterTokenBucket tokenBucketForScope(String scope) {
return scopeToTokenBucket.get(scope);
}
@@ -56,6 +67,7 @@ public static RateLimiterTokenBucketStore.Builder builder() {
public static class Builder implements CopyableBuilder {
private RateLimiterClock clock;
+ private ScheduledExecutorService scheduler;
Builder() {
this.clock = DEFAULT_CLOCK;
@@ -63,6 +75,7 @@ public static class Builder implements CopyableBuilder buckets = new ArrayList<>(entries);
+
+ ScheduledExecutorService scheduler = mock(ScheduledExecutorService.class);
+ RateLimiterTokenBucketStore store = RateLimiterTokenBucketStore.builder()
+ .clock(new SystemClock())
+ .executor(scheduler)
+ .build();
+
+ List> futures = new ArrayList<>(entries);
+
+ for (int i = 0; i < entries; ++i) {
+ RateLimiterTokenBucket bucket = store.tokenBucketForScope(Integer.toString(i));
+ buckets.add(bucket);
+
+ // enable throttling so futures from acquireAsync get queued
+ bucket.updateRateAfterThrottling();
+ futures.add(bucket.acquireAsync());
+ }
+
+ store.close();
+
+ // New acquires from the closed bucket should fail
+ assertThat(buckets).allSatisfy(b -> assertThatThrownBy(b.acquireAsync()::join)
+ .hasMessageContaining("Rate limiter bucket is closed"));
+
+ // All pending futures should be failed
+ assertThat(futures).allSatisfy(cf -> assertThatThrownBy(cf::join)
+ .hasMessageContaining("Rate limiter bucket is closed"));
+ }
+}
diff --git a/core/retries/src/test/java/software/amazon/awssdk/retries/internal/ratelimiter/RateLimiterTokenBucketTest.java b/core/retries/src/test/java/software/amazon/awssdk/retries/internal/ratelimiter/RateLimiterTokenBucketTest.java
index 14021ebde812..806826a20d07 100644
--- a/core/retries/src/test/java/software/amazon/awssdk/retries/internal/ratelimiter/RateLimiterTokenBucketTest.java
+++ b/core/retries/src/test/java/software/amazon/awssdk/retries/internal/ratelimiter/RateLimiterTokenBucketTest.java
@@ -16,44 +16,206 @@
package software.amazon.awssdk.retries.internal.ratelimiter;
import static org.assertj.core.api.Assertions.assertThat;
+import static org.assertj.core.api.Assertions.assertThatThrownBy;
import static org.assertj.core.api.AssertionsForClassTypes.within;
+import static org.mockito.ArgumentMatchers.any;
+import static org.mockito.ArgumentMatchers.anyLong;
+import static org.mockito.ArgumentMatchers.eq;
+import static org.mockito.Mockito.doThrow;
+import static org.mockito.Mockito.mock;
+import static org.mockito.Mockito.verify;
+import static org.mockito.Mockito.verifyNoInteractions;
+import static org.mockito.Mockito.when;
import java.util.Arrays;
import java.util.Collection;
-import org.junit.jupiter.api.BeforeAll;
-import org.junit.jupiter.params.ParameterizedTest;
-import org.junit.jupiter.params.provider.MethodSource;
+import java.util.List;
+import java.util.concurrent.CompletableFuture;
+import java.util.concurrent.RejectedExecutionException;
+import java.util.concurrent.ScheduledExecutorService;
+import java.util.concurrent.TimeUnit;
+import java.util.stream.Collectors;
+import java.util.stream.IntStream;
+import org.junit.jupiter.api.BeforeEach;
+import org.junit.jupiter.api.Test;
class RateLimiterTokenBucketTest {
- private static MutableClock clock = null;
- private static RateLimiterTokenBucket tokenBucket = null;
private static final double EPSILON = 0.0001;
+ private MutableClock clock = null;
+ private ScheduledExecutorService scheduler = null;
+ private RateLimiterTokenBucket tokenBucket = null;
- @BeforeAll
- static void setup() {
+ @BeforeEach
+ void setup() {
clock = new MutableClock();
- tokenBucket = new RateLimiterTokenBucket(clock);
+ scheduler = mock(ScheduledExecutorService.class);
+ tokenBucket = new RateLimiterTokenBucket(clock, scheduler);
}
- @ParameterizedTest
- @MethodSource("parameters")
- void testCase(TestCase testCase) {
+ @Test
+ void acquireAsync_bucketClosed_futureCompletedExceptionally() {
+ tokenBucket.close();
+ CompletableFuture f = tokenBucket.acquireAsync();
+ assertThatThrownBy(f::join).satisfies(t -> {
+ Throwable cause = t.getCause();
+ assertThat(cause).isExactlyInstanceOf(IllegalStateException.class);
+ assertThat(cause).hasMessage("Rate limiter bucket is closed");
+ });
+ }
+
+ @Test
+ void acquireAsync_notEnabled_doesNotScheduleTask() {
+ CompletableFuture f = tokenBucket.acquireAsync();
+ assertThat(f).isCompleted();
+ verifyNoInteractions(scheduler);
+ }
+
+ @Test
+ void acquireAsync_enabled_schedulesTask() {
+ tokenBucket.updateRateAfterThrottling();
+
+ tokenBucket.acquireAsync();
+ verify(scheduler).schedule(any(Runnable.class), anyLong(), any(TimeUnit.class));
+ }
+
+
+ @Test
+ void acquireAsync_scheduleFails_completesFutureExceptionally() {
+ tokenBucket.updateRateAfterThrottling();
+
+ doThrow(new RejectedExecutionException("no")).when(scheduler).schedule(any(Runnable.class),
+ anyLong(),
+ any(TimeUnit.class));
+
+ CompletableFuture f = tokenBucket.acquireAsync();
+ assertThatThrownBy(f::join).satisfies(t -> {
+ Throwable cause = t.getCause();
+ assertThat(cause).hasMessage("Rate limiter bucket is closed");
+ });
+ }
+
+ @Test
+ void acquireAsync_scheduleFails_futureNotInWaitingDeque() {
+ tokenBucket.updateRateAfterThrottling();
+
+ doThrow(new RejectedExecutionException("no")).when(scheduler).schedule(any(Runnable.class),
+ anyLong(),
+ any(TimeUnit.class));
+
+ CompletableFuture f = tokenBucket.acquireAsync();
+ assertThat(f).isCompletedExceptionally();
+ assertThat(tokenBucket.waiting()).isEmpty();
+ }
+
+ @Test
+ void acquireAsync_scheduleFails_closesBucket() {
+ tokenBucket.updateRateAfterThrottling();
+
+ doThrow(new RejectedExecutionException("no")).when(scheduler).schedule(any(Runnable.class),
+ anyLong(),
+ any(TimeUnit.class));
+
+ CompletableFuture f = tokenBucket.acquireAsync();
+ assertThat(f).isCompletedExceptionally();
+ assertThat(tokenBucket.isClosed()).isTrue();
+ }
+
+ @Test
+ void close_completesAllPendingFutures() {
+ // enable throttling so futures actually get queued instead of being completed immediately
+ tokenBucket.updateRateAfterThrottling();
+
+ List> futures = IntStream.range(0, 10)
+ .mapToObj(i -> tokenBucket.acquireAsync())
+ .collect(Collectors.toList());
+
+ tokenBucket.close();
+
+ assertThat(futures).allSatisfy(f -> {
+ assertThatThrownBy(f::join).satisfies(t -> {
+ Throwable cause = t.getCause();
+ assertThat(cause).isExactlyInstanceOf(IllegalStateException.class);
+ assertThat(cause).hasMessage("Rate limiter bucket is closed");
+ });
+ });
+
+ assertThat(tokenBucket.waiting()).isEmpty();
+ }
+
+ @Test
+ void close_doesShutDownExecutor() {
+ tokenBucket.close();
+ verifyNoInteractions(scheduler);
+ }
+
+ @Test
+ void doNotify_scheduleRejected_failsFuture() {
+ tokenBucket.updateRateAfterThrottling();
+
+ // Empty bucket at default rate of 0.5 tokens per second should be 2seconds
+ when(scheduler.schedule(any(Runnable.class), eq(2000L), eq(TimeUnit.MILLISECONDS)))
+ .thenThrow(new RejectedExecutionException("no"));
+
+ // 0L is the initial schedule from acquireAsync, capture the doNotify schedule and execute that.
+ when(scheduler.schedule(any(Runnable.class), eq(0L), any(TimeUnit.class))).thenAnswer(i -> {
+ Runnable r = i.getArgument(0);
+ r.run();
+ return null;
+ });
+
+ tokenBucket.acquireAsync();
+ assertThat(tokenBucket.isClosed()).isTrue();
+ }
+
+ @Test
+ void doNotify_scheduleRejected_closesBucket() {
+ tokenBucket.updateRateAfterThrottling();
+
+ // Empty bucket at default rate of 0.5 tokens per second should be 2seconds
+ when(scheduler.schedule(any(Runnable.class), eq(2000L), eq(TimeUnit.MILLISECONDS)))
+ .thenThrow(new RejectedExecutionException("no"));
+
+ // 0L is the initial schedule from acquireAsync, capture the doNotify schedule and execute that.
+ when(scheduler.schedule(any(Runnable.class), eq(0L), any(TimeUnit.class))).thenAnswer(i -> {
+ Runnable r = i.getArgument(0);
+ r.run();
+ return null;
+ });
+
+ CompletableFuture future = tokenBucket.acquireAsync();
+ assertThatThrownBy(future::join).hasRootCauseInstanceOf(RejectedExecutionException.class);
+ assertThat(tokenBucket.isClosed()).isTrue();
+ }
+
+ @Test
+ void sendingRateEndToEndTest() {
+ for (TestCase sendingRateTestCase : sendingRateTestCases()) {
+ assertSendingRateTestCase(sendingRateTestCase);
+ }
+ }
+
+ void assertSendingRateTestCase(TestCase testCase) {
clock.setCurrent(testCase.givenTimestamp);
RateLimiterUpdateResponse res;
- tokenBucket.tryAcquire();
+
if (testCase.throttleResponse) {
res = tokenBucket.updateRateAfterThrottling();
} else {
res = tokenBucket.updateRateAfterSuccess();
}
double measuredTxRate = res.measuredTxRate();
- assertThat(measuredTxRate).isCloseTo(testCase.expectMeasuredTxRate, within(EPSILON));
+ assertThat(measuredTxRate)
+ .as("%s: Measured TX rate", testCase)
+ .isCloseTo(testCase.expectMeasuredTxRate, within(EPSILON));
double fillRate = res.fillRate();
- assertThat(fillRate).isCloseTo(testCase.expectFillRate, within(EPSILON));
+ assertThat(fillRate)
+ .as("%s: Fill rate", testCase)
+ .isCloseTo(testCase.expectFillRate, within(EPSILON));
}
-
- static Collection parameters() {
+ static Collection sendingRateTestCases() {
+ // Note: Test cases are not independent. Each case depends on the state of the bucket being correctly updated from the
+ // previous test.
return Arrays.asList(
new TestCase()
.givenSuccessResponse()
@@ -174,6 +336,15 @@ TestCase expectFillRate(double expectFillRate) {
return this;
}
+ @Override
+ public String toString() {
+ return "TestCase{" +
+ "throttleResponse=" + throttleResponse +
+ ", givenTimestamp=" + givenTimestamp +
+ ", expectMeasuredTxRate=" + expectMeasuredTxRate +
+ ", expectFillRate=" + expectFillRate +
+ '}';
+ }
}
static class MutableClock implements RateLimiterClock {
diff --git a/utils/src/main/java/software/amazon/awssdk/utils/cache/lru/LruCache.java b/utils/src/main/java/software/amazon/awssdk/utils/cache/lru/LruCache.java
index df7bc222d261..eff5892b01ab 100644
--- a/utils/src/main/java/software/amazon/awssdk/utils/cache/lru/LruCache.java
+++ b/utils/src/main/java/software/amazon/awssdk/utils/cache/lru/LruCache.java
@@ -15,6 +15,8 @@
package software.amazon.awssdk.utils.cache.lru;
+import java.util.ArrayList;
+import java.util.List;
import java.util.Map;
import java.util.Objects;
import java.util.concurrent.ConcurrentHashMap;
@@ -77,6 +79,20 @@ public V get(K key) {
}
}
+ public List evictAll() {
+ List evicted = new ArrayList<>(cache.size());
+ synchronized (listLock) {
+ while (true) {
+ CacheEntry evictedEntry = evict();
+ if (evictedEntry == null) {
+ break;
+ }
+ evicted.add(evictedEntry.value);
+ }
+ return evicted;
+ }
+ }
+
private CacheEntry newEntry(K key) {
V value = valueSupplier.apply(key);
return new CacheEntry<>(key, value);
@@ -147,11 +163,18 @@ private void addToQueue(CacheEntry entry) {
/**
* Removes the least recently used entry from the cache, marks it as evicted and removes it from the queue.
*/
- private void evict() {
- leastRecentlyUsed.isEvicted(true);
- closeEvictedResourcesIfPossible(leastRecentlyUsed.value);
- cache.remove(leastRecentlyUsed.key());
- removeFromQueue(leastRecentlyUsed);
+ private CacheEntry evict() {
+ if (leastRecentlyUsed == null) {
+ return null;
+ }
+
+ CacheEntry entryToEvict = leastRecentlyUsed;
+
+ entryToEvict.isEvicted(true);
+ closeEvictedResourcesIfPossible(entryToEvict.value);
+ cache.remove(entryToEvict.key());
+ removeFromQueue(entryToEvict);
+ return entryToEvict;
}
private void closeEvictedResourcesIfPossible(V value) {
diff --git a/utils/src/test/java/software/amazon/awssdk/utils/cache/lru/LruCacheTest.java b/utils/src/test/java/software/amazon/awssdk/utils/cache/lru/LruCacheTest.java
index 2ee389b8849a..12b8dc3ef470 100644
--- a/utils/src/test/java/software/amazon/awssdk/utils/cache/lru/LruCacheTest.java
+++ b/utils/src/test/java/software/amazon/awssdk/utils/cache/lru/LruCacheTest.java
@@ -16,13 +16,16 @@
package software.amazon.awssdk.utils.cache.lru;
import static org.assertj.core.api.Assertions.assertThat;
+import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.times;
import static org.mockito.Mockito.verify;
import static software.amazon.awssdk.utils.FunctionalUtils.invokeSafely;
import java.util.ArrayList;
import java.util.Collections;
+import java.util.HashSet;
import java.util.List;
+import java.util.Set;
import java.util.concurrent.ExecutorService;
import java.util.concurrent.Executors;
import java.util.concurrent.Future;
@@ -40,6 +43,7 @@
import org.junit.jupiter.params.provider.MethodSource;
import org.mockito.Spy;
import org.mockito.junit.jupiter.MockitoExtension;
+import software.amazon.awssdk.utils.SdkAutoCloseable;
@ExtendWith(MockitoExtension.class)
public class LruCacheTest {
@@ -246,6 +250,39 @@ void when_multipleThreadsAreCallingCache_WorksAsExpected(Integer numThreads,
}
}
+ @Test
+ void evictAll_noEntries_returnsEmptyList() {
+ LruCache cache = simpleCache.get();
+ assertThat(cache.evictAll()).isEmpty();
+ }
+
+ @Test
+ void evictAll_maxEntries_returnsAllEntries() {
+ LruCache cache = simpleCache.get();
+ List expected = new ArrayList<>();
+ for (int i = 0; i < MAX_SIMPLE_CACHE_SIZE; ++i) {
+ expected.add(cache.get(i));
+ }
+ assertThat(cache.evictAll()).containsExactlyInAnyOrder(expected.toArray(new String[0]));
+ }
+
+ @Test
+ void evictAll_closesEvictedEntries() {
+ LruCache cache = LruCache.
+ builder(k -> mock(SdkAutoCloseable.class))
+ .maxSize(MAX_SIMPLE_CACHE_SIZE)
+ .build();
+
+ for (int i = 0; i < MAX_SIMPLE_CACHE_SIZE; ++i) {
+ cache.get(i);
+ }
+
+ List evicted = cache.evictAll();
+
+ assertThat(evicted).isNotEmpty();
+ assertThat(evicted).allSatisfy(e -> verify(e).close());
+ }
+
private static Stream concurrencyTestValues() {
// numThreads, numGetsPerThreads, sleepDurationMillis, cacheSize
return Stream.of(Arguments.of(1000, 5000, false, 5),