From a811fc041a8091cbe4092f3600e3ee9ab813bbd6 Mon Sep 17 00:00:00 2001 From: Sylvester Kaczmarek <16242628+sylvesterkaczmarek@users.noreply.github.com> Date: Mon, 17 Aug 2026 12:05:53 +0100 Subject: [PATCH] Accept awaitable-returning async middleware callables --- src/anthropic/_middleware.py | 10 +-- ...est_async_middleware_awaitable_callable.py | 74 +++++++++++++++++++ 2 files changed, 79 insertions(+), 5 deletions(-) create mode 100644 tests/test_async_middleware_awaitable_callable.py diff --git a/src/anthropic/_middleware.py b/src/anthropic/_middleware.py index c4df340f5..5138d7500 100644 --- a/src/anthropic/_middleware.py +++ b/src/anthropic/_middleware.py @@ -119,8 +119,8 @@ def validate_async_middleware(middleware: Iterable[MiddlewareInput]) -> None: ) elif not callable(entry): raise TypeError(f"middleware {_middleware_name(entry)} is not callable") - elif not _is_async_callable(entry): - raise TypeError( - f"middleware {_middleware_name(entry)} is not an async function; " - "the asynchronous client requires async middleware functions" - ) + # Function-style async middleware is typed as Callable[..., Awaitable], + # not specifically as a coroutine function. A synchronous wrapper may + # therefore validly return an awaitable; the async middleware chain + # awaits that result at invocation time. Do not reject that supported + # shape based only on inspect.iscoroutinefunction(). diff --git a/tests/test_async_middleware_awaitable_callable.py b/tests/test_async_middleware_awaitable_callable.py new file mode 100644 index 000000000..26ba420f3 --- /dev/null +++ b/tests/test_async_middleware_awaitable_callable.py @@ -0,0 +1,74 @@ +from __future__ import annotations + +from typing import Any, Awaitable + +import anyio +import httpx +import pytest + +from anthropic import AsyncAnthropic +from anthropic._middleware import AsyncCallNext +from anthropic._request import APIRequest + + +def test_async_client_accepts_sync_wrapper_returning_awaitable() -> None: + seen: list[str] = [] + + def middleware(request: APIRequest, call_next: AsyncCallNext) -> Awaitable[Any]: + seen.append(request.url) + return call_next(request) + + async def run() -> None: + async def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response(200, json={"ok": True}, request=request) + + client = AsyncAnthropic( + api_key="test", + base_url="https://example.test", + middleware=[middleware], + http_client=httpx.AsyncClient(transport=httpx.MockTransport(handler)), + ) + try: + result = await client.get("/probe", cast_to=object) + finally: + await client.close() + + assert result == {"ok": True} + + anyio.run(run) + assert seen == ["/probe"] + + +def test_async_client_accepts_sync_callable_object_returning_awaitable() -> None: + class MiddlewareWrapper: + def __init__(self) -> None: + self.calls = 0 + + def __call__(self, request: APIRequest, call_next: AsyncCallNext) -> Awaitable[Any]: + self.calls += 1 + return call_next(request) + + wrapper = MiddlewareWrapper() + + async def run() -> None: + async def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response(200, json={"ok": True}, request=request) + + client = AsyncAnthropic( + api_key="test", + base_url="https://example.test", + middleware=[wrapper], + http_client=httpx.AsyncClient(transport=httpx.MockTransport(handler)), + ) + try: + assert await client.get("/probe", cast_to=object) == {"ok": True} + finally: + await client.close() + + anyio.run(run) + assert wrapper.calls == 1 + + +def test_async_client_still_rejects_non_callable_middleware() -> None: + with pytest.raises(TypeError, match="is not callable"): + AsyncAnthropic(api_key="test", middleware=[object()]) # type: ignore[list-item]