diff options
| author | Jordan Cook <jordan.cook@pioneer.com> | 2022-04-01 16:29:13 -0500 |
|---|---|---|
| committer | Jordan Cook <jordan.cook@pioneer.com> | 2022-04-01 17:29:22 -0500 |
| commit | 0d2d9c690a787f8894bb81fec25d65a4b774ad43 (patch) | |
| tree | 4154088435cdbe152e271974215479cf50fff02c /requests_cache | |
| parent | 026b627c63124c885ff734d3b30a15464f9b0c93 (diff) | |
| download | requests-cache-0d2d9c690a787f8894bb81fec25d65a4b774ad43.tar.gz | |
Add an intermediate wrapper class, OriginalResponse, to provide type hints for extra attributes set on requests.Response objects
Diffstat (limited to 'requests_cache')
| -rw-r--r-- | requests_cache/backends/base.py | 8 | ||||
| -rw-r--r-- | requests_cache/models/__init__.py | 4 | ||||
| -rwxr-xr-x | requests_cache/models/response.py | 59 | ||||
| -rw-r--r-- | requests_cache/session.py | 31 | ||||
| -rw-r--r-- | requests_cache/settings.py | 8 |
5 files changed, 78 insertions, 32 deletions
diff --git a/requests_cache/backends/base.py b/requests_cache/backends/base.py index 3cc9413..f85c175 100644 --- a/requests_cache/backends/base.py +++ b/requests_cache/backends/base.py @@ -12,9 +12,11 @@ from datetime import datetime from logging import getLogger from typing import Iterable, Iterator, Optional, Tuple, Union +from requests import PreparedRequest, Response + from ..cache_keys import create_key, redact_response from ..expiration import ExpirationTime -from ..models import AnyRequest, AnyResponse, CachedResponse +from ..models import CachedResponse from ..serializers import init_serializer from ..settings import DEFAULT_CACHE_NAME, CacheSettings @@ -77,7 +79,7 @@ class BaseCache: logger.debug(e, exc_info=True) return default - def save_response(self, response: AnyResponse, cache_key: str = None, expires: datetime = None): + def save_response(self, response: Response, cache_key: str = None, expires: datetime = None): """Save a response to the cache Args: @@ -105,7 +107,7 @@ class BaseCache: self.responses.clear() self.redirects.clear() - def create_key(self, request: AnyRequest = None, **kwargs) -> str: + def create_key(self, request: PreparedRequest = None, **kwargs) -> str: """Create a normalized cache key from a request object""" key_fn = self._settings.key_fn or create_key return key_fn( diff --git a/requests_cache/models/__init__.py b/requests_cache/models/__init__.py index 6ffc7ad..1824a6c 100644 --- a/requests_cache/models/__init__.py +++ b/requests_cache/models/__init__.py @@ -6,8 +6,8 @@ from requests import PreparedRequest, Request, Response from .raw_response import CachedHTTPResponse from .request import CachedRequest -from .response import CachedResponse, set_response_defaults +from .response import CachedResponse, OriginalResponse -AnyResponse = Union[Response, CachedResponse] +AnyResponse = Union[OriginalResponse, CachedResponse] AnyRequest = Union[Request, PreparedRequest, CachedRequest] AnyPreparedRequest = Union[PreparedRequest, CachedRequest] diff --git a/requests_cache/models/response.py b/requests_cache/models/response.py index 4ac24ce..b6ba3ae 100755 --- a/requests_cache/models/response.py +++ b/requests_cache/models/response.py @@ -1,6 +1,6 @@ from datetime import datetime, timedelta, timezone from logging import getLogger -from typing import TYPE_CHECKING, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, List, Optional, Tuple import attr from attr import define, field @@ -12,17 +12,53 @@ from urllib3._collections import HTTPHeaderDict from ..expiration import ExpirationTime, get_expiration_datetime from . import CachedHTTPResponse, CachedRequest +if TYPE_CHECKING: + from ..cache_control import CacheActions + DATETIME_FORMAT = '%Y-%m-%d %H:%M:%S %Z' # Format used for __str__ only HeaderList = List[Tuple[str, str]] logger = getLogger(__name__) @define(auto_attribs=False, slots=False) -class CachedResponse(Response): - """A class that emulates :py:class:`requests.Response`, with some additional optimizations - for serialization. +class BaseResponse(Response): + """Wrapper class for responses returned by :py:class:`.CachedSession`. This mainly exists to + provide type hints for extra cache-related attributes that are added to non-cached responses. """ + cache_key: Optional[str] = None + created_at: datetime = field(factory=datetime.utcnow) + expires: Optional[datetime] = field(default=None) + + @property + def from_cache(self) -> bool: + return False + + @property + def is_expired(self) -> bool: + return False + + +@define(auto_attribs=False, repr=False, slots=False) +class OriginalResponse(BaseResponse): + """Wrapper class for non-cached responses returned by :py:class:`.CachedSession`""" + + @classmethod + def wrap_response(cls, response: Response, actions: 'CacheActions'): + """Modify a response object in-place and add extra cache-related attributes""" + if not isinstance(response, cls): + response.__class__ = cls + # Add expires and cache_key only if the response was written to the cache + response.expires = None if actions.skip_write else actions.expires # type: ignore + response.cache_key = None if actions.skip_write else actions.cache_key # type: ignore + response.created_at = datetime.utcnow() # type: ignore + return response + + +@define(auto_attribs=False, slots=False) +class CachedResponse(BaseResponse): + """A class that emulates :py:class:`requests.Response`, optimized for serialization""" + _content: bytes = field(default=None) _next: Optional[CachedRequest] = field(default=None) cache_key: Optional[str] = None # Not serialized; set by BaseCache.get_response() @@ -155,18 +191,3 @@ def format_file_size(n_bytes: int) -> str: if TYPE_CHECKING: return _format(unit) - - -def set_response_defaults( - response: Union[Response, CachedResponse], cache_key: str = None -) -> Union[Response, CachedResponse]: - """Set some default CachedResponse values on a requests.Response object, so they can be - expected to always be present - """ - if not isinstance(response, CachedResponse): - response.cache_key = cache_key # type: ignore - response.created_at = None # type: ignore - response.expires = None # type: ignore - response.from_cache = False # type: ignore - response.is_expired = False # type: ignore - return response diff --git a/requests_cache/session.py b/requests_cache/session.py index 63a2cb4..6abd360 100644 --- a/requests_cache/session.py +++ b/requests_cache/session.py @@ -27,7 +27,7 @@ from ._utils import get_valid_kwargs from .backends import BackendSpecifier, init_backend from .cache_control import REFRESH_TEMP_HEADER, CacheActions, append_directive from .expiration import ExpirationTime, get_expiration_seconds -from .models import AnyResponse, CachedResponse, set_response_defaults +from .models import AnyResponse, CachedResponse, OriginalResponse from .serializers import SerializerPipeline from .settings import ( DEFAULT_CACHE_NAME, @@ -107,6 +107,31 @@ class CacheMixin(MIXIN_BASE): def expire_after(self, value: ExpirationTime): self.settings.expire_after = value + # Wrapper methods to add return type hints + def get(self, url: str, **kwargs) -> AnyResponse: # type: ignore + kwargs.setdefault('allow_redirects', True) + return self.request('GET', url, **kwargs) + + def options(self, url: str, **kwargs) -> AnyResponse: # type: ignore + kwargs.setdefault('allow_redirects', True) + return self.request('OPTIONS', url, **kwargs) + + def head(self, url: str, **kwargs) -> AnyResponse: # type: ignore + kwargs.setdefault('allow_redirects', False) + return self.request('HEAD', url, **kwargs) + + def post(self, url: str, **kwargs) -> AnyResponse: # type: ignore + return self.request('POST', url, **kwargs) + + def put(self, url: str, **kwargs) -> AnyResponse: # type: ignore + return self.request('PUT', url, **kwargs) + + def patch(self, url: str, **kwargs) -> AnyResponse: # type: ignore + return self.request('PATCH', url, **kwargs) + + def delete(self, url: str, **kwargs) -> AnyResponse: # type: ignore + return self.request('DELETE', url, **kwargs) + def request( # type: ignore self, method: str, @@ -151,7 +176,7 @@ class CacheMixin(MIXIN_BASE): kwargs['headers'] = headers with patch_form_boundary(**kwargs): - return super().request(method, url, *args, **kwargs) + return super().request(method, url, *args, **kwargs) # type: ignore def send(self, request: PreparedRequest, **kwargs) -> AnyResponse: """Send a prepared request, with caching. See :py:meth:`requests.Session.send` for base @@ -219,7 +244,7 @@ class CacheMixin(MIXIN_BASE): return cached_response else: logger.debug(f'Skipping cache write for URL: {request.url}') - return set_response_defaults(response, actions.cache_key) + return OriginalResponse.wrap_response(response, actions) def _resend( self, diff --git a/requests_cache/settings.py b/requests_cache/settings.py index 7621a86..e9c6006 100644 --- a/requests_cache/settings.py +++ b/requests_cache/settings.py @@ -1,20 +1,18 @@ -from typing import TYPE_CHECKING, Callable, Dict, Iterable, Union +from typing import Callable, Dict, Iterable, Union from attr import asdict, define, field +from requests import Response from ._utils import get_valid_kwargs from .expiration import ExpirationTime -if TYPE_CHECKING: - from .models import AnyResponse - ALL_METHODS = ('GET', 'HEAD', 'OPTIONS', 'POST', 'PUT', 'PATCH', 'DELETE') DEFAULT_CACHE_NAME = 'http_cache' DEFAULT_METHODS = ('GET', 'HEAD') DEFAULT_STATUS_CODES = (200,) # Signatures for user-provided callbacks -FilterCallback = Callable[['AnyResponse'], bool] +FilterCallback = Callable[[Response], bool] KeyCallback = Callable[..., str] |
