mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(http): stop pooled clients persisting cookies on the aiohttp jar too (#36149)
#35978 stopped the pooled A2A client replaying one upstream's Set-Cookie to another by installing a blocking policy on that client's httpx cookie jar. That covers only one of the two jars on the request path. AiohttpTransport is the default transport unless it is explicitly disabled, and the aiohttp ClientSession behind it keeps its own cookie jar which no httpx-level assertion can observe, so the leak is still live on the default path: a live proxy on that commit still delivers agent-alpha's session cookie to agent-beta's card fetch and JSON-RPC call. The reason it looked fixed is that aiohttp's default CookieJar is built with unsafe=False and refuses to store cookies for IP hosts, so a proof addressed to 127.0.0.1 comes back clean whether or not that jar is blocked. Cookie persistence is now blocked where the clients are built rather than at one call site: blocked_cookie_jar() gives every httpx client, async and sync, a jar whose DefaultCookiePolicy(allowed_domains=()) rejects every domain in both directions, and both ClientSession constructions litellm owns, the transport's session factory and the proxy's shared startup session, get a DummyCookieJar. LiteLLM reads a response cookie nowhere, and an explicitly supplied Cookie header still goes out, so passthrough forwarding and an agent's extra_headers are unaffected. The A2A-scoped policy #35978 added is removed, since it is now dead. The two suites that drive the aiohttp session factory synchronously mock ClientSession because a real one needs a running event loop; DummyCookieJar has the same requirement, so they mock it for the same reason.
This commit is contained in:
parent
6ba744b340
commit
ae1d1cb05e
7 changed files with 116 additions and 38 deletions
|
|
@ -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()))
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)",
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue