diff --git a/litellm/a2a_protocol/main.py b/litellm/a2a_protocol/main.py index fe0e6837586..322393cd9c4 100644 --- a/litellm/a2a_protocol/main.py +++ b/litellm/a2a_protocol/main.py @@ -13,7 +13,6 @@ import asyncio import datetime import uuid from collections.abc import AsyncIterator, Coroutine -from http.cookiejar import DefaultCookiePolicy from typing import TYPE_CHECKING, Any, Final, Optional, cast import litellm @@ -80,8 +79,6 @@ from litellm.a2a_protocol.exceptions import A2ALocalhostURLError # Use our custom resolver instead of the default A2A SDK resolver A2ACardResolver: Final = LiteLLMA2ACardResolver -_BLOCK_ALL_COOKIES: Final = DefaultCookiePolicy(allowed_domains=()) - def _set_usage_on_logging_obj( kwargs: dict[str, Any], @@ -770,7 +767,6 @@ async def create_a2a_client( params={"timeout": timeout}, ) httpx_client: Final = _async_handler.client - httpx_client.cookies.jar.set_policy(_BLOCK_ALL_COOKIES) if extra_headers: verbose_proxy_logger.debug("A2A client created with extra_headers=%s", list(extra_headers.keys())) diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index d586156b625..9ada3674d33 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -8,11 +8,12 @@ import sys import threading import time from collections.abc import Callable, Mapping +from http.cookiejar import CookieJar, DefaultCookiePolicy from typing import TYPE_CHECKING, Any, Final, Optional import certifi import httpx -from aiohttp import ClientSession, TCPConnector +from aiohttp import ClientSession, DummyCookieJar, TCPConnector from httpx import USE_CLIENT_DEFAULT, AsyncHTTPTransport, HTTPTransport from httpx._types import RequestFiles @@ -144,6 +145,15 @@ def _handler_may_close_client(client_refcount: int, owns_client: bool) -> bool: return owns_client and client_refcount <= _CLIENT_REFCOUNT_WHEN_HANDLER_IS_SOLE_REFERRER +def blocked_cookie_jar() -> CookieJar: + """A jar that stores no response cookie and sends none, for httpx clients. + + LiteLLM's outbound clients are pooled and shared by every caller, so a cookie one + upstream sets would be replayed to every other upstream on a matching domain. + """ + return CookieJar(policy=DefaultCookiePolicy(allowed_domains=())) + + _STREAMING_ERROR_BODY_READ_TIMEOUT_SECONDS: Final = 5.0 _STREAMING_ERROR_BODY_READ_EXECUTOR: Final = concurrent.futures.ThreadPoolExecutor( max_workers=50, @@ -587,6 +597,7 @@ class AsyncHTTPHandler: verify=ssl_config, cert=cert, headers=default_headers, + cookies=blocked_cookie_jar(), follow_redirects=True, ) @@ -1063,6 +1074,7 @@ class AsyncHTTPHandler: def session_factory() -> ClientSession: return ClientSession( connector=TCPConnector(**transport_connector_kwargs), + cookie_jar=DummyCookieJar(), trust_env=trust_env, ) @@ -1132,6 +1144,7 @@ class HTTPHandler: verify=ssl_config, cert=cert, headers=default_headers, + cookies=blocked_cookie_jar(), follow_redirects=True, ) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 84f8685f5dc..d016eb57dd4 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -871,7 +871,7 @@ async def proxy_shutdown_event(): async def _initialize_shared_aiohttp_session(): """Initialize shared aiohttp session for connection reuse with connection limits.""" try: - from aiohttp import ClientSession, TCPConnector + from aiohttp import ClientSession, DummyCookieJar, TCPConnector from litellm.llms.custom_httpx.http_handler import ( _build_aiohttp_keepalive_socket_factory, @@ -892,7 +892,7 @@ async def _initialize_shared_aiohttp_session(): connector_kwargs["socket_factory"] = socket_factory connector: Final = TCPConnector(**connector_kwargs) - session: Final = ClientSession(connector=connector) + session: Final = ClientSession(connector=connector, cookie_jar=DummyCookieJar()) verbose_proxy_logger.info( "SESSION REUSE: Created shared aiohttp session for connection pooling (ID: %s, limit=%s, limit_per_host=%s)", diff --git a/tests/test_litellm/a2a_protocol/test_main.py b/tests/test_litellm/a2a_protocol/test_main.py index 59cb8c2c438..08f6b9f25bb 100644 --- a/tests/test_litellm/a2a_protocol/test_main.py +++ b/tests/test_litellm/a2a_protocol/test_main.py @@ -174,21 +174,15 @@ _RPC_REPLY = { _AGENT_A_HEADERS = {"x-agent-token": "token-for-a", "x-tenant": "tenant-a"} _AGENT_B_HEADERS = {"x-agent-token": "token-for-b", "x-tenant": "tenant-b"} -_UPSTREAM_SESSION_COOKIE = "a2a_session=only-agent-a-may-hold-this; Path=/" class _RequestRecorder: - """Records the headers httpx put on the wire, per outbound request. + """Records the headers httpx put on the wire, per outbound request.""" - ``cookie_from_tenant`` makes that tenant's agent answer with a Set-Cookie, standing in - for an upstream that issues a session cookie. - """ - - def __init__(self, cookie_from_tenant: str | None = None): + def __init__(self): self.card_requests = [] self.rpc_requests = [] self.client = None - self.cookie_from_tenant = cookie_from_tenant def __call__(self, request: httpx.Request) -> httpx.Response: headers = {k.lower(): v for k, v in request.headers.items()} @@ -196,8 +190,6 @@ class _RequestRecorder: self.card_requests.append(headers) return httpx.Response(200, json=_AGENT_CARD) self.rpc_requests.append(headers) - if self.cookie_from_tenant is not None and headers.get("x-tenant") == self.cookie_from_tenant: - return httpx.Response(200, json=_RPC_REPLY, headers={"set-cookie": _UPSTREAM_SESSION_COOKIE}) return httpx.Response(200, json=_RPC_REPLY) @@ -205,15 +197,14 @@ def _a2a_client_cache_key(timeout: float) -> str: return "async_httpx_client" + f"timeout_{timeout}" + httpxSpecialProvider.A2AProvider -async def _seed_shared_a2a_client(cookie_from_tenant: str | None = None) -> _RequestRecorder: +async def _seed_shared_a2a_client() -> _RequestRecorder: """Put the one A2A client the cache will hand out behind a mock transport. Seeding has to happen on the test's own event loop, because the client cache keys on it. The injected client is a real httpx.AsyncClient, so the merge of per-request - headers over client defaults, and httpx's own cookie handling, which is what these - tests are about, stay real. + headers over client defaults, which is what these tests are about, stays real. """ - recorder = _RequestRecorder(cookie_from_tenant=cookie_from_tenant) + recorder = _RequestRecorder() handler = AsyncHTTPHandler(timeout=DEFAULT_A2A_AGENT_TIMEOUT) owned_client = handler.client handler.client = httpx.AsyncClient(transport=httpx.MockTransport(recorder)) @@ -333,17 +324,21 @@ async def test_agent_card_fetch_carries_the_callers_headers(isolated_client_cach @pytest.mark.asyncio -async def test_one_agents_session_cookie_never_reaches_another_agent(isolated_client_cache): - """One pooled client is also one httpx cookie jar. httpx stores every Set-Cookie on the - client and replays it on any later request to a matching domain, so an agent's session - cookie would ride along on a different agent's call to the same host.""" - recorder = await _seed_shared_a2a_client(cookie_from_tenant="tenant-a") +async def test_the_pooled_a2a_client_arrives_with_cookie_persistence_disabled(isolated_client_cache): + """create_a2a_client takes its client from the shared builder rather than building one, + and the builder is what refuses to persist cookies. This pins the join between those + two facts, so the A2A path cannot quietly start acquiring a client that keeps a jar. - client_a = await create_a2a_client(base_url="http://127.0.0.1:9", extra_headers=_AGENT_A_HEADERS) - await _send_message(client_a, _send_request("a")) - client_b = await create_a2a_client(base_url="http://127.0.0.1:9", extra_headers=_AGENT_B_HEADERS) - await _send_message(client_b, _send_request("b")) + test_callers_with_different_headers_reuse_one_pooled_client pins the other half, that + create_a2a_client hands back exactly this cached client.""" + handler = get_async_httpx_client( + llm_provider=httpxSpecialProvider.A2AProvider, + params={"timeout": DEFAULT_A2A_AGENT_TIMEOUT}, + ) + request = httpx.Request("GET", "https://agent-a.example.com/") + handler.client.cookies.extract_cookies( + httpx.Response(200, headers={"set-cookie": "SESSION=only-agent-a-may-hold-this"}, request=request) + ) - assert dict(recorder.client.cookies) == {}, "the shared client kept an agent's session cookie" - assert "cookie" not in recorder.card_requests[-1] - assert "cookie" not in recorder.rpc_requests[-1] + assert dict(handler.client.cookies) == {}, "the pooled A2A client kept an upstream's cookie" + await handler.close() diff --git a/tests/test_litellm/llms/custom_httpx/test_aiohttp_cleanup_closed.py b/tests/test_litellm/llms/custom_httpx/test_aiohttp_cleanup_closed.py index f279acfd60c..82010e82cea 100644 --- a/tests/test_litellm/llms/custom_httpx/test_aiohttp_cleanup_closed.py +++ b/tests/test_litellm/llms/custom_httpx/test_aiohttp_cleanup_closed.py @@ -13,7 +13,7 @@ def test_create_aiohttp_transport_sets_enable_cleanup_closed_when_needed(monkeyp ) as mock_tcp_connector: with patch.object( http_handler_module, "ClientSession", return_value=session_mock - ): + ), patch.object(http_handler_module, "DummyCookieJar", return_value=MagicMock(name="cookie_jar")): transport = http_handler_module.AsyncHTTPHandler._create_aiohttp_transport( shared_session=None ) @@ -36,7 +36,7 @@ def test_create_aiohttp_transport_omits_enable_cleanup_closed_when_not_needed( ) as mock_tcp_connector: with patch.object( http_handler_module, "ClientSession", return_value=session_mock - ): + ), patch.object(http_handler_module, "DummyCookieJar", return_value=MagicMock(name="cookie_jar")): transport = http_handler_module.AsyncHTTPHandler._create_aiohttp_transport( shared_session=None ) diff --git a/tests/test_litellm/llms/custom_httpx/test_aiohttp_so_keepalive.py b/tests/test_litellm/llms/custom_httpx/test_aiohttp_so_keepalive.py index 5a37e681c4a..0065bf8f4ef 100644 --- a/tests/test_litellm/llms/custom_httpx/test_aiohttp_so_keepalive.py +++ b/tests/test_litellm/llms/custom_httpx/test_aiohttp_so_keepalive.py @@ -34,7 +34,7 @@ def test_socket_factory_omitted_when_disabled(monkeypatch): ) as mock_tcp_connector: with patch.object( http_handler_module, "ClientSession", return_value=session_mock - ): + ), patch.object(http_handler_module, "DummyCookieJar", return_value=MagicMock(name="cookie_jar")): _invoke_connector_factory(http_handler_module) assert mock_tcp_connector.call_count >= 1 @@ -55,7 +55,7 @@ def test_socket_factory_attached_when_enabled(monkeypatch): ) as mock_tcp_connector: with patch.object( http_handler_module, "ClientSession", return_value=session_mock - ): + ), patch.object(http_handler_module, "DummyCookieJar", return_value=MagicMock(name="cookie_jar")): _invoke_connector_factory(http_handler_module) assert mock_tcp_connector.call_count >= 1 @@ -77,7 +77,7 @@ def test_socket_factory_skipped_on_old_aiohttp(monkeypatch): ) as mock_tcp_connector: with patch.object( http_handler_module, "ClientSession", return_value=session_mock - ): + ), patch.object(http_handler_module, "DummyCookieJar", return_value=MagicMock(name="cookie_jar")): _invoke_connector_factory(http_handler_module) assert mock_tcp_connector.call_count >= 1 diff --git a/tests/test_litellm/llms/custom_httpx/test_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_http_handler.py index 86d097c6123..fa1c7308c6f 100644 --- a/tests/test_litellm/llms/custom_httpx/test_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_http_handler.py @@ -1171,3 +1171,77 @@ async def test_client_handed_out_by_async_cache_survives_eviction_and_collection assert not consumer_client.is_closed await consumer_client.aclose() + + +_SET_COOKIE = "SESSION=upstream-a-secret; Path=/" + + +def _cookie_recorder(): + """A transport that hands out a Set-Cookie once, and records what comes back.""" + seen = [] + + def handler(request: httpx.Request) -> httpx.Response: + seen.append(request.headers.get("cookie")) + if request.url.path == "/set": + return httpx.Response(200, headers={"set-cookie": _SET_COOKIE}) + return httpx.Response(200) + + return handler, seen + + +@pytest.mark.asyncio +async def test_async_client_never_replays_one_upstreams_cookie_to_another(): + """LiteLLM's async clients are pooled and shared by every caller, so a cookie one + upstream sets would be attached to every later request on a matching domain, reaching + a different tenant's upstream. The client must persist no response cookie.""" + handler, seen = _cookie_recorder() + http_handler = AsyncHTTPHandler() + client = http_handler.client + client._transport = httpx.MockTransport(handler) + + await client.get("https://upstream-a.example.com/set") + await client.get("https://upstream-b.example.com/rpc") + await client.aclose() + + assert dict(client.cookies) == {}, "the shared client stored an upstream's cookie" + assert seen == [None, None] + + +def test_sync_client_never_replays_one_upstreams_cookie_to_another(): + """Same invariant on the sync client, which is pooled the same way.""" + handler, seen = _cookie_recorder() + http_handler = HTTPHandler() + client = http_handler.client + client._transport = httpx.MockTransport(handler) + + client.get("https://upstream-a.example.com/set") + client.get("https://upstream-b.example.com/rpc") + client.close() + + assert dict(client.cookies) == {} + assert seen == [None, None] + + +@pytest.mark.asyncio +async def test_aiohttp_session_never_replays_one_upstreams_cookie_to_another(): + """The httpx jar is not the only one. AiohttpTransport is litellm's default transport + and the aiohttp ClientSession keeps its own cookie jar, which httpx-level assertions + cannot see, so blocking only the httpx jar leaves the leak intact on the real path. + + aiohttp's default jar refuses cookies for IP hosts, so this drives a hostname. An + IP-addressed check passes whether or not the session jar is blocked.""" + from aiohttp import DummyCookieJar + from yarl import URL + + http_handler = AsyncHTTPHandler(timeout=61.0) + transport = http_handler.client._transport + assert isinstance(transport, LiteLLMAiohttpTransport), "aiohttp is no longer the default transport" + + session = transport.client() if callable(transport.client) else transport.client + jar = session.cookie_jar + assert isinstance(jar, DummyCookieJar) + + jar.update_cookies({"SESSION": "upstream-a-secret"}, URL("https://upstream-a.example.com")) + assert len(jar) == 0 + assert dict(jar.filter_cookies(URL("https://upstream-a.example.com"))) == {} + await session.close()