mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge ee44c3dd36 into a2bf67a037
This commit is contained in:
commit
4e10527075
9 changed files with 787 additions and 3 deletions
41
cookbook/guardrails/ismalicious/README.md
Normal file
41
cookbook/guardrails/ismalicious/README.md
Normal file
|
|
@ -0,0 +1,41 @@
|
|||
# IsMalicious MCP guardrail
|
||||
|
||||
Inspect normalized MCP argument strings before tool execution and text/structured-content strings before a tool result is returned. IsMalicious checks URL reputation for standalone HTTP(S) URL arguments, then scans the complete normalized text list for content injection and links
|
||||
|
||||
## Configuration
|
||||
|
||||
Obtain an API key and secret from [your account](https://ismalicious.com/app/account). In your secret manager, set `ISMALICIOUS_ENCODED_API_KEY` to the Base64 encoding of `apiKey:apiSecret`, with no newline. Pass the secret explicitly through LiteLLM's `api_key` configuration as shown below; this provider does not read an ambient credential fallback. Do not log that value or put it directly into YAML
|
||||
|
||||
Merge [`config.yaml`](config.yaml) into the proxy configuration containing your MCP servers. It sets both MCP modes and `default_on: true`. MCP subcalls do not necessarily inherit a parent chat request's guardrail selection, so relying only on a `guardrails` field on the parent request does not establish MCP enforcement
|
||||
|
||||
```yaml
|
||||
guardrails:
|
||||
- guardrail_name: ismalicious-mcp
|
||||
litellm_params:
|
||||
guardrail: ismalicious
|
||||
mode: [pre_mcp_call, post_mcp_call]
|
||||
default_on: true
|
||||
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. 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
|
||||
|
||||
Top-level `warn` and `block` refuse. Timeout, quota429, redirects, invalid JSON/schema/verdict, incomplete link inspection and oversized serialized UTF-8 request bodies also refuse. A valid `allow` returns the original inputs without masking or sanitization. It means the service did not block under its current rules, not that content is proven benign. An unknown inner link reputation is a valid response and can still have an allow decision
|
||||
|
||||
The content request serializes the complete normalized string list into `content` and uses `mode: fast`. The serialized HTTP body must fit within 1 MiB; it is never truncated. Results in `structuredContent`, including labels and numeric fields emitted by the native handler, are included when that handler exposes them. These calls use the separate scan quota, not the indicator lookup quota. URL reputation does not fetch the destination
|
||||
|
||||
## Scope
|
||||
|
||||
This provider uses LiteLLM's native MCP guardrail translation. It supports scannable text fields only. The native handler can skip a response that has no scannable text and can omit binary or multimodal blocks alongside text. This integration therefore does not enforce binary/multimodal/streaming content policy. Route only text tools through this policy and exclude unsupported tools at the gateway. Streaming LLM/tool effects, model-native search, hidden tool network calls and direct network access are outside its guarantees
|
||||
|
||||
The normalized list is not the complete MCP envelope: transport metadata, content-block metadata, annotations, binary fields and fields not exposed by LiteLLM's string extraction are not inspected. Scanning this list must not be described as scanning every field in the original response
|
||||
|
||||
The result can have existed in process memory or logging structures before the post-call inspection. Disable message/content logging and prompt storage as shown in the sample, do not enable detailed/debug logging or install callbacks that expose raw tool output, and review your tracing configuration. The native logging decorator records a decision summary on success and the fixed refusal on failure, but native debug logging can include allowed content. Tool side effects cannot be undone. The failure message omits the raw text, URLs, credentials and upstream exception details. Normal MCP error responses may retain HTTP200 while indicating `isError`; clients must inspect the JSON-RPC/tool result rather than treating HTTP200 as an allow decision
|
||||
|
||||
## Validation
|
||||
|
||||
The mapped unit tests execute native pre/post MCP translation with an injected HTTP fixture. They check identity-preserving allowed inputs, original URL query/comma/fragment encoding, both refusal stages, structured-only results, warn, invalid verdicts, incomplete inspection, errors and UTF-8 limits without truncation. Synthetic replies validate integration policy, not the detector's accuracy
|
||||
|
||||
An authenticated live proxy run against the actual IsMalicious API and a real LLM provider is still required before asking for maintainer review under this repository's contribution rules. The local MCP/unit tests must not be described as that proof. No pricing, plan counts or detection-accuracy claim is embedded in this example
|
||||
13
cookbook/guardrails/ismalicious/config.yaml
Normal file
13
cookbook/guardrails/ismalicious/config.yaml
Normal file
|
|
@ -0,0 +1,13 @@
|
|||
guardrails:
|
||||
- guardrail_name: ismalicious-mcp
|
||||
litellm_params:
|
||||
guardrail: ismalicious
|
||||
mode: [pre_mcp_call, post_mcp_call]
|
||||
default_on: true
|
||||
api_key: os.environ/ISMALICIOUS_ENCODED_API_KEY
|
||||
|
||||
litellm_settings:
|
||||
turn_off_message_logging: true
|
||||
|
||||
general_settings:
|
||||
store_prompts_in_spend_logs: false
|
||||
|
|
@ -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(),
|
||||
|
|
@ -1750,7 +1755,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(),
|
||||
|
|
|
|||
|
|
@ -0,0 +1,25 @@
|
|||
from typing import Final
|
||||
|
||||
import litellm
|
||||
from litellm.types.guardrails import Guardrail, GuardrailEventHooks, LitellmParams, SupportedGuardrailIntegrations
|
||||
|
||||
from .ismalicious import IsMaliciousGuardrail
|
||||
|
||||
|
||||
def initialize_guardrail(litellm_params: LitellmParams, guardrail: Guardrail) -> IsMaliciousGuardrail:
|
||||
modes: Final = [litellm_params.mode] if isinstance(litellm_params.mode, str) else litellm_params.mode
|
||||
if not isinstance(modes, list):
|
||||
raise ValueError("IsMalicious requires explicit MCP modes")
|
||||
callback: Final = IsMaliciousGuardrail(
|
||||
api_key=litellm_params.api_key,
|
||||
api_base=litellm_params.api_base,
|
||||
guardrail_name=guardrail["guardrail_name"],
|
||||
event_hook=[GuardrailEventHooks(mode) for mode in modes],
|
||||
default_on=litellm_params.default_on is True,
|
||||
)
|
||||
litellm.logging_callback_manager.add_litellm_callback(callback) # pyright: ignore[reportUnknownMemberType] # manager exposes an untyped callable union
|
||||
return callback
|
||||
|
||||
|
||||
guardrail_initializer_registry: Final = {SupportedGuardrailIntegrations.ISMALICIOUS.value: initialize_guardrail}
|
||||
guardrail_class_registry: Final = {SupportedGuardrailIntegrations.ISMALICIOUS.value: IsMaliciousGuardrail}
|
||||
|
|
@ -0,0 +1,240 @@
|
|||
import base64
|
||||
import binascii
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Final, Literal, NoReturn
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict, Field, ValidationError, model_validator
|
||||
from typing_extensions import Self
|
||||
|
||||
from litellm.exceptions import GuardrailRaisedException
|
||||
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
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
|
||||
_API_BASE: Final = "https://api.ismalicious.com"
|
||||
_MAX_BODY_BYTES: Final = 1024 * 1024
|
||||
_REFUSAL: Final = "IsMalicious could not allow the inspected MCP content"
|
||||
|
||||
|
||||
class _Response(BaseModel):
|
||||
model_config = ConfigDict(extra="allow", strict=True)
|
||||
verdict: Literal["allow", "warn", "block"]
|
||||
latency_ms: int = Field(ge=0, le=2**63 - 1)
|
||||
|
||||
|
||||
class _Link(BaseModel):
|
||||
model_config = ConfigDict(extra="allow", strict=True)
|
||||
url: str
|
||||
entity: str
|
||||
verdict: Literal["clean", "suspicious", "malicious", "unknown"]
|
||||
sources: int = Field(ge=0)
|
||||
|
||||
|
||||
class _Span(BaseModel):
|
||||
model_config = ConfigDict(extra="allow", strict=True)
|
||||
start: int = Field(ge=0)
|
||||
end: int = Field(ge=0)
|
||||
family: str
|
||||
|
||||
@model_validator(mode="after")
|
||||
def ordered_interval(self) -> Self:
|
||||
if self.end < self.start:
|
||||
raise ValueError("Invalid injection span interval")
|
||||
return self
|
||||
|
||||
|
||||
class _Injection(BaseModel):
|
||||
model_config = ConfigDict(extra="allow", strict=True)
|
||||
score: float = Field(ge=0, le=1)
|
||||
families: list[str]
|
||||
spans: list[_Span]
|
||||
|
||||
|
||||
class _Scan(_Response):
|
||||
injection: _Injection
|
||||
links: list[_Link]
|
||||
links_truncated: bool
|
||||
mode: Literal["fast", "thorough"]
|
||||
source: _Link | None = None
|
||||
sanitized_content: str | None = None
|
||||
|
||||
|
||||
class _Url(_Response):
|
||||
url: str
|
||||
entity: str
|
||||
sources: int = Field(ge=0)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Failure:
|
||||
blocked_content: bool = False
|
||||
|
||||
|
||||
def _verdict(response: httpx.Response, url: str | None) -> _Failure | None:
|
||||
try:
|
||||
response.raise_for_status()
|
||||
parsed: Final = (_Scan if url is None else _Url).model_validate_json(response.content)
|
||||
except (httpx.HTTPError, ValidationError, ValueError):
|
||||
return _Failure()
|
||||
if parsed.verdict != "allow":
|
||||
return _Failure(blocked_content=True)
|
||||
if isinstance(parsed, _Scan) and parsed.links_truncated:
|
||||
return _Failure()
|
||||
if isinstance(parsed, _Url) and parsed.url != url:
|
||||
return _Failure()
|
||||
return None
|
||||
|
||||
|
||||
def _encoded_body(texts: list[str]) -> bytes | _Failure:
|
||||
try:
|
||||
content: Final = json.dumps(texts, ensure_ascii=False, separators=(",", ":"))
|
||||
body: Final = json.dumps({"content": content, "mode": "fast"}, ensure_ascii=False).encode("utf-8")
|
||||
except (UnicodeError, ValueError):
|
||||
return _Failure()
|
||||
return body if len(body) <= _MAX_BODY_BYTES else _Failure()
|
||||
|
||||
|
||||
def _url_argument(text: str) -> str | None | _Failure:
|
||||
if not text.startswith(("http://", "https://")):
|
||||
return None
|
||||
try:
|
||||
parsed: Final = urlsplit(text)
|
||||
except ValueError:
|
||||
return _Failure()
|
||||
if not parsed.hostname or parsed.username or parsed.password or any(ord(char) < 33 for char in text):
|
||||
return _Failure()
|
||||
return text
|
||||
|
||||
|
||||
def _verified_transport() -> httpx.AsyncBaseTransport:
|
||||
return httpx.AsyncHTTPTransport(verify=True, retries=0, trust_env=False)
|
||||
|
||||
|
||||
class IsMaliciousGuardrail(CustomGuardrail):
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
guardrail_name: str | None = None,
|
||||
event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | None = None,
|
||||
default_on: bool = False,
|
||||
transport: httpx.AsyncBaseTransport | None = None,
|
||||
) -> None:
|
||||
resolved_key: Final = api_key
|
||||
if not resolved_key:
|
||||
raise ValueError("IsMalicious requires a Base64 API key and secret pair")
|
||||
try:
|
||||
decoded_key: Final = base64.b64decode(resolved_key, validate=True).decode("utf-8")
|
||||
except (binascii.Error, UnicodeError):
|
||||
raise ValueError("IsMalicious requires a Base64 API key and secret pair") from None
|
||||
if ":" not in decoded_key or not all(decoded_key.split(":", 1)):
|
||||
raise ValueError("IsMalicious requires a Base64 API key and secret pair")
|
||||
if api_base is not None and api_base.rstrip("/") != _API_BASE:
|
||||
raise ValueError("IsMalicious credentials may only be sent to its HTTPS API")
|
||||
modes: Final = event_hook if isinstance(event_hook, list) else [event_hook]
|
||||
if not modes or any(mode not in self.get_supported_event_hooks() for mode in modes):
|
||||
raise ValueError("IsMalicious requires pre_mcp_call or post_mcp_call")
|
||||
self._headers = {"X-API-KEY": resolved_key, "Content-Type": "application/json"}
|
||||
self._transport = transport
|
||||
super().__init__( # pyright: ignore[reportUnknownMemberType] # base guardrail accepts untyped kwargs
|
||||
guardrail_name=guardrail_name,
|
||||
supported_event_hooks=self.get_supported_event_hooks(),
|
||||
event_hook=event_hook,
|
||||
default_on=default_on,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
|
||||
return [GuardrailEventHooks.pre_mcp_call, GuardrailEventHooks.post_mcp_call]
|
||||
|
||||
@staticmethod
|
||||
def get_config_model() -> type["GuardrailConfigModel[BaseModel]"] | None:
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.ismalicious import IsMaliciousConfigModel
|
||||
|
||||
return IsMaliciousConfigModel
|
||||
|
||||
def _raise_failure(self, failure: _Failure) -> NoReturn:
|
||||
raise GuardrailRaisedException(
|
||||
guardrail_name=self.guardrail_name,
|
||||
message=_REFUSAL,
|
||||
blocked_content=failure.blocked_content,
|
||||
)
|
||||
|
||||
async def _inspect(
|
||||
self, client: httpx.AsyncClient, *, body: bytes | None = None, url: str | None = None
|
||||
) -> _Failure | None:
|
||||
try:
|
||||
response: Final = (
|
||||
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(
|
||||
f"{_API_BASE}/gate/scan",
|
||||
content=body,
|
||||
headers=self._headers,
|
||||
timeout=15,
|
||||
follow_redirects=False,
|
||||
)
|
||||
)
|
||||
except httpx.HTTPError:
|
||||
return _Failure()
|
||||
return _verdict(response, url)
|
||||
|
||||
async def _inspect_url_argument(self, client: httpx.AsyncClient, text: str) -> None:
|
||||
url: Final = _url_argument(text)
|
||||
if isinstance(url, _Failure):
|
||||
self._raise_failure(url)
|
||||
if isinstance(url, str):
|
||||
failure: Final = await self._inspect(client, url=url)
|
||||
if failure is not None:
|
||||
self._raise_failure(failure)
|
||||
|
||||
@log_guardrail_information
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict[str, object],
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: "LiteLLMLoggingObj | None" = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
texts: Final = inputs.get("texts", [])
|
||||
if inputs.get("images") or not texts:
|
||||
self._raise_failure(_Failure())
|
||||
body: Final = _encoded_body(texts)
|
||||
if isinstance(body, _Failure):
|
||||
self._raise_failure(body)
|
||||
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
|
||||
|
|
@ -146,6 +146,7 @@ class SupportedGuardrailIntegrations(Enum):
|
|||
ALICE = "alice"
|
||||
AGENT_365 = "agent_365"
|
||||
CONDUCT = "conduct"
|
||||
ISMALICIOUS = "ismalicious"
|
||||
|
||||
|
||||
class Role(Enum):
|
||||
|
|
|
|||
|
|
@ -0,0 +1,11 @@
|
|||
from pydantic import BaseModel, Field
|
||||
|
||||
from .base import GuardrailConfigModel
|
||||
|
||||
|
||||
class IsMaliciousConfigModel(GuardrailConfigModel[BaseModel]):
|
||||
api_key: str | None = Field(default=None, description="Base64 encoding of the IsMalicious API key and secret pair")
|
||||
|
||||
@staticmethod
|
||||
def ui_friendly_name() -> str:
|
||||
return "IsMalicious"
|
||||
|
|
@ -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
|
||||
|
|
@ -1875,3 +1876,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)),
|
||||
)
|
||||
|
|
|
|||
369
tests/unit/proxy/guardrails/guardrail_hooks/test_ismalicious.py
Normal file
369
tests/unit/proxy/guardrails/guardrail_hooks/test_ismalicious.py
Normal file
|
|
@ -0,0 +1,369 @@
|
|||
import base64
|
||||
import json
|
||||
|
||||
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, LitellmParams, Mode
|
||||
|
||||
URL = "https://example.com/path?q=a,b&x=1#part"
|
||||
KEY = base64.b64encode(b"test-key:test-secret").decode()
|
||||
|
||||
|
||||
def service_response(request, verdict="allow", truncated=False):
|
||||
if request.url.path == "/gate/url":
|
||||
body = {
|
||||
"url": request.url.params["u"],
|
||||
"entity": "example.com",
|
||||
"verdict": verdict,
|
||||
"sources": 0,
|
||||
"latency_ms": 1,
|
||||
}
|
||||
else:
|
||||
body = {
|
||||
"verdict": verdict,
|
||||
"injection": {"score": 0.0, "families": [], "spans": []},
|
||||
"links": [{"url": URL, "entity": "example.com", "verdict": "unknown", "sources": 0}],
|
||||
"links_truncated": truncated,
|
||||
"mode": "fast",
|
||||
"latency_ms": 1,
|
||||
}
|
||||
return httpx.Response(200, json=body)
|
||||
|
||||
|
||||
def guardrail(handler=service_response):
|
||||
return IsMaliciousGuardrail(
|
||||
api_key=KEY,
|
||||
guardrail_name="ismalicious",
|
||||
event_hook=[GuardrailEventHooks.pre_mcp_call, GuardrailEventHooks.post_mcp_call],
|
||||
default_on=True,
|
||||
transport=httpx.MockTransport(handler),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_pre_handler_preserves_original_url_and_arguments():
|
||||
requests = []
|
||||
|
||||
def handler(request):
|
||||
requests.append(request)
|
||||
return service_response(request)
|
||||
|
||||
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
|
||||
async def test_native_post_handler_preserves_text_and_structured_content_identity():
|
||||
payload = CallToolResult(
|
||||
content=[TextContent(type="text", text="Allowed å🚀 content")],
|
||||
structuredContent={"context": "Additional untrusted context", "count": 3},
|
||||
)
|
||||
requests = []
|
||||
|
||||
def handler(request):
|
||||
requests.append(request)
|
||||
return service_response(request)
|
||||
|
||||
original_content = payload.content
|
||||
original_structured = payload.structured_content
|
||||
result = await MCPGuardrailTranslationHandler().process_output_response(payload, guardrail(handler))
|
||||
assert result is payload
|
||||
assert result.content is original_content
|
||||
assert result.structured_content is original_structured
|
||||
inspected = json.loads(json.loads(requests[0].content)["content"])
|
||||
assert "Allowed å🚀 content" in inspected
|
||||
assert "Additional untrusted context" in inspected
|
||||
assert "context" in inspected and "count" in inspected and "3" in inspected
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("phase", ["url", "scan"])
|
||||
@pytest.mark.parametrize("verdict", ["warn", "block", "unknown"])
|
||||
async def test_native_pre_handler_does_not_return_a_blocked_request(phase, verdict):
|
||||
def handler(request):
|
||||
return service_response(request, verdict if request.url.path.endswith(phase) else "allow")
|
||||
|
||||
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) and KEY not in str(exc.value)
|
||||
assert exc.value.blocked_content is (verdict in {"warn", "block"})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("verdict", ["warn", "block", "unknown"])
|
||||
async def test_native_post_handler_does_not_return_blocked_text(verdict):
|
||||
payload = CallToolResult(content=[TextContent(type="text", text="Private untrusted result")])
|
||||
|
||||
def handler(request):
|
||||
return service_response(request, verdict)
|
||||
|
||||
with pytest.raises(GuardrailRaisedException) as exc:
|
||||
await MCPGuardrailTranslationHandler().process_output_response(payload, guardrail(handler))
|
||||
assert "Private untrusted result" not in str(exc.value)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("failure", ["429", "302", "500", "bad-json", "missing", "timeout", "truncated"])
|
||||
async def test_native_post_handler_fails_closed_without_retry(failure):
|
||||
calls = []
|
||||
|
||||
def handler(request):
|
||||
calls.append(request)
|
||||
if failure.isdigit():
|
||||
return httpx.Response(int(failure), headers={"Location": "https://attacker.example"})
|
||||
if failure == "bad-json":
|
||||
return httpx.Response(200, content=b"not-json")
|
||||
if failure == "missing":
|
||||
return httpx.Response(200, json={"verdict": "allow"})
|
||||
if failure == "timeout":
|
||||
raise httpx.ReadTimeout("Sensitive detail", request=request)
|
||||
return service_response(request, truncated=True)
|
||||
|
||||
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 not exc.value.blocked_content
|
||||
assert "Sensitive detail" not in str(exc.value)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_serialized_utf8_request_body_limit_refuses_without_network():
|
||||
def handler(request):
|
||||
pytest.fail("An oversized request must not be sent")
|
||||
|
||||
with pytest.raises(GuardrailRaisedException):
|
||||
await guardrail(handler).apply_guardrail(
|
||||
inputs={"texts": ["🚀" * (1024 * 1024 // 4)]}, request_data={}, input_type="response"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_handler_scans_structured_only_untrusted_content():
|
||||
def handler(request):
|
||||
assert "Hidden instruction" in json.loads(request.content)["content"]
|
||||
return service_response(request, "block")
|
||||
|
||||
with pytest.raises(GuardrailRaisedException):
|
||||
await MCPGuardrailTranslationHandler().process_output_response(
|
||||
CallToolResult(content=[], structuredContent={"instruction": "Hidden instruction"}), guardrail(handler)
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_return_inputs_are_identity_without_rewriting():
|
||||
inputs = {"texts": ["Allowed text"]}
|
||||
assert await guardrail().apply_guardrail(inputs, {}, "response") is inputs
|
||||
|
||||
|
||||
@pytest.mark.parametrize("mode", [GuardrailEventHooks.pre_call, GuardrailEventHooks.during_mcp_call, None, []])
|
||||
def test_unsupported_modes_fail_at_startup(mode):
|
||||
with pytest.raises(ValueError, match="requires pre_mcp_call"):
|
||||
IsMaliciousGuardrail(api_key=KEY, event_hook=mode)
|
||||
|
||||
|
||||
@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)
|
||||
assert api_key not in str(exc.value)
|
||||
|
||||
|
||||
def test_credentials_never_redirect_to_a_custom_endpoint():
|
||||
with pytest.raises(ValueError, match="HTTPS API"):
|
||||
IsMaliciousGuardrail(
|
||||
api_key=KEY, api_base="https://attacker.example", event_hook=GuardrailEventHooks.pre_mcp_call
|
||||
)
|
||||
|
||||
|
||||
def test_missing_explicit_credentials_fail_at_startup():
|
||||
with pytest.raises(ValueError, match="requires a Base64"):
|
||||
IsMaliciousGuardrail(event_hook=GuardrailEventHooks.pre_mcp_call)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("verdict", ["allow", "block"])
|
||||
async def test_native_decision_logging_excludes_content_and_credentials(verdict):
|
||||
def handler(request):
|
||||
return service_response(request, verdict)
|
||||
|
||||
data = {"request_id": "test"}
|
||||
inputs = {"texts": ["Private untrusted result"]}
|
||||
if verdict == "block":
|
||||
with pytest.raises(GuardrailRaisedException):
|
||||
await guardrail(handler).apply_guardrail(inputs=inputs, request_data=data, input_type="response")
|
||||
else:
|
||||
await guardrail(handler).apply_guardrail(inputs=inputs, request_data=data, input_type="response")
|
||||
recorded = json.dumps(data)
|
||||
assert "standard_logging_guardrail_information" in recorded
|
||||
assert "Private untrusted result" not in recorded
|
||||
assert URL not in recorded and KEY not in recorded
|
||||
|
||||
|
||||
def test_mcp_subcalls_without_guardrail_metadata_use_explicit_default_on():
|
||||
gate = guardrail()
|
||||
assert gate.should_run_guardrail(data={}, event_type=GuardrailEventHooks.pre_mcp_call)
|
||||
assert gate.should_run_guardrail(data={}, event_type=GuardrailEventHooks.post_mcp_call)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_images_passed_to_the_provider_are_not_silently_allowed():
|
||||
with pytest.raises(GuardrailRaisedException):
|
||||
await guardrail().apply_guardrail({"texts": ["caption"], "images": ["image"]}, {}, "response")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("invalid", ["backward-span", "latency-overflow"])
|
||||
async def test_malformed_response_semantics_fail_closed(invalid):
|
||||
def handler(request):
|
||||
body = service_response(request).json()
|
||||
if invalid == "backward-span":
|
||||
body["injection"]["spans"] = [{"start": 9, "end": 2, "family": "instruction"}]
|
||||
else:
|
||||
body["latency_ms"] = 2**63
|
||||
return httpx.Response(200, json=body)
|
||||
|
||||
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