diff options
| author | Andy McCurdy <andy@andymccurdy.com> | 2014-03-28 10:06:45 -0700 |
|---|---|---|
| committer | Andy McCurdy <andy@andymccurdy.com> | 2014-03-28 10:06:45 -0700 |
| commit | 169cc0241a4f8824f35add9a872c067b70b1e95c (patch) | |
| tree | b896d9524d00cae8d728b8916f82870bddc127f6 | |
| parent | 8a2f3698576b8a67a8b3629173517c56d744e150 (diff) | |
| download | redis-py-169cc0241a4f8824f35add9a872c067b70b1e95c.tar.gz | |
fixes PubSub.subscribe once and for all.
| -rw-r--r-- | redis/client.py | 97 | ||||
| -rw-r--r-- | tests/conftest.py | 21 | ||||
| -rw-r--r-- | tests/test_pubsub.py | 152 |
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) |
