Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
60 changes: 36 additions & 24 deletions src/unihttp/bind_method.py
Original file line number Diff line number Diff line change
@@ -1,47 +1,52 @@
import functools
import inspect
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING, Any, Generic, ParamSpec, TypeVar, cast, overload
from typing import TYPE_CHECKING, Any, Generic, ParamSpec, TypeVar, overload

from unihttp.method import BaseMethod
from unihttp.http.stream import AsyncChunkStream, ChunkStream
from unihttp.method import BaseMethod, StreamMethod

if TYPE_CHECKING:
from unihttp.clients.base import BaseAsyncClient, BaseSyncClient


MethodParamSpec = ParamSpec("MethodParamSpec")
MethodResultT = TypeVar("MethodResultT")
SyncResultT = TypeVar("SyncResultT")
AsyncResultT = TypeVar("AsyncResultT")


class MethodBinder(Generic[MethodParamSpec, MethodResultT]): # noqa: UP046
__slots__ = ("_method_tp",)
class MethodBinder(Generic[MethodParamSpec, SyncResultT, AsyncResultT]): # noqa: UP046
__slots__ = ("_call_name", "_method_tp")

def __init__(
self,
method_tp: Callable[MethodParamSpec, BaseMethod[MethodResultT]],
method_tp: Callable[MethodParamSpec, Any],
call_name: str,
) -> None:
self._method_tp = method_tp
self._call_name = call_name

@overload
def __get__(
self,
instance: None,
owner: type,
) -> "MethodBinder[MethodParamSpec, MethodResultT]": ...
) -> "MethodBinder[MethodParamSpec, SyncResultT, AsyncResultT]": ...

@overload
def __get__(
self,
instance: "BaseSyncClient",
owner: type,
) -> Callable[MethodParamSpec, MethodResultT]: ...
) -> Callable[MethodParamSpec, SyncResultT]: ...

@overload
def __get__(
self,
instance: "BaseAsyncClient",
owner: type,
) -> Callable[MethodParamSpec, Awaitable[MethodResultT]]: ...
) -> Callable[MethodParamSpec, Awaitable[AsyncResultT]]: ...

def __get__(
self,
Expand All @@ -51,42 +56,49 @@ def __get__(
if instance is None:
return self

if not hasattr(instance, "call_method"):
call_name = self._call_name
if not hasattr(instance, call_name):
raise RuntimeError(
"`bind_method` is available only for classes with `call_method`",
f"`bind_method` is available only for classes with `{call_name}`",
)

call_method = instance.call_method
call = getattr(instance, call_name)
method_tp = self._method_tp

if inspect.iscoroutinefunction(call_method):
if inspect.iscoroutinefunction(call):

@functools.wraps(method_tp)
async def async_wrapper(
*args: MethodParamSpec.args,
**kwargs: MethodParamSpec.kwargs,
) -> MethodResultT:
return cast(
MethodResultT,
await call_method(method_tp(*args, **kwargs)),
)
) -> Any:
return await call(method_tp(*args, **kwargs))

return async_wrapper

@functools.wraps(method_tp)
def sync_wrapper(
*args: MethodParamSpec.args,
**kwargs: MethodParamSpec.kwargs,
) -> MethodResultT:
return cast(
MethodResultT,
call_method(method_tp(*args, **kwargs)),
)
) -> Any:
return call(method_tp(*args, **kwargs))

return sync_wrapper


@overload
def bind_method( # noqa: UP047
method_tp: Callable[MethodParamSpec, BaseMethod[MethodResultT]],
) -> MethodBinder[MethodParamSpec, MethodResultT]:
return MethodBinder(method_tp)
) -> MethodBinder[MethodParamSpec, MethodResultT, MethodResultT]: ...


@overload
def bind_method( # noqa: UP047
method_tp: Callable[MethodParamSpec, StreamMethod],
) -> MethodBinder[MethodParamSpec, ChunkStream, AsyncChunkStream]: ...


def bind_method(method_tp: Callable[..., Any]) -> Any:
if isinstance(method_tp, type) and issubclass(method_tp, StreamMethod):
return MethodBinder(method_tp, "call_method_stream")
return MethodBinder(method_tp, "call_method")
93 changes: 64 additions & 29 deletions src/unihttp/clients/aiohttp.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,17 +4,36 @@
from urllib.parse import urljoin

import aiohttp
from aiohttp import ClientSession, FormData
from aiohttp import ClientResponse, ClientSession, FormData

from unihttp.clients.base import BaseAsyncClient
from unihttp.exceptions import NetworkError, RequestTimeoutError
from unihttp.http import UploadFile
from unihttp.http.request import HTTPRequest
from unihttp.http.response import HTTPResponse
from unihttp.http.stream import AsyncChunkStream
from unihttp.middlewares.base import AsyncMiddleware
from unihttp.serialize import RequestDumper, ResponseLoader


class _AiohttpChunkStream(AsyncChunkStream):
def __init__(self, response: ClientResponse, chunk_size: int) -> None:
super().__init__()
self._response = response
self._iter = response.content.iter_chunked(chunk_size)

async def _fetch_chunk(self) -> bytes:
try:
return await anext(self._iter)
except aiohttp.ClientConnectionError as e:
raise NetworkError(str(e)) from e
except TimeoutError as e:
raise RequestTimeoutError(str(e)) from e

async def _close_response(self) -> None:
self._response.close()


class AiohttpAsyncClient(BaseAsyncClient):
def __init__(
self,
Expand Down Expand Up @@ -70,50 +89,66 @@ def _build_form_data(self, request: HTTPRequest) -> FormData:

return form_data

async def make_request(self, request: HTTPRequest) -> HTTPResponse:
data: FormData | str | None = None

def _build_data(self, request: HTTPRequest) -> FormData | str | bytes | None:
"""Resolve the request payload: raw, then multipart/form, then JSON body."""
if request.raw is not None:
return request.raw
if request.form or request.file:
data = self._build_form_data(request)

return self._build_form_data(request)
if request.body:
if request.form or request.file:
raise ValueError(
"Cannot use Body with Form or File. "
"Use Form for fields in multipart requests."
)

data = self.json_dumps(request.body)
if "Content-Type" not in request.header:
request.header["Content-Type"] = "application/json"
return data
return None

async def _do_request(self, request: HTTPRequest) -> ClientResponse:
data = self._build_data(request)

try:
async with self._session.request(
return await self._session.request(
method=request.method,
url=urljoin(self.base_url, request.url),
headers=request.header,
params=request.query,
data=data,
) as response:
response_data: Any = None
content = await response.read()
if content:
try:
response_data = self.json_loads(content)
except (ValueError, TypeError):
response_data = content

return HTTPResponse(
status_code=response.status,
headers=response.headers,
cookies=response.cookies,
data=response_data,
raw_response=response,
)
)
except aiohttp.ClientConnectionError as e:
raise NetworkError(str(e)) from e
except TimeoutError as e:
raise RequestTimeoutError(str(e)) from e

async def make_request(self, request: HTTPRequest) -> HTTPResponse:
response = await self._do_request(request)

response_data: Any = None
content = await response.read()
if content:
try:
response_data = self.json_loads(content)
except (ValueError, TypeError):
response_data = content

return HTTPResponse(
status_code=response.status,
headers=response.headers,
cookies=response.cookies,
data=response_data,
raw_response=response,
)

async def stream_make_request(
self, request: HTTPRequest, chunk_size: int = 65536
) -> HTTPResponse[AsyncChunkStream]:
response = await self._do_request(request)

return HTTPResponse(
status_code=response.status,
headers=response.headers,
cookies=response.cookies,
data=_AiohttpChunkStream(response, chunk_size),
raw_response=response,
)

async def close(self) -> None:
await self._session.close()
Loading