Merge pull request #864 from niphlod/fix/redis_sessions

fixes issue with deleting sessions
This commit is contained in:
mdipierro
2015-03-22 00:36:47 -05:00
+49 -32
View File
@@ -59,8 +59,7 @@ class RedisClient(object):
self.app = '' self.app = ''
self.r_server = redis.Redis(host=host, port=port, db=self.db, password=self.password) self.r_server = redis.Redis(host=host, port=port, db=self.db, password=self.password)
if with_lock: if with_lock:
RedisClient._release_script = \ RedisClient._release_script = self.r_server.register_script(_LUA_RELEASE_LOCK)
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 self.with_lock = with_lock
@@ -87,7 +86,7 @@ class RedisClient(object):
return q return q
def commit(self): def commit(self):
#this is only called by session2trash.py # this is only called by session2trash.py
pass pass
@@ -97,22 +96,23 @@ class MockTable(object):
self.db = db self.db = db
self.r_server = r_server self.r_server = r_server
self.tablename = tablename self.tablename = tablename
#set the namespace for sessions of this app # set the namespace for sessions of this app
self.keyprefix = 'w2p:sess:%s' % tablename.replace( self.keyprefix = 'w2p:sess:%s' % tablename.replace(
'web2py_session_', '') 'web2py_session_', '')
#fast auto-increment id (needed for session handling) # fast auto-increment id (needed for session handling)
self.serial = "%s:serial" % self.keyprefix self.serial = "%s:serial" % self.keyprefix
#index of all the session keys of this app # index of all the session keys of this app
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 self.with_lock = with_lock
def __call__(self, record_id, unique_key=None): def __call__(self, record_id, unique_key=None):
# Support DAL shortcut query: table(record_id) # Support DAL shortcut query: table(record_id)
q = self.id # This will call the __getattr__ below # This will call the __getattr__ below
# returning a MockQuery # returning a MockQuery
q = self.id
# Instructs MockQuery, to behave as db(table.id == record_id) # Instructs MockQuery, to behave as db(table.id == record_id)
q.op = 'eq' q.op = 'eq'
@@ -124,29 +124,31 @@ class MockTable(object):
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, self.query = MockQuery(
prefix=self.keyprefix, session_expiry=self.session_expiry, field='id', db=self.r_server,
with_lock=self.with_lock, unique_key=self.unique_key) prefix=self.keyprefix, session_expiry=self.session_expiry,
with_lock=self.with_lock, unique_key=self.unique_key
)
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
return self.db return self.db
def insert(self, **kwargs): def insert(self, **kwargs):
#usually kwargs would be a Storage with several keys: # usually kwargs would be a Storage with several keys:
#'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 = str(self.r_server.incr(self.serial)) newid = str(self.r_server.incr(self.serial))
key = self.keyprefix + ':' + newid key = self.keyprefix + ':' + newid
if self.with_lock: if self.with_lock:
key_lock = key + ':lock' key_lock = key + ':lock'
acquire_lock(self.r_server, key_lock, newid) acquire_lock(self.r_server, key_lock, newid)
with self.r_server.pipeline() as pipe: with self.r_server.pipeline() as pipe:
#add it to the index # add it to the index
pipe.sadd(self.id_idx, key) pipe.sadd(self.id_idx, key)
#set a hash key with the Storage # set a hash key with the Storage
pipe.hmset(key, kwargs) pipe.hmset(key, kwargs)
if self.session_expiry: if self.session_expiry:
pipe.expire(key, self.session_expiry) pipe.expire(key, self.session_expiry)
@@ -155,12 +157,13 @@ class MockTable(object):
release_lock(self.r_server, key_lock, newid) release_lock(self.r_server, key_lock, newid)
return newid return newid
class MockQuery(object): 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, unique_key=None): with_lock=False, unique_key=None):
self.field = field self.field = field
self.value = None self.value = None
self.db = db self.db = db
@@ -180,34 +183,34 @@ 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
key = self.keyprefix + ':' + str(self.value) key = self.keyprefix + ':' + str(self.value)
if self.with_lock: if self.with_lock:
acquire_lock(self.db, key + ':lock', self.value) acquire_lock(self.db, key + ':lock', self.value)
rtn = self.db.hgetall(key) rtn = self.db.hgetall(key)
if rtn: if rtn:
if self.unique_key: if self.unique_key:
#make sure the id and unique_key are correct # make sure the id and unique_key are correct
if rtn['unique_key'] == self.unique_key: if rtn['unique_key'] == self.unique_key:
rtn['update_record'] = self.update # update record support rtn['update_record'] = self.update # update record support
else: else:
rtn = None rtn = None
return [Storage(rtn)] if rtn else [] return [Storage(rtn)] if rtn else []
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 = []
id_idx = "%s:id_idx" % self.keyprefix id_idx = "%s:id_idx" % self.keyprefix
#find all session keys of this app # find all session keys of this app
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 not val: 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
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(
self.db, sess, self.keyprefix) self.db, sess, self.keyprefix)
rtn.append(val) rtn.append(val)
@@ -216,9 +219,11 @@ class MockQuery(object):
raise Exception("Operation not supported") raise Exception("Operation not supported")
def update(self, **kwargs): def update(self, **kwargs):
#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 = self.keyprefix + ':' + str(self.value) key = self.keyprefix + ':' + str(self.value)
if not self.db.exists(key):
return None
with self.db.pipeline() as pipe: with self.db.pipeline() as pipe:
pipe.hmset(key, kwargs) pipe.hmset(key, kwargs)
if self.session_expiry: if self.session_expiry:
@@ -228,6 +233,17 @@ class MockQuery(object):
release_lock(self.db, key + ':lock', self.value) release_lock(self.db, key + ':lock', self.value)
return rtn return rtn
def delete(self, **kwargs):
# means that we want this session to be deleted
if self.op == 'eq' and self.field == 'id' and self.value:
id_idx = "%s:id_idx" % self.keyprefix
key = self.keyprefix + ':' + str(self.value)
with self.db.pipeline() as pipe:
pipe.delete(key)
pipe.srem(id_idx, key)
rtn = pipe.execute()
return rtn[1]
class RecordDeleter(object): class RecordDeleter(object):
"""Dumb record deleter to support sessions2trash.py""" """Dumb record deleter to support sessions2trash.py"""
@@ -237,9 +253,9 @@ class RecordDeleter(object):
def __call__(self): def __call__(self):
id_idx = "%s:id_idx" % self.keyprefix id_idx = "%s:id_idx" % self.keyprefix
#remove from the index # remove from the index
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)
@@ -261,5 +277,6 @@ end
def release_lock(conn, lockname, identifier): def release_lock(conn, lockname, identifier):
return RedisClient._release_script(keys=[lockname], args=[identifier], return RedisClient._release_script(
client=conn) keys=[lockname], args=[identifier],
client=conn)