This commit is contained in:
Mateo Wang 2026-09-12 09:41:15 -07:00 • committed by GitHub
commit 48f7ca1b97
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 835 additions and 125 deletions

View file

@ -1,5 +1,9 @@
import importlib
import importlib.util
import os
from collections.abc import Iterable, Iterator, Mapping
from dataclasses import dataclass
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
from litellm._logging import verbose_logger
@ -80,93 +84,181 @@ def get_cost_for_web_search_request(custom_llm_provider: str, usage: "Usage", mo
return None
def discover_guardrail_translation_mappings() -> dict[CallTypes, type["BaseTranslation"]]:
_GUARDRAIL_TRANSLATION_PACKAGE: Final = "guardrail_translation"
_MCP_GUARDRAIL_TRANSLATION_MODULE: Final = "litellm.proxy._experimental.mcp_server.guardrail_translation"
_NO_MAPPINGS: Final[Mapping[CallTypes, type["BaseTranslation"]]] = MappingProxyType({})
@dataclass(frozen=True, slots=True)
class GuardrailTranslationDiscovery:
"""
The outcome of one scan for guardrail translation handlers.
unavailable maps each package that failed to import to the reason, which is what tells a complete
result apart from one that is missing handlers and therefore has to be retried.
"""
mappings: Mapping[CallTypes, type["BaseTranslation"]]
unavailable: Mapping[str, str]
def _bundled_guardrail_translation_modules() -> Iterator[str]:
"""Yield the import path of every guardrail_translation package shipped under litellm/llms."""
llms_dir: Final = os.path.dirname(__file__)
for root, dirs, files in os.walk(llms_dir):
dirs[:] = tuple(d for d in dirs if not d.startswith("__") and d != "base_llm")
if os.path.basename(root) == _GUARDRAIL_TRANSLATION_PACKAGE and "__init__.py" in files:
yield "litellm." + os.path.relpath(root, os.path.dirname(llms_dir)).replace(os.sep, ".")
@dataclass(frozen=True, slots=True)
class _UnavailablePackage:
"""
Why one guardrail_translation package could not be imported.
missing_dependency is set when the package asked for a module this install does not have at all, which is
the one failure that says the package is absent rather than momentarily unimportable.
"""
reason: str
missing_dependency: bool
def _is_absent(module_name: str) -> bool:
"""Whether this install has no module of that name, as opposed to one that is present but failed to import."""
root: Final = module_name.partition(".")[0]
try:
return importlib.util.find_spec(root) is None
except (ImportError, ValueError):
return False
def _import_guardrail_translations(
module_path: str,
) -> Mapping[CallTypes, type["BaseTranslation"]] | _UnavailablePackage:
"""Import one guardrail_translation package, reporting why that failed instead of raising."""
try:
module: Final = importlib.import_module(module_path)
except Exception as e: # noqa: BLE001 # a package failing at import time for any reason is unavailable, not fatal
return _UnavailablePackage(
reason=f"{type(e).__name__}: {e}",
missing_dependency=isinstance(e, ModuleNotFoundError) and _is_absent(e.name or ""),
)
mappings: Final = getattr(module, "guardrail_translation_mappings", None)
if not isinstance(mappings, dict):
return _NO_MAPPINGS
declared: Final[Mapping[CallTypes, type[BaseTranslation]]] = mappings
return declared
def _guardrail_translations_from(
module_path: str,
) -> Mapping[CallTypes, type["BaseTranslation"]] | _UnavailablePackage:
"""
Import one package's handlers, tolerating an install that does not ship the optional MCP server.
litellm ships every package under llms, so a failure there is a gap to retry rather than a fact about the
install. The MCP package instead arrives with the proxy extra, and a dependency this install does not have
at all means it serves no MCP endpoints for a guardrail to scan, so there is nothing to retry or report. A
dependency that is installed and still fails to import is a broken install, which is reported and retried.
"""
result: Final = _import_guardrail_translations(module_path)
if (
module_path == _MCP_GUARDRAIL_TRANSLATION_MODULE
and isinstance(result, _UnavailablePackage)
and result.missing_dependency
):
verbose_logger.debug("%s is not installed: %s", module_path, result.reason)
return _NO_MAPPINGS
return result
def _guardrail_translation_modules() -> Iterator[str]:
"""Yield every module that can declare guardrail translation handlers, the optional MCP one last."""
yield from _bundled_guardrail_translation_modules()
yield _MCP_GUARDRAIL_TRANSLATION_MODULE
def _discover(
module_paths: Iterable[str], already_found: Mapping[CallTypes, type["BaseTranslation"]]
) -> GuardrailTranslationDiscovery:
imported: Final = tuple((module_path, _guardrail_translations_from(module_path)) for module_path in module_paths)
found: Final = (
already_found,
*(result for _, result in imported if not isinstance(result, _UnavailablePackage)),
)
return GuardrailTranslationDiscovery(
mappings=MappingProxyType(
{call_type: handler for mappings in found for call_type, handler in mappings.items()}
),
unavailable=MappingProxyType(
{module_path: result.reason for module_path, result in imported if isinstance(result, _UnavailablePackage)}
),
)
def discover_guardrail_translations() -> GuardrailTranslationDiscovery:
"""
Scan the llms tree, plus the optional MCP package, for guardrail translation handlers.
Returns:
GuardrailTranslationDiscovery: the handlers found, and the packages that failed to import
"""
return _discover(_guardrail_translation_modules(), already_found=_NO_MAPPINGS)
def discover_guardrail_translation_mappings() -> Mapping[CallTypes, type["BaseTranslation"]]:
"""
Discover guardrail translation mappings by scanning the llms directory structure.
Scans for modules with guardrail_translation_mappings dictionaries and aggregates them.
Returns:
Dict[CallTypes, Type[BaseTranslation]]: A dictionary mapping call types to their translation handler classes
Mapping[CallTypes, Type[BaseTranslation]]: the call types that have a translation handler class
"""
discovered_mappings: Final[dict[CallTypes, type[BaseTranslation]]] = {}
return discover_guardrail_translations().mappings
try:
# Get the path to the llms directory
current_dir: Final = os.path.dirname(__file__)
llms_dir: Final = current_dir
if not os.path.exists(llms_dir):
verbose_logger.debug("llms directory not found")
return discovered_mappings
# Recursively scan for guardrail_translation directories
for root, dirs, files in os.walk(llms_dir):
# Skip __pycache__ and base_llm directories
dirs[:] = [d for d in dirs if not d.startswith("__") and d != "base_llm"]
# Check if this is a guardrail_translation directory with __init__.py
if os.path.basename(root) == "guardrail_translation" and "__init__.py" in files:
# Build the module path relative to litellm
rel_path = os.path.relpath(root, os.path.dirname(llms_dir))
module_path = "litellm." + rel_path.replace(os.sep, ".")
try:
# Import the module
verbose_logger.debug("Discovering guardrail translations in: %s", module_path)
module = importlib.import_module(module_path)
# Check for guardrail_translation_mappings dictionary
if hasattr(module, "guardrail_translation_mappings"):
mappings = getattr(module, "guardrail_translation_mappings")
if isinstance(mappings, dict):
discovered_mappings.update(mappings)
verbose_logger.debug(
"Found guardrail_translation_mappings in %s: %s", module_path, list(mappings.keys())
)
except ImportError as e:
verbose_logger.error("Could not import %s: %s", module_path, e)
continue
except Exception as e:
verbose_logger.error("Error processing %s: %s", module_path, e)
continue
try:
from litellm.proxy._experimental.mcp_server.guardrail_translation import (
guardrail_translation_mappings as mcp_guardrail_translation_mappings,
)
discovered_mappings.update(mcp_guardrail_translation_mappings)
verbose_logger.debug(
"Loaded MCP guardrail translation mappings: %s",
list(mcp_guardrail_translation_mappings.keys()),
)
except ImportError:
verbose_logger.debug("MCP guardrail translation mappings not available; skipping")
verbose_logger.debug(
"Discovered %s guardrail translation mappings: %s",
len(discovered_mappings),
list(discovered_mappings.keys()),
def _announce_discovery(previous: GuardrailTranslationDiscovery | None, current: GuardrailTranslationDiscovery) -> None:
if previous is None and not current.unavailable:
return
if previous is None:
verbose_logger.error(
"Could not import guardrail translation handlers from %s; guardrails cannot run for their call types "
"until the import succeeds, which every lookup retries. %s",
", ".join(current.unavailable),
"; ".join(f"{module_path}: {reason}" for module_path, reason in current.unavailable.items()),
)
except Exception as e:
verbose_logger.error("Error discovering guardrail translation mappings: %s", e)
return discovered_mappings
return
recovered: Final = tuple(
module_path for module_path in previous.unavailable if module_path not in current.unavailable
)
if not recovered:
return
verbose_logger.warning("Guardrail translation handlers from %s are available again.", ", ".join(recovered))
# Cache the discovered mappings
endpoint_guardrail_translation_mappings: dict[CallTypes, type["BaseTranslation"]] | None = None
guardrail_translation_discovery: GuardrailTranslationDiscovery | None = None
def load_guardrail_translation_mappings():
global endpoint_guardrail_translation_mappings
if endpoint_guardrail_translation_mappings is None:
endpoint_guardrail_translation_mappings = discover_guardrail_translation_mappings()
return endpoint_guardrail_translation_mappings
def load_guardrail_translation_mappings() -> Mapping[CallTypes, type["BaseTranslation"]]:
"""
Return the guardrail translation handlers, retrying any bundled package that could not be imported last time.
Serving an incomplete scan as if it were complete would silently strip the missing call types off every
guardrail for the rest of the process, so the packages that failed are imported again on each lookup, and
only the part that succeeded is kept.
"""
global guardrail_translation_discovery
cached: Final = guardrail_translation_discovery
if cached is not None and not cached.unavailable:
return cached.mappings
discovery: Final = (
discover_guardrail_translations()
if cached is None
else _discover(cached.unavailable, already_found=cached.mappings)
)
_announce_discovery(previous=cached, current=discovery)
guardrail_translation_discovery = discovery
return discovery.mappings
def get_guardrail_translation_mapping(call_type: CallTypes) -> type["BaseTranslation"]:
@ -182,18 +274,10 @@ def get_guardrail_translation_mapping(call_type: CallTypes) -> type["BaseTransla
Raises:
ValueError: If no translation mapping exists for the given call type
"""
global endpoint_guardrail_translation_mappings
# Lazy load the mappings on first access
if endpoint_guardrail_translation_mappings is None:
endpoint_guardrail_translation_mappings = discover_guardrail_translation_mappings()
# Get the translation handler class for the call type
if call_type not in endpoint_guardrail_translation_mappings:
mappings: Final = load_guardrail_translation_mappings()
if call_type not in mappings:
raise ValueError(
f"No guardrail translation mapping found for call_type: {call_type}. "
f"Available mappings: {list(endpoint_guardrail_translation_mappings.keys())}"
f"Available mappings: {list(mappings.keys())}"
)
# Return the handler class directly
return endpoint_guardrail_translation_mappings[call_type]
return mappings[call_type]

View file

@ -9,6 +9,7 @@ Unified Guardrail, leveraging LiteLLM's /applyGuardrail endpoint
import copy
import json
from collections.abc import AsyncGenerator, AsyncIterable, Awaitable, Callable, Mapping, Sequence
from functools import lru_cache
from typing import TYPE_CHECKING, Any, Final, Protocol
from fastapi import HTTPException
@ -163,6 +164,72 @@ def _ensure_litellm_metadata(data: dict, user_api_key_dict: UserAPIKeyAuth) -> N
data["litellm_metadata"] = user_metadata
_UNSCANNED_WARNING_KEYS: Final = 4096
def _resolved_call_type(call_type: str | None) -> CallTypes | None:
"""Return the CallTypes member a route's call type names, or None when the enum has no member for it."""
try:
return CallTypes(call_type)
except ValueError:
return None
@lru_cache(maxsize=_UNSCANNED_WARNING_KEYS)
def _warn_left_unscanned_once(
guardrail_name: str | None,
request_route: str | None,
call_type: str | None,
consequence: str,
) -> None:
if _resolved_call_type(call_type) is not None:
verbose_proxy_logger.warning(
"Guardrail '%s' selected for route '%s' but call type '%s' has no guardrail translation handler; %s. "
"Add a guardrail translation handler for that call type.",
guardrail_name,
request_route,
call_type,
consequence,
)
return
unscannable, remedy = (
(
f"call type '{call_type}' is not one litellm can scan",
"Map the route to a CallTypes member in API_ROUTE_TO_CALL_TYPES",
)
if call_type
else ("its call type could not be resolved", "Add the route to API_ROUTE_TO_CALL_TYPES")
)
verbose_proxy_logger.warning(
"Guardrail '%s' selected for route '%s' but %s, so no guardrail can run on that route; %s. %s.",
guardrail_name,
request_route,
unscannable,
consequence,
remedy,
)
def _warn_left_unscanned(
guardrail_to_apply: "CustomGuardrail",
user_api_key_dict: UserAPIKeyAuth,
call_type: str | None,
consequence: str,
) -> None:
"""
Say that a selected guardrail could not scan this request, once per route and reason.
Most proxy routes have no translation handler and never will, so a line per request would bury the
outage it is meant to surface, and repeating it adds nothing an operator can act on twice.
"""
_warn_left_unscanned_once(
guardrail_name=guardrail_to_apply.guardrail_name,
request_route=user_api_key_dict.request_route,
call_type=call_type,
consequence=consequence,
)
class UnifiedLLMGuardrails(CustomLogger):
def __init__(
self,
@ -213,14 +280,17 @@ class UnifiedLLMGuardrails(CustomLogger):
return data
mappings: Final = load_guardrail_translation_mappings()
resolved_call_type: Final = _resolved_call_type(call_type)
if resolved_call_type is None or resolved_call_type not in mappings:
_warn_left_unscanned(
guardrail_to_apply=guardrail_to_apply,
user_api_key_dict=user_api_key_dict,
call_type=call_type,
consequence="skipping pre-call scanning",
)
return data
try:
if CallTypes(call_type) not in mappings:
return data
except ValueError:
return data # handle unmapped call types
endpoint_translation: Final = _as_endpoint_translation(mappings[CallTypes(call_type)]())
endpoint_translation: Final = _as_endpoint_translation(mappings[resolved_call_type]())
_ensure_litellm_metadata(data, user_api_key_dict)
@ -263,10 +333,17 @@ class UnifiedLLMGuardrails(CustomLogger):
return data
mappings: Final = load_guardrail_translation_mappings()
if call_type is not None and CallTypes(call_type) not in mappings:
resolved_call_type: Final = _resolved_call_type(call_type)
if resolved_call_type is None or resolved_call_type not in mappings:
_warn_left_unscanned(
guardrail_to_apply=guardrail_to_apply,
user_api_key_dict=user_api_key_dict,
call_type=call_type,
consequence="skipping during-call scanning",
)
return data
endpoint_translation: Final = _as_endpoint_translation(mappings[CallTypes(call_type)]())
endpoint_translation: Final = _as_endpoint_translation(mappings[resolved_call_type]())
_ensure_litellm_metadata(data, user_api_key_dict)
@ -311,7 +388,7 @@ class UnifiedLLMGuardrails(CustomLogger):
if user_api_key_dict.request_route is not None:
call_types: Final = get_call_types_for_route(user_api_key_dict.request_route)
if call_types is not None and len(call_types) > 0:
call_type = call_types[0]
call_type = call_types[0].value
if call_type is None:
call_type = _infer_call_type(call_type=None, completion_response=response)
@ -327,28 +404,18 @@ class UnifiedLLMGuardrails(CustomLogger):
):
call_type = logging_call_type
if call_type is None:
verbose_proxy_logger.warning(
"Guardrail '%s' selected for route '%s' but its call type could not be resolved; "
"skipping post-call scanning. Add the route to API_ROUTE_TO_CALL_TYPES.",
guardrail_to_apply.guardrail_name,
user_api_key_dict.request_route,
)
return response
mappings: Final = load_guardrail_translation_mappings()
if CallTypes(call_type) not in mappings:
verbose_proxy_logger.warning(
"Guardrail '%s' selected for route '%s' but call type '%s' has no guardrail translation handler; "
"skipping post-call scanning.",
guardrail_to_apply.guardrail_name,
user_api_key_dict.request_route,
call_type,
resolved_call_type: Final = _resolved_call_type(call_type)
if resolved_call_type is None or resolved_call_type not in mappings:
_warn_left_unscanned(
guardrail_to_apply=guardrail_to_apply,
user_api_key_dict=user_api_key_dict,
call_type=call_type,
consequence="skipping post-call scanning",
)
return response
endpoint_translation: Final = _as_endpoint_translation(mappings[CallTypes(call_type)]())
endpoint_translation: Final = _as_endpoint_translation(mappings[resolved_call_type]())
try:
response = await endpoint_translation.process_output_response(
@ -450,7 +517,7 @@ class UnifiedLLMGuardrails(CustomLogger):
has no in-stream error frame) the exception is re-raised so the proxy
can report it with a real HTTP status.
"""
if call_type is not None and CallTypes(call_type) in A2A_CALL_TYPES:
if _resolved_call_type(call_type) in A2A_CALL_TYPES:
yield _a2a_jsonrpc_error_chunk(exc, _get_a2a_request_id(responses_so_far, request_data))
return
if stream_started and endpoint_translation is not None:
@ -1046,7 +1113,13 @@ class UnifiedLLMGuardrails(CustomLogger):
call_type = _infer_call_type(call_type=None, completion_response=item)
# If call type not supported, just pass through all chunks
if call_type is None or CallTypes(call_type) not in mappings:
if _resolved_call_type(call_type) not in mappings:
_warn_left_unscanned(
guardrail_to_apply=guardrail_to_apply,
user_api_key_dict=user_api_key_dict,
call_type=call_type,
consequence="streaming this response to the client unscanned",
)
yield item
async for remaining_item in response:
yield remaining_item
@ -1152,14 +1225,15 @@ class UnifiedLLMGuardrails(CustomLogger):
yield item
# Stream has ended - do final processing with all collected chunks
if call_type is not None and CallTypes(call_type) in mappings:
final_call_type: Final = _resolved_call_type(call_type)
if final_call_type is not None and final_call_type in mappings:
verbose_proxy_logger.debug(
"Processing final streaming response with all %s chunks for guardrail %s",
len(responses_so_far),
guardrail_to_apply.guardrail_name,
)
endpoint_translation = mappings[CallTypes(call_type)]()
endpoint_translation = mappings[final_call_type]()
# When buffering, snapshot the original chunks before moderation.
# A shallow copy suffices: end-of-stream

View file

@ -0,0 +1,210 @@
import importlib.abc
import importlib.util
import logging
import sys
from collections.abc import Iterator
from contextlib import contextmanager
from types import ModuleType
import pytest
import litellm.llms as llms_package
from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.guardrail_translation.handler import (
MCPGuardrailTranslationHandler,
)
from litellm.types.utils import CallTypes
OPENAI_CHAT_TRANSLATION_MODULE = "litellm.llms.openai.chat.guardrail_translation"
MCP_TRANSLATION_MODULE = "litellm.proxy._experimental.mcp_server.guardrail_translation"
class RaisingLoader(importlib.abc.MetaPathFinder, importlib.abc.Loader):
"""Serves one module path, and raises the given error when Python executes it."""
def __init__(self, module_path: str, error: BaseException) -> None:
self.module_path = module_path
self.error = error
def find_spec(self, fullname: str, path=None, target=None):
if fullname != self.module_path:
return None
return importlib.util.spec_from_loader(fullname, self)
def create_module(self, spec) -> ModuleType | None:
return None
def exec_module(self, module: ModuleType) -> None:
raise self.error
@contextmanager
def raising_on_import(module_path: str, error: BaseException) -> Iterator[None]:
with pytest.MonkeyPatch.context() as mp:
mp.delitem(sys.modules, module_path, raising=False)
mp.setattr(sys, "meta_path", [RaisingLoader(module_path, error), *sys.meta_path])
yield
@contextmanager
def unimportable(module_path: str) -> Iterator[None]:
with pytest.MonkeyPatch.context() as mp:
mp.setitem(sys.modules, module_path, None)
yield
@contextmanager
def capturing(caplog: pytest.LogCaptureFixture, level: int) -> Iterator[None]:
with pytest.MonkeyPatch.context() as mp:
mp.setattr(verbose_logger, "propagate", True)
caplog.set_level(level, logger=verbose_logger.name)
yield
@pytest.fixture(autouse=True)
def reset_guardrail_translation_discovery():
saved = llms_package.guardrail_translation_discovery
llms_package.guardrail_translation_discovery = None
yield
llms_package.guardrail_translation_discovery = saved
def test_discovery_reports_the_handler_package_it_could_not_import():
with unimportable(OPENAI_CHAT_TRANSLATION_MODULE):
discovery = llms_package.discover_guardrail_translations()
assert tuple(discovery.unavailable) == (OPENAI_CHAT_TRANSLATION_MODULE,)
assert "None in sys.modules" in discovery.unavailable[OPENAI_CHAT_TRANSLATION_MODULE]
assert CallTypes.acompletion not in discovery.mappings
assert CallTypes.completion not in discovery.mappings
assert CallTypes.aembedding in discovery.mappings
def test_complete_discovery_reports_nothing_unavailable():
discovery = llms_package.discover_guardrail_translations()
assert not discovery.unavailable
assert CallTypes.acompletion in discovery.mappings
def test_the_next_lookup_retries_a_package_that_failed_to_import():
with unimportable(OPENAI_CHAT_TRANSLATION_MODULE):
partial = llms_package.load_guardrail_translation_mappings()
assert CallTypes.acompletion not in partial
assert CallTypes.aembedding in partial
recovered = llms_package.load_guardrail_translation_mappings()
assert CallTypes.acompletion in recovered
assert CallTypes.completion in recovered
assert CallTypes.aembedding in recovered
assert not llms_package.guardrail_translation_discovery.unavailable
def test_a_complete_discovery_is_not_scanned_again():
healthy = llms_package.load_guardrail_translation_mappings()
assert CallTypes.acompletion in healthy
with unimportable(OPENAI_CHAT_TRANSLATION_MODULE):
after_the_package_breaks = llms_package.load_guardrail_translation_mappings()
assert CallTypes.acompletion in after_the_package_breaks
assert not llms_package.guardrail_translation_discovery.unavailable
def test_a_package_that_keeps_failing_is_reported_once_and_its_recovery_announced(caplog):
with capturing(caplog, logging.INFO), unimportable(OPENAI_CHAT_TRANSLATION_MODULE):
for _ in range(3):
llms_package.load_guardrail_translation_mappings()
errors = [record for record in caplog.records if record.levelno == logging.ERROR]
assert len(errors) == 1, [record.getMessage() for record in errors]
assert OPENAI_CHAT_TRANSLATION_MODULE in errors[0].getMessage()
assert "None in sys.modules" in errors[0].getMessage()
caplog.clear()
with capturing(caplog, logging.INFO):
llms_package.load_guardrail_translation_mappings()
llms_package.load_guardrail_translation_mappings()
recoveries = [record for record in caplog.records if "available again" in record.getMessage()]
assert len(recoveries) == 1
assert OPENAI_CHAT_TRANSLATION_MODULE in recoveries[0].getMessage()
assert recoveries[0].levelno >= logging.WARNING
assert not [record for record in caplog.records if record.levelno == logging.ERROR]
def test_lookup_recovers_after_a_failed_discovery():
with unimportable(OPENAI_CHAT_TRANSLATION_MODULE):
with pytest.raises(ValueError, match="acompletion"):
llms_package.get_guardrail_translation_mapping(CallTypes.acompletion)
assert llms_package.get_guardrail_translation_mapping(CallTypes.acompletion) is not None
def test_an_mcp_package_that_fails_to_import_is_reported_and_retried():
with raising_on_import(MCP_TRANSLATION_MODULE, AttributeError("module 'mcp.types' has no attribute 'ToolCall'")):
partial = llms_package.load_guardrail_translation_mappings()
assert CallTypes.call_mcp_tool not in partial
assert CallTypes.acompletion in partial
assert tuple(llms_package.guardrail_translation_discovery.unavailable) == (MCP_TRANSLATION_MODULE,)
assert "AttributeError" in llms_package.guardrail_translation_discovery.unavailable[MCP_TRANSLATION_MODULE]
recovered = llms_package.load_guardrail_translation_mappings()
assert CallTypes.call_mcp_tool in recovered
assert CallTypes.acompletion in recovered
assert not llms_package.guardrail_translation_discovery.unavailable
def test_an_install_without_the_mcp_server_is_discovered_once_and_quietly(caplog):
absent = ModuleNotFoundError("No module named 'mcp'", name="mcp")
with capturing(caplog, logging.DEBUG), unimportable("mcp"), raising_on_import(MCP_TRANSLATION_MODULE, absent):
first = llms_package.load_guardrail_translation_mappings()
second = llms_package.load_guardrail_translation_mappings()
assert first is second
assert CallTypes.call_mcp_tool not in first
assert CallTypes.acompletion in first
assert not llms_package.guardrail_translation_discovery.unavailable
assert not [record for record in caplog.records if record.levelno >= logging.WARNING]
def test_a_broken_mcp_dependency_is_reported_and_retried(caplog):
"""An mcp the install has but cannot import is a broken install, not a lean one, so it must be loud."""
broken = ModuleNotFoundError("No module named 'mcp.types'", name="mcp.types")
with capturing(caplog, logging.DEBUG), raising_on_import(MCP_TRANSLATION_MODULE, broken):
partial = llms_package.load_guardrail_translation_mappings()
assert CallTypes.call_mcp_tool not in partial
assert CallTypes.acompletion in partial
assert tuple(llms_package.guardrail_translation_discovery.unavailable) == (MCP_TRANSLATION_MODULE,)
errors = [record for record in caplog.records if record.levelno >= logging.ERROR]
assert len(errors) == 1, [record.getMessage() for record in caplog.records]
assert "mcp.types" in errors[0].getMessage()
recovered = llms_package.load_guardrail_translation_mappings()
assert CallTypes.call_mcp_tool in recovered
assert not llms_package.guardrail_translation_discovery.unavailable
class _StandInMCPHandler:
"""A bundled package's handler that claims the call type the MCP package owns."""
def test_the_mcp_package_wins_a_handler_collision_with_a_bundled_package(monkeypatch):
"""MCP handlers are the specialised ones, so a bundled package declaring the same call type must not shadow them."""
colliding = ModuleType("litellm.llms.colliding_stub.guardrail_translation")
colliding.guardrail_translation_mappings = {CallTypes.call_mcp_tool: _StandInMCPHandler}
monkeypatch.setitem(sys.modules, colliding.__name__, colliding)
monkeypatch.setattr(llms_package, "_bundled_guardrail_translation_modules", lambda: iter((colliding.__name__,)))
discovery = llms_package.discover_guardrail_translations()
assert discovery.mappings[CallTypes.call_mcp_tool] is MCPGuardrailTranslationHandler

View file

@ -133,8 +133,8 @@ def restore_callbacks(monkeypatch):
monkeypatch.setattr(litellm, "callbacks", litellm.callbacks)
monkeypatch.setattr(
litellm_llms,
"endpoint_guardrail_translation_mappings",
litellm_llms.endpoint_guardrail_translation_mappings,
"guardrail_translation_discovery",
litellm_llms.guardrail_translation_discovery,
)
yield
ProxyLogging._callback_capabilities_cache.clear()

View file

@ -0,0 +1,13 @@
import pytest
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail import (
unified_guardrail as unified_module,
)
@pytest.fixture(autouse=True)
def _forget_unscanned_warnings():
"""The unscanned warning fires once per route and reason, so every guardrail test starts with nothing remembered."""
unified_module._warn_left_unscanned_once.cache_clear()
yield
unified_module._warn_left_unscanned_once.cache_clear()

View file

@ -1,14 +1,24 @@
import pytest
from unittest.mock import MagicMock, patch
import os
import sys
from unittest.mock import MagicMock, patch
import pytest
import litellm.llms as llms_package
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.openai.moderations import (
OpenAIModerationGuardrail,
)
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
from litellm.types.utils import ModelResponseStream, ModelResponse
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.utils import (
CallTypes,
Delta,
ModelResponse,
ModelResponseStream,
StreamingChoices,
)
@pytest.mark.asyncio
@ -516,3 +526,74 @@ async def test_openai_moderation_streaming_sampled_when_end_of_stream_only_disab
f"because chunk 6 already scanned the full text), "
f"got {patched_make_request.await_count}"
)
@pytest.fixture
def reset_guardrail_translation_discovery():
saved = llms_package.guardrail_translation_discovery
llms_package.guardrail_translation_discovery = None
yield llms_package
llms_package.guardrail_translation_discovery = saved
@pytest.mark.asyncio
async def test_moderation_still_runs_after_a_failed_translation_discovery(
reset_guardrail_translation_discovery,
):
"""
A guardrail translation discovery that could not import the chat handler must not silently
disable moderation for the rest of the process.
"""
with pytest.MonkeyPatch.context() as poison:
poison.setitem(sys.modules, "litellm.llms.openai.chat.guardrail_translation", None)
poisoned = llms_package.load_guardrail_translation_mappings()
assert CallTypes.acompletion not in poisoned
with patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}):
openai_guardrail = OpenAIModerationGuardrail(
guardrail_name="test-openai-moderation",
event_hook="post_call",
)
unified_guardrail = UnifiedLLMGuardrails()
mock_mod_response = MagicMock()
mock_mod_response.results = []
async def mock_stream():
chunks_data = ["Hello", " ", "world", "!", " Goodbye"]
for i, content in enumerate(chunks_data):
yield ModelResponseStream(
model="gpt-4",
choices=[
StreamingChoices(
index=0,
delta=Delta(content=content, role="assistant"),
finish_reason=(
"stop" if i == len(chunks_data) - 1 else None
),
)
],
)
with patch.object(
openai_guardrail, "async_make_request", return_value=mock_mod_response
) as patched_make_request:
chunks_received = 0
async for _ in unified_guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(
api_key="test", request_route="/chat/completions"
),
response=mock_stream(),
request_data={
"messages": [{"role": "user", "content": "hi"}],
"guardrail_to_apply": openai_guardrail,
"metadata": {"guardrails": ["test-openai-moderation"]},
},
):
chunks_received += 1
assert chunks_received == 5
assert patched_make_request.await_count > 0, (
"Moderation never ran: the failed discovery was cached and the streaming hook "
"passed every chunk through unscanned"
)

