summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorAndy McCurdy <andy@andymccurdy.com>2014-03-28 10:06:45 -0700
committerAndy McCurdy <andy@andymccurdy.com>2014-03-28 10:06:45 -0700
commit169cc0241a4f8824f35add9a872c067b70b1e95c (patch)
treeb896d9524d00cae8d728b8916f82870bddc127f6
parent8a2f3698576b8a67a8b3629173517c56d744e150 (diff)
downloadredis-py-169cc0241a4f8824f35add9a872c067b70b1e95c.tar.gz
fixes PubSub.subscribe once and for all.
-rw-r--r--redis/client.py97
-rw-r--r--tests/conftest.py21
-rw-r--r--tests/test_pubsub.py152
3 files changed, 170 insertions, 100 deletions
diff --git a/redis/client.py b/redis/client.py
index dd31e3f..7837dd3 100644
--- a/redis/client.py
+++ b/redis/client.py
@@ -1753,7 +1753,8 @@ class PubSub(object):
until a message arrives on one of the subscribed channels. That message
will be returned and it's safe to start listening again.
"""
- MESSAGE_TYPES = ('message', 'pmessage')
+ PUBLISH_MESSAGE_TYPES = ('message', 'pmessage')
+ UNSUBSCRIBE_MESSAGE_TYPES = ('unsubscribe', 'punsubscribe')
def __init__(self, connection_pool, shard_hint=None,
ignore_subscribe_messages=False):
@@ -1761,12 +1762,7 @@ class PubSub(object):
self.shard_hint = shard_hint
self.ignore_subscribe_messages = ignore_subscribe_messages
self.connection = None
- self.channels = {}
- self.patterns = {}
- self.subscription_count = 0
- self.subscribe_commands = set(
- ('subscribe', 'psubscribe', 'unsubscribe', 'punsubscribe')
- )
+ self.reset()
def __del__(self):
try:
@@ -1783,6 +1779,8 @@ class PubSub(object):
self.connection.clear_connect_callbacks()
self.connection_pool.release(self.connection)
self.connection = None
+ self.channels = {}
+ self.patterns = {}
def close(self):
self.reset()
@@ -1796,7 +1794,7 @@ class PubSub(object):
@property
def subscribed(self):
"Indicates if there are subscriptions to any channels or patterns"
- return bool(self.subscription_count or self.channels or self.patterns)
+ return bool(self.channels or self.patterns)
def execute_command(self, *args, **kwargs):
"Execute a publish/subscribe command"
@@ -1811,7 +1809,8 @@ class PubSub(object):
self.shard_hint
)
# initially connect here so we don't run our callback the first
- # time. It's primarily there for reconnection purposes.
+ # time. If we did, it would dupe the subscriptions, once from the
+ # callback and a second time from the actual command invocation
self.connection.connect()
self.connection.register_connect_callback(self.on_connect)
connection = self.connection
@@ -1825,8 +1824,9 @@ class PubSub(object):
# Connect manually here. If the Redis server is down, this will
# fail and raise a ConnectionError as desired.
connection.connect()
- # resubscribe to all channels and patterns before
- # resending the current command
+ # the ``on_connect`` callback should haven been called by the
+ # connection to resubscribe us to any channels and patterns we were
+ # previously listening to
return command(*args)
def parse_response(self, block=True):
@@ -1834,10 +1834,7 @@ class PubSub(object):
connection = self.connection
if not block and not connection.can_read():
return None
- response = self._execute(connection, connection.read_response)
- if nativestr(response[0]) in self.subscribe_commands:
- self.subscription_count = response[2]
- return response
+ return self._execute(connection, connection.read_response)
def psubscribe(self, *args, **kwargs):
"""
@@ -1862,13 +1859,6 @@ class PubSub(object):
"""
if args:
args = list_or_args(args[0], args[1:])
- for pattern in args:
- try:
- del self.patterns[pattern]
- except KeyError:
- pass
- if not args:
- self.patterns = {}
return self.execute_command('PUNSUBSCRIBE', *args)
def subscribe(self, *args, **kwargs):
@@ -1876,7 +1866,8 @@ class PubSub(object):
Subscribe to channels. Channels supplied as keyword arguments expect
a channel name as the key and a callable as the value. A channel's
callable will be invoked automatically when a message is received on
- that channel rather than producing a message via ``listen()``.
+ that channel rather than producing a message via ``listen()`` or
+ ``get_message()``.
"""
if args:
args = list_or_args(args[0], args[1:])
@@ -1893,13 +1884,6 @@ class PubSub(object):
"""
if args:
args = list_or_args(args[0], args[1:])
- for channel in args:
- try:
- del self.channels[channel]
- except KeyError:
- pass
- if not args:
- self.channels = {}
return self.execute_command('UNSUBSCRIBE', *args)
def listen(self):
@@ -1922,36 +1906,51 @@ class PubSub(object):
with a message handler, the handler is invoked instead of a parsed
message being returned.
"""
- msg_type = nativestr(response[0])
- handler = None
-
- # optionally ignore subscribe/unsubscribe messages
- s = ignore_subscribe_messages or self.ignore_subscribe_messages
- if s and msg_type not in self.MESSAGE_TYPES:
- return None
-
- if msg_type == 'pmessage':
- msg = {
- 'type': msg_type,
+ message_type = nativestr(response[0])
+ if message_type == 'pmessage':
+ message = {
+ 'type': message_type,
'pattern': nativestr(response[1]),
'channel': nativestr(response[2]),
'data': response[3]
}
- handler = self.patterns.get(msg['pattern'], None)
else:
- msg = {
- 'type': msg_type,
+ message = {
+ 'type': message_type,
'pattern': None,
'channel': nativestr(response[1]),
'data': response[2]
}
- handler = self.channels.get(msg['channel'], None)
- if handler:
- handler(msg)
- return None
+ if message_type in self.PUBLISH_MESSAGE_TYPES:
+ # if there's a message handler, invoke it
+ handler = None
+ if message_type == 'pmessage':
+ handler = self.patterns.get(message['pattern'], None)
+ else:
+ handler = self.channels.get(message['channel'], None)
+ if handler:
+ handler(message)
+ return None
else:
- return msg
+ # this is a subscribe/unsubscribe message. ignore if we don't
+ # want them
+ if ignore_subscribe_messages or self.ignore_subscribe_messages:
+ return None
+
+ # if this is an unsubscribe message, remove it from memory
+ if message_type in self.UNSUBSCRIBE_MESSAGE_TYPES:
+ subscribed_dict = None
+ if message_type == 'punsubscribe':
+ subscribed_dict = self.patterns
+ else:
+ subscribed_dict = self.channels
+ try:
+ del subscribed_dict[message['channel']]
+ except KeyError:
+ pass
+
+ return message
def run_in_thread(self, sleep_time=0):
for channel, handler in iteritems(self.channels):
diff --git a/tests/conftest.py b/tests/conftest.py
index 231081f..553838d 100644
--- a/tests/conftest.py
+++ b/tests/conftest.py
@@ -2,18 +2,35 @@ import pytest
import redis
+_REDIS_VERSIONS = {}
+
+
+def get_version(**kwargs):
+ params = {'host': 'localhost', 'port': 6379, 'db': 9}
+ params.update(kwargs)
+ key = '%s:%s' % (params['host'], params['port'])
+ if key not in _REDIS_VERSIONS:
+ client = redis.Redis(**params)
+ _REDIS_VERSIONS[key] = client.info()['redis_version']
+ client.connection_pool.disconnect()
+ return _REDIS_VERSIONS[key]
+
+
def _get_client(cls, request=None, **kwargs):
params = {'host': 'localhost', 'port': 6379, 'db': 9}
params.update(kwargs)
client = cls(**params)
client.flushdb()
if request:
- request.addfinalizer(client.flushdb)
+ def teardown():
+ client.flushdb()
+ client.connection_pool.disconnect()
+ request.addfinalizer(teardown)
return client
def skip_if_server_version_lt(min_version):
- version = _get_client(redis.Redis).info()['redis_version']
+ version = get_version()
c = "StrictVersion('%s') < StrictVersion('%s')" % (version, min_version)
return pytest.mark.skipif(c)
diff --git a/tests/test_pubsub.py b/tests/test_pubsub.py
index 52825af..9d47b62 100644
--- a/tests/test_pubsub.py
+++ b/tests/test_pubsub.py
@@ -28,75 +28,129 @@ def make_message(type, channel, data, pattern=None):
}
-class TestPubSubSubscribeUnsubscribe(object):
-
- def test_subscribe_unsubscribe(self, r):
- p = r.pubsub()
-
- assert p.subscribe('foo', 'bar') is None
-
- # should be 2 messages indicating that we've subscribed
- assert wait_for_message(p) == make_message('subscribe', 'foo', 1)
- assert wait_for_message(p) == make_message('subscribe', 'bar', 2)
-
- assert p.unsubscribe('foo', 'bar') is None
+def make_subscribe_test_data(pubsub, type):
+ if type == 'channel':
+ return {
+ 'p': pubsub,
+ 'sub_type': 'subscribe',
+ 'unsub_type': 'unsubscribe',
+ 'sub_func': pubsub.subscribe,
+ 'unsub_func': pubsub.unsubscribe,
+ 'keys': ['foo', 'bar']
+ }
+ elif type == 'pattern':
+ return {
+ 'p': pubsub,
+ 'sub_type': 'psubscribe',
+ 'unsub_type': 'punsubscribe',
+ 'sub_func': pubsub.psubscribe,
+ 'unsub_func': pubsub.punsubscribe,
+ 'keys': ['f*', 'b*']
+ }
+ assert False, 'invalid subscribe type: %s' % type
- # should be 2 messages indicating that we've unsubscribed
- assert wait_for_message(p) == make_message('unsubscribe', 'foo', 1)
- assert wait_for_message(p) == make_message('unsubscribe', 'bar', 0)
- def test_pattern_subscribe_unsubscribe(self, r):
- p = r.pubsub()
+class TestPubSubSubscribeUnsubscribe(object):
- assert p.psubscribe('f*', 'b*') is None
+ def _test_subscribe_unsubscribe(self, p, sub_type, unsub_type, sub_func,
+ unsub_func, keys):
+ assert sub_func(*keys) is None
# should be 2 messages indicating that we've subscribed
- assert wait_for_message(p) == make_message('psubscribe', 'f*', 1)
- assert wait_for_message(p) == make_message('psubscribe', 'b*', 2)
+ assert wait_for_message(p) == make_message(sub_type, keys[0], 1)
+ assert wait_for_message(p) == make_message(sub_type, keys[1], 2)
- assert p.punsubscribe('f*', 'b*') is None
+ assert unsub_func(*keys) is None
# should be 2 messages indicating that we've unsubscribed
- assert wait_for_message(p) == make_message('punsubscribe', 'f*', 1)
- assert wait_for_message(p) == make_message('punsubscribe', 'b*', 0)
+ assert wait_for_message(p) == make_message(unsub_type, keys[0], 1)
+ assert wait_for_message(p) == make_message(unsub_type, keys[1], 0)
- def test_resubscribe_to_channels_on_reconnection(self, r):
- channels = ['foo', 'bar']
- p = r.pubsub()
+ def test_channel_subscribe_unsubscribe(self, r):
+ kwargs = make_subscribe_test_data(r.pubsub(), 'channel')
+ self._test_subscribe_unsubscribe(**kwargs)
- assert p.subscribe(*channels) is None
+ def test_pattern_subscribe_unsubscribe(self, r):
+ kwargs = make_subscribe_test_data(r.pubsub(), 'pattern')
+ self._test_subscribe_unsubscribe(**kwargs)
- for i, channel in enumerate(channels):
+ def _test_resubscribe_on_reconnection(self, p, sub_type, unsub_type,
+ sub_func, unsub_func, keys):
+ assert sub_func(*keys) is None
+
+ for i, key in enumerate(keys):
i += 1 # enumerate is 0 index, but we want 1 based indexing
- assert wait_for_message(p) == make_message('subscribe', channel, i)
+ assert wait_for_message(p) == make_message(sub_type, key, i)
# manually disconnect
p.connection.disconnect()
# calling get_message again reconnects and resubscribes
- for i, channel in enumerate(channels):
- i += 1 # enumerate is 0 index, but we want 1 based indexing
- assert wait_for_message(p) == make_message('subscribe', channel, i)
-
- def test_resubscribe_to_patterns_on_reconnection(self, r):
- patterns = ['f*', 'b*']
- p = r.pubsub()
-
- assert p.psubscribe(*patterns) is None
-
- for i, pattern in enumerate(patterns):
+ for i, key in enumerate(keys):
i += 1 # enumerate is 0 index, but we want 1 based indexing
- assert wait_for_message(p) == make_message(
- 'psubscribe', pattern, i)
+ assert wait_for_message(p) == make_message(sub_type, key, i)
- # manually disconnect
- p.connection.disconnect()
+ def test_resubscribe_to_channels_on_reconnection(self, r):
+ kwargs = make_subscribe_test_data(r.pubsub(), 'channel')
+ self._test_resubscribe_on_reconnection(**kwargs)
- # calling get_message again reconnects and resubscribes
- for i, pattern in enumerate(patterns):
- i += 1 # enumerate is 0 index, but we want 1 based indexing
- assert wait_for_message(p) == make_message(
- 'psubscribe', pattern, i)
+ def test_resubscribe_to_patterns_on_reconnection(self, r):
+ kwargs = make_subscribe_test_data(r.pubsub(), 'pattern')
+ self._test_resubscribe_on_reconnection(**kwargs)
+
+ def _test_subscribed_property(self, p, sub_type, unsub_type, sub_func,
+ unsub_func, keys):
+
+ assert p.subscribed is False
+ sub_func(keys[0])
+ # we're now subscribed even though we haven't processed the
+ # reply from the server just yet
+ assert p.subscribed is True
+ assert wait_for_message(p) == make_message(sub_type, keys[0], 1)
+ # we're still subscribed
+ assert p.subscribed is True
+
+ # unsubscribe from all channels
+ unsub_func()
+ # we're still technically subscribed until we process the
+ # response messages from the server
+ assert p.subscribed is True
+ assert wait_for_message(p) == make_message(unsub_type, keys[0], 0)
+ # now we're no longer subscribed as no more messages can be delivered
+ # to any channels we were listening to
+ assert p.subscribed is False
+
+ # subscribing again flips the flag back
+ sub_func(keys[0])
+ assert p.subscribed is True
+ assert wait_for_message(p) == make_message(sub_type, keys[0], 1)
+
+ # unsubscribe again
+ unsub_func()
+ assert p.subscribed is True
+ # subscribe to another channel before reading the unsubscribe response
+ sub_func(keys[1])
+ assert p.subscribed is True
+ # read the unsubscribe for key1
+ assert wait_for_message(p) == make_message(unsub_type, keys[0], 0)
+ # we're still subscribed to key2, so subscribed should still be True
+ assert p.subscribed is True
+ # read the key2 subscribe message
+ assert wait_for_message(p) == make_message(sub_type, keys[1], 1)
+ unsub_func()
+ # haven't read the message yet, so we're still subscribed
+ assert p.subscribed is True
+ assert wait_for_message(p) == make_message(unsub_type, keys[1], 0)
+ # now we're finally unsubscribed
+ assert p.subscribed is False
+
+ def test_subscribe_property_with_channels(self, r):
+ kwargs = make_subscribe_test_data(r.pubsub(), 'channel')
+ self._test_subscribed_property(**kwargs)
+
+ def test_subscribe_property_with_patterns(self, r):
+ kwargs = make_subscribe_test_data(r.pubsub(), 'pattern')
+ self._test_subscribed_property(**kwargs)
def test_ignore_all_subscribe_messages(self, r):
p = r.pubsub(ignore_subscribe_messages=True)