Polishing.

Refine unlocking by checking whether the lock was actually applied.

Reduce allocations, refine test assertions to check for concurrency.

See #1686
Original pull request: #2879
This commit is contained in:
Mark Paluch
2024-04-19 08:57:23 +02:00
parent 3cf7cbfa0c
commit 0948d0dd67
2 changed files with 61 additions and 24 deletions

View File

@@ -32,6 +32,7 @@ import org.springframework.data.redis.connection.ReactiveRedisConnectionFactory;
import org.springframework.data.redis.connection.ReactiveStringCommands; import org.springframework.data.redis.connection.ReactiveStringCommands;
import org.springframework.data.redis.connection.RedisConnection; import org.springframework.data.redis.connection.RedisConnection;
import org.springframework.data.redis.connection.RedisConnectionFactory; import org.springframework.data.redis.connection.RedisConnectionFactory;
import org.springframework.data.redis.connection.RedisStringCommands;
import org.springframework.data.redis.connection.RedisStringCommands.SetOption; import org.springframework.data.redis.connection.RedisStringCommands.SetOption;
import org.springframework.data.redis.core.types.Expiration; import org.springframework.data.redis.core.types.Expiration;
import org.springframework.data.redis.util.ByteUtils; import org.springframework.data.redis.util.ByteUtils;
@@ -219,8 +220,10 @@ class DefaultRedisCacheWriter implements RedisCacheWriter {
return execute(name, connection -> { return execute(name, connection -> {
boolean wasLocked = false;
if (isLockingCacheWriter()) { if (isLockingCacheWriter()) {
doLock(name, key, value, connection); doLock(name, key, value, connection);
wasLocked = true;
} }
try { try {
@@ -242,7 +245,7 @@ class DefaultRedisCacheWriter implements RedisCacheWriter {
return connection.stringCommands().get(key); return connection.stringCommands().get(key);
} finally { } finally {
if (isLockingCacheWriter()) { if (isLockingCacheWriter() && wasLocked) {
doUnlock(name, connection); doUnlock(name, connection);
} }
} }
@@ -319,15 +322,17 @@ class DefaultRedisCacheWriter implements RedisCacheWriter {
execute(name, connection -> doLock(name, name, null, connection)); execute(name, connection -> doLock(name, name, null, connection));
} }
@Nullable boolean doLock(String name, Object contextualKey, @Nullable Object contextualValue, RedisConnection connection) {
protected Boolean doLock(String name, Object contextualKey, @Nullable Object contextualValue,
RedisConnection connection) {
RedisStringCommands commands = connection.stringCommands();
Expiration expiration = Expiration.from(this.lockTtl.getTimeToLive(contextualKey, contextualValue)); Expiration expiration = Expiration.from(this.lockTtl.getTimeToLive(contextualKey, contextualValue));
byte[] cacheLockKey = createCacheLockKey(name);
while (!ObjectUtils.nullSafeEquals(connection.stringCommands().set(createCacheLockKey(name), new byte[0], expiration, SetOption.SET_IF_ABSENT),true)) { while (!ObjectUtils.nullSafeEquals(commands.set(cacheLockKey, new byte[0], expiration, SetOption.SET_IF_ABSENT),
true)) {
checkAndPotentiallyWaitUntilUnlocked(name, connection); checkAndPotentiallyWaitUntilUnlocked(name, connection);
} }
return true; return true;
} }
@@ -341,7 +346,7 @@ class DefaultRedisCacheWriter implements RedisCacheWriter {
} }
@Nullable @Nullable
private Long doUnlock(String name, RedisConnection connection) { Long doUnlock(String name, RedisConnection connection) {
return connection.keyCommands().del(createCacheLockKey(name)); return connection.keyCommands().del(createCacheLockKey(name));
} }
@@ -489,8 +494,7 @@ class DefaultRedisCacheWriter implements RedisCacheWriter {
Mono<?> cacheLockCheck = isLockingCacheWriter() ? waitForLock(connection, name) : Mono.empty(); Mono<?> cacheLockCheck = isLockingCacheWriter() ? waitForLock(connection, name) : Mono.empty();
ReactiveStringCommands stringCommands = connection.stringCommands(); ReactiveStringCommands stringCommands = connection.stringCommands();
Mono<ByteBuffer> get = shouldExpireWithin(ttl) Mono<ByteBuffer> get = shouldExpireWithin(ttl) ? stringCommands.getEx(wrappedKey, Expiration.from(ttl))
? stringCommands.getEx(wrappedKey, Expiration.from(ttl))
: stringCommands.get(wrappedKey); : stringCommands.get(wrappedKey);
return cacheLockCheck.then(get).map(ByteUtils::getBytes).toFuture(); return cacheLockCheck.then(get).map(ByteUtils::getBytes).toFuture();
@@ -502,8 +506,7 @@ class DefaultRedisCacheWriter implements RedisCacheWriter {
return doWithConnection(connection -> { return doWithConnection(connection -> {
Mono<?> mono = isLockingCacheWriter() Mono<?> mono = isLockingCacheWriter() ? doStoreWithLocking(name, key, value, ttl, connection)
? doStoreWithLocking(name, key, value, ttl, connection)
: doStore(key, value, ttl, connection); : doStore(key, value, ttl, connection);
return mono.then().toFuture(); return mono.then().toFuture();
@@ -531,7 +534,6 @@ class DefaultRedisCacheWriter implements RedisCacheWriter {
} }
} }
private Mono<Object> doLock(String name, Object contextualKey, @Nullable Object contextualValue, private Mono<Object> doLock(String name, Object contextualKey, @Nullable Object contextualValue,
ReactiveRedisConnection connection) { ReactiveRedisConnection connection) {

View File

@@ -21,14 +21,19 @@ import static org.springframework.data.redis.cache.RedisCacheWriter.*;
import java.nio.charset.Charset; import java.nio.charset.Charset;
import java.nio.charset.StandardCharsets; import java.nio.charset.StandardCharsets;
import java.time.Duration; import java.time.Duration;
import java.util.ArrayList;
import java.util.Collection; import java.util.Collection;
import java.util.List;
import java.util.concurrent.CompletableFuture;
import java.util.concurrent.CountDownLatch; import java.util.concurrent.CountDownLatch;
import java.util.concurrent.ExecutionException; import java.util.concurrent.ExecutionException;
import java.util.concurrent.TimeUnit; import java.util.concurrent.TimeUnit;
import java.util.concurrent.atomic.AtomicLong;
import java.util.concurrent.atomic.AtomicReference; import java.util.concurrent.atomic.AtomicReference;
import java.util.function.Consumer; import java.util.function.Consumer;
import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.BeforeEach;
import org.springframework.data.redis.connection.RedisConnection; import org.springframework.data.redis.connection.RedisConnection;
import org.springframework.data.redis.connection.RedisConnectionFactory; import org.springframework.data.redis.connection.RedisConnectionFactory;
import org.springframework.data.redis.connection.RedisStringCommands.SetOption; import org.springframework.data.redis.connection.RedisStringCommands.SetOption;
@@ -421,43 +426,73 @@ public class DefaultRedisCacheWriterTests {
assertThat(stats.getPuts()).isZero(); assertThat(stats.getPuts()).isZero();
} }
@ParameterizedRedisTest @ParameterizedRedisTest // GH-1686
void doLockShouldGetLock() throws InterruptedException { void doLockShouldGetLock() throws InterruptedException {
int threadCount = 3; int threadCount = 3;
CountDownLatch beforeWrite = new CountDownLatch(threadCount); CountDownLatch beforeWrite = new CountDownLatch(threadCount);
CountDownLatch afterWrite = new CountDownLatch(threadCount); CountDownLatch afterWrite = new CountDownLatch(threadCount);
AtomicLong concurrency = new AtomicLong();
DefaultRedisCacheWriter cw = new DefaultRedisCacheWriter(connectionFactory, Duration.ofMillis(50), DefaultRedisCacheWriter cw = new DefaultRedisCacheWriter(connectionFactory, Duration.ofMillis(10),
BatchStrategies.keys()){ BatchStrategies.keys()) {
@Nullable
protected Boolean doLock(String name, Object contextualKey, @Nullable Object contextualValue, boolean doLock(String name, Object contextualKey, @Nullable Object contextualValue, RedisConnection connection) {
RedisConnection connection) {
Boolean doLock = super.doLock(name, contextualKey, contextualValue, connection); boolean doLock = super.doLock(name, contextualKey, contextualValue, connection);
assertThat(doLock).isTrue();
// any concurrent access (aka not waiting until the lock is acquired) will result in a concurrency greater 1
assertThat(concurrency.incrementAndGet()).isOne();
return doLock; return doLock;
} }
@Nullable
@Override
Long doUnlock(String name, RedisConnection connection) {
try {
return super.doUnlock(name, connection);
} finally {
concurrency.decrementAndGet();
}
}
}; };
cw.lock(CACHE_NAME); cw.lock(CACHE_NAME);
// introduce concurrency
List<CompletableFuture<?>> completions = new ArrayList<>();
for (int i = 0; i < threadCount; i++) { for (int i = 0; i < threadCount; i++) {
CompletableFuture<?> completion = new CompletableFuture<>();
completions.add(completion);
Thread th = new Thread(() -> { Thread th = new Thread(() -> {
beforeWrite.countDown(); beforeWrite.countDown();
cw.putIfAbsent(CACHE_NAME, binaryCacheKey, binaryCacheValue, Duration.ZERO); try {
cw.putIfAbsent(CACHE_NAME, binaryCacheKey, binaryCacheValue, Duration.ZERO);
completion.complete(null);
} catch (Throwable e) {
completion.completeExceptionally(e);
}
afterWrite.countDown(); afterWrite.countDown();
}); });
th.start(); th.start();
} }
beforeWrite.await(); assertThat(beforeWrite.await(5, TimeUnit.SECONDS)).isTrue();
Thread.sleep(100);
Thread.sleep(200);
cw.unlock(CACHE_NAME); cw.unlock(CACHE_NAME);
afterWrite.await(); assertThat(afterWrite.await(5, TimeUnit.SECONDS)).isTrue();
for (CompletableFuture<?> completion : completions) {
assertThat(completion).isCompleted().isCompletedWithValue(null);
}
doWithConnection(conn -> {
assertThat(conn.exists("default-redis-cache-writer-tests~lock".getBytes())).isFalse();
});
} }
private void doWithConnection(Consumer<RedisConnection> callback) { private void doWithConnection(Consumer<RedisConnection> callback) {