summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorKonstantin Merenkov <kmerenkov@yandex-team.ru>2010-11-02 16:06:27 +0300
committerKonstantin Merenkov <kmerenkov@yandex-team.ru>2010-11-02 16:06:27 +0300
commit36c73d49a6255b1c3f86a257d6bbc97ba898c77d (patch)
tree3982987c05cbd952d720797a4b1f28c72f0c83fd
parentab623e03ca84d22313af06744a1f887ed436b9f0 (diff)
parent6eeee751a90779d65d45165279dc21dae399ac7a (diff)
downloadredis-py-36c73d49a6255b1c3f86a257d6bbc97ba898c77d.tar.gz
Merge branch 'master' of http://github.com/andymccurdy/redis-py
-rw-r--r--README.md463
-rw-r--r--redis/__init__.py2
-rw-r--r--redis/client.py264
-rw-r--r--redis/exceptions.py14
-rwxr-xr-xrun_tests9
-rw-r--r--setup.py2
-rw-r--r--tests/__init__.py2
-rw-r--r--tests/lock.py53
-rw-r--r--tests/server_commands.py184
9 files changed, 927 insertions, 66 deletions
diff --git a/README.md b/README.md
index 5a086e6..d1971ba 100644
--- a/README.md
+++ b/README.md
@@ -17,6 +17,469 @@ Usage
For a complete list of commands, check out the list of Redis commands here:
http://code.google.com/p/redis/wiki/CommandReference
+Installation
+------------
+
+ $ sudo easy-install redis
+
+alternatively:
+
+ $ sudo pip install redis
+
+From sources:
+
+ $ sudo python setup.py install
+
+Versioning scheme
+-----------------
+
+redis-py is versioned after Redis. So, for example, redis-py 2.0.0 should
+support all the commands available in Redis 2.0.0.
+
+API Reference
+-------------
+
+### append(self, key, value)
+ Appends the string _value_ to the value at _key_. If _key_
+ doesn't already exist, create it with a value of _value_.
+ Returns the new length of the value at _key_.
+
+### bgrewriteaof(self)
+ Tell the Redis server to rewrite the AOF file from data in memory.
+
+### bgsave(self)
+ Tell the Redis server to save its data to disk. Unlike save(),
+ this method is asynchronous and returns immediately.
+
+### blpop(self, keys, timeout=0)
+ LPOP a value off of the first non-empty list
+ named in the _keys_ list.
+
+ If none of the lists in _keys_ has a value to LPOP, then block
+ for _timeout_ seconds, or until a value gets pushed on to one
+ of the lists.
+
+ If timeout is 0, then block indefinitely.
+
+### brpop(self, keys, timeout=0)
+ RPOP a value off of the first non-empty list
+ named in the _keys_ list.
+
+ If none of the lists in _keys_ has a value to LPOP, then block
+ for _timeout_ seconds, or until a value gets pushed on to one
+ of the lists.
+
+ If timeout is 0, then block indefinitely.
+
+### dbsize(self)
+ Returns the number of keys in the current database
+
+### decr(self, name, amount=1)
+ Decrements the value of _key_ by _amount_. If no key exists,
+ the value will be initialized as 0 - _amount_
+
+### delete(self, *names)
+ Delete one or more keys specified by _names_
+
+### encode(self, value)
+ Encode _value_ using the instance's charset
+
+### execute_command(self, *args, **options)
+ Sends the command to the redis server and returns it's response
+
+### exists(self, name)
+ Returns a boolean indicating whether key _name_ exists
+
+### expire(self, name, time)
+ Set an expire flag on key _name_ for _time_ seconds
+
+### expireat(self, name, when)
+ Set an expire flag on key _name_. _when_ can be represented
+ as an integer indicating unix time or a Python datetime object.
+
+### flush(self, all_dbs=False)
+
+### flushall(self)
+ Delete all keys in all databases on the current host
+
+### flushdb(self)
+ Delete all keys in the current database
+
+### get(self, name)
+ Return the value at key _name_, or None of the key doesn't exist
+
+### get_connection(self, host, port, db, password, socket_timeout)
+ Returns a connection object
+
+### getset(self, name, value)
+ Set the value at key _name_ to _value_ if key doesn't exist
+ Return the value at key _name_ atomically
+
+### hdel(self, name, key)
+ Delete _key_ from hash _name_
+
+### hexists(self, name, key)
+ Returns a boolean indicating if _key_ exists within hash _name_
+
+### hget(self, name, key)
+ Return the value of _key_ within the hash _name_
+
+### hgetall(self, name)
+ Return a Python dict of the hash's name/value pairs
+
+### hincrby(self, name, key, amount=1)
+ Increment the value of _key_ in hash _name_ by _amount_
+
+### hkeys(self, name)
+ Return the list of keys within hash _name_
+
+### hlen(self, name)
+ Return the number of elements in hash _name_
+
+### hmget(self, name, keys)
+ Returns a list of values ordered identically to _keys_
+
+### hmset(self, name, mapping)
+ Sets each key in the _mapping_ dict to its corresponding value
+ in the hash _name_
+
+### hset(self, name, key, value)
+ Set _key_ to _value_ within hash _name_
+ Returns 1 if HSET created a new field, otherwise 0
+
+### hsetnx(self, name, key, value)
+ Set _key_ to _value_ within hash _name_ if _key_ does not
+ exist. Returns 1 if HSETNX created a field, otherwise 0.
+
+### hvals(self, name)
+ Return the list of values within hash _name_
+
+### incr(self, name, amount=1)
+ Increments the value of _key_ by _amount_. If no key exists,
+ the value will be initialized as _amount_
+
+### info(self)
+ Returns a dictionary containing information about the Redis server
+
+### keys(self, pattern='*')
+ Returns a list of keys matching _pattern_
+
+### lastsave(self)
+ Return a Python datetime object representing the last time the
+ Redis database was saved to disk
+
+### lindex(self, name, index)
+ Return the item from list _name_ at position _index_
+
+ Negative indexes are supported and will return an item at the
+ end of the list
+
+### listen(self)
+ Listen for messages on channels this client has been subscribed to
+
+### llen(self, name)
+ Return the length of the list _name_
+
+###lock(self, name, timeout=None, sleep=0.10000000000000001)
+ Return a new Lock object using key _name_ that mimics
+ the behavior of threading.Lock.
+
+ If specified, _timeout_ indicates a maximum life for the lock.
+ By default, it will remain locked until release() is called.
+
+ _sleep_ indicates the amount of time to sleep per loop iteration
+ when the lock is in blocking mode and another client is currently
+ holding the lock.
+
+### lpop(self, name)
+ Remove and return the first item of the list _name_
+
+### lpush(self, name, value)
+ Push _value_ onto the head of the list _name_
+
+### lrange(self, name, start, end)
+ Return a slice of the list _name_ between
+ position _start_ and _end_
+
+ _start_ and _end_ can be negative numbers just like
+ Python slicing notation
+
+### lrem(self, name, value, num=0)
+ Remove the first _num_ occurrences of _value_ from list _name_
+
+ If _num_ is 0, then all occurrences will be removed
+
+### lset(self, name, index, value)
+ Set _position_ of list _name_ to _value_
+
+### ltrim(self, name, start, end)
+ Trim the list _name_, removing all values not within the slice
+ between _start_ and _end_
+
+ _start_ and _end_ can be negative numbers just like
+ Python slicing notation
+
+### mget(self, keys, *args)
+ Returns a list of values ordered identically to _keys_
+
+ * Passing *args to this method has been deprecated *
+
+### move(self, name, db)
+ Moves the key _name_ to a different Redis database _db_
+
+### mset(self, mapping)
+ Sets each key in the _mapping_ dict to its corresponding value
+
+### msetnx(self, mapping)
+ Sets each key in the _mapping_ dict to its corresponding value if
+ none of the keys are already set
+
+### parse_response(self, command_name, catch_errors=False, **options)
+ Parses a response from the Redis server
+
+### ping(self)
+ Ping the Redis server
+
+### pipeline(self, transaction=True)
+ Return a new pipeline object that can queue multiple commands for
+ later execution. _transaction_ indicates whether all commands
+ should be executed atomically. Apart from multiple atomic operations,
+ pipelines are useful for batch loading of data as they reduce the
+ number of back and forth network operations between client and server.
+
+### pop(self, name, tail=False)
+ Pop and return the first or last element of list _name_
+
+ This method has been deprecated, use _Redis.lpop_ or _Redis.rpop_ instead.
+
+### psubscribe(self, patterns)
+ Subscribe to all channels matching any pattern in _patterns_
+
+### publish(self, channel, message)
+ Publish _message_ on _channel_.
+ Returns the number of subscribers the message was delivered to.
+
+### punsubscribe(self, patterns=[])
+ Unsubscribe from any channel matching any pattern in _patterns_.
+ If empty, unsubscribe from all channels.
+
+### push(self, name, value, head=False)
+ Push _value_ onto list _name_.
+
+ This method has been deprecated, use __Redis.lpush__ or __Redis.rpush__ instead.
+
+### randomkey(self)
+ Returns the name of a random key
+
+### rename(self, src, dst, **kwargs)
+ Rename key _src_ to _dst_
+
+ * The following flags have been deprecated *
+ If _preserve_ is True, rename the key only if the destination name
+ doesn't already exist
+
+### renamenx(self, src, dst)
+ Rename key _src_ to _dst_ if _dst_ doesn't already exist
+
+### rpop(self, name)
+ Remove and return the last item of the list _name_
+
+### rpoplpush(self, src, dst)
+ RPOP a value off of the _src_ list and atomically LPUSH it
+ on to the _dst_ list. Returns the value.
+
+### rpush(self, name, value)
+ Push _value_ onto the tail of the list _name_
+
+### sadd(self, name, value)
+ Add _value_ to set _name_
+
+### save(self)
+ Tell the Redis server to save its data to disk,
+ blocking until the save is complete
+
+### scard(self, name)
+ Return the number of elements in set _name_
+
+### sdiff(self, keys, *args)
+ Return the difference of sets specified by _keys_
+
+### sdiffstore(self, dest, keys, *args)
+ Store the difference of sets specified by _keys_ into a new
+ set named _dest_. Returns the number of keys in the new set.
+
+### select(self, db, host=None, port=None, password=None, socket_timeout=None)
+ Switch to a different Redis connection.
+
+ If the host and port aren't provided and there's an existing
+ connection, use the existing connection's host and port instead.
+
+ Note this method actually replaces the underlying connection object
+ prior to issuing the SELECT command. This makes sure we protect
+ the thread-safe connections
+
+### set(self, name, value, **kwargs)
+ Set the value at key _name_ to _value_
+
+ * The following flags have been deprecated *
+ If _preserve_ is True, set the value only if key doesn't already
+ exist
+ If _getset_ is True, set the value only if key doesn't already exist
+ and return the resulting value of key
+
+### setex(self, name, value, time)
+ Set the value of key _name_ to _value_
+ that expires in _time_ seconds
+
+### setnx(self, name, value)
+ Set the value of key _name_ to _value_ if key doesn't exist
+
+### sinter(self, keys, *args)
+ Return the intersection of sets specified by _keys_
+
+### sinterstore(self, dest, keys, *args)
+ Store the intersection of sets specified by _keys_ into a new
+ set named _dest_. Returns the number of keys in the new set.
+
+### sismember(self, name, value)
+ Return a boolean indicating if _value_ is a member of set _name_
+
+### smembers(self, name)
+ Return all members of the set _name_
+
+### smove(self, src, dst, value)
+ Move _value_ from set _src_ to set _dst_ atomically
+
+### sort(self, name, start=None, num=None, by=None, get=None, desc=False, alpha=False, store=None)
+ Sort and return the list, set or sorted set at _name_.
+
+ _start_ and _num_ allow for paging through the sorted data
+
+ _by_ allows using an external key to weight and sort the items.
+ Use an "*" to indicate where in the key the item value is located
+
+ _get_ allows for returning items from external keys rather than the
+ sorted data itself. Use an "*" to indicate where int he key
+ the item value is located
+
+ _desc_ allows for reversing the sort
+
+ _alpha_ allows for sorting lexicographically rather than numerically
+
+ _store_ allows for storing the result of the sort into
+ the key _store_
+
+### spop(self, name)
+ Remove and return a random member of set _name_
+
+### srandmember(self, name)
+ Return a random member of set _name_
+
+### srem(self, name, value)
+ Remove _value_ from set _name_
+
+### subscribe(self, channels)
+ Subscribe to _channels_, waiting for messages to be published
+
+### substr(self, name, start, end=-1)
+ Return a substring of the string at key _name_. _start_ and _end_
+ are 0-based integers specifying the portion of the string to return.
+
+### sunion(self, keys, *args)
+ Return the union of sets specifiued by _keys_
+
+### sunionstore(self, dest, keys, *args)
+ Store the union of sets specified by _keys_ into a new
+ set named _dest_. Returns the number of keys in the new set.
+
+### ttl(self, name)
+ Returns the number of seconds until the key _name_ will expire
+
+### type(self, name)
+ Returns the type of key _name_
+
+### unsubscribe(self, channels=[])
+ Unsubscribe from _channels_. If empty, unsubscribe
+ from all channels
+
+### zadd(self, name, value, score)
+ Add member _value_ with score _score_ to sorted set _name_
+
+### zcard(self, name)
+ Return the number of elements in the sorted set _name_
+
+### zincr(self, key, member, value=1)
+ This has been deprecated, use zincrby instead
+
+### zincrby(self, name, value, amount=1)
+ Increment the score of _value_ in sorted set _name_ by _amount_
+
+### zinter(self, dest, keys, aggregate=None)
+
+###zinterstore(self, dest, keys, aggregate=None)
+ Intersect multiple sorted sets specified by _keys_ into
+ a new sorted set, _dest_. Scores in the destination will be
+ aggregated based on the _aggregate_, or SUM if none is provided.
+
+### zrange(self, name, start, end, desc=False, withscores=False)
+ Return a range of values from sorted set _name_ between
+ _start_ and _end_ sorted in ascending order.
+
+ _start_ and _end_ can be negative, indicating the end of the range.
+
+ _desc_ indicates to sort in descending order.
+
+ _withscores_ indicates to return the scores along with the values.
+ The return type is a list of (value, score) pairs
+
+### zrangebyscore(self, name, min, max, start=None, num=None, withscores=False)
+ Return a range of values from the sorted set _name_ with scores
+ between _min_ and _max_.
+
+ If _start_ and _num_ are specified, then return a slice of the range.
+
+ _withscores_ indicates to return the scores along with the values.
+ The return type is a list of (value, score) pairs
+
+### zrank(self, name, value)
+ Returns a 0-based value indicating the rank of _value_ in sorted set
+ _name_
+
+### zrem(self, name, value)
+ Remove member _value_ from sorted set _name_
+
+### zremrangebyrank(self, name, min, max)
+ Remove all elements in the sorted set _name_ with ranks between
+ _min_ and _max_. Values are 0-based, ordered from smallest score
+ to largest. Values can be negative indicating the highest scores.
+ Returns the number of elements removed
+
+### zremrangebyscore(self, name, min, max)
+ Remove all elements in the sorted set _name_ with scores
+ between _min_ and _max_. Returns the number of elements removed.
+
+### zrevrange(self, name, start, num, withscores=False)
+ Return a range of values from sorted set _name_ between
+ _start_ and _num_ sorted in descending order.
+
+ _start_ and _num_ can be negative, indicating the end of the range.
+
+ _withscores_ indicates to return the scores along with the values
+ as a dictionary of value => score
+
+### zrevrank(self, name, value)
+ Returns a 0-based value indicating the descending rank of
+ _value_ in sorted set _name_
+
+### zscore(self, name, value)
+ Return the score of element _value_ in sorted set _name_
+
+### zunion(self, dest, keys, aggregate=None)
+
+### zunionstore(self, dest, keys, aggregate=None)
+ Union multiple sorted sets specified by _keys_ into
+ a new sorted set, _dest_. Scores in the destination will be
+ aggregated based on the _aggregate_, or SUM if none is provided.
Author
------
diff --git a/redis/__init__.py b/redis/__init__.py
index 93155fb..8850091 100644
--- a/redis/__init__.py
+++ b/redis/__init__.py
@@ -3,6 +3,8 @@ from redis.client import Redis, ConnectionPool
from redis.exceptions import RedisError, ConnectionError, AuthenticationError
from redis.exceptions import ResponseError, InvalidResponse, InvalidData
+__version__ = '2.0.1'
+
__all__ = [
'Redis', 'ConnectionPool',
'RedisError', 'ConnectionError', 'ResponseError', 'AuthenticationError'
diff --git a/redis/client.py b/redis/client.py
index d5e33ab..1274228 100644
--- a/redis/client.py
+++ b/redis/client.py
@@ -4,8 +4,8 @@ import socket
import threading
import time
import warnings
-from itertools import chain
-from redis.exceptions import ConnectionError, ResponseError, InvalidResponse
+from itertools import chain, imap
+from redis.exceptions import ConnectionError, ResponseError, InvalidResponse, WatchError
from redis.exceptions import RedisError, AuthenticationError
@@ -207,8 +207,9 @@ class Redis(threading.local):
bool
),
string_keys_to_dict(
- 'DECRBY HLEN INCRBY LLEN SCARD SDIFFSTORE SINTERSTORE '
- 'SUNIONSTORE ZCARD ZREMRANGEBYSCORE ZREVRANK',
+ 'DECRBY HLEN INCRBY LINSERT LLEN LPUSHX RPUSHX SCARD SDIFFSTORE '
+ 'SINTERSTORE SUNIONSTORE ZCARD ZREMRANGEBYRANK ZREMRANGEBYSCORE '
+ 'ZREVRANK',
int
),
string_keys_to_dict(
@@ -219,11 +220,12 @@ class Redis(threading.local):
string_keys_to_dict('ZSCORE ZINCRBY', float_or_none),
string_keys_to_dict(
'FLUSHALL FLUSHDB LSET LTRIM MSET RENAME '
- 'SAVE SELECT SET SHUTDOWN',
+ 'SAVE SELECT SET SHUTDOWN WATCH UNWATCH',
lambda r: r == 'OK'
),
+ string_keys_to_dict('BLPOP BRPOP', lambda r: r and tuple(r) or None),
string_keys_to_dict('SDIFF SINTER SMEMBERS SUNION',
- lambda r: set(r)
+ lambda r: r and set(r) or set()
),
string_keys_to_dict('ZRANGE ZRANGEBYSCORE ZREVRANGE', zset_score_pairs),
{
@@ -241,7 +243,9 @@ class Redis(threading.local):
)
# commands that should NOT pull data off the network buffer when executed
- SUBSCRIPTION_COMMANDS = set(['SUBSCRIBE', 'UNSUBSCRIBE'])
+ SUBSCRIPTION_COMMANDS = set([
+ 'SUBSCRIBE', 'UNSUBSCRIBE', 'PSUBSCRIBE', 'PUNSUBSCRIBE'
+ ])
def __init__(self, host='localhost', port=6379,
db=0, password=None, socket_timeout=None,
@@ -282,6 +286,19 @@ class Redis(threading.local):
self.errors
)
+ def lock(self, name, timeout=None, sleep=0.1):
+ """
+ Return a new Lock object using key ``name`` that mimics
+ the behavior of threading.Lock.
+
+ If specified, ``timeout`` indicates a maximum life for the lock.
+ By default, it will remain locked until release() is called.
+
+ ``sleep`` indicates the amount of time to sleep per loop iteration
+ when the lock is in blocking mode and another client is currently
+ holding the lock.
+ """
+ return Lock(self, name, timeout=timeout, sleep=sleep)
#### COMMAND EXECUTION AND PROTOCOL PARSING ####
def _execute_command(self, command_name, command, **options):
@@ -303,14 +320,11 @@ class Redis(threading.local):
def execute_command(self, *args, **options):
"Sends the command to the redis server and returns it's response"
- cmd_count = len(args)
- cmds = []
- for i in args:
- enc_value = self.encode(i)
- cmds.append('$%s\r\n%s\r\n' % (len(enc_value), enc_value))
+ cmds = ['$%s\r\n%s\r\n' % (len(enc_value), enc_value)
+ for enc_value in imap(self.encode, args)]
return self._execute_command(
args[0],
- '*%s\r\n%s' % (cmd_count, ''.join(cmds)),
+ '*%s\r\n%s' % (len(cmds), ''.join(cmds)),
**options
)
@@ -432,6 +446,17 @@ class Redis(threading.local):
self.connection = self.get_connection(
host, port, db, password, socket_timeout)
+ def shutdown(self):
+ "Shutdown the server"
+ if self.subscribed:
+ raise RedisError("Can't call 'shutdown' from a pipeline'")
+ try:
+ self.execute_command('SHUTDOWN')
+ except ConnectionError:
+ # a ConnectionError here is expected
+ return
+ raise RedisError("SHUTDOWN seems to have failed.")
+
#### SERVER INFORMATION ####
def bgrewriteaof(self):
@@ -563,7 +588,8 @@ class Redis(threading.local):
def mset(self, mapping):
"Sets each key in the ``mapping`` dict to its corresponding value"
items = []
- [items.extend(pair) for pair in mapping.iteritems()]
+ for pair in mapping.iteritems():
+ items.extend(pair)
return self.execute_command('MSET', *items)
def msetnx(self, mapping):
@@ -572,7 +598,8 @@ class Redis(threading.local):
none of the keys are already set
"""
items = []
- [items.extend(pair) for pair in mapping.iteritems()]
+ for pair in mapping.iteritems():
+ items.extend(pair)
return self.execute_command('MSETNX', *items)
def move(self, name, db):
@@ -657,6 +684,23 @@ class Redis(threading.local):
"Returns the type of key ``name``"
return self.execute_command('TYPE', name)
+ def watch(self, name):
+ """
+ Watches the value at key ``name``, or None of the key doesn't exist
+ """
+ if self.subscribed:
+ raise RedisError("Can't call 'watch' from a pipeline'")
+
+ return self.execute_command('WATCH', name)
+
+ def unwatch(self):
+ """
+ Unwatches the value at key ``name``, or None of the key doesn't exist
+ """
+ if self.subscribed:
+ raise RedisError("Can't call 'unwatch' from a pipeline'")
+
+ return self.execute_command('UNWATCH')
#### LIST COMMANDS ####
def blpop(self, keys, timeout=0):
@@ -670,7 +714,12 @@ class Redis(threading.local):
If timeout is 0, then block indefinitely.
"""
- keys = list(keys)
+ if timeout is None:
+ timeout = 0
+ if isinstance(keys, basestring):
+ keys = [keys]
+ else:
+ keys = list(keys)
keys.append(timeout)
return self.execute_command('BLPOP', *keys)
@@ -685,7 +734,12 @@ class Redis(threading.local):
If timeout is 0, then block indefinitely.
"""
- keys = list(keys)
+ if timeout is None:
+ timeout = 0
+ if isinstance(keys, basestring):
+ keys = [keys]
+ else:
+ keys = list(keys)
keys.append(timeout)
return self.execute_command('BRPOP', *keys)
@@ -697,7 +751,17 @@ class Redis(threading.local):
end of the list
"""
return self.execute_command('LINDEX', name, index)
-
+
+ def linsert(self, name, where, refvalue, value):
+ """
+ Insert ``value`` in list ``name`` either immediately before or after
+ [``where``] ``refvalue``
+
+ Returns the new length of the list on success or -1 if ``refvalue``
+ is not in the list.
+ """
+ return self.execute_command('LINSERT', name, where, refvalue, value)
+
def llen(self, name):
"Return the length of the list ``name``"
return self.execute_command('LLEN', name)
@@ -709,6 +773,10 @@ class Redis(threading.local):
def lpush(self, name, value):
"Push ``value`` onto the head of the list ``name``"
return self.execute_command('LPUSH', name, value)
+
+ def lpushx(self, name, value):
+ "Push ``value`` onto the head of the list ``name`` if ``name`` exists"
+ return self.execute_command('LPUSHX', name, value)
def lrange(self, name, start, end):
"""
@@ -785,8 +853,12 @@ class Redis(threading.local):
"Push ``value`` onto the tail of the list ``name``"
return self.execute_command('RPUSH', name, value)
+ def rpushx(self, name, value):
+ "Push ``value`` onto the tail of the list ``name`` if ``name`` exists"
+ return self.execute_command('RPUSHX', name, value)
+
def sort(self, name, start=None, num=None, by=None, get=None,
- desc=False, alpha=False, store=None):
+ desc=False, alpha=False, store=None):
"""
Sort and return the list, set or sorted set at ``name``.
@@ -819,8 +891,17 @@ class Redis(threading.local):
pieces.append(start)
pieces.append(num)
if get is not None:
- pieces.append('GET')
- pieces.append(get)
+ # If get is a string assume we want to get a single value.
+ # Otherwise assume it's an interable and we want to get multiple
+ # values. We can't just iterate blindly because strings are
+ # iterable.
+ if isinstance(get, basestring):
+ pieces.append('GET')
+ pieces.append(get)
+ else:
+ for g in get:
+ pieces.append('GET')
+ pieces.append(g)
if desc:
pieces.append('DESC')
if alpha:
@@ -913,6 +994,9 @@ class Redis(threading.local):
"Return the number of elements in the sorted set ``name``"
return self.execute_command('ZCARD', name)
+ def zcount(self, name, min, max):
+ return self.execute_command('ZCOUNT', name, min, max)
+
def zincr(self, key, member, value=1):
"This has been deprecated, use zincrby instead"
warnings.warn(DeprecationWarning(
@@ -925,6 +1009,12 @@ class Redis(threading.local):
return self.execute_command('ZINCRBY', name, amount, value)
def zinter(self, dest, keys, aggregate=None):
+ warnings.warn(DeprecationWarning(
+ "Redis.zinter has been deprecated, use Redis.zinterstore instead"
+ ))
+ return self.zinterstore(dest, keys, aggregate)
+
+ def zinterstore(self, dest, keys, aggregate=None):
"""
Intersect multiple sorted sets specified by ``keys`` into
a new sorted set, ``dest``. Scores in the destination will be
@@ -983,10 +1073,19 @@ class Redis(threading.local):
"Remove member ``value`` from sorted set ``name``"
return self.execute_command('ZREM', name, value)
+ def zremrangebyrank(self, name, min, max):
+ """
+ Remove all elements in the sorted set ``name`` with ranks between
+ ``min`` and ``max``. Values are 0-based, ordered from smallest score
+ to largest. Values can be negative indicating the highest scores.
+ Returns the number of elements removed
+ """
+ return self.execute_command('ZREMRANGEBYRANK', name, min, max)
+
def zremrangebyscore(self, name, min, max):
"""
Remove all elements in the sorted set ``name`` with scores
- between ``min`` and ``max``
+ between ``min`` and ``max``. Returns the number of elements removed.
"""
return self.execute_command('ZREMRANGEBYSCORE', name, min, max)
@@ -1017,6 +1116,12 @@ class Redis(threading.local):
return self.execute_command('ZSCORE', name, value)
def zunion(self, dest, keys, aggregate=None):
+ warnings.warn(DeprecationWarning(
+ "Redis.zunion has been deprecated, use Redis.zunionstore instead"
+ ))
+ return self.zunionstore(dest, keys, aggregate)
+
+ def zunionstore(self, dest, keys, aggregate=None):
"""
Union multiple sorted sets specified by ``keys`` into
a new sorted set, ``dest``. Scores in the destination will be
@@ -1077,13 +1182,21 @@ class Redis(threading.local):
"""
return self.execute_command('HSET', name, key, value)
+ def hsetnx(self, name, key, value):
+ """
+ Set ``key`` to ``value`` within hash ``name`` if ``key`` does not
+ exist. Returns 1 if HSETNX created a field, otherwise 0.
+ """
+ return self.execute_command("HSETNX", name, key, value)
+
def hmset(self, name, mapping):
"""
Sets each key in the ``mapping`` dict to its corresponding value
in the hash ``name``
"""
items = []
- [items.extend(pair) for pair in mapping.iteritems()]
+ for pair in mapping.iteritems():
+ items.extend(pair)
return self.execute_command('HMSET', name, *items)
def hmget(self, name, keys):
@@ -1145,10 +1258,23 @@ class Redis(threading.local):
"Listen for messages on channels this client has been subscribed to"
while self.subscribed:
r = self.parse_response('LISTEN')
- message_type, channel, message = r[0], r[1], r[2]
- yield (message_type, channel, message)
- if message_type == 'unsubscribe' and message == 0:
+ if r[0] == 'pmessage':
+ msg = {
+ 'type': r[0],
+ 'pattern': r[1],
+ 'channel': r[2],
+ 'data': r[3]
+ }
+ else:
+ msg = {
+ 'type': r[0],
+ 'pattern': None,
+ 'channel': r[1],
+ 'data': r[2]
+ }
+ if r[0] == 'unsubscribe' and r[2] == 0:
self.subscribed = False
+ yield msg
class Pipeline(Redis):
@@ -1217,6 +1343,10 @@ class Pipeline(Redis):
_ = self.parse_response('_')
# parse the EXEC. we want errors returned as items in the response
response = self.parse_response('_', catch_errors=True)
+
+ if response is None:
+ raise WatchError("Watched variable changed.")
+
if len(response) != len(commands):
raise ResponseError("Wrong number of response items from "
"pipeline execution")
@@ -1257,3 +1387,85 @@ class Pipeline(Redis):
def select(self, *args, **kwargs):
raise RedisError("Cannot select a different database from a pipeline")
+
+class Lock(object):
+ """
+ A shared, distributed Lock. Using Redis for locking allows the Lock
+ to be shared across processes and/or machines.
+
+ It's left to the user to resolve deadlock issues and make sure
+ multiple clients play nicely together.
+ """
+
+ LOCK_FOREVER = 2**31+1 # 1 past max unix time
+
+ def __init__(self, redis, name, timeout=None, sleep=0.1):
+ """
+ Create a new Lock instnace named ``name`` using the Redis client
+ supplied by ``redis``.
+
+ ``timeout`` indicates a maximum life for the lock.
+ By default, it will remain locked until release() is called.
+
+ ``sleep`` indicates the amount of time to sleep per loop iteration
+ when the lock is in blocking mode and another client is currently
+ holding the lock.
+
+ Note: If using ``timeout``, you should make sure all the hosts
+ that are running clients are within the same timezone and are using
+ a network time service like ntp.
+ """
+ self.redis = redis
+ self.name = name
+ self.acquired_until = None
+ self.timeout = timeout
+ self.sleep = sleep
+
+ def __enter__(self):
+ return self.acquire()
+
+ def __exit__(self, exc_type, exc_value, traceback):
+ self.release()
+
+ def acquire(self, blocking=True):
+ """
+ Use Redis to hold a shared, distributed lock named ``name``.
+ Returns True once the lock is acquired.
+
+ If ``blocking`` is False, always return immediately. If the lock
+ was acquired, return True, otherwise return False.
+ """
+ sleep = self.sleep
+ timeout = self.timeout
+ while 1:
+ unixtime = int(time.time())
+ if timeout:
+ timeout_at = unixtime + timeout
+ else:
+ timeout_at = Lock.LOCK_FOREVER
+ if self.redis.setnx(self.name, timeout_at):
+ self.acquired_until = timeout_at
+ return True
+ # We want blocking, but didn't acquire the lock
+ # check to see if the current lock is expired
+ existing = long(self.redis.get(self.name) or 1)
+ if existing < unixtime:
+ # the previous lock is expired, attempt to overwrite it
+ existing = long(self.redis.getset(self.name, timeout_at) or 1)
+ if existing < unixtime:
+ # we successfully acquired the lock
+ self.acquired_until = timeout_at
+ return True
+ if not blocking:
+ return False
+ time.sleep(sleep)
+
+ def release(self):
+ "Releases the already acquired lock"
+ if self.acquired_until is None:
+ raise ValueError("Cannot release an unlocked lock")
+ existing = long(self.redis.get(self.name) or 1)
+ # if the lock time is in the future, delete the lock
+ if existing >= self.acquired_until:
+ self.redis.delete(self.name)
+ self.acquired_until = None
diff --git a/redis/exceptions.py b/redis/exceptions.py
index d3449b6..b3257ac 100644
--- a/redis/exceptions.py
+++ b/redis/exceptions.py
@@ -2,19 +2,21 @@
class RedisError(Exception):
pass
-
+
class AuthenticationError(RedisError):
pass
-
+
class ConnectionError(RedisError):
pass
-
+
class ResponseError(RedisError):
pass
-
+
class InvalidResponse(RedisError):
pass
-
+
class InvalidData(RedisError):
pass
- \ No newline at end of file
+
+class WatchError(RedisError):
+ pass
diff --git a/run_tests b/run_tests
new file mode 100755
index 0000000..2d629c6
--- /dev/null
+++ b/run_tests
@@ -0,0 +1,9 @@
+#!/usr/bin/env python
+
+import unittest
+from tests import all_tests
+
+
+if __name__ == "__main__":
+ tests = all_tests()
+ results = unittest.TextTestRunner().run(tests)
diff --git a/setup.py b/setup.py
index 0660371..4bdbd63 100644
--- a/setup.py
+++ b/setup.py
@@ -7,7 +7,7 @@
@brief Setuptools configuration for redis client
"""
-version = '1.36'
+version = '2.0.1'
sdict = {
'name' : 'redis',
diff --git a/tests/__init__.py b/tests/__init__.py
index 45e55b0..8931b07 100644
--- a/tests/__init__.py
+++ b/tests/__init__.py
@@ -2,10 +2,12 @@ import unittest
from server_commands import ServerCommandsTestCase
from connection_pool import ConnectionPoolTestCase
from pipeline import PipelineTestCase
+from lock import LockTestCase
def all_tests():
suite = unittest.TestSuite()
suite.addTest(unittest.makeSuite(ServerCommandsTestCase))
suite.addTest(unittest.makeSuite(ConnectionPoolTestCase))
suite.addTest(unittest.makeSuite(PipelineTestCase))
+ suite.addTest(unittest.makeSuite(LockTestCase))
return suite
diff --git a/tests/lock.py b/tests/lock.py
new file mode 100644
index 0000000..a6e447d
--- /dev/null
+++ b/tests/lock.py
@@ -0,0 +1,53 @@
+from __future__ import with_statement
+import redis
+import time
+import unittest
+from redis.client import Lock
+
+class LockTestCase(unittest.TestCase):
+ def setUp(self):
+ self.client = redis.Redis(host='localhost', port=6379, db=9)
+ self.client.flushdb()
+
+ def tearDown(self):
+ self.client.flushdb()
+
+ def test_lock(self):
+ lock = self.client.lock('foo')
+ self.assert_(lock.acquire())
+ self.assertEquals(self.client['foo'], str(Lock.LOCK_FOREVER))
+ lock.release()
+ self.assertEquals(self.client['foo'], None)
+
+ def test_competing_locks(self):
+ lock1 = self.client.lock('foo')
+ lock2 = self.client.lock('foo')
+ self.assert_(lock1.acquire())
+ self.assertFalse(lock2.acquire(blocking=False))
+ lock1.release()
+ self.assert_(lock2.acquire())
+ self.assertFalse(lock1.acquire(blocking=False))
+ lock2.release()
+
+ def test_timeouts(self):
+ lock1 = self.client.lock('foo', timeout=1)
+ lock2 = self.client.lock('foo')
+ self.assert_(lock1.acquire())
+ self.assertEquals(lock1.acquired_until, long(time.time()) + 1)
+ self.assertEquals(lock1.acquired_until, long(self.client['foo']))
+ self.assertFalse(lock2.acquire(blocking=False))
+ time.sleep(2) # need to wait up to 2 seconds for lock to timeout
+ self.assert_(lock2.acquire(blocking=False))
+ lock2.release()
+
+ def test_non_blocking(self):
+ lock1 = self.client.lock('foo')
+ self.assert_(lock1.acquire(blocking=False))
+ self.assert_(lock1.acquired_until)
+ lock1.release()
+ self.assert_(lock1.acquired_until is None)
+
+ def test_context_manager(self):
+ with self.client.lock('foo'):
+ self.assertEquals(self.client['foo'], str(Lock.LOCK_FOREVER))
+ self.assertEquals(self.client['foo'], None)
diff --git a/tests/server_commands.py b/tests/server_commands.py
index 056188a..5455174 100644
--- a/tests/server_commands.py
+++ b/tests/server_commands.py
@@ -16,6 +16,8 @@ class ServerCommandsTestCase(unittest.TestCase):
def tearDown(self):
self.client.flushdb()
+ for c in self.client.connection_pool.get_all_connections():
+ c.disconnect()
# GENERAL SERVER COMMANDS
def test_dbsize(self):
@@ -211,6 +213,25 @@ class ServerCommandsTestCase(unittest.TestCase):
self.client.zadd('a', '1', 1)
self.assertEquals(self.client.type('a'), 'zset')
+ def test_watch(self):
+ self.client.set("a", 1)
+
+ self.client.watch("a")
+ pipeline = self.client.pipeline()
+ pipeline.set("a", 2)
+ self.assertEquals(pipeline.execute(), [True])
+
+ self.client.set("b", 1)
+ self.client.watch("b")
+ self.get_client().set("b", 2)
+ pipeline = self.client.pipeline()
+ pipeline.set("b", 3)
+
+ self.assertRaises(redis.exceptions.WatchError, pipeline.execute)
+
+ def test_unwatch(self):
+ self.assertEquals(self.client.unwatch(), True)
+
# LISTS
def make_list(self, name, l):
for i in l:
@@ -219,20 +240,24 @@ class ServerCommandsTestCase(unittest.TestCase):
def test_blpop(self):
self.make_list('a', 'ab')
self.make_list('b', 'cd')
- self.assertEquals(self.client.blpop(['b', 'a'], timeout=1), ['b', 'c'])
- self.assertEquals(self.client.blpop(['b', 'a'], timeout=1), ['b', 'd'])
- self.assertEquals(self.client.blpop(['b', 'a'], timeout=1), ['a', 'a'])
- self.assertEquals(self.client.blpop(['b', 'a'], timeout=1), ['a', 'b'])
+ self.assertEquals(self.client.blpop(['b', 'a'], timeout=1), ('b', 'c'))
+ self.assertEquals(self.client.blpop(['b', 'a'], timeout=1), ('b', 'd'))
+ self.assertEquals(self.client.blpop(['b', 'a'], timeout=1), ('a', 'a'))
+ self.assertEquals(self.client.blpop(['b', 'a'], timeout=1), ('a', 'b'))
self.assertEquals(self.client.blpop(['b', 'a'], timeout=1), None)
+ self.make_list('c', 'a')
+ self.assertEquals(self.client.blpop('c', timeout=1), ('c', 'a'))
def test_brpop(self):
self.make_list('a', 'ab')
self.make_list('b', 'cd')
- self.assertEquals(self.client.brpop(['b', 'a'], timeout=1), ['b', 'd'])
- self.assertEquals(self.client.brpop(['b', 'a'], timeout=1), ['b', 'c'])
- self.assertEquals(self.client.brpop(['b', 'a'], timeout=1), ['a', 'b'])
- self.assertEquals(self.client.brpop(['b', 'a'], timeout=1), ['a', 'a'])
+ self.assertEquals(self.client.brpop(['b', 'a'], timeout=1), ('b', 'd'))
+ self.assertEquals(self.client.brpop(['b', 'a'], timeout=1), ('b', 'c'))
+ self.assertEquals(self.client.brpop(['b', 'a'], timeout=1), ('a', 'b'))
+ self.assertEquals(self.client.brpop(['b', 'a'], timeout=1), ('a', 'a'))
self.assertEquals(self.client.brpop(['b', 'a'], timeout=1), None)
+ self.make_list('c', 'a')
+ self.assertEquals(self.client.brpop('c', timeout=1), ('c', 'a'))
def test_lindex(self):
# no key
@@ -247,6 +272,24 @@ class ServerCommandsTestCase(unittest.TestCase):
self.assertEquals(self.client.lindex('a', '1'), 'b')
self.assertEquals(self.client.lindex('a', '2'), 'c')
+ def test_linsert(self):
+ # no key
+ self.assertEquals(self.client.linsert('a', 'after', 'x', 'y'), 0)
+ # key is not a list
+ self.client['a'] = 'b'
+ self.assertRaises(
+ redis.ResponseError, self.client.linsert, 'a', 'after', 'x', 'y'
+ )
+ del self.client['a']
+ # real logic
+ self.make_list('a', 'abc')
+ self.assertEquals(self.client.linsert('a', 'after', 'b', 'b1'), 4)
+ self.assertEquals(self.client.lrange('a', 0, -1),
+ ['a', 'b', 'b1', 'c'])
+ self.assertEquals(self.client.linsert('a', 'before', 'b', 'a1'), 5)
+ self.assertEquals(self.client.lrange('a', 0, -1),
+ ['a', 'a1', 'b', 'b1', 'c'])
+
def test_llen(self):
# no key
self.assertEquals(self.client.llen('a'), 0)
@@ -288,6 +331,18 @@ class ServerCommandsTestCase(unittest.TestCase):
self.assertEquals(self.client.lindex('a', 0), 'a')
self.assertEquals(self.client.lindex('a', 1), 'b')
+ def test_lpushx(self):
+ # key is not a list
+ self.client['a'] = 'b'
+ self.assertRaises(redis.ResponseError, self.client.lpushx, 'a', 'a')
+ del self.client['a']
+ # real logic
+ self.assertEquals(self.client.lpushx('a', 'b'), 0)
+ self.assertEquals(self.client.lrange('a', 0, -1), [])
+ self.make_list('a', 'abc')
+ self.assertEquals(self.client.lpushx('a', 'd'), 4)
+ self.assertEquals(self.client.lrange('a', 0, -1), ['d', 'a', 'b', 'c'])
+
def test_lrange(self):
# no key
self.assertEquals(self.client.lrange('a', 0, 1), [])
@@ -411,6 +466,18 @@ class ServerCommandsTestCase(unittest.TestCase):
self.assertEquals(self.client.lindex('a', 0), 'a')
self.assertEquals(self.client.lindex('a', 1), 'b')
+ def test_rpushx(self):
+ # key is not a list
+ self.client['a'] = 'b'
+ self.assertRaises(redis.ResponseError, self.client.rpushx, 'a', 'a')
+ del self.client['a']
+ # real logic
+ self.assertEquals(self.client.rpushx('a', 'b'), 0)
+ self.assertEquals(self.client.lrange('a', 0, -1), [])
+ self.make_list('a', 'abc')
+ self.assertEquals(self.client.rpushx('a', 'd'), 4)
+ self.assertEquals(self.client.lrange('a', 0, -1), ['a', 'b', 'c', 'd'])
+
# Set commands
def make_set(self, name, l):
for i in l:
@@ -609,6 +676,17 @@ class ServerCommandsTestCase(unittest.TestCase):
self.make_zset('a', {'a1': 1, 'a2': 2, 'a3': 3})
self.assertEquals(self.client.zcard('a'), 3)
+ def test_zcount(self):
+ # key is not a zset
+ self.client['a'] = 'a'
+ self.assertRaises(redis.ResponseError, self.client.zcount, 'a', 0, 0)
+ del self.client['a']
+ # real logic
+ self.make_zset('a', {'a1': 1, 'a2': 2, 'a3': 3})
+ self.assertEquals(self.client.zcount('a', '-inf', '+inf'), 3)
+ self.assertEquals(self.client.zcount('a', 1, 2), 2)
+ self.assertEquals(self.client.zcount('a', 10, 20), 0)
+
def test_zincrby(self):
# key is not a zset
self.client['a'] = 'a'
@@ -621,32 +699,34 @@ class ServerCommandsTestCase(unittest.TestCase):
self.assertEquals(self.client.zscore('a', 'a2'), 3.0)
self.assertEquals(self.client.zscore('a', 'a3'), 8.0)
- def test_zinter(self):
+ def test_zinterstore(self):
self.make_zset('a', {'a1': 1, 'a2': 1, 'a3': 1})
self.make_zset('b', {'a1': 2, 'a3': 2, 'a4': 2})
self.make_zset('c', {'a1': 6, 'a3': 5, 'a4': 4})
# sum, no weight
- self.assert_(self.client.zinter('z', ['a', 'b', 'c']))
+ self.assert_(self.client.zinterstore('z', ['a', 'b', 'c']))
self.assertEquals(
self.client.zrange('z', 0, -1, withscores=True),
[('a3', 8), ('a1', 9)]
)
# max, no weight
- self.assert_(self.client.zinter('z', ['a', 'b', 'c'], aggregate='MAX'))
+ self.assert_(
+ self.client.zinterstore('z', ['a', 'b', 'c'], aggregate='MAX')
+ )
self.assertEquals(
self.client.zrange('z', 0, -1, withscores=True),
[('a3', 5), ('a1', 6)]
)
# with weight
- self.assert_(self.client.zinter('z', {'a': 1, 'b': 2, 'c': 3}))
+ self.assert_(self.client.zinterstore('z', {'a': 1, 'b': 2, 'c': 3}))
self.assertEquals(
self.client.zrange('z', 0, -1, withscores=True),
[('a3', 20), ('a1', 23)]
)
-
+
def test_zrange(self):
# key is not a zset
@@ -709,6 +789,17 @@ class ServerCommandsTestCase(unittest.TestCase):
self.assertEquals(self.client.zrem('a', 'b'), False)
self.assertEquals(self.client.zrange('a', 0, 5), ['a1', 'a3'])
+ def test_zremrangebyrank(self):
+ # key is not a zset
+ self.client['a'] = 'a'
+ self.assertRaises(redis.ResponseError, self.client.zremrangebyscore,
+ 'a', 0, 1)
+ del self.client['a']
+ # real logic
+ self.make_zset('a', {'a1': 1, 'a2': 2, 'a3': 3, 'a4': 4, 'a5': 5})
+ self.assertEquals(self.client.zremrangebyrank('a', 1, 3), 3)
+ self.assertEquals(self.client.zrange('a', 0, 5), ['a1', 'a5'])
+
def test_zremrangebyscore(self):
# key is not a zset
self.client['a'] = 'a'
@@ -764,27 +855,29 @@ class ServerCommandsTestCase(unittest.TestCase):
# test a non-existant member
self.assertEquals(self.client.zscore('a', 'a4'), None)
- def test_zunion(self):
+ def test_zunionstore(self):
self.make_zset('a', {'a1': 1, 'a2': 1, 'a3': 1})
self.make_zset('b', {'a1': 2, 'a3': 2, 'a4': 2})
self.make_zset('c', {'a1': 6, 'a4': 5, 'a5': 4})
-
+
# sum, no weight
- self.assert_(self.client.zunion('z', ['a', 'b', 'c']))
+ self.assert_(self.client.zunionstore('z', ['a', 'b', 'c']))
self.assertEquals(
self.client.zrange('z', 0, -1, withscores=True),
[('a2', 1), ('a3', 3), ('a5', 4), ('a4', 7), ('a1', 9)]
)
# max, no weight
- self.assert_(self.client.zunion('z', ['a', 'b', 'c'], aggregate='MAX'))
+ self.assert_(
+ self.client.zunionstore('z', ['a', 'b', 'c'], aggregate='MAX')
+ )
self.assertEquals(
self.client.zrange('z', 0, -1, withscores=True),
[('a2', 1), ('a3', 2), ('a5', 4), ('a4', 5), ('a1', 6)]
)
# with weight
- self.assert_(self.client.zunion('z', {'a': 1, 'b': 2, 'c': 3}))
+ self.assert_(self.client.zunionstore('z', {'a': 1, 'b': 2, 'c': 3}))
self.assertEquals(
self.client.zrange('z', 0, -1, withscores=True),
[('a2', 1), ('a3', 5), ('a5', 12), ('a4', 19), ('a1', 23)]
@@ -815,6 +908,14 @@ class ServerCommandsTestCase(unittest.TestCase):
# key inside of hash that doesn't exist returns null value
self.assertEquals(self.client.hget('a', 'b'), None)
+ def test_hsetnx(self):
+ # Initially set the hash field
+ self.client.hsetnx('a', 'a1', 1)
+ self.assertEqual(self.client.hget('a', 'a1'), '1')
+ # Try and set the existing hash field to a different value
+ self.client.hsetnx('a', 'a1', 2)
+ self.assertEqual(self.client.hget('a', 'a1'), '1')
+
def test_hmset(self):
d = {'a': '1', 'b': '2', 'c': '3'}
self.assert_(self.client.hmset('foo', d))
@@ -970,6 +1071,14 @@ class ServerCommandsTestCase(unittest.TestCase):
self.assertEquals(self.client.sort('a', get='user:*'),
['u1', 'u2', 'u3'])
+ def test_sort_get_multi(self):
+ self.client['user:1'] = 'u1'
+ self.client['user:2'] = 'u2'
+ self.client['user:3'] = 'u3'
+ self.make_list('a', '231')
+ self.assertEquals(self.client.sort('a', get=('user:*', '#')),
+ ['u1', '1', 'u2', '2', 'u3', '3'])
+
def test_sort_desc(self):
self.make_list('a', '231')
self.assertEquals(self.client.sort('a', desc=True), ['3', '2', '1'])
@@ -1018,6 +1127,9 @@ class ServerCommandsTestCase(unittest.TestCase):
channels = ('a1', 'a2', 'a3')
for c in channels:
r.subscribe(c)
+ # state variable should be flipped
+ self.assertEquals(r.subscribed, True)
+
channels_to_publish_to = channels + ('a4',)
messages_per_channel = 4
def publish():
@@ -1025,34 +1137,40 @@ class ServerCommandsTestCase(unittest.TestCase):
for c in channels_to_publish_to:
self.client.publish(c, 'a message')
time.sleep(0.01)
- t = threading.Thread(target=publish)
+ for c in channels_to_publish_to:
+ self.client.publish(c, 'unsubscribe')
+ time.sleep(0.01)
+
messages = []
- # should receive a message for each subscribe command
+ # should receive a message for each subscribe/unsubscribe command
# plus a message for each iteration of the loop * num channels
- num_messages_to_expect = len(channels) + \
+ # we hide the data messages that tell the client to unsubscribe
+ num_messages_to_expect = len(channels)*2 + \
(messages_per_channel*len(channels))
- thread_started = False
+ t = threading.Thread(target=publish)
+ t.start()
for msg in r.listen():
- if not thread_started:
- # start the thread delayed so that we are intermingling
- # publish commands with pulling messsages off the socket
- # with subscribe
- thread_started = True
- t.start()
- messages.append(msg)
- if len(messages) == num_messages_to_expect:
- break
+ if msg['data'] == 'unsubscribe':
+ r.unsubscribe(msg['channel'])
+ else:
+ messages.append(msg)
+
+ self.assertEquals(r.subscribed, False)
+ self.assertEquals(len(messages), num_messages_to_expect)
sent_types, sent_channels = {}, {}
- for msg_type, channel, _ in messages:
+ for msg in messages:
+ msg_type = msg['type']
+ channel = msg['channel']
sent_types.setdefault(msg_type, 0)
sent_types[msg_type] += 1
if msg_type == 'message':
sent_channels.setdefault(channel, 0)
sent_channels[channel] += 1
for channel in channels:
+ self.assert_(channel in sent_channels)
self.assertEquals(sent_channels[channel], messages_per_channel)
- self.assert_(channel in channels)
self.assertEquals(sent_types['subscribe'], len(channels))
+ self.assertEquals(sent_types['unsubscribe'], len(channels))
self.assertEquals(sent_types['message'],
len(channels) * messages_per_channel)