diff options
| author | Jordan Cook <jordan.cook@pioneer.com> | 2022-04-20 16:50:40 -0500 |
|---|---|---|
| committer | Jordan Cook <jordan.cook@pioneer.com> | 2022-04-22 17:28:33 -0500 |
| commit | 0dbd82d4d28875f2c0a592dfc89f50bf1c63cb2b (patch) | |
| tree | f3982671c81005c29c39fbd6da241e79509cc0cf /requests_cache/backends | |
| parent | 57579af3a5c4e683f2dd96f493471077808d1d39 (diff) | |
| download | requests-cache-0dbd82d4d28875f2c0a592dfc89f50bf1c63cb2b.tar.gz | |
Merge *PickleDict storage classes into parent classes
Diffstat (limited to 'requests_cache/backends')
| -rw-r--r-- | requests_cache/backends/__init__.py | 26 | ||||
| -rw-r--r-- | requests_cache/backends/base.py | 13 | ||||
| -rw-r--r-- | requests_cache/backends/dynamodb.py | 41 | ||||
| -rw-r--r-- | requests_cache/backends/filesystem.py | 4 | ||||
| -rw-r--r-- | requests_cache/backends/gridfs.py | 8 | ||||
| -rw-r--r-- | requests_cache/backends/mongodb.py | 39 | ||||
| -rw-r--r-- | requests_cache/backends/redis.py | 13 | ||||
| -rw-r--r-- | requests_cache/backends/sqlite.py | 44 |
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 |
