From 7c6b143964ced894c600e9be29dcf49401195ae2 Mon Sep 17 00:00:00 2001 From: Vedran Pavic Date: Thu, 13 Sep 2018 18:23:21 +0200 Subject: [PATCH] Ensure RedisHttpSessionConfiguration handles events for configured database At present, RedisHttpSessionConfiguration doesn't take into account database index when handlng events. In situations where multiple apps use Spring Session with same Redis instance, but different database, this results in invalid session events. This commits improves event handling in RedisHttpSessionConfiguration to ensure currently used database is considered. Closes gh-1193 --- ...edisOperationsSessionRepositoryITests.java | 15 ++- .../RedisOperationsSessionRepository.java | 70 +++++++++-- .../http/RedisHttpSessionConfiguration.java | 24 +++- ...RedisOperationsSessionRepositoryTests.java | 110 ++++++++++++++---- 4 files changed, 179 insertions(+), 40 deletions(-) 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; }