This commit is contained in:
JVQ 2026-10-04 23:11:26 +08:00 • committed by GitHub
commit 4e10527075
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 787 additions and 3 deletions

View 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

View 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

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(),
@ -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(),

View file

@ -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}

View file

@ -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

View file

@ -146,6 +146,7 @@ class SupportedGuardrailIntegrations(Enum):
ALICE = "alice"
AGENT_365 = "agent_365"
CONDUCT = "conduct"
ISMALICIOUS = "ismalicious"
class Role(Enum):

View file

@ -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"

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
@ -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)),
)

View 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