mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(guardrails): reuse native clients with isolated transports
This commit is contained in:
parent
7e2df548d3
commit
ee44c3dd36
5 changed files with 257 additions and 24 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue