Skip to content
Merged
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
8 changes: 6 additions & 2 deletions python/packages/core/agent_framework/_mcp.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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))

Expand Down Expand Up @@ -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:
Expand Down
23 changes: 23 additions & 0 deletions python/packages/core/tests/core/test_mcp.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading