summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorAndy McCurdy <andy@andymccurdy.com>2011-05-16 11:45:53 -0700
committerAndy McCurdy <andy@andymccurdy.com>2011-05-16 11:45:53 -0700
commitc650073a51c37e80815b6d5db901e7ae3c893411 (patch)
treecc3f7ec19a966b8ccaf71c605c0002ba9d503f80
parentf64c4ad7ed9a2d87a92415c3490edf8e9b757d84 (diff)
downloadredis-py-c650073a51c37e80815b6d5db901e7ae3c893411.tar.gz
all tests passing with new connection pool
-rw-r--r--redis/client.py200
-rw-r--r--redis/connection.py78
-rw-r--r--tests/connection_pool.py49
-rw-r--r--tests/server_commands.py3
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):