summaryrefslogtreecommitdiff
path: root/requests_cache/backends
diff options
context:
space:
mode:
authorJordan Cook <jordan.cook@pioneer.com>2022-04-20 16:50:40 -0500
committerJordan Cook <jordan.cook@pioneer.com>2022-04-22 17:28:33 -0500
commit0dbd82d4d28875f2c0a592dfc89f50bf1c63cb2b (patch)
treef3982671c81005c29c39fbd6da241e79509cc0cf /requests_cache/backends
parent57579af3a5c4e683f2dd96f493471077808d1d39 (diff)
downloadrequests-cache-0dbd82d4d28875f2c0a592dfc89f50bf1c63cb2b.tar.gz
Merge *PickleDict storage classes into parent classes
Diffstat (limited to 'requests_cache/backends')
-rw-r--r--requests_cache/backends/__init__.py26
-rw-r--r--requests_cache/backends/base.py13
-rw-r--r--requests_cache/backends/dynamodb.py41
-rw-r--r--requests_cache/backends/filesystem.py4
-rw-r--r--requests_cache/backends/gridfs.py8
-rw-r--r--requests_cache/backends/mongodb.py39
-rw-r--r--requests_cache/backends/redis.py13
-rw-r--r--requests_cache/backends/sqlite.py44
8 files changed, 80 insertions, 108 deletions
diff --git a/requests_cache/backends/__init__.py b/requests_cache/backends/__init__.py
index 7695b8f..9dd206a 100644
--- a/requests_cache/backends/__init__.py
+++ b/requests_cache/backends/__init__.py
@@ -15,35 +15,35 @@ logger = getLogger(__name__)
# Import all backend classes for which dependencies are installed
try:
- from .dynamodb import DynamoDbCache, DynamoDbDict, DynamoDbDocumentDict
+ from .dynamodb import DynamoDbCache, DynamoDbDict
except ImportError as e:
- DynamoDbCache = DynamoDbDict = DynamoDbDocumentDict = get_placeholder_class(e) # type: ignore
+ DynamoDbCache = DynamoDbDict = get_placeholder_class(e) # type: ignore
+
try:
- from .gridfs import GridFSCache, GridFSPickleDict
+ from .gridfs import GridFSCache, GridFSDict
except ImportError as e:
- GridFSCache = GridFSPickleDict = get_placeholder_class(e) # type: ignore
+ GridFSCache = GridFSDict = get_placeholder_class(e) # type: ignore
+
try:
- from .mongodb import MongoCache, MongoDict, MongoDocumentDict
+ from .mongodb import MongoCache, MongoDict
except ImportError as e:
- MongoCache = MongoDict = MongoDocumentDict = get_placeholder_class(e) # type: ignore
+ MongoCache = MongoDict = get_placeholder_class(e) # type: ignore
+
try:
from .redis import RedisCache, RedisDict, RedisHashDict
except ImportError as e:
RedisCache = RedisDict = RedisHashDict = get_placeholder_class(e) # type: ignore
+
try:
- # Note: Heroku doesn't support SQLite due to ephemeral storage
- from .sqlite import SQLiteCache, SQLiteDict, SQLitePickleDict
+ from .sqlite import SQLiteCache, SQLiteDict
except ImportError as e:
- SQLiteCache = SQLiteDict = SQLitePickleDict = get_placeholder_class(e) # type: ignore
+ SQLiteCache = SQLiteDict = get_placeholder_class(e) # type: ignore
+
try:
from .filesystem import FileCache, FileDict
except ImportError as e:
FileCache = FileDict = get_placeholder_class(e) # type: ignore
-# Aliases for backwards-compatibility
-DbCache = SQLiteCache
-DbDict = SQLiteDict
-DbPickleDict = SQLitePickleDict
BACKEND_CLASSES = {
'dynamodb': DynamoDbCache,
diff --git a/requests_cache/backends/base.py b/requests_cache/backends/base.py
index 0a551ea..e837671 100644
--- a/requests_cache/backends/base.py
+++ b/requests_cache/backends/base.py
@@ -263,8 +263,8 @@ class BaseStorage(MutableMapping, ABC):
kwargs: Additional backend-specific keyword arguments
"""
- def __init__(self, serializer=None, **kwargs):
- self.serializer = init_serializer(serializer)
+ def __init__(self, **kwargs):
+ self.serializer = init_serializer(kwargs.get('serializer', 'pickle'))
logger.debug(f'Initializing {type(self).__name__} with serializer: {self.serializer}')
def bulk_delete(self, keys: Iterable[str]):
@@ -281,6 +281,14 @@ class BaseStorage(MutableMapping, ABC):
def close(self):
"""Close any open backend connections"""
+ def serialize(self, value):
+ """Serialize value, if a serializer is available"""
+ return self.serializer.dumps(value) if self.serializer else value
+
+ def deserialize(self, value):
+ """Deserialize value, if a serializer is available"""
+ return self.serializer.loads(value) if self.serializer else value
+
def __str__(self):
return str(list(self.keys()))
@@ -297,7 +305,6 @@ class DictStorage(UserDict, BaseStorage):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
- self._serializer = None
self.serializer = None
def __getitem__(self, key):
diff --git a/requests_cache/backends/dynamodb.py b/requests_cache/backends/dynamodb.py
index d6a5fdb..f39d267 100644
--- a/requests_cache/backends/dynamodb.py
+++ b/requests_cache/backends/dynamodb.py
@@ -19,6 +19,7 @@ from . import BaseCache, BaseStorage
class DynamoDbCache(BaseCache):
"""DynamoDB cache backend.
+ By default, responses are only partially serialized into a DynamoDB-compatible document format.
Args:
table_name: DynamoDB table name
@@ -34,14 +35,25 @@ class DynamoDbCache(BaseCache):
table_name: str = 'http_cache',
ttl: bool = True,
connection: ServiceResource = None,
+ serializer=None,
**kwargs,
):
super().__init__(cache_name=table_name, **kwargs)
- self.responses = DynamoDbDocumentDict(
- table_name, 'responses', ttl=ttl, connection=connection, **kwargs
+ self.responses = DynamoDbDict(
+ table_name,
+ 'responses',
+ ttl=ttl,
+ serializer=serializer or dynamodb_document_serializer,
+ connection=connection,
+ **kwargs,
)
self.redirects = DynamoDbDict(
- table_name, 'redirects', ttl=False, connection=self.responses.connection, **kwargs
+ table_name,
+ 'redirects',
+ ttl=False,
+ connection=self.responses.connection,
+ serializer=None,
+ **kwargs,
)
@@ -132,14 +144,15 @@ class DynamoDbDict(BaseStorage):
# With a custom serializer, the value may be a Binary object
raw_value = result['Item']['value']
- return raw_value.value if isinstance(raw_value, Binary) else raw_value
+ value = raw_value.value if isinstance(raw_value, Binary) else raw_value
+ return self.deserialize(value)
def __setitem__(self, key, value):
- item = {**self._composite_key(key), 'value': value}
+ item = {**self._composite_key(key), 'value': self.serialize(value)}
# If enabled, set TTL value as a timestamp in unix format
if self.ttl and getattr(value, 'ttl', None):
- item['ttl'] = int(time() + value.ttl)
+ item['ttl'] = round(time() + value.ttl)
self._table.put_item(Item=item)
@@ -169,19 +182,3 @@ class DynamoDbDict(BaseStorage):
def clear(self):
self.bulk_delete((k for k in self))
-
-
-class DynamoDbDocumentDict(DynamoDbDict):
- """Same as :class:`DynamoDbDict`, but serializes values before saving.
-
- By default, responses are only partially serialized into a DynamoDB-compatible document format.
- """
-
- def __init__(self, *args, serializer=None, **kwargs):
- super().__init__(*args, serializer=serializer or dynamodb_document_serializer, **kwargs)
-
- def __getitem__(self, key):
- return self.serializer.loads(super().__getitem__(key))
-
- def __setitem__(self, key, item):
- super().__setitem__(key, self.serializer.dumps(item))
diff --git a/requests_cache/backends/filesystem.py b/requests_cache/backends/filesystem.py
index 834b0ab..c1cb77e 100644
--- a/requests_cache/backends/filesystem.py
+++ b/requests_cache/backends/filesystem.py
@@ -91,7 +91,7 @@ class FileDict(BaseStorage):
mode = 'rb' if self.is_binary else 'r'
with self._try_io():
with self._path(key).open(mode) as f:
- return self.serializer.loads(f.read())
+ return self.deserialize(f.read())
def __delitem__(self, key):
with self._try_io():
@@ -100,7 +100,7 @@ class FileDict(BaseStorage):
def __setitem__(self, key, value):
with self._try_io():
with self._path(key).open(mode='wb' if self.is_binary else 'w') as f:
- f.write(self.serializer.dumps(value))
+ f.write(self.serialize(value))
def __iter__(self):
yield from self.keys()
diff --git a/requests_cache/backends/gridfs.py b/requests_cache/backends/gridfs.py
index eb139d5..dc66323 100644
--- a/requests_cache/backends/gridfs.py
+++ b/requests_cache/backends/gridfs.py
@@ -29,7 +29,7 @@ class GridFSCache(BaseCache):
def __init__(self, db_name: str, **kwargs):
super().__init__(cache_name=db_name, **kwargs)
- self.responses = GridFSPickleDict(db_name, **kwargs)
+ self.responses = GridFSDict(db_name, **kwargs)
self.redirects = MongoDict(
db_name, collection_name='redirects', connection=self.responses.connection, **kwargs
)
@@ -39,7 +39,7 @@ class GridFSCache(BaseCache):
return super().remove_expired_responses(*args, **kwargs)
-class GridFSPickleDict(BaseStorage):
+class GridFSDict(BaseStorage):
"""A dictionary-like interface for a GridFS database
Args:
@@ -63,13 +63,13 @@ class GridFSPickleDict(BaseStorage):
result = self.fs.find_one({'_id': key})
if result is None:
raise KeyError
- return self.serializer.loads(result.read())
+ return self.deserialize(result.read())
except CorruptGridFile as e:
logger.warning(e, exc_info=True)
raise KeyError
def __setitem__(self, key, item):
- value = self.serializer.dumps(item)
+ value = self.serialize(item)
encoding = None if isinstance(value, bytes) else 'utf-8'
with self._lock:
diff --git a/requests_cache/backends/mongodb.py b/requests_cache/backends/mongodb.py
index d0fe79c..9c881be 100644
--- a/requests_cache/backends/mongodb.py
+++ b/requests_cache/backends/mongodb.py
@@ -21,6 +21,7 @@ logger = getLogger(__name__)
class MongoCache(BaseCache):
"""MongoDB cache backend.
+ By default, responses are only partially serialized into a MongoDB-compatible document format.
Args:
db_name: Database name
@@ -28,18 +29,22 @@ class MongoCache(BaseCache):
kwargs: Additional keyword arguments for :py:class:`pymongo.mongo_client.MongoClient`
"""
- def __init__(self, db_name: str = 'http_cache', connection: MongoClient = None, **kwargs):
+ def __init__(
+ self, db_name: str = 'http_cache', connection: MongoClient = None, serializer=None, **kwargs
+ ):
super().__init__(cache_name=db_name, **kwargs)
- self.responses: MongoDict = MongoDocumentDict(
+ self.responses: MongoDict = MongoDict(
db_name,
collection_name='responses',
connection=connection,
+ serializer=serializer or bson_document_serializer,
**kwargs,
)
self.redirects: MongoDict = MongoDict(
db_name,
collection_name='redirects',
connection=self.responses.connection,
+ serializer=None,
**kwargs,
)
@@ -107,15 +112,17 @@ class MongoDict(BaseStorage):
result = self.collection.find_one({'_id': key})
if result is None:
raise KeyError
- return result['data'] if 'data' in result else result
+ value = result['data'] if 'data' in result else result
+ return self.deserialize(value)
- def __setitem__(self, key, item):
- """If ``item`` is already a dict, its values will be stored under top-level keys.
+ def __setitem__(self, key, value):
+ """If ``value`` is already a dict, its values will be stored under top-level keys.
Otherwise, it will be stored under a 'data' key.
"""
- if not isinstance(item, Mapping):
- item = {'data': item}
- self.collection.replace_one({'_id': key}, item, upsert=True)
+ value = self.serialize(value)
+ if not isinstance(value, Mapping):
+ value = {'data': value}
+ self.collection.replace_one({'_id': key}, value, upsert=True)
def __delitem__(self, key):
result = self.collection.find_one_and_delete({'_id': key}, {'_id': True})
@@ -138,19 +145,3 @@ class MongoDict(BaseStorage):
def close(self):
self.connection.close()
-
-
-class MongoDocumentDict(MongoDict):
- """Same as :class:`MongoDict`, but serializes values before saving.
-
- By default, responses are only partially serialized into a MongoDB-compatible document format.
- """
-
- def __init__(self, *args, serializer=None, **kwargs):
- super().__init__(*args, serializer=serializer or bson_document_serializer, **kwargs)
-
- def __getitem__(self, key):
- return self.serializer.loads(super().__getitem__(key))
-
- def __setitem__(self, key, item):
- super().__setitem__(key, self.serializer.dumps(item))
diff --git a/requests_cache/backends/redis.py b/requests_cache/backends/redis.py
index 0697023..76c83c1 100644
--- a/requests_cache/backends/redis.py
+++ b/requests_cache/backends/redis.py
@@ -75,14 +75,15 @@ class RedisDict(BaseStorage):
result = self.connection.get(self._bkey(key))
if result is None:
raise KeyError
- return self.serializer.loads(result)
+ return self.deserialize(result)
def __setitem__(self, key, item):
"""Save an item to the cache, optionally with TTL"""
- if self.ttl and getattr(item, 'ttl', None):
- self.connection.setex(self._bkey(key), item.ttl, self.serializer.dumps(item))
+ ttl_seconds = getattr(item, 'ttl', None)
+ if self.ttl and ttl_seconds and ttl_seconds > 0:
+ self.connection.setex(self._bkey(key), round(ttl_seconds), self.serialize(item))
else:
- self.connection.set(self._bkey(key), self.serializer.dumps(item))
+ self.connection.set(self._bkey(key), self.serialize(item))
def __delitem__(self, key):
if not self.connection.delete(self._bkey(key)):
@@ -141,10 +142,10 @@ class RedisHashDict(BaseStorage):
result = self.connection.hget(self._hash_key, encode(key))
if result is None:
raise KeyError
- return self.serializer.loads(result)
+ return self.deserialize(result)
def __setitem__(self, key, item):
- self.connection.hset(self._hash_key, encode(key), self.serializer.dumps(item))
+ self.connection.hset(self._hash_key, encode(key), self.serialize(item))
def __delitem__(self, key):
if not self.connection.hdel(self._hash_key, encode(key)):
diff --git a/requests_cache/backends/sqlite.py b/requests_cache/backends/sqlite.py
index eb9b712..cb15697 100644
--- a/requests_cache/backends/sqlite.py
+++ b/requests_cache/backends/sqlite.py
@@ -43,10 +43,14 @@ class SQLiteCache(BaseCache):
kwargs: Additional keyword arguments for :py:func:`sqlite3.connect`
"""
- def __init__(self, db_path: AnyPath = 'http_cache', **kwargs):
+ def __init__(self, db_path: AnyPath = 'http_cache', serializer=None, **kwargs):
super().__init__(cache_name=str(db_path), **kwargs)
- self.responses: SQLiteDict = SQLitePickleDict(db_path, table_name='responses', **kwargs)
- self.redirects: SQLiteDict = SQLiteDict(db_path, table_name='redirects', **kwargs)
+ self.responses: SQLiteDict = SQLiteDict(
+ db_path, table_name='responses', serializer=serializer or 'pickle', **kwargs
+ )
+ self.redirects: SQLiteDict = SQLiteDict(
+ db_path, table_name='redirects', serializer=None, **kwargs
+ )
@property
def db_path(self) -> AnyPath:
@@ -211,7 +215,8 @@ class SQLiteDict(BaseStorage):
# raise error after the with block, otherwise the connection will be locked
if not row:
raise KeyError
- return row[0]
+
+ return self.deserialize(row[0])
def __setitem__(self, key, value):
self._insert(key, value)
@@ -288,36 +293,13 @@ class SQLiteDict(BaseStorage):
f' ORDER BY {key} {direction} {limit_expr}',
params,
):
- yield row[0]
+ yield self.deserialize(row[0])
def vacuum(self):
with self.connection(commit=True) as con:
con.execute('VACUUM')
-class SQLitePickleDict(SQLiteDict):
- """Same as :class:`SQLiteDict`, but serializes values before saving"""
-
- def __setitem__(self, key, value: CachedResponse):
- serialized_value = self.serializer.dumps(value)
- if isinstance(serialized_value, bytes):
- serialized_value = sqlite3.Binary(serialized_value)
- super()._insert(key, serialized_value, getattr(value, 'expires', None))
-
- def __getitem__(self, key):
- return self.serializer.loads(super().__getitem__(key))
-
- def sorted(
- self,
- key: str = 'expires',
- reversed: bool = False,
- limit: int = None,
- exclude_expired: bool = False,
- ):
- for value in super().sorted(key, reversed, limit, exclude_expired):
- yield self.serializer.loads(value)
-
-
def _format_sequence(values: Collection) -> Tuple[str, List]:
"""Get SQL parameter marks for a sequence-based query"""
return ','.join(['?'] * len(values)), list(values)
@@ -372,9 +354,3 @@ def sqlite_template(
uri: bool = False,
):
"""Template function to get an accurate signature for the builtin :py:func:`sqlite3.connect`"""
-
-
-# Aliases for backwards-compatibility
-DbCache = SQLiteCache
-DbDict = SQLiteDict
-DbPickeDict = SQLitePickleDict