View file

@ -1,6 +1,7 @@
"""Tests for unified guardrail."""
import logging
from contextlib import contextmanager
import pytest
@ -2289,3 +2290,250 @@ class TestTranslationMappingsAreReadLive:
for name, value in vars(unified_module).items()
if isinstance(value, dict) and CallTypes.aocr in value
]
class TestUnscannedStreamIsAnnounced:
"""The streaming hook must never forward a whole response unscanned without saying so."""
@staticmethod
async def _drive(caplog, monkeypatch, request_route, mappings, response_chunks):
_patch_translation_mappings(monkeypatch, mappings)
async def stream():
for chunk in response_chunks:
yield chunk
caplog.set_level(logging.WARNING, logger="LiteLLM Proxy")
unified_module.verbose_proxy_logger.addHandler(caplog.handler)
try:
chunks = [
chunk
async for chunk in UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(
api_key="test", request_route=request_route
),
response=stream(),
request_data={
"guardrail_to_apply": RecordingGuardrail(),
"metadata": {"guardrails": ["recording-guardrail"]},
},
)
]
finally:
unified_module.verbose_proxy_logger.removeHandler(caplog.handler)
return chunks, [
record.getMessage()
for record in caplog.records
if record.levelno >= logging.WARNING
]
@pytest.mark.asyncio
async def test_warns_when_the_call_type_has_no_translation_handler(
self, caplog, monkeypatch
):
chunks, warnings = await self._drive(
caplog,
monkeypatch,
request_route="/chat/completions",
mappings={CallTypes.aembedding: _NoopTranslation},
response_chunks=[
ModelResponseStream(
choices=[StreamingChoices(index=0, delta=Delta(content=content))]
)
for content in ("a", "b", "c")
],
)
assert len(chunks) == 3
assert any(
"no guardrail translation handler" in message
and "Add a guardrail translation handler for that call type." in message
and "recording-guardrail" in message
and "/chat/completions" in message
for message in warnings
), warnings
@pytest.mark.asyncio
async def test_warns_when_the_call_type_cannot_be_resolved(self, caplog, monkeypatch):
chunks, warnings = await self._drive(
caplog,
monkeypatch,
request_route="/v1/not-a-mapped-route",
mappings=load_guardrail_translation_mappings(),
response_chunks=[{"event": "delta", "text": content} for content in ("a", "b", "c")],
)
assert len(chunks) == 3
assert any(
"call type could not be resolved" in message
and "Add the route to API_ROUTE_TO_CALL_TYPES." in message
and "recording-guardrail" in message
for message in warnings
), warnings
class TestUnscannedRequestIsAnnounced:
"""A request hook that cannot scan must say so instead of passing the request through in silence."""
@staticmethod
def _request(guardrail):
return {
"guardrail_to_apply": guardrail,
"model": "gpt-4o",
"messages": [{"role": "user", "content": "hello world"}],
}
@staticmethod
@contextmanager
def _capturing(caplog):
caplog.set_level(logging.WARNING, logger="LiteLLM Proxy")
unified_module.verbose_proxy_logger.addHandler(caplog.handler)
try:
yield
finally:
unified_module.verbose_proxy_logger.removeHandler(caplog.handler)
@staticmethod
def _warnings(caplog):
return [record.getMessage() for record in caplog.records if record.levelno >= logging.WARNING]
@staticmethod
def _distinct_warnings(caplog):
"""caplog holds every record twice here, once through the handler above and once through propagation."""
by_record = {id(record): record for record in caplog.records if record.levelno >= logging.WARNING}
return [record.getMessage() for record in by_record.values()]
@pytest.mark.asyncio
async def test_pre_call_warns_when_the_call_type_has_no_translation_handler(self, caplog, monkeypatch):
_patch_translation_mappings(monkeypatch, {CallTypes.aembedding: _NoopTranslation})
guardrail = RecordingGuardrail()
data = self._request(guardrail)
with self._capturing(caplog):
returned = await UnifiedLLMGuardrails().async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/moderations"),
cache=DualCache(),
data=data,
call_type=CallTypes.acompletion.value,
)
assert guardrail.apply_calls == []
assert returned["messages"] == [{"role": "user", "content": "hello world"}]
assert any(
"no guardrail translation handler" in message
and "skipping pre-call scanning" in message
and "recording-guardrail" in message
and "/v1/moderations" in message
and "acompletion" in message
for message in self._warnings(caplog)
), self._warnings(caplog)
@pytest.mark.asyncio
async def test_during_call_warns_when_the_call_type_has_no_translation_handler(self, caplog, monkeypatch):
_patch_translation_mappings(monkeypatch, {CallTypes.aembedding: _NoopTranslation})
guardrail = RecordingGuardrail()
data = self._request(guardrail)
with self._capturing(caplog):
returned = await UnifiedLLMGuardrails().async_moderation_hook(
data=data,
user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/moderations"),
call_type=CallTypes.acompletion.value,
)
assert guardrail.apply_calls == []
assert returned["messages"] == [{"role": "user", "content": "hello world"}]
assert any(
"skipping during-call scanning" in message and "recording-guardrail" in message
for message in self._warnings(caplog)
), self._warnings(caplog)
@pytest.mark.asyncio
async def test_pre_call_warns_instead_of_raising_on_a_call_type_outside_the_enum(self, caplog, monkeypatch):
_patch_translation_mappings(monkeypatch, load_guardrail_translation_mappings())
guardrail = RecordingGuardrail()
data = self._request(guardrail)
with self._capturing(caplog):
returned = await UnifiedLLMGuardrails().async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/moderations"),
cache=DualCache(),
data=data,
call_type="not_a_call_type",
)
assert guardrail.apply_calls == []
assert returned["messages"] == [{"role": "user", "content": "hello world"}]
assert any(
"call type 'not_a_call_type' is not one litellm can scan" in message
and "Map the route to a CallTypes member in API_ROUTE_TO_CALL_TYPES." in message
and "skipping pre-call scanning" in message
for message in self._warnings(caplog)
), self._warnings(caplog)
@pytest.mark.asyncio
async def test_during_call_warns_instead_of_raising_on_a_call_type_outside_the_enum(self, caplog, monkeypatch):
_patch_translation_mappings(monkeypatch, load_guardrail_translation_mappings())
guardrail = RecordingGuardrail()
data = self._request(guardrail)
with self._capturing(caplog):
returned = await UnifiedLLMGuardrails().async_moderation_hook(
data=data,
user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/moderations"),
call_type="not_a_call_type",
)
assert guardrail.apply_calls == []
assert returned["messages"] == [{"role": "user", "content": "hello world"}]
assert any(
"call type 'not_a_call_type' is not one litellm can scan" in message
and "Map the route to a CallTypes member in API_ROUTE_TO_CALL_TYPES." in message
and "skipping during-call scanning" in message
for message in self._warnings(caplog)
), self._warnings(caplog)
@pytest.mark.asyncio
async def test_a_route_with_no_handler_warns_once_instead_of_once_per_request(self, caplog, monkeypatch):
_patch_translation_mappings(monkeypatch, {CallTypes.aembedding: _NoopTranslation})
guardrail = RecordingGuardrail()
with self._capturing(caplog):
for _ in range(5):
await UnifiedLLMGuardrails().async_moderation_hook(
data=self._request(guardrail),
user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/responses/resp_123"),
call_type="aget_responses",
)
await UnifiedLLMGuardrails().async_moderation_hook(
data=self._request(guardrail),
user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/moderations"),
call_type="aget_responses",
)
unscanned = [
message for message in self._distinct_warnings(caplog) if "skipping during-call scanning" in message
]
assert len(unscanned) == 2, unscanned
assert "/v1/responses/resp_123" in unscanned[0]
assert "aget_responses" in unscanned[0]
assert "/v1/moderations" in unscanned[1]
@pytest.mark.asyncio
async def test_post_call_names_the_call_type_the_route_maps_to(self, caplog, monkeypatch):
_patch_translation_mappings(monkeypatch, {CallTypes.aembedding: _NoopTranslation})
guardrail = RecordingGuardrail()
with self._capturing(caplog):
await UnifiedLLMGuardrails().async_post_call_success_hook(
data={"guardrail_to_apply": guardrail, "model": "gpt-4o"},
user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/responses"),
response=litellm.ModelResponse(),
)
assert guardrail.apply_calls == []
unscanned = [message for message in self._warnings(caplog) if "skipping post-call scanning" in message]
assert unscanned, self._warnings(caplog)
assert "call type 'aresponses'" in unscanned[0], unscanned[0]
assert "CallTypes." not in unscanned[0], unscanned[0]