diff --git a/spring-session-data-redis/src/integration-test/java/org/springframework/session/data/redis/RedisOperationsSessionRepositoryITests.java b/spring-session-data-redis/src/integration-test/java/org/springframework/session/data/redis/RedisOperationsSessionRepositoryITests.java index c3e13e05..8e95658f 100644 --- a/spring-session-data-redis/src/integration-test/java/org/springframework/session/data/redis/RedisOperationsSessionRepositoryITests.java +++ b/spring-session-data-redis/src/integration-test/java/org/springframework/session/data/redis/RedisOperationsSessionRepositoryITests.java @@ -16,6 +16,7 @@ package org.springframework.session.data.redis; +import java.nio.charset.StandardCharsets; import java.util.Map; import java.util.UUID; @@ -190,9 +191,10 @@ public class RedisOperationsSessionRepositoryITests extends AbstractRedisITests String body = "RedisOperationsSessionRepositoryITests:sessions:expires:" + toSave.getId(); - String channel = ":expired"; - DefaultMessage message = new DefaultMessage(channel.getBytes("UTF-8"), - body.getBytes("UTF-8")); + String channel = "__keyevent@0__:expired"; + DefaultMessage message = new DefaultMessage( + channel.getBytes(StandardCharsets.UTF_8), + body.getBytes(StandardCharsets.UTF_8)); byte[] pattern = new byte[] {}; this.repository.onMessage(message, pattern); @@ -358,9 +360,10 @@ public class RedisOperationsSessionRepositoryITests extends AbstractRedisITests String body = "RedisOperationsSessionRepositoryITests:sessions:expires:" + toSave.getId(); - String channel = ":expired"; - DefaultMessage message = new DefaultMessage(channel.getBytes("UTF-8"), - body.getBytes("UTF-8")); + String channel = "__keyevent@0__:expired"; + DefaultMessage message = new DefaultMessage( + channel.getBytes(StandardCharsets.UTF_8), + body.getBytes(StandardCharsets.UTF_8)); byte[] pattern = new byte[] {}; this.repository.onMessage(message, pattern); diff --git a/spring-session-data-redis/src/main/java/org/springframework/session/data/redis/RedisOperationsSessionRepository.java b/spring-session-data-redis/src/main/java/org/springframework/session/data/redis/RedisOperationsSessionRepository.java index 63abfb32..28aecd1a 100644 --- a/spring-session-data-redis/src/main/java/org/springframework/session/data/redis/RedisOperationsSessionRepository.java +++ b/spring-session-data-redis/src/main/java/org/springframework/session/data/redis/RedisOperationsSessionRepository.java @@ -254,6 +254,11 @@ public class RedisOperationsSessionRepository implements static PrincipalNameResolver PRINCIPAL_NAME_RESOLVER = new PrincipalNameResolver(); + /** + * The default Redis database used by Spring Session. + */ + public static final int DEFAULT_DATABASE = 0; + /** * The default namespace for each key and channel in Redis used by Spring Session. */ @@ -286,11 +291,19 @@ public class RedisOperationsSessionRepository implements */ static final String SESSION_ATTR_PREFIX = "sessionAttr:"; + private int database = RedisOperationsSessionRepository.DEFAULT_DATABASE; + /** * The namespace for every key used by Spring Session in Redis. */ private String namespace = DEFAULT_NAMESPACE + ":"; + private String sessionCreatedChannelPrefix; + + private String sessionDeletedChannel; + + private String sessionExpiredChannel; + private final RedisOperations sessionRedisOperations; private final RedisSessionExpirationPolicy expirationPolicy; @@ -327,6 +340,7 @@ public class RedisOperationsSessionRepository implements this.sessionRedisOperations = sessionRedisOperations; this.expirationPolicy = new RedisSessionExpirationPolicy(sessionRedisOperations, this::getExpirationsKey, this::getSessionKey); + configureSessionChannels(); } /** @@ -377,6 +391,27 @@ public class RedisOperationsSessionRepository implements this.redisFlushMode = redisFlushMode; } + /** + * Sets the database index to use. Defaults to {@link #DEFAULT_DATABASE}. + * @param database the database index to use + */ + public void setDatabase(int database) { + this.database = database; + configureSessionChannels(); + } + + private void configureSessionChannels() { + this.sessionCreatedChannelPrefix = this.namespace + "event:" + this.database + + ":created:"; + this.sessionDeletedChannel = "__keyevent@" + this.database + "__:del"; + this.sessionExpiredChannel = "__keyevent@" + this.database + "__:expired"; + } + + /** + * Returns the {@link RedisOperations} used for sessions. + * @return the {@link RedisOperations} used for sessions + * @since 2.0.0 + */ public RedisOperations getSessionRedisOperations() { return this.sessionRedisOperations; } @@ -497,7 +532,7 @@ public class RedisOperationsSessionRepository implements String channel = new String(messageChannel); - if (channel.startsWith(getSessionCreatedChannelPrefix())) { + if (channel.startsWith(this.sessionCreatedChannelPrefix)) { // TODO: is this thread safe? Map loaded = (Map) this.defaultSerializer .deserialize(message.getBody()); @@ -510,8 +545,8 @@ public class RedisOperationsSessionRepository implements return; } - boolean isDeleted = channel.endsWith(":del"); - if (isDeleted || channel.endsWith(":expired")) { + boolean isDeleted = channel.equals(this.sessionDeletedChannel); + if (isDeleted || channel.equals(this.sessionExpiredChannel)) { int beginIndex = body.lastIndexOf(":") + 1; int endIndex = body.length(); String sessionId = body.substring(beginIndex, endIndex); @@ -574,6 +609,7 @@ public class RedisOperationsSessionRepository implements public void setRedisKeyNamespace(String namespace) { Assert.hasText(namespace, "namespace cannot be null or empty"); this.namespace = namespace.trim() + ":"; + configureSessionChannels(); } /** @@ -605,17 +641,33 @@ public class RedisOperationsSessionRepository implements } private String getExpiredKeyPrefix() { - return this.namespace + "sessions:" + "expires:"; + return this.namespace + "sessions:expires:"; } /** - * Gets the prefix for the channel that SessionCreatedEvent are published to. The - * suffix is the session id of the session that was created. - * - * @return the prefix for the channel that SessionCreatedEvent are published to + * Gets the prefix for the channel that {@link SessionCreatedEvent}s are published to. + * The suffix is the session id of the session that was created. + * @return the prefix for the channel that {@link SessionCreatedEvent}s are published + * to */ public String getSessionCreatedChannelPrefix() { - return this.namespace + "event:created:"; + return this.sessionCreatedChannelPrefix; + } + + /** + * Gets the name of the channel that {@link SessionDeletedEvent}s are published to. + * @return the name for the channel that {@link SessionDeletedEvent}s are published to + */ + public String getSessionDeletedChannel() { + return this.sessionDeletedChannel; + } + + /** + * Gets the name of the channel that {@link SessionExpiredEvent}s are published to. + * @return the name for the channel that {@link SessionExpiredEvent}s are published to + */ + public String getSessionExpiredChannel() { + return this.sessionExpiredChannel; } /** diff --git a/spring-session-data-redis/src/main/java/org/springframework/session/data/redis/config/annotation/web/http/RedisHttpSessionConfiguration.java b/spring-session-data-redis/src/main/java/org/springframework/session/data/redis/config/annotation/web/http/RedisHttpSessionConfiguration.java index 50888ff4..16c4ffb3 100644 --- a/spring-session-data-redis/src/main/java/org/springframework/session/data/redis/config/annotation/web/http/RedisHttpSessionConfiguration.java +++ b/spring-session-data-redis/src/main/java/org/springframework/session/data/redis/config/annotation/web/http/RedisHttpSessionConfiguration.java @@ -37,7 +37,10 @@ import org.springframework.core.annotation.AnnotationAttributes; import org.springframework.core.type.AnnotationMetadata; import org.springframework.data.redis.connection.RedisConnection; import org.springframework.data.redis.connection.RedisConnectionFactory; +import org.springframework.data.redis.connection.jedis.JedisConnectionFactory; +import org.springframework.data.redis.connection.lettuce.LettuceConnectionFactory; import org.springframework.data.redis.core.RedisTemplate; +import org.springframework.data.redis.listener.ChannelTopic; import org.springframework.data.redis.listener.PatternTopic; import org.springframework.data.redis.listener.RedisMessageListenerContainer; import org.springframework.data.redis.serializer.RedisSerializer; @@ -54,6 +57,7 @@ import org.springframework.session.data.redis.config.ConfigureRedisAction; import org.springframework.session.data.redis.config.annotation.SpringSessionRedisConnectionFactory; import org.springframework.session.web.http.SessionRepositoryFilter; import org.springframework.util.Assert; +import org.springframework.util.ClassUtils; import org.springframework.util.StringUtils; import org.springframework.util.StringValueResolver; @@ -115,6 +119,8 @@ public class RedisHttpSessionConfiguration extends SpringHttpSessionConfiguratio sessionRepository.setRedisKeyNamespace(this.redisNamespace); } sessionRepository.setRedisFlushMode(this.redisFlushMode); + int database = resolveDatabase(); + sessionRepository.setDatabase(database); return sessionRepository; } @@ -128,9 +134,9 @@ public class RedisHttpSessionConfiguration extends SpringHttpSessionConfiguratio if (this.redisSubscriptionExecutor != null) { container.setSubscriptionExecutor(this.redisSubscriptionExecutor); } - container.addMessageListener(sessionRepository(), - Arrays.asList(new PatternTopic("__keyevent@*:del"), - new PatternTopic("__keyevent@*:expired"))); + container.addMessageListener(sessionRepository(), Arrays.asList( + new ChannelTopic(sessionRepository().getSessionDeletedChannel()), + new ChannelTopic(sessionRepository().getSessionExpiredChannel()))); container.addMessageListener(sessionRepository(), Collections.singletonList(new PatternTopic( sessionRepository().getSessionCreatedChannelPrefix() + "*"))); @@ -256,6 +262,18 @@ public class RedisHttpSessionConfiguration extends SpringHttpSessionConfiguratio return redisTemplate; } + private int resolveDatabase() { + if (ClassUtils.isPresent("io.lettuce.core.RedisClient", null) + && this.redisConnectionFactory instanceof LettuceConnectionFactory) { + return ((LettuceConnectionFactory) this.redisConnectionFactory).getDatabase(); + } + if (ClassUtils.isPresent("redis.clients.jedis.Jedis", null) + && this.redisConnectionFactory instanceof JedisConnectionFactory) { + return ((JedisConnectionFactory) this.redisConnectionFactory).getDatabase(); + } + return RedisOperationsSessionRepository.DEFAULT_DATABASE; + } + /** * Ensures that Redis is configured to send keyspace notifications. This is important * to ensure that expiration and deletion of sessions trigger SessionDestroyedEvents. diff --git a/spring-session-data-redis/src/test/java/org/springframework/session/data/redis/RedisOperationsSessionRepositoryTests.java b/spring-session-data-redis/src/test/java/org/springframework/session/data/redis/RedisOperationsSessionRepositoryTests.java index 3730d497..0bde7976 100644 --- a/spring-session-data-redis/src/test/java/org/springframework/session/data/redis/RedisOperationsSessionRepositoryTests.java +++ b/spring-session-data-redis/src/test/java/org/springframework/session/data/redis/RedisOperationsSessionRepositoryTests.java @@ -16,6 +16,7 @@ package org.springframework.session.data.redis; +import java.nio.charset.StandardCharsets; import java.time.Duration; import java.time.Instant; import java.time.temporal.ChronoUnit; @@ -522,14 +523,15 @@ public class RedisOperationsSessionRepositoryTests { } @Test - public void onMessageCreated() throws Exception { + public void onMessageCreated() { MapSession session = this.cached; - byte[] pattern = "".getBytes("UTF-8"); - String channel = "spring:session:event:created:" + session.getId(); + byte[] pattern = "".getBytes(StandardCharsets.UTF_8); + String channel = "spring:session:event:0:created:" + session.getId(); JdkSerializationRedisSerializer defaultSerailizer = new JdkSerializationRedisSerializer(); this.redisRepository.setDefaultSerializer(defaultSerailizer); byte[] body = defaultSerailizer.serialize(new HashMap()); - DefaultMessage message = new DefaultMessage(channel.getBytes("UTF-8"), body); + DefaultMessage message = new DefaultMessage( + channel.getBytes(StandardCharsets.UTF_8), body); this.redisRepository.setApplicationEventPublisher(this.publisher); @@ -539,16 +541,16 @@ public class RedisOperationsSessionRepositoryTests { assertThat(this.event.getValue().getSessionId()).isEqualTo(session.getId()); } - // gh-309 - @Test - public void onMessageCreatedCustomSerializer() throws Exception { + @Test // gh-309 + public void onMessageCreatedCustomSerializer() { MapSession session = this.cached; - byte[] pattern = "".getBytes("UTF-8"); + byte[] pattern = "".getBytes(StandardCharsets.UTF_8); byte[] body = new byte[0]; - String channel = "spring:session:event:created:" + session.getId(); + String channel = "spring:session:event:0:created:" + session.getId(); given(this.defaultSerializer.deserialize(body)) .willReturn(new HashMap()); - DefaultMessage message = new DefaultMessage(channel.getBytes("UTF-8"), body); + DefaultMessage message = new DefaultMessage( + channel.getBytes(StandardCharsets.UTF_8), body); this.redisRepository.setApplicationEventPublisher(this.publisher); this.redisRepository.onMessage(message, pattern); @@ -559,7 +561,7 @@ public class RedisOperationsSessionRepositoryTests { } @Test - public void onMessageDeletedSessionFound() throws Exception { + public void onMessageDeletedSessionFound() { String deletedId = "deleted-id"; given(this.redisOperations.boundHashOps(getKey(deletedId))) .willReturn(this.boundHashOperations); @@ -570,10 +572,12 @@ public class RedisOperationsSessionRepositoryTests { String channel = "__keyevent@0__:del"; String body = "spring:session:sessions:expires:" + deletedId; - DefaultMessage message = new DefaultMessage(channel.getBytes("UTF-8"), body.getBytes("UTF-8")); + DefaultMessage message = new DefaultMessage( + channel.getBytes(StandardCharsets.UTF_8), + body.getBytes(StandardCharsets.UTF_8)); this.redisRepository.setApplicationEventPublisher(this.publisher); - this.redisRepository.onMessage(message, "".getBytes("UTF-8")); + this.redisRepository.onMessage(message, "".getBytes(StandardCharsets.UTF_8)); verify(this.redisOperations).boundHashOps(eq(getKey(deletedId))); verify(this.boundHashOperations).entries(); @@ -586,7 +590,7 @@ public class RedisOperationsSessionRepositoryTests { } @Test - public void onMessageDeletedSessionNotFound() throws Exception { + public void onMessageDeletedSessionNotFound() { String deletedId = "deleted-id"; given(this.redisOperations.boundHashOps(getKey(deletedId))) .willReturn(this.boundHashOperations); @@ -594,10 +598,12 @@ public class RedisOperationsSessionRepositoryTests { String channel = "__keyevent@0__:del"; String body = "spring:session:sessions:expires:" + deletedId; - DefaultMessage message = new DefaultMessage(channel.getBytes("UTF-8"), body.getBytes("UTF-8")); + DefaultMessage message = new DefaultMessage( + channel.getBytes(StandardCharsets.UTF_8), + body.getBytes(StandardCharsets.UTF_8)); this.redisRepository.setApplicationEventPublisher(this.publisher); - this.redisRepository.onMessage(message, "".getBytes("UTF-8")); + this.redisRepository.onMessage(message, "".getBytes(StandardCharsets.UTF_8)); verify(this.redisOperations).boundHashOps(eq(getKey(deletedId))); verify(this.boundHashOperations).entries(); @@ -608,7 +614,7 @@ public class RedisOperationsSessionRepositoryTests { } @Test - public void onMessageExpiredSessionFound() throws Exception { + public void onMessageExpiredSessionFound() { String expiredId = "expired-id"; given(this.redisOperations.boundHashOps(getKey(expiredId))) .willReturn(this.boundHashOperations); @@ -619,10 +625,12 @@ public class RedisOperationsSessionRepositoryTests { String channel = "__keyevent@0__:expired"; String body = "spring:session:sessions:expires:" + expiredId; - DefaultMessage message = new DefaultMessage(channel.getBytes("UTF-8"), body.getBytes("UTF-8")); + DefaultMessage message = new DefaultMessage( + channel.getBytes(StandardCharsets.UTF_8), + body.getBytes(StandardCharsets.UTF_8)); this.redisRepository.setApplicationEventPublisher(this.publisher); - this.redisRepository.onMessage(message, "".getBytes("UTF-8")); + this.redisRepository.onMessage(message, "".getBytes(StandardCharsets.UTF_8)); verify(this.redisOperations).boundHashOps(eq(getKey(expiredId))); verify(this.boundHashOperations).entries(); @@ -635,7 +643,7 @@ public class RedisOperationsSessionRepositoryTests { } @Test - public void onMessageExpiredSessionNotFound() throws Exception { + public void onMessageExpiredSessionNotFound() { String expiredId = "expired-id"; given(this.redisOperations.boundHashOps(getKey(expiredId))) .willReturn(this.boundHashOperations); @@ -643,10 +651,12 @@ public class RedisOperationsSessionRepositoryTests { String channel = "__keyevent@0__:expired"; String body = "spring:session:sessions:expires:" + expiredId; - DefaultMessage message = new DefaultMessage(channel.getBytes("UTF-8"), body.getBytes("UTF-8")); + DefaultMessage message = new DefaultMessage( + channel.getBytes(StandardCharsets.UTF_8), + body.getBytes(StandardCharsets.UTF_8)); this.redisRepository.setApplicationEventPublisher(this.publisher); - this.redisRepository.onMessage(message, "".getBytes("UTF-8")); + this.redisRepository.onMessage(message, "".getBytes(StandardCharsets.UTF_8)); verify(this.redisOperations).boundHashOps(eq(getKey(expiredId))); verify(this.boundHashOperations).entries(); @@ -881,6 +891,62 @@ public class RedisOperationsSessionRepositoryTests { assertThat(session.getAttributeNames()).isEmpty(); } + @Test + public void onMessageCreatedInOtherDatabase() { + JdkSerializationRedisSerializer serializer = new JdkSerializationRedisSerializer(); + this.redisRepository.setApplicationEventPublisher(this.publisher); + this.redisRepository.setDefaultSerializer(serializer); + + MapSession session = this.cached; + String channel = "spring:session:event:created:1:" + session.getId(); + byte[] body = serializer.serialize(new HashMap()); + DefaultMessage message = new DefaultMessage( + channel.getBytes(StandardCharsets.UTF_8), body); + + this.redisRepository.onMessage(message, "".getBytes(StandardCharsets.UTF_8)); + + assertThat(this.event.getAllValues()).isEmpty(); + verifyZeroInteractions(this.publisher); + } + + @Test + public void onMessageDeletedInOtherDatabase() { + JdkSerializationRedisSerializer serializer = new JdkSerializationRedisSerializer(); + this.redisRepository.setApplicationEventPublisher(this.publisher); + this.redisRepository.setDefaultSerializer(serializer); + + MapSession session = this.cached; + String channel = "__keyevent@1__:del"; + String body = "spring:session:sessions:expires:" + session.getId(); + DefaultMessage message = new DefaultMessage( + channel.getBytes(StandardCharsets.UTF_8), + body.getBytes(StandardCharsets.UTF_8)); + + this.redisRepository.onMessage(message, "".getBytes(StandardCharsets.UTF_8)); + + assertThat(this.event.getAllValues()).isEmpty(); + verifyZeroInteractions(this.publisher); + } + + @Test + public void onMessageExpiredInOtherDatabase() { + JdkSerializationRedisSerializer serializer = new JdkSerializationRedisSerializer(); + this.redisRepository.setApplicationEventPublisher(this.publisher); + this.redisRepository.setDefaultSerializer(serializer); + + MapSession session = this.cached; + String channel = "__keyevent@1__:expired"; + String body = "spring:session:sessions:expires:" + session.getId(); + DefaultMessage message = new DefaultMessage( + channel.getBytes(StandardCharsets.UTF_8), + body.getBytes(StandardCharsets.UTF_8)); + + this.redisRepository.onMessage(message, "".getBytes(StandardCharsets.UTF_8)); + + assertThat(this.event.getAllValues()).isEmpty(); + verifyZeroInteractions(this.publisher); + } + private String getKey(String id) { return "spring:session:sessions:" + id; }