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),