diff --git a/python/packages/core/agent_framework/_mcp.py b/python/packages/core/agent_framework/_mcp.py index 0197d9a353..83410bd9d5 100644 --- a/python/packages/core/agent_framework/_mcp.py +++ b/python/packages/core/agent_framework/_mcp.py @@ -58,7 +58,7 @@ from typing_extensions import Self # pragma: no cover if TYPE_CHECKING: - from httpx import AsyncClient + from httpx import AsyncClient, Request, Response from mcp import types from mcp.client.session import ClientSession from mcp.shared.context import RequestContext @@ -437,6 +437,10 @@ def _tagged_kwargs(self, kwargs: dict[str, Any]) -> dict[str, Any]: def stream(self, *args: Any, **kwargs: Any) -> Any: return self._client.stream(*args, **self._tagged_kwargs(kwargs)) + async def send(self, request: Request, **kwargs: Any) -> Response: + request.extensions[_MCP_HEADER_OWNER_EXTENSION] = self._owner + return await self._client.send(request, **kwargs) + async def delete(self, *args: Any, **kwargs: Any) -> Any: return await self._client.delete(*args, **self._tagged_kwargs(kwargs)) @@ -3682,7 +3686,7 @@ def get_mcp_client(self) -> _AsyncGeneratorContextManager[Any, None]: Returns: An async context manager for the streamable HTTP client transport. """ - from httpx import URL, AsyncClient, Request, Timeout + from httpx import URL, AsyncClient, Timeout http_client = self._httpx_client if self._header_provider is not None: diff --git a/python/packages/core/tests/core/test_mcp.py b/python/packages/core/tests/core/test_mcp.py index 063ffc288a..7102dbd7e3 100644 --- a/python/packages/core/tests/core/test_mcp.py +++ b/python/packages/core/tests/core/test_mcp.py @@ -7464,6 +7464,29 @@ async def delayed_close() -> None: await user_client.aclose() +async def test_mcp_header_scoped_client_tags_send_requests(): + """The transport wrapper must identify requests sent through AsyncClient.send.""" + import httpx + + from agent_framework._mcp import _MCP_HEADER_OWNER_EXTENSION, _MCPHeaderScopedClient + + owner = object() + observed_owners: list[object | None] = [] + + async def handle(request: httpx.Request) -> httpx.Response: + observed_owners.append(request.extensions.get(_MCP_HEADER_OWNER_EXTENSION)) + return httpx.Response(200) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handle)) as user_client: + wrapper = _MCPHeaderScopedClient(user_client, owner) + request = user_client.build_request("POST", "http://example.com/mcp") + + response = await wrapper.send(request) + + assert response.status_code == 200 + assert observed_owners == [owner] + + async def test_mcp_header_scoped_client_delegates_unwrapped_attributes(): """The transport wrapper must stay a drop-in for the caller's httpx client.""" import httpx