Merge pull request #96 from rpedroso/redis-session

lock in redis session
This commit is contained in:
mdipierro
2013-05-19 17:12:53 -07:00
+68 -30
View File
@@ -42,8 +42,10 @@ class RedisClient(object):
meta_storage = {} meta_storage = {}
MAX_RETRIES = 5 MAX_RETRIES = 5
RETRIES = 0 RETRIES = 0
_release_script = None
def __init__(self, server='localhost:6379', db=None, debug=False, session_expiry=False): def __init__(self, server='localhost:6379', db=None, debug=False,
session_expiry=False, with_lock=False):
"""session_expiry can be an integer, in seconds, to set the default expiration """session_expiry can be an integer, in seconds, to set the default expiration
of sessions. The corresponding record will be deleted from the redis instance, of sessions. The corresponding record will be deleted from the redis instance,
and there's virtually no need to run sessions2trash.py and there's virtually no need to run sessions2trash.py
@@ -58,8 +60,12 @@ class RedisClient(object):
else: else:
self.app = '' self.app = ''
self.r_server = redis.Redis(host=host, port=port, db=self.db) self.r_server = redis.Redis(host=host, port=port, db=self.db)
if with_lock:
RedisClient._release_script = \
self.r_server.register_script(_LUA_RELEASE_LOCK)
self.tablename = None self.tablename = None
self.session_expiry = session_expiry self.session_expiry = session_expiry
self.with_lock = with_lock
def get(self, what, default): def get(self, what, default):
return self.tablename return self.tablename
@@ -71,7 +77,8 @@ class RedisClient(object):
def define_table(self, tablename, *fields, **args): def define_table(self, tablename, *fields, **args):
if not self.tablename: if not self.tablename:
self.tablename = MockTable( self.tablename = MockTable(
self, self.r_server, tablename, self.session_expiry) self, self.r_server, tablename, self.session_expiry,
self.with_lock)
return self.tablename return self.tablename
def __getitem__(self, key): def __getitem__(self, key):
@@ -88,7 +95,7 @@ class RedisClient(object):
class MockTable(object): class MockTable(object):
def __init__(self, db, r_server, tablename, session_expiry): def __init__(self, db, r_server, tablename, session_expiry, with_lock=False):
self.db = db self.db = db
self.r_server = r_server self.r_server = r_server
self.tablename = tablename self.tablename = tablename
@@ -101,15 +108,14 @@ class MockTable(object):
self.id_idx = "%s:id_idx" % self.keyprefix self.id_idx = "%s:id_idx" % self.keyprefix
#remember the session_expiry setting #remember the session_expiry setting
self.session_expiry = session_expiry self.session_expiry = session_expiry
self.with_lock = with_lock
def getserial(self):
#return an auto-increment id
return "%s" % self.r_server.incr(self.serial, 1)
def __getattr__(self, key): def __getattr__(self, key):
if key == 'id': if key == 'id':
#return a fake query. We need to query it just by id for normal operations #return a fake query. We need to query it just by id for normal operations
self.query = MockQuery(field='id', db=self.r_server, prefix=self.keyprefix, session_expiry=self.session_expiry) self.query = MockQuery(field='id', db=self.r_server,
prefix=self.keyprefix, session_expiry=self.session_expiry,
with_lock=self.with_lock)
return self.query return self.query
elif key == '_db': elif key == '_db':
#needed because of the calls in sessions2trash.py and globals.py #needed because of the calls in sessions2trash.py and globals.py
@@ -120,14 +126,21 @@ class MockTable(object):
#'locked', 'client_ip','created_datetime','modified_datetime' #'locked', 'client_ip','created_datetime','modified_datetime'
#'unique_key', 'session_data' #'unique_key', 'session_data'
#retrieve a new key #retrieve a new key
newid = self.getserial() newid = str(self.r_server.incr(self.serial))
key = "%s:%s" % (self.keyprefix, newid) key = self.keyprefix + ':' + newid
#add it to the index if self.with_lock:
self.r_server.sadd(self.id_idx, key) key_lock = key + ':lock'
#set a hash key with the Storage acquire_lock(self.r_server, key_lock, newid)
self.r_server.hmset(key, kwargs) with self.r_server.pipeline() as pipe:
if self.session_expiry: #add it to the index
self.r_server.expire(key, self.session_expiry) pipe.sadd(self.id_idx, key)
#set a hash key with the Storage
pipe.hmset(key, kwargs)
if self.session_expiry:
pipe.expire(key, self.session_expiry)
pipe.execute()
if self.with_lock:
release_lock(self.r_server, key_lock, newid)
return newid return newid
@@ -135,13 +148,15 @@ class MockQuery(object):
"""a fake Query object that supports querying by id """a fake Query object that supports querying by id
and listing all keys. No other operation is supported and listing all keys. No other operation is supported
""" """
def __init__(self, field=None, db=None, prefix=None, session_expiry=False): def __init__(self, field=None, db=None, prefix=None, session_expiry=False,
with_lock=False):
self.field = field self.field = field
self.value = None self.value = None
self.db = db self.db = db
self.keyprefix = prefix self.keyprefix = prefix
self.op = None self.op = None
self.session_expiry = session_expiry self.session_expiry = session_expiry
self.with_lock = with_lock
def __eq__(self, value, op='eq'): def __eq__(self, value, op='eq'):
self.value = value self.value = value
@@ -154,12 +169,11 @@ class MockQuery(object):
def select(self): def select(self):
if self.op == 'eq' and self.field == 'id' and self.value: if self.op == 'eq' and self.field == 'id' and self.value:
#means that someone wants to retrieve the key self.value #means that someone wants to retrieve the key self.value
rtn = self.db.hgetall("%s:%s" % (self.keyprefix, self.value)) key = self.keyprefix + ':' + self.value
if rtn == dict(): if self.with_lock:
#return an empty resultset for non existing key acquire_lock(self.db, key + ':lock', self.value)
return [] rtn = self.db.hgetall(key)
else: return [Storage(rtn)] if rtn else []
return [Storage(rtn)]
elif self.op == 'ge' and self.field == 'id' and self.value == 0: elif self.op == 'ge' and self.field == 'id' and self.value == 0:
#means that someone wants the complete list #means that someone wants the complete list
rtn = [] rtn = []
@@ -168,13 +182,11 @@ class MockQuery(object):
allkeys = self.db.smembers(id_idx) allkeys = self.db.smembers(id_idx)
for sess in allkeys: for sess in allkeys:
val = self.db.hgetall(sess) val = self.db.hgetall(sess)
if val == dict(): if not val:
if self.session_expiry: if self.session_expiry:
#clean up the idx, because the key expired #clean up the idx, because the key expired
self.db.srem(id_idx, sess) self.db.srem(id_idx, sess)
continue continue
else:
continue
val = Storage(val) val = Storage(val)
#add a delete_record method (necessary for sessions2trash.py) #add a delete_record method (necessary for sessions2trash.py)
val.delete_record = RecordDeleter( val.delete_record = RecordDeleter(
@@ -188,9 +200,13 @@ class MockQuery(object):
#means that the session has been found and needs an update #means that the session has been found and needs an update
if self.op == 'eq' and self.field == 'id' and self.value: if self.op == 'eq' and self.field == 'id' and self.value:
key = "%s:%s" % (self.keyprefix, self.value) key = "%s:%s" % (self.keyprefix, self.value)
rtn = self.db.hmset(key, kwargs) with self.db.pipeline() as pipe:
if self.session_expiry: pipe.hmset(key, kwargs)
self.db.expire(key, self.session_expiry) if self.session_expiry:
pipe.expire(key, self.session_expiry)
rtn = pipe.execute()[0]
if self.with_lock:
release_lock(self.db, key + ':lock', self.value)
return rtn return rtn
@@ -206,3 +222,25 @@ class RecordDeleter(object):
self.db.srem(id_idx, self.key) self.db.srem(id_idx, self.key)
#remove the key itself #remove the key itself
self.db.delete(self.key) self.db.delete(self.key)
def acquire_lock(conn, lockname, identifier, ltime=10):
while True:
if conn.set(lockname, identifier, ex=ltime, nx=True):
return identifier
time.sleep(.01)
_LUA_RELEASE_LOCK = """
if redis.call("get", KEYS[1]) == ARGV[1]
then
return redis.call("del", KEYS[1])
else
return 0
end
"""
def release_lock(conn, lockname, identifier):
return RedisClient._release_script(keys=[lockname], args=[identifier],
client=conn)