From 0086dd830b693ca9a1277c93836c96210295face Mon Sep 17 00:00:00 2001 From: Rossen Stoyanchev Date: Tue, 13 Oct 2015 12:28:27 -0400 Subject: [PATCH] Enforce cacheLimit in DefaultSubscriptionRegistry When the cacheLimit is reached and there is an eviction from the updateCache, the accessCache is now also updated. This change also ensures that adding a destination to the cache is protected with synchronization on the updateCache. Issue: SPR-13555 --- .../broker/AbstractSubscriptionRegistry.java | 8 ++- .../broker/DefaultSubscriptionRegistry.java | 55 ++++++++++--------- .../DefaultSubscriptionRegistryTests.java | 38 +++++++++---- 3 files changed, 62 insertions(+), 39 deletions(-) diff --git a/spring-messaging/src/main/java/org/springframework/messaging/simp/broker/AbstractSubscriptionRegistry.java b/spring-messaging/src/main/java/org/springframework/messaging/simp/broker/AbstractSubscriptionRegistry.java index 3a9a04ece1..cbaa802131 100644 --- a/spring-messaging/src/main/java/org/springframework/messaging/simp/broker/AbstractSubscriptionRegistry.java +++ b/spring-messaging/src/main/java/org/springframework/messaging/simp/broker/AbstractSubscriptionRegistry.java @@ -23,6 +23,8 @@ import org.springframework.messaging.Message; import org.springframework.messaging.MessageHeaders; import org.springframework.messaging.simp.SimpMessageHeaderAccessor; import org.springframework.messaging.simp.SimpMessageType; +import org.springframework.util.CollectionUtils; +import org.springframework.util.LinkedMultiValueMap; import org.springframework.util.MultiValueMap; /** @@ -35,6 +37,10 @@ import org.springframework.util.MultiValueMap; */ public abstract class AbstractSubscriptionRegistry implements SubscriptionRegistry { + private static MultiValueMap EMPTY_MAP = + CollectionUtils.unmodifiableMultiValueMap(new LinkedMultiValueMap(0)); + + protected final Log logger = LogFactory.getLog(getClass()); @@ -104,7 +110,7 @@ public abstract class AbstractSubscriptionRegistry implements SubscriptionRegist String destination = SimpMessageHeaderAccessor.getDestination(headers); if (destination == null) { logger.error("No destination in " + message); - return null; + return EMPTY_MAP; } return findSubscriptionsInternal(destination, message); diff --git a/spring-messaging/src/main/java/org/springframework/messaging/simp/broker/DefaultSubscriptionRegistry.java b/spring-messaging/src/main/java/org/springframework/messaging/simp/broker/DefaultSubscriptionRegistry.java index 38035c602d..a80a92b7de 100644 --- a/spring-messaging/src/main/java/org/springframework/messaging/simp/broker/DefaultSubscriptionRegistry.java +++ b/spring-messaging/src/main/java/org/springframework/messaging/simp/broker/DefaultSubscriptionRegistry.java @@ -115,24 +115,7 @@ public class DefaultSubscriptionRegistry extends AbstractSubscriptionRegistry { @Override protected MultiValueMap findSubscriptionsInternal(String destination, Message message) { - MultiValueMap result = this.destinationCache.getSubscriptions(destination); - if (result != null) { - return result; - } - result = new LinkedMultiValueMap(); - for (SessionSubscriptionInfo info : this.subscriptionRegistry.getAllSubscriptions()) { - for (String destinationPattern : info.getDestinations()) { - if (this.pathMatcher.match(destinationPattern, destination)) { - for (String subscriptionId : info.getSubscriptions(destinationPattern)) { - result.add(info.sessionId, subscriptionId); - } - } - } - } - if (!result.isEmpty()) { - this.destinationCache.addSubscriptions(destination, result); - } - return result; + return this.destinationCache.getSubscriptions(destination, message); } @Override @@ -157,20 +140,38 @@ public class DefaultSubscriptionRegistry extends AbstractSubscriptionRegistry { new LinkedHashMap>(DEFAULT_CACHE_LIMIT, 0.75f, true) { @Override protected boolean removeEldestEntry(Map.Entry> eldest) { - return size() > getCacheLimit(); + if (size() > getCacheLimit()) { + accessCache.remove(eldest.getKey()); + return true; + } + else { + return false; + } } }; - public MultiValueMap getSubscriptions(String destination) { - return this.accessCache.get(destination); - } - - public void addSubscriptions(String destination, MultiValueMap subscriptions) { - synchronized (this.updateCache) { - this.updateCache.put(destination, deepCopy(subscriptions)); - this.accessCache.put(destination, subscriptions); + public MultiValueMap getSubscriptions(String destination, Message message) { + MultiValueMap result = this.accessCache.get(destination); + if (result == null) { + synchronized (this.updateCache) { + result = new LinkedMultiValueMap(); + for (SessionSubscriptionInfo info : subscriptionRegistry.getAllSubscriptions()) { + for (String destinationPattern : info.getDestinations()) { + if (getPathMatcher().match(destinationPattern, destination)) { + for (String subscriptionId : info.getSubscriptions(destinationPattern)) { + result.add(info.sessionId, subscriptionId); + } + } + } + } + if (!result.isEmpty()) { + this.updateCache.put(destination, deepCopy(result)); + this.accessCache.put(destination, result); + } + } } + return result; } public void updateAfterNewSubscription(String destination, String sessionId, String subsId) { diff --git a/spring-messaging/src/test/java/org/springframework/messaging/simp/broker/DefaultSubscriptionRegistryTests.java b/spring-messaging/src/test/java/org/springframework/messaging/simp/broker/DefaultSubscriptionRegistryTests.java index 130835e342..9c6cd90a04 100644 --- a/spring-messaging/src/test/java/org/springframework/messaging/simp/broker/DefaultSubscriptionRegistryTests.java +++ b/spring-messaging/src/test/java/org/springframework/messaging/simp/broker/DefaultSubscriptionRegistryTests.java @@ -69,7 +69,7 @@ public class DefaultSubscriptionRegistryTests { MultiValueMap actual = this.registry.findSubscriptions(message(dest)); assertEquals("Expected one element " + actual, 1, actual.size()); - assertEquals(Arrays.asList(subsId), actual.get(sessId)); + assertEquals(Collections.singletonList(subsId), actual.get(sessId)); } @Test @@ -116,7 +116,7 @@ public class DefaultSubscriptionRegistryTests { MultiValueMap actual = this.registry.findSubscriptions(message(dest)); assertEquals("Expected one element " + actual, 1, actual.size()); - assertEquals(Arrays.asList(subsId), actual.get(sessId)); + assertEquals(Collections.singletonList(subsId), actual.get(sessId)); } @Test // SPR-11657 @@ -143,13 +143,13 @@ public class DefaultSubscriptionRegistryTests { actual = this.registry.findSubscriptions(message("/topic/PRICE.STOCK.NASDAQ.IBM")); assertEquals(2, actual.size()); assertEquals(Arrays.asList(subs2, subs1), actual.get(sess1)); - assertEquals(Arrays.asList(subs1), actual.get(sess2)); + assertEquals(Collections.singletonList(subs1), actual.get(sess2)); this.registry.unregisterAllSubscriptions(sess1); actual = this.registry.findSubscriptions(message("/topic/PRICE.STOCK.NASDAQ.IBM")); assertEquals(1, actual.size()); - assertEquals(Arrays.asList(subs1), actual.get(sess2)); + assertEquals(Collections.singletonList(subs1), actual.get(sess2)); this.registry.registerSubscription(subscribeMessage(sess1, subs1, "/topic/PRICE.STOCK.*.IBM")); this.registry.registerSubscription(subscribeMessage(sess1, subs2, "/topic/PRICE.STOCK.NASDAQ.IBM")); @@ -157,20 +157,20 @@ public class DefaultSubscriptionRegistryTests { actual = this.registry.findSubscriptions(message("/topic/PRICE.STOCK.NASDAQ.IBM")); assertEquals(2, actual.size()); assertEquals(Arrays.asList(subs1, subs2), actual.get(sess1)); - assertEquals(Arrays.asList(subs1), actual.get(sess2)); + assertEquals(Collections.singletonList(subs1), actual.get(sess2)); this.registry.unregisterSubscription(unsubscribeMessage(sess1, subs2)); actual = this.registry.findSubscriptions(message("/topic/PRICE.STOCK.NASDAQ.IBM")); assertEquals(2, actual.size()); - assertEquals(Arrays.asList(subs1), actual.get(sess1)); - assertEquals(Arrays.asList(subs1), actual.get(sess2)); + assertEquals(Collections.singletonList(subs1), actual.get(sess1)); + assertEquals(Collections.singletonList(subs1), actual.get(sess2)); this.registry.unregisterSubscription(unsubscribeMessage(sess1, subs1)); actual = this.registry.findSubscriptions(message("/topic/PRICE.STOCK.NASDAQ.IBM")); assertEquals(1, actual.size()); - assertEquals(Arrays.asList(subs1), actual.get(sess2)); + assertEquals(Collections.singletonList(subs1), actual.get(sess2)); this.registry.unregisterSubscription(unsubscribeMessage(sess2, subs1)); @@ -222,13 +222,13 @@ public class DefaultSubscriptionRegistryTests { MultiValueMap actual = this.registry.findSubscriptions(message); assertEquals("Expected one element " + actual, 1, actual.size()); - assertEquals(Arrays.asList(subsId), actual.get(sessId)); + assertEquals(Collections.singletonList(subsId), actual.get(sessId)); message = message("/topic/PRICE.STOCK.NASDAQ.MSFT"); actual = this.registry.findSubscriptions(message); assertEquals("Expected one element " + actual, 1, actual.size()); - assertEquals(Arrays.asList(subsId), actual.get(sessId)); + assertEquals(Collections.singletonList(subsId), actual.get(sessId)); message = message("/topic/PRICE.STOCK.NASDAQ.VMW"); actual = this.registry.findSubscriptions(message); @@ -249,7 +249,7 @@ public class DefaultSubscriptionRegistryTests { actual = this.registry.findSubscriptions(message("/foo")); assertEquals("Expected 1 element", 1, actual.size()); - assertEquals(Arrays.asList("subs02"), actual.get("sess01")); + assertEquals(Collections.singletonList("subs02"), actual.get("sess01")); this.registry.unregisterSubscription(unsubscribeMessage("sess01", "subs02")); @@ -345,6 +345,22 @@ public class DefaultSubscriptionRegistryTests { // no ConcurrentModificationException } + @Test // SPR-13555 + public void cacheLimitExceeded() throws Exception { + this.registry.setCacheLimit(1); + this.registry.registerSubscription(subscribeMessage("sess1", "1", "/foo")); + this.registry.registerSubscription(subscribeMessage("sess1", "2", "/bar")); + + assertEquals(1, this.registry.findSubscriptions(message("/foo")).size()); + assertEquals(1, this.registry.findSubscriptions(message("/bar")).size()); + + this.registry.registerSubscription(subscribeMessage("sess2", "1", "/foo")); + this.registry.registerSubscription(subscribeMessage("sess2", "2", "/bar")); + + assertEquals(2, this.registry.findSubscriptions(message("/foo")).size()); + assertEquals(2, this.registry.findSubscriptions(message("/bar")).size()); + } + private Message subscribeMessage(String sessionId, String subscriptionId, String destination) { SimpMessageHeaderAccessor headers = SimpMessageHeaderAccessor.create(SimpMessageType.SUBSCRIBE);