diff options
| author | Andy McCurdy <andy@andymccurdy.com> | 2010-04-28 11:53:21 -0700 |
|---|---|---|
| committer | Andy McCurdy <andy@andymccurdy.com> | 2010-04-28 11:53:21 -0700 |
| commit | f4c8393113e2205eb8eeddeed10d42ad70d648ef (patch) | |
| tree | 20b2af0c235382422a4fc90011fd864f4c45b63d | |
| parent | 60063e8ef022ea500788a572d7b3c07080557c54 (diff) | |
| download | redis-py-f4c8393113e2205eb8eeddeed10d42ad70d648ef.tar.gz | |
added support for zinter and zunion
| -rw-r--r-- | redis/client.py | 32 | ||||
| -rw-r--r-- | tests/server_commands.py | 54 |
2 files changed, 86 insertions, 0 deletions
diff --git a/redis/client.py b/redis/client.py index 16e357e..d29df67 100644 --- a/redis/client.py +++ b/redis/client.py @@ -896,6 +896,14 @@ class Redis(threading.local): "Increment the score of ``value`` in sorted set ``name`` by ``amount``" return self.execute_command('ZINCRBY', name, amount, value) + def zinter(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. + """ + return self._zaggregate('ZINTER', dest, keys, aggregate) + def zrange(self, name, start, end, desc=False, withscores=False): """ Return a range of values from sorted set ``name`` between @@ -980,6 +988,30 @@ class Redis(threading.local): "Return the score of element ``value`` in sorted set ``name``" return self.execute_command('ZSCORE', name, value) + def zunion(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. + """ + return self._zaggregate('ZUNION', dest, keys, aggregate) + + def _zaggregate(self, command, dest, keys, aggregate=None): + pieces = [command, dest, len(keys)] + if isinstance(keys, dict): + items = keys.items() + keys = [i[0] for i in items] + weights = [i[1] for i in items] + else: + weights = None + pieces.extend(keys) + if weights: + pieces.append('WEIGHTS') + pieces.extend(weights) + if aggregate: + pieces.append('AGGREGATE') + pieces.append(aggregate) + return self.execute_command(*pieces) #### HASH COMMANDS #### def hdel(self, name, key): diff --git a/tests/server_commands.py b/tests/server_commands.py index 360b763..da20a97 100644 --- a/tests/server_commands.py +++ b/tests/server_commands.py @@ -611,6 +611,33 @@ 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): + 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.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.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.assertEquals( + self.client.zrange('z', 0, -1, withscores=True), + [('a3', 20), ('a1', 23)] + ) + + def test_zrange(self): # key is not a zset self.client['a'] = 'a' @@ -724,6 +751,33 @@ class ServerCommandsTestCase(unittest.TestCase): # test a non-existant member self.assertEquals(self.client.zscore('a', 'a4'), None) + def test_zunion(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.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.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.assertEquals( + self.client.zrange('z', 0, -1, withscores=True), + [('a2', 1), ('a3', 5), ('a5', 12), ('a4', 19), ('a1', 23)] + ) + + # HASHES def make_hash(self, key, d): for k,v in d.iteritems(): |
