diff options
| author | Andy McCurdy <andy@andymccurdy.com> | 2011-05-16 11:45:53 -0700 |
|---|---|---|
| committer | Andy McCurdy <andy@andymccurdy.com> | 2011-05-16 11:45:53 -0700 |
| commit | c650073a51c37e80815b6d5db901e7ae3c893411 (patch) | |
| tree | cc3f7ec19a966b8ccaf71c605c0002ba9d503f80 | |
| parent | f64c4ad7ed9a2d87a92415c3490edf8e9b757d84 (diff) | |
| download | redis-py-c650073a51c37e80815b6d5db901e7ae3c893411.tar.gz | |
all tests passing with new connection pool
| -rw-r--r-- | redis/client.py | 200 | ||||
| -rw-r--r-- | redis/connection.py | 78 | ||||
| -rw-r--r-- | tests/connection_pool.py | 49 | ||||
| -rw-r--r-- | tests/server_commands.py | 3 |
4 files changed, 165 insertions, 165 deletions
diff --git a/redis/client.py b/redis/client.py index d014bd1..a5b2f9b 100644 --- a/redis/client.py +++ b/redis/client.py @@ -177,14 +177,21 @@ class Redis(threading.local): db=0, password=None, socket_timeout=None, connection_pool=None, charset='utf-8', errors='strict'): - self.encoding = charset - self.errors = errors - self.connection = None self.subscribed = False - self.connection_pool = connection_pool and connection_pool or ConnectionPool() - self.select(db, host, port, password, socket_timeout) - - #### Legacty accessors of connection information #### + if connection_pool: + self.connection_pool = connection_pool + else: + self.connection_pool = ConnectionPool( + host=host, + port=port, + db=db, + password=password, + socket_timeout=socket_timeout, + encoding=charset, + encoding_errors=errors + ) + + #### Legacy accessors of connection information #### def _get_host(self): return self.connection.host host = property(_get_host) @@ -205,12 +212,7 @@ class Redis(threading.local): pipelines are useful for batch loading of data as they reduce the number of back and forth network operations between client and server. """ - return Pipeline( - self.connection, - transaction, - self.encoding, - self.errors - ) + return Pipeline(self.connection_pool, transaction) def lock(self, name, timeout=None, sleep=0.1): """ @@ -228,77 +230,38 @@ class Redis(threading.local): #### COMMAND EXECUTION AND PROTOCOL PARSING #### def execute_command(self, *args, **options): - command_name = args[0] - subscription_command = command_name in self.SUBSCRIPTION_COMMANDS - if self.subscribed and not subscription_command: - raise PubSubError("Cannot issue commands other than SUBSCRIBE and " - "UNSUBSCRIBE while channels are open") + connection = self.connection_pool.get_connection() try: - self.connection.send_command(*args) - if subscription_command: - return None - return self.parse_response(command_name, **options) - except ConnectionError: - self.connection.disconnect() - self.connection.send_command(*args) - if subscription_command: - return None - return self.parse_response(command_name, **options) - - def parse_response(self, command_name, **options): + command_name = args[0] + subscription_command = command_name in self.SUBSCRIPTION_COMMANDS + if self.subscribed and not subscription_command: + raise PubSubError("Cannot issue commands other than SUBSCRIBE " + "and UNSUBSCRIBE while channels are open") + try: + connection.send_command(*args) + if subscription_command: + return None + return self.parse_response(connection, command_name, **options) + except ConnectionError: + connection.disconnect() + connection.send_command(*args) + if subscription_command: + return None + return self.parse_response(connection, command_name, **options) + finally: + self.connection_pool.release(connection) + + def parse_response(self, connection, command_name, **options): "Parses a response from the Redis server" - response = self.connection.read_response() + response = connection.read_response() if command_name in self.RESPONSE_CALLBACKS: return self.RESPONSE_CALLBACKS[command_name](response, **options) return response #### CONNECTION HANDLING #### - def get_connection(self, host, port, db, password, socket_timeout): - "Returns a connection object" - conn = self.connection_pool.get_connection( - host, port, db, password, socket_timeout) - # if for whatever reason the connection gets a bad password, make - # sure a subsequent attempt with the right password makes its way - # to the connection - conn.password = password - return conn - - def _setup_connection(self): - """ - After successfully opening a socket to the Redis server, the - connection object calls this method to authenticate and select - the appropriate database. - """ - self.subscribed = False - if self.connection.password: - if not self.execute_command('AUTH', self.connection.password): - raise AuthenticationError("Invalid Password") - self.execute_command('SELECT', self.connection.db) - - def 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 - """ - if host is None: - if self.connection is None: - raise RedisError("A valid hostname or IP address " - "must be specified") - host = self.connection.host - if port is None: - if self.connection is None: - raise RedisError("A valid port must be specified") - port = self.connection.port - - self.connection = self.get_connection( - host, port, db, password, socket_timeout) + def select(self, db): + "SELECT a differnet Redis database." + return self.execute_command('SELECT', db) def shutdown(self): "Shutdown the server" @@ -1246,11 +1209,9 @@ class Pipeline(Redis): ResponseError exceptions, such as those raised when issuing a command on a key of a different datatype. """ - def __init__(self, connection, transaction, charset, errors): - self.connection = connection + def __init__(self, connection_pool, transaction): + self.connection_pool = connection_pool self.transaction = transaction - self.encoding = charset - self.errors = errors self.subscribed = False # NOTE not in use, but necessary self.reset() @@ -1275,42 +1236,51 @@ class Pipeline(Redis): return self def _execute_transaction(self, commands): - all_cmds = ''.join(starmap(self.connection.pack_command, - [args for args, options in commands])) - self.connection.send_packed_command(all_cmds) - # we don't care about the multi/exec any longer - commands = commands[1:-1] - # parse off the response for MULTI and all commands prior to EXEC - # the only data we care about is the response the EXEC, the last command - for i in range(len(commands)+1): - _ = self.parse_response('_') - # parse the EXEC. - response = self.parse_response('_') - - if response is None: - raise WatchError("Watched variable changed.") - - if len(response) != len(commands): - raise ResponseError("Wrong number of response items from " - "pipeline execution") - # We have to run response callbacks manually - data = [] - for r, cmd in izip(response, commands): - if not isinstance(r, Exception): - args, options = cmd - command_name = args[0] - if command_name in self.RESPONSE_CALLBACKS: - r = self.RESPONSE_CALLBACKS[command_name](r, **options) - data.append(r) - return data + connection = self.connection_pool.get_connection() + try: + all_cmds = ''.join(starmap(connection.pack_command, + [args for args, options in commands])) + connection.send_packed_command(all_cmds) + # we don't care about the multi/exec any longer + commands = commands[1:-1] + # parse off the response for MULTI and all commands prior to EXEC. + # the only data we care about is the response the EXEC + # which is the last command + for i in range(len(commands)+1): + _ = self.parse_response(connection, '_') + # parse the EXEC. + response = self.parse_response(connection, '_') + + if response is None: + raise WatchError("Watched variable changed.") + + if len(response) != len(commands): + raise ResponseError("Wrong number of response items from " + "pipeline execution") + # We have to run response callbacks manually + data = [] + for r, cmd in izip(response, commands): + if not isinstance(r, Exception): + args, options = cmd + command_name = args[0] + if command_name in self.RESPONSE_CALLBACKS: + r = self.RESPONSE_CALLBACKS[command_name](r, **options) + data.append(r) + return data + finally: + self.connection_pool.release(connection) def _execute_pipeline(self, commands): # build up all commands into a single request to increase network perf - all_cmds = ''.join(starmap(self.connection.pack_command, - [args for args, options in commands])) - self.connection.send_packed_command(all_cmds) - return [self.parse_response(args[0], **options) - for args, options in commands] + connection = self.connection_pool.get_connection() + try: + all_cmds = ''.join(starmap(connection.pack_command, + [args for args, options in commands])) + connection.send_packed_command(all_cmds) + return [self.parse_response(connection, args[0], **options) + for args, options in commands] + finally: + self.connection_pool.release(connection) def execute(self): "Execute all the commands in the current pipeline" @@ -1324,7 +1294,7 @@ class Pipeline(Redis): try: return execute(stack) except ConnectionError: - self.connection.disconnect() + connection.disconnect() return execute(stack) def select(self, *args, **kwargs): diff --git a/redis/connection.py b/redis/connection.py index d5e4165..47c4f2f 100644 --- a/redis/connection.py +++ b/redis/connection.py @@ -116,9 +116,7 @@ class Connection(object): if self._sock: return try: - sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - sock.settimeout(self.socket_timeout) - sock.connect((self.host, self.port)) + sock = self._connect() except socket.error, e: # args for socket.error can either be (errno, "message") # or just "message" @@ -133,6 +131,13 @@ class Connection(object): self._sock = sock self.on_connect() + def _connect(self): + "Create a TCP socket connection" + sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + sock.settimeout(self.socket_timeout) + sock.connect((self.host, self.port)) + return sock + def on_connect(self): "Initialize the connection, authenticate and select a database" self._parser.on_connect(self) @@ -199,7 +204,7 @@ class Connection(object): return response def encode(self, value): - "Return a bytestring of the value" + "Return a bytestring representation of the value" if isinstance(value, unicode): return value.encode(self.encoding, self.encoding_errors) return str(value) @@ -210,24 +215,47 @@ class Connection(object): for enc_value in imap(self.encode, args)] return '*%s\r\n%s' % (len(command), ''.join(command)) -class ConnectionPool(threading.local): - "Manages a list of connections on the local thread" - def __init__(self, connection_class=None): - self.connections = {} - self.connection_class = connection_class or Connection - - def make_connection_key(self, host, port, db): - "Create a unique key for the specified host, port and db" - return '%s:%s:%s' % (host, port, db) - - def get_connection(self, host, port, db, password, socket_timeout): - "Return a specific connection for the specified host, port and db" - key = self.make_connection_key(host, port, db) - if key not in self.connections: - self.connections[key] = self.connection_class( - host, port, db, password, socket_timeout) - return self.connections[key] - - def get_all_connections(self): - "Return a list of all connection objects the manager knows about" - return self.connections.values() +class UnixDomainSocketConnection(Connection): + def __init__(self, path='', db=0, password=None, + socket_timeout=None, encoding='utf-8', + encoding_errors='strict', parser_class=DefaultParser): + self.path = path + self.db = db + self.password = password + self.socket_timeout = socket_timeout + self.encoding = encoding + self.encoding_errors = encoding_errors + self._sock = None + self._parser = parser_class() + + def _connect(self): + "Create a Unix domain socket connection" + sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) + sock.settimeout(self.socket_timeout) + sock.connect(self.path) + return sock + +class ConnectionPool(object): + """ + A connection pool that maintains only one connection. Great for + single-threaded apps with no sharding + """ + def __init__(self, connection_class=Connection, **kwargs): + self.connection_class = connection_class + self.kwargs = kwargs + self._connection = None + + def get_connection(self, *args, **kwargs): + "Get a connection from the pool" + if not self._connection: + self._connection = self.connection_class(**self.kwargs) + return self._connection + + def release(self, connection): + "Releases the connection back to the pool" + pass + + def disconnect(self): + "Disconnects all connections in the pool" + if self._connection: + self._connection.disconnect() diff --git a/tests/connection_pool.py b/tests/connection_pool.py index 56f5f43..f8f669e 100644 --- a/tests/connection_pool.py +++ b/tests/connection_pool.py @@ -4,29 +4,32 @@ import time import unittest class ConnectionPoolTestCase(unittest.TestCase): - def test_multiple_connections(self): - # 2 clients to the same host/port/db/pool should use the same connection - pool = redis.ConnectionPool() - r1 = redis.Redis(host='localhost', port=6379, db=9, connection_pool=pool) - r2 = redis.Redis(host='localhost', port=6379, db=9, connection_pool=pool) - self.assertEquals(r1.connection, r2.connection) - - # if one of them switches, they should have - # separate conncetion objects - r2.select(db=10, host='localhost', port=6379) - self.assertNotEqual(r1.connection, r2.connection) - - conns = [r1.connection, r2.connection] - conns.sort() - - # but returning to the original state shares the object again - r2.select(db=9, host='localhost', port=6379) - self.assertEquals(r1.connection, r2.connection) - - # the connection manager should still have just 2 connections - mgr_conns = pool.get_all_connections() - mgr_conns.sort() - self.assertEquals(conns, mgr_conns) + # TODO: + # THIS TEST IS INVALID WITH THE DEFAULT CONNECTIONPOOL + # + # def test_multiple_connections(self): + # # 2 clients to the same host/port/db/pool should use the same connection + # pool = redis.ConnectionPool() + # r1 = redis.Redis(host='localhost', port=6379, db=9, connection_pool=pool) + # r2 = redis.Redis(host='localhost', port=6379, db=9, connection_pool=pool) + # self.assertEquals(r1.connection, r2.connection) + + # # if one of them switches, they should have + # # separate conncetion objects + # r2.select(db=10, host='localhost', port=6379) + # self.assertNotEqual(r1.connection, r2.connection) + + # conns = [r1.connection, r2.connection] + # conns.sort() + + # # but returning to the original state shares the object again + # r2.select(db=9, host='localhost', port=6379) + # self.assertEquals(r1.connection, r2.connection) + + # # the connection manager should still have just 2 connections + # mgr_conns = pool.get_all_connections() + # mgr_conns.sort() + # self.assertEquals(conns, mgr_conns) def test_threaded_workers(self): r = redis.Redis(host='localhost', port=6379, db=9) diff --git a/tests/server_commands.py b/tests/server_commands.py index a68c4fb..1501441 100644 --- a/tests/server_commands.py +++ b/tests/server_commands.py @@ -16,8 +16,7 @@ class ServerCommandsTestCase(unittest.TestCase): def tearDown(self): self.client.flushdb() - for c in self.client.connection_pool.get_all_connections(): - c.disconnect() + self.client.connection_pool.disconnect() # GENERAL SERVER COMMANDS def test_dbsize(self): |
