from __future__ import annotations
import asyncio
import typing
from eggfetch.compat.httpx._request import Request
from eggfetch.compat.httpx._response import Response
from eggfetch.compat.httpx._transports import AsyncBaseTransport, BaseTransport
if typing.TYPE_CHECKING:
from typing import Any, Callable
class MockTransport(AsyncBaseTransport, BaseTransport):
def __init__(self, handler: Callable[[Request], Response]) -> None:
self._handler = handler
self._is_closed = False
def handle_request(self, request: Request) -> Response:
if self._is_closed:
raise RuntimeError("MockTransport is closed")
if asyncio.iscoroutinefunction(self._handler):
raise RuntimeError(
"Cannot use an async handler with a synchronous client. "
"Use AsyncClient instead."
)
response = self._handler(request)
if not isinstance(response, Response):
raise TypeError(
f"Handler must return a Response, got {type(response)}"
)
try:
response.request
except RuntimeError:
response._request = request return response
async def handle_async_request(self, request: Request) -> Response:
if self._is_closed:
raise RuntimeError("MockTransport is closed")
if asyncio.iscoroutinefunction(self._handler):
response = await self._handler(request)
else:
response = self._handler(request)
if not isinstance(response, Response):
raise TypeError(
f"Handler must return a Response, got {type(response)}"
)
try:
response.request
except RuntimeError:
response._request = request return response
def close(self) -> None:
self._is_closed = True
async def aclose(self) -> None:
self._is_closed = True
def __enter__(self) -> MockTransport:
return self
def __exit__(self, *args: Any) -> None:
self.close()
async def __aenter__(self) -> MockTransport:
return self
async def __aexit__(self, *args: Any) -> None:
await self.aclose()
def _build_response(
status_code: int = 200,
*,
headers: dict | None = None,
content: bytes | None = None,
text: str | None = None,
json: Any = None,
stream: Any = None,
) -> Response:
if stream is not None:
return Response(status_code, headers=headers, stream=stream)
if text is not None:
return Response(status_code, headers=headers, text=text)
if json is not None:
return Response(status_code, headers=headers, json=json)
return Response(status_code, headers=headers, content=content or b"")