from abc import abstractmethod from typing import Any, Protocol, Sequence from aiohttp import ClientConnectionError from aioresponses import aioresponses from aioresponses.core import merge_params, normalize_url from responses import RequestsMock __all__ = ["RequestManager", "SyncRequestManager", "AsyncRequestManager"] class RequestManager(Protocol): @property @abstractmethod def call_count(self) -> int: ... @property @abstractmethod def last_request_data(self) -> dict[str, Any]: ... @abstractmethod def __enter__(self): ... @abstractmethod def __exit__(self, exc_type, exc_val, exc_tb): ... @abstractmethod def add(self, *args, **kwargs): ... def add_many(self, requests: Sequence[dict[str, Any]]) -> None: for request in requests: self.add(**request) class SyncRequestManager(RequestsMock, RequestManager): @property def call_count(self): return len(self._calls) @property def last_request_data(self): request = self._calls[-1].request return {"method": request.method, "url": request.url, "headers": request.headers, "data": request.body} class AsyncRequestManager(aioresponses, RequestManager): def __init__(self, **kwargs): super().__init__(**kwargs) self._last_request_data = {} self._call_count = 0 @property def call_count(self): return self._call_count @property def last_request_data(self): return self._last_request_data def _save_last_request_data(self, method, url, **kwargs): self._last_request_data = { "method": method, "url": url, "headers": kwargs.get("headers", {}), "data": kwargs.get("data", {}), } async def _request_mock(self, orig_self, method, url, *args, **kwargs): """Override to add 'call_args'.""" if orig_self.closed: raise RuntimeError("Session is closed") url_origin = url url = normalize_url(merge_params(url, kwargs.get("params"))) url_str = str(url) for prefix in self._passthrough: if url_str.startswith(prefix): return await self.patcher.temp_original(orig_self, method, url_origin, *args, **kwargs) key = (method, url) self.requests.setdefault(key, []) request_call = self._build_request_call(method, *args, **kwargs) self.requests[key].append(request_call) self._call_count += 1 self._save_last_request_data(method, url, **kwargs) response = await self.match(method, url, **kwargs) if response is None: raise ClientConnectionError("Connection refused: {} {}".format(method, url)) self._responses.append(response) raise_for_status = kwargs.get("raise_for_status") if raise_for_status is None: raise_for_status = getattr(orig_self, "_raise_for_status", False) if raise_for_status: response.raise_for_status() return response