fix(guardrails): reuse native clients with isolated transports

This commit is contained in:
JVQ 2026-10-03 14:32:16 +02:00
parent 7e2df548d3
commit ee44c3dd36
No known key found for this signature in database
GPG key ID: 9365B9C958929949
5 changed files with 257 additions and 24 deletions

View file

@ -18,7 +18,7 @@ guardrails:
api_key: os.environ/ISMALICIOUS_ENCODED_API_KEY
```
Requests only send credentials to `https://api.ismalicious.com`, with TLS verification, no redirects, a 15-second timeout and no retries. A custom API base is refused. Incoming request headers, provider credentials, user metadata and model prompts are not added to the inspection body by this provider
Requests only send credentials to `https://api.ismalicious.com`, with TLS verification, no redirects, a 15-second timeout and no retries. A custom API base is refused. Clients reuse LiteLLM’s native cache, which retains ownership of closure. Each client generation gets its own transport, so closing a retired client does not close the replacement. API credentials are passed per request and are excluded from the cache key. Incoming request headers, provider credentials, user metadata and model prompts are not added to the inspection body by this provider
## Decisions

View file

@ -616,12 +616,16 @@ class AsyncHTTPHandler:
shared_session: Optional["ClientSession"] = None,
transport: httpx.AsyncBaseTransport | None = None,
follow_redirects: bool = True,
transport_factory: Callable[[], httpx.AsyncBaseTransport] | None = None,
):
if transport is not None and transport_factory is not None:
raise ValueError("transport and transport_factory are mutually exclusive")
self.timeout = timeout
self.event_hooks = event_hooks
self.ssl_verify = ssl_verify
self.shared_session = shared_session
self.transport = transport
self.transport_factory = transport_factory
self.follow_redirects = follow_redirects
self._owns_client = True
self._client = self.create_client(
@ -655,9 +659,10 @@ class AsyncHTTPHandler:
ssl_verify: VerifyTypes | None = None,
shared_session: Optional["ClientSession"] = None,
) -> httpx.AsyncClient:
if self.transport is not None:
explicit_transport: Final = self.transport_factory() if self.transport_factory is not None else self.transport
if explicit_transport is not None:
return httpx.AsyncClient(
transport=self.transport,
transport=explicit_transport,
event_hooks=event_hooks,
timeout=timeout if timeout is not None else _DEFAULT_TIMEOUT,
headers=get_default_headers(),
@ -1748,7 +1753,9 @@ def get_async_httpx_client(
# Filter out params that are only used for cache key, not for AsyncHTTPHandler.__init__
handler_params: Final = {k: v for k, v in params.items() if k != "disable_aiohttp_transport"}
handler_params["shared_session"] = shared_session
_new_client = AsyncHTTPHandler(**handler_params)
_new_client = AsyncHTTPHandler(
**handler_params, # pyright: ignore[reportUnknownArgumentType] # native cache forwards an untyped constructor-parameter dictionary
)
else:
_new_client = AsyncHTTPHandler(
timeout=_default_cached_client_timeout(),

View file

@ -14,6 +14,9 @@ from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information, # pyright: ignore[reportUnknownVariableType] # native logging decorator has an untyped signature
)
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # native factory uses an untyped parameter dictionary
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import GenericGuardrailAPIInputs
@ -116,6 +119,10 @@ def _url_argument(text: str) -> str | None | _Failure:
return text
def _verified_transport() -> httpx.AsyncBaseTransport:
return httpx.AsyncHTTPTransport(verify=True, retries=0, trust_env=False)
class IsMaliciousGuardrail(CustomGuardrail):
def __init__(
self,
@ -171,9 +178,21 @@ class IsMaliciousGuardrail(CustomGuardrail):
) -> _Failure | None:
try:
response: Final = (
await client.get("/gate/url", params={"u": url})
await client.get(
f"{_API_BASE}/gate/url",
params={"u": url},
headers=self._headers,
timeout=15,
follow_redirects=False,
)
if url is not None
else await client.post("/gate/scan", content=body)
else await client.post(
f"{_API_BASE}/gate/scan",
content=body,
headers=self._headers,
timeout=15,
follow_redirects=False,
)
)
except httpx.HTTPError:
return _Failure()
@ -202,19 +221,20 @@ class IsMaliciousGuardrail(CustomGuardrail):
body: Final = _encoded_body(texts)
if isinstance(body, _Failure):
self._raise_failure(body)
async with httpx.AsyncClient(
base_url=_API_BASE,
headers=self._headers,
transport=self._transport,
timeout=15,
follow_redirects=False,
verify=True,
trust_env=False,
) as client:
if input_type == "request":
for text in texts:
await self._inspect_url_argument(client, text)
scan_failure: Final = await self._inspect(client, body=body)
if scan_failure is not None:
self._raise_failure(scan_failure)
params: Final = {
"timeout": 15,
"follow_redirects": False,
**(
{"transport": self._transport}
if self._transport is not None
else {"transport_factory": _verified_transport}
),
}
client: Final = get_async_httpx_client("ismalicious", params=params).client
if input_type == "request":
for text in texts:
await self._inspect_url_argument(client, text)
scan_failure: Final = await self._inspect(client, body=body)
if scan_failure is not None:
self._raise_failure(scan_failure)
return inputs

View file

@ -24,6 +24,7 @@ from litellm.llms.custom_httpx.http_handler import (
HTTPHandler,
MaskedHTTPStatusError,
_get_httpx_client,
get_async_httpx_client,
get_ssl_configuration,
)
from litellm.types.llms.custom_http import VerifyTypes
@ -1828,3 +1829,79 @@ async def test_http2_disabled_by_default(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(litellm, "disable_aiohttp_transport", False)
assert AsyncHTTPHandler._should_use_aiohttp_transport() is True
class _FactoryTransport(httpx.MockTransport):
def __init__(self, generation: int) -> None:
self.closed = False
super().__init__(lambda request: httpx.Response(200, json={"generation": generation}))
async def handle_async_request(self, request: httpx.Request) -> httpx.Response:
if self.closed:
raise httpx.ConnectError("Transport closed", request=request)
return await super().handle_async_request(request)
async def aclose(self) -> None:
self.closed = True
def _transport_factory() -> tuple[Callable[[], httpx.AsyncBaseTransport], list[_FactoryTransport]]:
transports: Final[list[_FactoryTransport]] = []
def create() -> httpx.AsyncBaseTransport:
transport: Final = _FactoryTransport(len(transports))
transports.append(transport)
return transport
return create, transports
@pytest.mark.asyncio
async def test_transport_factory_reuses_cached_client_and_refreshes_closed_generation() -> None:
factory, transports = _transport_factory()
params: Final = {"transport_factory": factory, "follow_redirects": False, "timeout": 15}
handler: Final = get_async_httpx_client("factory-fixture", params=params)
cached: Final = get_async_httpx_client("factory-fixture", params=params)
first: Final = handler.client
assert cached is handler
assert (await first.get("https://fixture.example/first")).json() == {"generation": 0}
assert len(transports) == 1
await first.aclose()
refreshed: Final = cached.client
try:
assert (await refreshed.get("https://fixture.example/next")).json() == {"generation": 1}
assert refreshed is not first
assert refreshed.timeout == httpx.Timeout(15)
assert not refreshed.follow_redirects
assert transports[0].closed
assert not transports[1].closed
finally:
await handler.close()
@pytest.mark.asyncio
async def test_transport_factory_isolates_replacement_from_retired_cached_client() -> None:
factory, transports = _transport_factory()
params: Final = {"transport_factory": factory}
retired: Final = get_async_httpx_client("factory-fixture", params=params)
assert (await retired.client.get("https://fixture.example/first")).json() == {"generation": 0}
litellm.in_memory_llm_clients_cache.flush_cache()
replacement: Final = get_async_httpx_client("factory-fixture", params=params)
await retired.close()
try:
assert replacement is not retired
assert transports[0].closed
assert not transports[1].closed
assert (await replacement.client.get("https://fixture.example/next")).json() == {"generation": 1}
finally:
await replacement.close()
def test_transport_factory_and_transport_are_mutually_exclusive() -> None:
with pytest.raises(ValueError, match="mutually exclusive"):
AsyncHTTPHandler(
transport=httpx.MockTransport(lambda request: httpx.Response(200)),
transport_factory=lambda: httpx.MockTransport(lambda request: httpx.Response(200)),
)

View file

@ -5,10 +5,12 @@ import httpx
import pytest
from mcp.types import CallToolResult, TextContent
import litellm
from litellm.exceptions import GuardrailRaisedException
from litellm.proxy._experimental.mcp_server.guardrail_translation.handler import MCPGuardrailTranslationHandler
from litellm.proxy.guardrails.guardrail_hooks.ismalicious import initialize_guardrail
from litellm.proxy.guardrails.guardrail_hooks.ismalicious.ismalicious import IsMaliciousGuardrail
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.guardrails import GuardrailEventHooks, LitellmParams, Mode
URL = "https://example.com/path?q=a,b&x=1#part"
KEY = base64.b64encode(b"test-key:test-secret").decode()
@ -53,12 +55,19 @@ async def test_native_pre_handler_preserves_original_url_and_arguments():
requests.append(request)
return service_response(request)
data = {"mcp_tool_name": "fetch", "mcp_arguments": {"url": URL, "other": "unchanged"}}
data = {
"mcp_tool_name": "fetch",
"mcp_arguments": {"url": URL, "other": "unchanged"},
"headers": {"Authorization": "Incoming private token"},
}
result = await MCPGuardrailTranslationHandler().process_input_messages(data, guardrail(handler))
assert result is data
assert requests[0].url.params["u"] == URL
assert json.loads(json.loads(requests[1].content)["content"]) == [URL, "unchanged"]
assert data["mcp_arguments"]["url"] == URL
assert all(request.headers["X-API-KEY"] == KEY for request in requests)
assert all("Authorization" not in request.headers for request in requests)
assert all(request.extensions["timeout"]["read"] == 15 for request in requests)
@pytest.mark.asyncio
@ -174,7 +183,9 @@ def test_unsupported_modes_fail_at_startup(mode):
IsMaliciousGuardrail(api_key=KEY, event_hook=mode)
@pytest.mark.parametrize("api_key", ["not-base64", "dXNlcg==", "OnNlY3JldA=="])
@pytest.mark.parametrize(
"api_key", ["not-base64", "dXNlcg==", "OnNlY3JldA==", base64.b64encode(b"\xff:secret").decode()]
)
def test_invalid_credentials_are_not_echoed(api_key):
with pytest.raises(ValueError, match="requires a Base64") as exc:
IsMaliciousGuardrail(api_key=api_key, event_hook=GuardrailEventHooks.pre_mcp_call)
@ -238,3 +249,121 @@ async def test_malformed_response_semantics_fail_closed(invalid):
with pytest.raises(GuardrailRaisedException) as exc:
await guardrail(handler).apply_guardrail({"texts": ["Untrusted"]}, {}, "response")
assert not exc.value.blocked_content
@pytest.mark.parametrize("mode", ["pre_mcp_call", ["pre_mcp_call", "post_mcp_call"]])
@pytest.mark.parametrize("default_on", [False, True])
def test_native_config_roundtrip_registers_requested_policy(mode, default_on):
config_model = IsMaliciousGuardrail.get_config_model()
config = config_model.model_validate({"api_key": KEY})
params = LitellmParams(
guardrail="ismalicious", mode=mode, default_on=default_on, **config.model_dump(exclude_none=True)
)
callback = initialize_guardrail(params, {"guardrail_name": "selected-gate", "litellm_params": params})
assert any(item is callback for item in litellm.callbacks)
assert callback.should_run_guardrail(data={}, event_type=GuardrailEventHooks.pre_mcp_call) is default_on
assert callback.should_run_guardrail(data={}, event_type=GuardrailEventHooks.post_mcp_call) is (
default_on and isinstance(mode, list)
)
def test_native_initializer_refuses_conditional_mcp_modes():
params = LitellmParams(guardrail="ismalicious", api_key=KEY, mode=Mode(tags={}, default="pre_mcp_call"))
with pytest.raises(ValueError, match="explicit MCP modes"):
initialize_guardrail(params, {"guardrail_name": "selected-gate", "litellm_params": params})
@pytest.mark.asyncio
async def test_valid_ordered_span_allows_original_content():
def handler(request):
body = service_response(request).json()
body["injection"]["spans"] = [{"start": 0, "end": 1, "family": "fixture"}]
return httpx.Response(200, json=body)
payload = CallToolResult(content=[TextContent(type="text", text="Allowed result")])
assert await MCPGuardrailTranslationHandler().process_output_response(payload, guardrail(handler)) is payload
@pytest.mark.asyncio
@pytest.mark.parametrize("url", ["https://[invalid", "https://user:password@example.com/", "https://example.com/\n"])
async def test_native_pre_handler_refuses_malformed_urls_before_network(url):
def handler(request):
pytest.fail("A malformed URL must not be sent")
with pytest.raises(GuardrailRaisedException) as exc:
await MCPGuardrailTranslationHandler().process_input_messages(
{"mcp_tool_name": "fetch", "mcp_arguments": {"url": url}}, guardrail(handler)
)
assert url not in str(exc.value)
@pytest.mark.asyncio
async def test_native_pre_handler_refuses_a_service_response_for_a_different_url():
def handler(request):
body = service_response(request).json()
body["url"] = "https://example.com/path"
return httpx.Response(200, json=body)
with pytest.raises(GuardrailRaisedException):
await MCPGuardrailTranslationHandler().process_input_messages(
{"mcp_tool_name": "fetch", "mcp_arguments": {"url": URL}}, guardrail(handler)
)
@pytest.mark.asyncio
async def test_non_utf8_content_is_refused_before_network():
def handler(request):
pytest.fail("Invalid UTF-8 must not be sent")
with pytest.raises(GuardrailRaisedException):
await guardrail(handler).apply_guardrail({"texts": ["\ud800"]}, {}, "response")
@pytest.mark.asyncio
@pytest.mark.parametrize("error", [httpx.ConnectError, httpx.RemoteProtocolError])
async def test_native_transport_connection_errors_are_not_retried(error):
calls = []
def handler(request):
calls.append(request)
raise error("Private upstream details", request=request)
with pytest.raises(GuardrailRaisedException) as exc:
await MCPGuardrailTranslationHandler().process_output_response(
CallToolResult(content=[TextContent(type="text", text="Untrusted result")]), guardrail(handler)
)
assert len(calls) == 1
assert "Private upstream details" not in str(exc.value)
class ClosureAwareTransport(httpx.MockTransport):
def __init__(self, handler):
super().__init__(handler)
self.closed = False
async def handle_async_request(self, request):
if self.closed:
raise httpx.ConnectError("Transport already closed", request=request)
return await super().handle_async_request(request)
async def aclose(self):
self.closed = True
await super().aclose()
@pytest.mark.asyncio
async def test_native_pool_survives_calls_without_sharing_credentials():
requests = []
def handler(request):
requests.append(request)
return service_response(request)
transport = ClosureAwareTransport(handler)
other_key = base64.b64encode(b"other-key:other-secret").decode()
for key in (KEY, other_key):
callback = IsMaliciousGuardrail(api_key=key, event_hook=GuardrailEventHooks.post_mcp_call, transport=transport)
payload = CallToolResult(content=[TextContent(type="text", text="Allowed result")])
assert await MCPGuardrailTranslationHandler().process_output_response(payload, callback) is payload
assert [request.headers["X-API-KEY"] for request in requests] == [KEY, other_key]
assert not transport.closed