This commit is contained in:
Caduri 2026-09-30 10:31:04 +00:00 • committed by GitHub
commit 15d9e6fd13
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
28 changed files with 1847 additions and 191 deletions

View file

@ -80,7 +80,12 @@ def get_call_types_for_route(route: str) -> Sequence[CallTypes] | None:
return None
def get_routes_for_call_type(call_type: CallTypes) -> list:
def get_primary_call_type_for_route(route: str | None) -> CallTypes | None:
call_types: Final = get_call_types_for_route(route) if route else None
return call_types[0] if call_types else None
def get_routes_for_call_type(call_type: CallTypes) -> list[str]:
"""
Get all routes that use a specific CallType.
@ -90,7 +95,7 @@ def get_routes_for_call_type(call_type: CallTypes) -> list:
Returns:
List of routes that use this CallType
"""
routes: Final = []
routes: Final[list[str]] = []
for route, types in API_ROUTE_TO_CALL_TYPES.items():
if call_type in types:
routes.append(route)

View file

@ -12,7 +12,9 @@ import litellm
from litellm._logging import redact_internal_details_from_client_message, verbose_logger
from litellm.constants import REALTIME_SESSION_FAILURE_LOGGED_KEY, REALTIME_SESSION_SUCCESS_LOGGED_KEY
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig, RealtimeBackend
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.llms.openai import (
OpenAIRealtimeEvents,
OpenAIRealtimeOutputItemDone,
@ -123,6 +125,10 @@ DefaultLoggedRealTimeEventTypes: Final = [
]
def _as_user_api_key_auth(user_api_key_dict: object) -> UserAPIKeyAuth | None:
return user_api_key_dict if isinstance(user_api_key_dict, UserAPIKeyAuth) else None
class RealTimeStreaming:
def __init__(
self,
@ -852,7 +858,12 @@ class RealTimeStreaming:
try:
await callback.apply_guardrail(
inputs={"texts": [transcript], "images": []},
request_data={"user_api_key_dict": self.user_api_key_dict},
request_data={
"user_api_key_dict": self.user_api_key_dict,
"litellm_metadata": BaseTranslation.transform_user_api_key_dict_to_metadata(
_as_user_api_key_auth(self.user_api_key_dict)
),
},
input_type="request",
)
except Exception as e:

View file

@ -81,43 +81,18 @@ class BaseTranslation(ABC):
@staticmethod
def transform_user_api_key_dict_to_metadata(
user_api_key_dict: Any | None,
user_api_key_dict: "UserAPIKeyAuth | None",
) -> dict[str, object]:
"""
Transform user_api_key_dict to a metadata dict with prefixed keys.
Converts keys like 'user_id' to 'user_api_key_user_id' to clearly indicate
the source of the metadata.
Args:
user_api_key_dict: UserAPIKeyAuth object or dict with user information
Returns:
Dict with keys prefixed with 'user_api_key_'
"""
"""The authenticated key's identity as prefixed metadata, an allowlist safe to hand to guardrail vendors."""
if user_api_key_dict is None:
return {}
# Lazy: `import litellm` loads this module before litellm.Router exists, and litellm_pre_call_utils imports it
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
# Convert to dict if it's a Pydantic object
user_dict = user_api_key_dict.model_dump() if hasattr(user_api_key_dict, "model_dump") else user_api_key_dict
if not isinstance(user_dict, dict):
return {}
# Transform keys to be prefixed with 'user_api_key_'
transformed: Final[dict[str, object]] = {}
for key, value in user_dict.items():
# Skip None values and internal fields
if value is None or key.startswith("_"):
continue
# If key already has the prefix, use as-is, otherwise add prefix
if key.startswith("user_api_key_"):
transformed[key] = value
else:
transformed[f"user_api_key_{key}"] = value
return transformed
return {
**LiteLLMProxyRequestSetup.get_authenticated_identity_metadata(user_api_key_dict),
"user_api_key_key_alias": user_api_key_dict.key_alias,
}
@staticmethod
def merge_user_api_key_metadata_into_request(

View file

@ -20,6 +20,10 @@ if TYPE_CHECKING:
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging
_PROXY_OWNED_PAYLOAD_KEYS: Final = frozenset(
{"metadata", "litellm_metadata", "litellm_logging_obj", "proxy_server_request"}
)
class PassThroughEndpointHandler(BaseTranslation):
"""
@ -80,7 +84,7 @@ class PassThroughEndpointHandler(BaseTranslation):
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
payload_to_check: Final = {
k: v for k, v in data.items() if not k.startswith("_") and k not in ("metadata", "litellm_logging_obj")
k: v for k, v in data.items() if not k.startswith("_") and k not in _PROXY_OWNED_PAYLOAD_KEYS
}
verbose_proxy_logger.debug("PassThroughEndpointHandler: Using full payload for guardrail")
return safe_dumps(payload_to_check)

View file

@ -34,6 +34,7 @@ from litellm.proxy.guardrails.guardrail_hooks.custom_code.sandbox import (
)
from litellm.proxy.guardrails.guardrail_registry import GuardrailRegistry
from litellm.proxy.guardrails.usage_endpoints import router as guardrails_usage_router
from litellm.proxy.litellm_pre_call_utils import caller_metadata_with_authenticated_identity
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
from litellm.repositories.prisma_protocols import TableActions
from litellm.repositories.table_repositories import GuardrailsRepository
@ -2404,9 +2405,14 @@ async def apply_guardrail(
if litellm_logging_obj is not None:
_patch_logging_obj_for_guardrail(litellm_logging_obj, request)
processed_metadata: Final = data.get("metadata")
inbound_headers: Final = processed_metadata.get("headers") if isinstance(processed_metadata, dict) else None
request_data: Final[dict] = {
**({"messages": request.messages} if request.messages is not None else {}),
**({"metadata": request.metadata} if request.metadata is not None else {}),
"metadata": {
**caller_metadata_with_authenticated_identity(request.metadata, user_api_key_dict),
**({"headers": inbound_headers} if inbound_headers is not None else {}),
},
}
_input_type: Final = _resolve_guardrail_input_type(active_guardrail, request.input_type)
guardrailed_inputs: Final = await active_guardrail.apply_guardrail(

View file

@ -39,6 +39,8 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
streaming_end_of_stream_only=_get_config_value(litellm_params, optional_params, "streaming_end_of_stream_only"),
streaming_sampling_rate=_get_config_value(litellm_params, optional_params, "streaming_sampling_rate"),
streaming_transform_mode=_get_config_value(litellm_params, optional_params, "streaming_transform_mode"),
run_only_on_call_types=_get_config_value(litellm_params, optional_params, "run_only_on_call_types"),
skip_call_types=_get_config_value(litellm_params, optional_params, "skip_call_types"),
)
litellm.logging_callback_manager.add_litellm_callback(_generic_guardrail_api_callback)

View file

@ -0,0 +1,114 @@
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from typing import TYPE_CHECKING, Final
from typing_extensions import Self
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.api_route_to_call_types import get_primary_call_type_for_route
from litellm.types.utils import CallTypes
from .config_parsing import config_values
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
_KNOWN_CALL_TYPES: Final = frozenset(call_type.value for call_type in CallTypes)
_UNRESOLVABLE_CALL_TYPES: Final = frozenset({CallTypes.call_mcp_tool.value})
def _unknown_call_types_error(option_name: str, unknown: Sequence[str]) -> ValueError:
renamed: Final = tuple(
f"{value!r} -> {CallTypes[value].value!r}" for value in unknown if value in CallTypes.__members__
)
hint: Final = f" CallTypes member names map to these values: {', '.join(renamed)}." if renamed else ""
return ValueError(
f"{option_name} contains unknown call type(s) {list(unknown)}. Use CallTypes values, the strings "
f"logged as call_type (e.g. acompletion, aembedding, anthropic_messages, pass_through_endpoint).{hint}"
)
def _as_call_types(raw: Sequence[str] | None, *, option_name: str) -> frozenset[str]:
values: Final = config_values(raw, option_name=option_name)
unknown: Final = tuple(value for value in values if value not in _KNOWN_CALL_TYPES)
if unknown:
raise _unknown_call_types_error(option_name, unknown)
unresolvable: Final = tuple(value for value in values if value in _UNRESOLVABLE_CALL_TYPES)
if unresolvable:
raise ValueError(
f"{option_name} cannot filter {list(unresolvable)}: MCP tool calls reach the guardrail without a "
"call type, so the filter could never match them"
)
return frozenset(values)
@dataclass(frozen=True, slots=True)
class CallTypeFilter:
run_only_on: frozenset[str] | None = None
skip: frozenset[str] = frozenset()
@classmethod
def from_config(
cls,
*,
run_only_on_call_types: Sequence[str] | None,
skip_call_types: Sequence[str] | None,
guardrail_name: str | None,
) -> Self:
run_only_on: Final = (
_as_call_types(run_only_on_call_types, option_name="run_only_on_call_types")
if run_only_on_call_types
else None
)
skip: Final = _as_call_types(skip_call_types, option_name="skip_call_types")
if run_only_on is not None and skip:
verbose_proxy_logger.warning(
"Generic Guardrail API (%s): both run_only_on_call_types and skip_call_types are set. "
"The allowlist wins; skip_call_types=%s is ignored.",
guardrail_name,
sorted(skip),
)
return cls(run_only_on=run_only_on, skip=skip)
def allows(self, call_type: str | None) -> bool:
if call_type is None:
return True
if self.run_only_on is not None:
return call_type in self.run_only_on
return call_type not in self.skip
def skip_reason(
self,
*,
request_data: Mapping[str, object],
logging_obj: "LiteLLMLoggingObj | None",
) -> str | None:
if self.run_only_on is None and not self.skip:
return None
call_type: Final = resolve_call_type(request_data=request_data, logging_obj=logging_obj)
if self.allows(call_type):
return None
rule: Final = "not in run_only_on_call_types" if self.run_only_on is not None else "in skip_call_types"
return f"skipped: call type {call_type} {rule}"
def _non_empty_str(value: object) -> str | None:
return value if isinstance(value, str) and value else None
def _metadata_route(metadata: object) -> str | None:
return _non_empty_str(metadata.get("user_api_key_request_route")) if isinstance(metadata, Mapping) else None
def _request_route(request_data: Mapping[str, object]) -> str | None:
return _metadata_route(request_data.get("litellm_metadata")) or _metadata_route(request_data.get("metadata"))
def resolve_call_type(
request_data: Mapping[str, object],
logging_obj: "LiteLLMLoggingObj | None",
) -> str | None:
route_call_type: Final = get_primary_call_type_for_route(_request_route(request_data))
if route_call_type is not None:
return route_call_type.value
return _non_empty_str(logging_obj.call_type) if logging_obj is not None else None

View file

@ -0,0 +1,15 @@
import re
from collections.abc import Sequence
def config_values(raw: Sequence[str] | None, *, option_name: str) -> tuple[str, ...]:
if isinstance(raw, str):
raise ValueError(f"{option_name} must be a list of strings, got the single string {raw!r}")
return tuple(raw or ())
def compile_patterns(raw: Sequence[str] | None, *, option_name: str) -> tuple[re.Pattern[str], ...]:
try:
return tuple(re.compile(pattern) for pattern in config_values(raw, option_name=option_name))
except re.error as e:
raise ValueError(f"{option_name} contains an invalid regex: {e}") from e

View file

@ -20,6 +20,7 @@ from litellm.integrations.custom_guardrail import (
log_guardrail_information,
)
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
get_async_httpx_client,
httpxSpecialProvider,
)
@ -33,6 +34,8 @@ from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import
)
from litellm.types.utils import GenericGuardrailAPIInputs
from .call_type_filter import CallTypeFilter
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
@ -170,6 +173,10 @@ def _structured_rows_to_write_back(
)
def _passthrough_inputs(inputs: GenericGuardrailAPIInputs) -> GenericGuardrailAPIInputs:
return GenericGuardrailAPIInputs(**inputs)
class GenericGuardrailAPI(CustomGuardrail):
"""
Generic Guardrail API integration for LiteLLM.
@ -204,9 +211,14 @@ class GenericGuardrailAPI(CustomGuardrail):
streaming_end_of_stream_only: bool | None = None,
streaming_sampling_rate: int | None = None,
streaming_transform_mode: Literal["block_only", "incremental_diff"] | None = None,
run_only_on_call_types: Sequence[str] | None = None,
skip_call_types: Sequence[str] | None = None,
async_handler: AsyncHTTPHandler | None = None,
**kwargs,
):
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
self.async_handler = async_handler or get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback
)
self.headers = headers or {}
self.extra_headers = extra_headers or []
@ -251,6 +263,12 @@ class GenericGuardrailAPI(CustomGuardrail):
"block_only" if streaming_transform_mode is None else streaming_transform_mode
)
self.call_type_filter: Final = CallTypeFilter.from_config(
run_only_on_call_types=run_only_on_call_types,
skip_call_types=skip_call_types,
guardrail_name=kwargs.get("guardrail_name"),
)
# Set supported event hooks
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
@ -292,9 +310,8 @@ class GenericGuardrailAPI(CustomGuardrail):
if value is not None:
result_metadata[field_name] = value
# handle user_api_key_token = user_api_key_hash
if metadata_dict.get("user_api_key_token") is not None:
result_metadata["user_api_key_hash"] = metadata_dict.get("user_api_key_token")
if litellm_metadata.get("user_api_key_token") is not None and "user_api_key_hash" not in result_metadata:
result_metadata["user_api_key_hash"] = litellm_metadata["user_api_key_token"]
verbose_proxy_logger.debug(
"Generic Guardrail API: Extracted user metadata: %s",
@ -432,6 +449,16 @@ class GenericGuardrailAPI(CustomGuardrail):
if request_data is None:
request_data = {}
skip_reason: Final = self.call_type_filter.skip_reason(request_data=request_data, logging_obj=logging_obj)
if skip_reason is not None:
verbose_proxy_logger.debug("Generic Guardrail API: %s (input_type=%s)", skip_reason, input_type)
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_json_response=skip_reason,
request_data=request_data,
guardrail_status="not_run",
)
return _passthrough_inputs(inputs)
request_body: Final = request_data.get("body") or {}
# Merge additional provider specific params from config and dynamic params

View file

@ -1798,16 +1798,6 @@ class PanwPrismaAirsHandler(CustomGuardrail):
call_id,
_mcp_tool,
)
elif not request_data and logging_obj is None and input_type == "request":
# Direct /apply_guardrail endpoint — empty request_data, no
# logging_obj. Existing behavior: synthesize UUID.
call_id = str(uuid.uuid4())
request_data["litellm_call_id"] = call_id
verbose_proxy_logger.warning(
"PANW Prisma AIRS: litellm_call_id missing from empty "
"request_data, synthesized %s (direct /apply_guardrail?)",
call_id,
)
else:
call_id = str(uuid.uuid4())
request_data["litellm_call_id"] = call_id

View file

@ -18,9 +18,11 @@ from litellm.caching.caching import DualCache
from litellm.cost_calculator import _infer_call_type
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route
from litellm.litellm_core_utils.api_route_to_call_types import get_primary_call_type_for_route
from litellm.llms import get_guardrail_translation_mapping, load_guardrail_translation_mappings
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation, StreamingScanKey
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import (
MCP_GUARDRAIL_CALL_TYPES,
@ -35,10 +37,6 @@ if TYPE_CHECKING:
# Imported lazily at runtime (inside the streaming hook) to avoid a
# module-level cyclic import with litellm.integrations.custom_guardrail.
from litellm.integrations.custom_guardrail import ModifyResponseException
from litellm.llms.base_llm.guardrail_translation.base_translation import (
BaseTranslation,
StreamingScanKey,
)
# Call types that stream JSON-RPC events (A2A); guardrail HTTPException is emitted as in-stream error
A2A_CALL_TYPES: Final = (CallTypes.asend_message, CallTypes.send_message)
@ -57,7 +55,7 @@ class _EndpointTranslation(Protocol):
def process_output_streaming_response(self) -> "Callable[..., Awaitable[object]]": ...
@property
def get_streaming_scan_key(self) -> "Callable[[Sequence[object]], StreamingScanKey | None]": ...
def get_streaming_scan_key(self) -> Callable[[Sequence[object]], StreamingScanKey | None]: ...
@property
def build_block_sse_chunks(self) -> "Callable[..., Sequence[bytes] | None]": ...
@ -72,19 +70,17 @@ def _as_endpoint_translation(translation: _EndpointTranslation) -> _EndpointTran
def resolve_endpoint_translation(
user_api_key_dict: UserAPIKeyAuth, first_response_item: object | None
) -> "tuple[str, BaseTranslation] | None":
) -> tuple[str, BaseTranslation] | None:
"""
Resolve the endpoint guardrail translation for a streamed response: the
request route wins, falling back to inferring the call type from the first
response chunk (the same resolution order the streaming iterator hook uses).
Returns None when the call type is unresolvable or has no translation.
"""
route_call_types: Final = (
get_call_types_for_route(user_api_key_dict.request_route) if user_api_key_dict.request_route else None
)
route_call_type: Final = get_primary_call_type_for_route(user_api_key_dict.request_route)
call_type: Final = (
route_call_types[0].value
if route_call_types
route_call_type.value
if route_call_type is not None
else (
_infer_call_type(call_type=None, completion_response=first_response_item)
if first_response_item is not None
@ -109,7 +105,7 @@ def _held_choices(held_chars_per_choice: Mapping[int, int]) -> frozenset[int]:
return frozenset(idx for idx, held in held_chars_per_choice.items() if held > 0)
def _is_redundant_scan(scan_key: "StreamingScanKey | None", last_scan_key: "StreamingScanKey | None") -> bool:
def _is_redundant_scan(scan_key: StreamingScanKey | None, last_scan_key: StreamingScanKey | None) -> bool:
if scan_key is None:
return False
return scan_key == last_scan_key or scan_key.has_nothing_to_scan
@ -156,16 +152,20 @@ def _a2a_jsonrpc_error_chunk(exc: HTTPException, request_id: str | None) -> Mapp
}
def _ensure_litellm_metadata(data: dict, user_api_key_dict: UserAPIKeyAuth) -> None:
"""Populate data['litellm_metadata'] from user_api_key_dict if absent."""
if "litellm_metadata" not in data:
from litellm.llms.base_llm.guardrail_translation.base_translation import (
BaseTranslation,
)
_PROXY_ENRICHED_IDENTITY_FIELDS: Final = frozenset({"user_api_key_auth_metadata"})
user_metadata: Final = BaseTranslation.transform_user_api_key_dict_to_metadata(user_api_key_dict)
if user_metadata:
data["litellm_metadata"] = user_metadata
def _ensure_litellm_metadata(data: dict, user_api_key_dict: UserAPIKeyAuth) -> None:
"""Overwrite the identity fields of data['litellm_metadata'] from the authenticated key, in place."""
existing: Final = data.get("litellm_metadata")
if isinstance(existing, dict):
identity: Final = LiteLLMProxyRequestSetup.get_authenticated_identity_metadata(user_api_key_dict)
existing.update({key: value for key, value in identity.items() if key not in _PROXY_ENRICHED_IDENTITY_FIELDS})
existing.pop("user_api_key_token", None)
return
user_metadata: Final = BaseTranslation.transform_user_api_key_dict_to_metadata(user_api_key_dict)
if user_metadata:
data["litellm_metadata"] = user_metadata
class UnifiedLLMGuardrails(CustomLogger):
@ -312,11 +312,7 @@ class UnifiedLLMGuardrails(CustomLogger):
verbose_proxy_logger.debug("async_post_call_success_hook response: %s", response)
call_type: CallTypesLiteral | None = None
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: CallTypesLiteral | None = get_primary_call_type_for_route(user_api_key_dict.request_route)
if call_type is None:
call_type = _infer_call_type(call_type=None, completion_response=response)
@ -406,7 +402,7 @@ class UnifiedLLMGuardrails(CustomLogger):
@staticmethod
def _resolve_transform_call_type(
user_api_key_dict: UserAPIKeyAuth,
mappings: Mapping[CallTypes, type["BaseTranslation"]],
mappings: Mapping[CallTypes, type[BaseTranslation]],
) -> str | None:
"""Resolve the call type for the incremental_diff path, or None if the
route is unresolvable / unsupported.
@ -420,20 +416,13 @@ class UnifiedLLMGuardrails(CustomLogger):
OpenAIChatCompletionsHandler,
)
if user_api_key_dict.request_route is None:
route_call_type: Final = get_primary_call_type_for_route(user_api_key_dict.request_route)
if route_call_type is None:
return None
call_types: Final = get_call_types_for_route(user_api_key_dict.request_route)
if not call_types:
return None
call_type: Final = call_types[0].value
try:
mapped: Final = CallTypes(call_type)
except ValueError:
return None
handler_cls: Final = mappings.get(mapped)
handler_cls: Final = mappings.get(route_call_type)
if handler_cls is None or not issubclass(handler_cls, OpenAIChatCompletionsHandler):
return None
return call_type
return route_call_type.value
async def emit_streaming_http_error(
self,
@ -669,7 +658,7 @@ class UnifiedLLMGuardrails(CustomLogger):
call_type: str,
sampling_rate: int,
end_of_stream_only: bool,
mappings: Mapping[CallTypes, type["BaseTranslation"]],
mappings: Mapping[CallTypes, type[BaseTranslation]],
) -> AsyncGenerator[object, None]:
"""Emit guardrail text transformations as new deltas on the stream.
@ -1080,8 +1069,8 @@ class UnifiedLLMGuardrails(CustomLogger):
getattr(guardrail_to_apply, "guardrail_name", None),
)
# Infer call type from first chunk
call_type = None
route_call_type: Final = get_primary_call_type_for_route(user_api_key_dict.request_route)
call_type = route_call_type.value if route_call_type is not None else None
chunk_counter = 0
responses_so_far: Final[list[object]] = []
responses_yielded: Final[list[object]] = []
@ -1098,12 +1087,6 @@ class UnifiedLLMGuardrails(CustomLogger):
chunk_counter += 1
responses_so_far.append(item)
# Infer call type from first chunk if not already done
if call_type is None and user_api_key_dict.request_route is not None:
call_types = get_call_types_for_route(user_api_key_dict.request_route)
if call_types is not None:
call_type = call_types[0].value
if call_type is None:
call_type = _infer_call_type(call_type=None, completion_response=item)

View file

@ -496,6 +496,40 @@ def _strip_untrusted_request_header_controls(
headers.pop(header_name, None)
def is_untrusted_caller_metadata_key(key: str) -> bool:
return key.startswith("user_api_key_") or key in _UNTRUSTED_METADATA_CONTROL_FIELDS
def strip_untrusted_caller_metadata(
data: MutableMapping[str, object], *, allow_client_message_redaction_opt_out: bool
) -> None:
"""Remove, in place, the proxy-owned slots a caller put in either metadata bucket of a request body."""
for user_meta in (data.get("metadata"), data.get("litellm_metadata")):
if not isinstance(user_meta, dict):
continue
_strip_untrusted_request_header_controls(
user_meta.get("headers"),
allow_client_message_redaction_opt_out=allow_client_message_redaction_opt_out,
)
for untrusted_key in tuple(key for key in user_meta if is_untrusted_caller_metadata_key(key)):
user_meta.pop(untrusted_key, None)
_GUARDRAIL_UNTRUSTED_CALLER_METADATA_KEYS: Final = frozenset({"user_api_key", "headers"})
def caller_metadata_with_authenticated_identity(
caller_metadata: Mapping[str, object] | None, user_api_key_dict: UserAPIKeyAuth
) -> dict[str, object]:
"""Caller metadata minus proxy-owned slots, bare user_api_key and headers, with the key's identity on top."""
caller_fields: Final = {
key: value
for key, value in (caller_metadata or {}).items()
if not (is_untrusted_caller_metadata_key(key) or key in _GUARDRAIL_UNTRUSTED_CALLER_METADATA_KEYS)
}
return {**caller_fields, **LiteLLMProxyRequestSetup.get_authenticated_identity_metadata(user_api_key_dict)}
def _is_false_like(value: object) -> bool:
if isinstance(value, bool):
return value is False
@ -523,7 +557,7 @@ def _key_or_team_allows_client_mock_response(
)
def _key_or_team_allows_client_message_redaction_opt_out(
def key_or_team_allows_client_message_redaction_opt_out(
user_api_key_dict: UserAPIKeyAuth,
) -> bool:
return _key_or_team_metadata_flag_is_true(
@ -1645,6 +1679,24 @@ class LiteLLMProxyRequestSetup:
)
return user_api_key_logged_metadata
@staticmethod
def get_key_scoped_metadata(user_api_key_dict: UserAPIKeyAuth) -> dict[str, object]:
return {
"user_api_key_metadata": strip_callback_config(user_api_key_dict.metadata),
"user_api_key_team_metadata": strip_callback_config(user_api_key_dict.team_metadata),
"user_api_key_object_permission_id": user_api_key_dict.object_permission_id,
"user_api_key_team_object_permission_id": user_api_key_dict.team_object_permission_id,
}
@staticmethod
def get_authenticated_identity_metadata(user_api_key_dict: UserAPIKeyAuth) -> dict[str, object]:
"""Identity fields derived from the authenticated key alone, for paths that skip the chat-path build."""
return {
**LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict),
"user_api_key": LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict),
**LiteLLMProxyRequestSetup.get_key_scoped_metadata(user_api_key_dict),
}
@staticmethod
def add_user_api_key_auth_to_request_metadata(
data: dict,
@ -2006,7 +2058,7 @@ async def add_litellm_data_to_request(
# These keys are injected by the proxy itself below — user-supplied values
# must not be trusted.
_allow_client_mock_response: Final = _key_or_team_allows_client_mock_response(user_api_key_dict)
_allow_client_message_redaction_opt_out = _key_or_team_allows_client_message_redaction_opt_out(user_api_key_dict)
_allow_client_message_redaction_opt_out = key_or_team_allows_client_message_redaction_opt_out(user_api_key_dict)
for _internal_key in _UNTRUSTED_ROOT_CONTROL_FIELDS:
if _allow_client_mock_response and _internal_key in _CLIENT_MOCK_CONTROL_FIELDS:
continue
@ -2188,31 +2240,12 @@ async def add_litellm_data_to_request(
# profile_id) don't see attacker-injected admin slots preserved in
# the deepcopy.
# Strip internal pipeline state and admin-injection slots from user input.
# Runs AFTER the string-to-dict parse above so JSON-string metadata (sent
# via multipart/form-data or extra_body) cannot smuggle admin fields past
# the isinstance(dict) guard.
#
# The proxy populates a family of ``user_api_key_*`` fields below
# (user_api_key_metadata, user_api_key_user_id, user_api_key_alias,
# user_api_key_spend, user_api_key_team_metadata, …) into
# data[_metadata_variable_name]. Because the proxy only writes to ONE of
# the two metadata dicts, a caller pre-populating any of these keys on
# the OTHER metadata dict would have their forged values surface in
# guardrails, spend tracking, audit logs, and identity resolution. Strip
# by prefix so new ``user_api_key_*`` fields added in the future are
# covered without per-key maintenance.
for _meta_key in ("metadata", "litellm_metadata"):
_user_meta = data.get(_meta_key)
if isinstance(_user_meta, dict):
_strip_untrusted_request_header_controls(
_user_meta.get("headers"),
allow_client_message_redaction_opt_out=(_allow_client_message_redaction_opt_out),
)
for _k in [
k for k in _user_meta if k.startswith("user_api_key_") or k in _UNTRUSTED_METADATA_CONTROL_FIELDS
]:
_user_meta.pop(_k, None)
strip_untrusted_caller_metadata(
data, allow_client_message_redaction_opt_out=_allow_client_message_redaction_opt_out
)
# Strip pricing overrides AFTER the litellm_metadata string-to-dict parse
# above, for the same reason as the user_api_key_* strip — JSON-string
@ -2389,14 +2422,7 @@ async def add_litellm_data_to_request(
data[_metadata_variable_name]["user_api_key_user_model_max_budget"] = user_model_budget # rebind-ok: out-param
data[_metadata_variable_name].update(carried_budget_metadata(user_api_key_dict))
data[_metadata_variable_name]["user_api_key_metadata"] = strip_callback_config(user_api_key_dict.metadata)
data[_metadata_variable_name]["user_api_key_team_metadata"] = strip_callback_config(user_api_key_dict.team_metadata)
data[_metadata_variable_name]["user_api_key_object_permission_id"] = getattr(
user_api_key_dict, "object_permission_id", None
)
data[_metadata_variable_name]["user_api_key_team_object_permission_id"] = getattr(
user_api_key_dict, "team_object_permission_id", None
)
data[_metadata_variable_name].update(LiteLLMProxyRequestSetup.get_key_scoped_metadata(user_api_key_dict))
data[_metadata_variable_name]["headers"] = _logging_safe_headers
data[_metadata_variable_name]["endpoint"] = str(request.url)
# Carry the proxy-receive instant via metadata (like `endpoint`) so the

View file

@ -23,7 +23,11 @@ from typing_extensions import assert_never
from litellm.exceptions import GuardrailRaisedException
from litellm.integrations.custom_guardrail import is_guardrail_intervention
from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route
from litellm.litellm_core_utils.api_route_to_call_types import (
get_call_types_for_route,
get_primary_call_type_for_route,
get_routes_for_call_type,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.llms.openai import BatchGuardrailRecord, BatchGuardrailReport
from litellm.types.utils import CallTypes, CallTypesLiteral
@ -56,7 +60,8 @@ _ROUTE_APPLIED_KEY: Final = "sensitive_data_routing_applied"
# online path either. `guardrails` is dropped because guardrail selection reads it ahead of the
# proxy-injected list, so leaving it would let a record's own body opt out of the chain its key
# and team selected; online that key can only add to the list, never replace it.
_INJECTED_KEYS: Final = frozenset({_SCAN_METADATA_KEY, "metadata", "guardrails"})
# `litellm_logging_obj` is dropped because guardrails read it as the proxy's own logging object.
_INJECTED_KEYS: Final = frozenset({_SCAN_METADATA_KEY, "metadata", "guardrails", "litellm_logging_obj"})
# Only what guardrail dispatch reads. The parent OTel span is deliberately left out: parenting one
# guardrail span per record would put tens of thousands of spans on a single upload's trace.
@ -280,22 +285,26 @@ def _iter_records(source: BinaryIO) -> Iterator[_ParsedRecord]:
yield _ParsedRecord(line_number=line_number, payload=json.loads(raw_line))
def _call_type_from_url(url: str) -> CallTypesLiteral | None:
def _url_path(url: str) -> str | None:
"""
Resolve the route a record names, tolerating how callers actually write it.
The route a record names, tolerating how callers actually write it.
An absolute url has to reduce to its path or nothing matches, and a record naming
``/v1/responses`` in full would fall through to its body, where ``input`` reads as an
embedding and the record gets scanned as the wrong call type rather than the right one.
"""
try:
path: Final = urlsplit(url).path.split("?")[0].rstrip("/")
return urlsplit(url).path.split("?")[0].rstrip("/")
except ValueError:
# urlsplit rejects a few malformed authorities outright, and the validation that ran
# before this only checks the key is present. An unreadable url is one we do not
# recognize, which is what falling back to the body shape already handles.
return None
call_types: Final = get_call_types_for_route(path)
def _call_type_from_url(url: str) -> CallTypesLiteral | None:
path: Final = _url_path(url)
call_types: Final = get_call_types_for_route(path) if path else None
if call_types is None:
return None
scannable: Final = next((c for c in call_types if c in _SCANNABLE_CALL_TYPES), None)
@ -319,6 +328,61 @@ def _scannable_call_type(url: object, body: Mapping[str, object]) -> CallTypesLi
return from_url if from_url is not None else _call_type_from_body(body)
def _route_runs_as(route: str, call_type: CallTypesLiteral) -> bool:
primary: Final = get_primary_call_type_for_route(route)
return primary is not None and primary.value == call_type
def _record_route(url: object, call_type: CallTypesLiteral) -> str | None:
"""
The endpoint route a record is scanned as, so guardrails that read the key's route see the
record's endpoint rather than the upload route (``/v1/files``), which names no scannable call.
"""
path: Final = _url_path(url) if isinstance(url, str) and url else None
candidates: Final = (*((path,) if path else ()), *get_routes_for_call_type(CallTypes(call_type)))
return next((route for route in candidates if _route_runs_as(route, call_type)), None)
def _executed_call_types(payload: Mapping[str, object]) -> frozenset[str]:
"""
The call types Bedrock and Vertex run this record as, from their own record classifiers. OpenAI and
Azure run the record's url, which is already the scan's call type whenever the url is recognized.
"""
from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
from litellm.llms.vertex_ai.files.transformation import (
_is_embeddings_batch_entry, # pyright: ignore[reportPrivateUsage] # Vertex's runtime rule
)
from litellm.types.llms.bedrock import BedrockBatchRecordKind
bedrock_call_types: Final = MappingProxyType(
{
BedrockBatchRecordKind.CHAT: CallTypes.acompletion,
BedrockBatchRecordKind.TEXT_COMPLETION: CallTypes.atext_completion,
BedrockBatchRecordKind.RESPONSES: CallTypes.aresponses,
BedrockBatchRecordKind.EMBEDDING: CallTypes.aembedding,
}
)
classify: Final = BedrockFilesConfig._classify_batch_record # pyright: ignore[reportPrivateUsage] # its own rule
bedrock_kind: Final = classify(payload) # pyright: ignore[reportArgumentType] # reads only url and body
vertex_embeds: Final = _is_embeddings_batch_entry(payload)
return frozenset(
{
bedrock_call_types[bedrock_kind].value,
(CallTypes.aembedding if vertex_embeds else CallTypes.acompletion).value,
}
)
def _filter_route(payload: Mapping[str, object], body: Mapping[str, object], call_type: CallTypesLiteral) -> str | None:
"""
The route call-type filters see for a record, or None when the scan's call type is not the one every
provider would run it as or disagrees with the body, so a relabeled record is scanned rather than skipped.
"""
if _call_type_from_body(body) != call_type or _executed_call_types(payload) != frozenset({call_type}):
return None
return _record_route(payload.get("url"), call_type)
def _custom_id_of(payload: Mapping[str, object]) -> str | None:
"""
The record's identifier, rendered as text.
@ -394,7 +458,9 @@ async def _scan_record(
# The chain hands back the body it produced, which may be a replacement for the dict it was
# given rather than that same dict mutated, so this is what gets compared.
scanned: Final[dict] = await proxy_logging_obj.pre_call_hook( # mutable-ok: the guardrails' own dict
user_api_key_dict=user_api_key_dict,
user_api_key_dict=user_api_key_dict.model_copy(
update={"request_route": _filter_route(record.payload, body, call_type)}
),
data=scan_input,
call_type=call_type,
guardrails_only=True,

View file

@ -26,6 +26,7 @@ from fastapi import (
status,
)
from fastapi.responses import StreamingResponse
from starlette.datastructures import Headers
from starlette.datastructures import UploadFile as StarletteUploadFile
from starlette.websockets import WebSocketState
from websockets.asyncio.client import connect
@ -65,6 +66,7 @@ from litellm.llms.base_llm.managed_resources.utils import (
)
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.passthrough import BasePassthroughUtils
from litellm.proxy._experimental.mcp_server.utils import upstream_credential_headers
from litellm.proxy._types import (
ConfigFieldInfo,
ConfigFieldUpdate,
@ -100,6 +102,10 @@ from litellm.proxy.common_utils.sse_keepalive import (
from litellm.proxy.litellm_pre_call_utils import (
LiteLLMProxyRequestSetup,
_get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage] # shared proxy helper, same import style as _read_request_body above
clean_headers,
key_or_team_allows_client_message_redaction_opt_out,
redact_credential_headers,
strip_untrusted_caller_metadata,
)
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
from litellm.proxy.utils import normalize_route_for_root_path
@ -1022,6 +1028,14 @@ from litellm.passthrough.timeout_utils import (
)
def _guardrail_request_headers(headers: Headers, litellm_key_header_name: str | None) -> Mapping[str, str]:
cleaned: Final = clean_headers(headers, litellm_key_header_name=litellm_key_header_name)
mcp_credential_headers: Final = upstream_credential_headers(cleaned)
return redact_credential_headers(
{name: value for name, value in cleaned.items() if name.lower() not in mcp_credential_headers}
)
async def pass_through_request(
request: Request,
target: str,
@ -1120,6 +1134,26 @@ async def pass_through_request(
_parsed_body = {}
else:
_parsed_body = await _read_request_body(request)
strip_untrusted_caller_metadata(
_parsed_body,
allow_client_message_redaction_opt_out=key_or_team_allows_client_message_redaction_opt_out(
user_api_key_dict
),
)
# Lazy: proxy_server imports this module
from litellm.proxy.proxy_server import (
general_settings as proxy_general_settings,
)
from litellm.proxy.proxy_server import (
general_settings_view,
)
# Guardrails forward these to vendors as the inbound request headers; all are popped before the upstream send.
_parsed_body.pop("proxy_server_request", None)
_parsed_body.pop("headers", None)
for _caller_bucket in (_parsed_body.get("metadata"), _parsed_body.get("litellm_metadata")):
if isinstance(_caller_bucket, dict):
_caller_bucket.pop("headers", None)
verbose_proxy_logger.debug(
"Pass through endpoint sending request to \nURL %s\nheaders: %s\nbody: %s\n",
url,
@ -1174,6 +1208,10 @@ async def pass_through_request(
if _parsed_body is None:
_parsed_body = {}
_parsed_body["litellm_logging_obj"] = logging_obj
guardrail_headers: Final = _guardrail_request_headers(
request.headers, litellm_key_header_name=proxy_general_settings.get("litellm_key_header_name")
)
_parsed_body["proxy_server_request"] = {"headers": guardrail_headers}
### CALL HOOKS ### - modify incoming data / reject request before calling the model
_parsed_body = await proxy_logging_obj.pre_call_hook(
@ -1222,13 +1260,6 @@ async def pass_through_request(
# provider IDs before forwarding upstream. Gated by feature flag and
# enterprise managed-files hook. Runs after pre_call_hook so
# guardrails have already seen the managed IDs.
from litellm.proxy.proxy_server import (
general_settings as proxy_general_settings,
)
from litellm.proxy.proxy_server import (
general_settings_view,
)
_managed_id_provider: Final = resolve_passthrough_managed_id_provider(custom_llm_provider)
if proxy_general_settings.get("passthrough_managed_object_ids", False) and _managed_id_provider is not None:
@ -1636,6 +1667,7 @@ async def pass_through_request(
**existing_metadata,
"guardrails": guardrails_to_run,
}
hook_data["proxy_server_request"] = {"headers": guardrail_headers}
post_call_guardrail_data = hook_data
response_body = await proxy_logging_obj.post_call_success_hook(
data=hook_data,

View file

@ -103,6 +103,31 @@ class GenericGuardrailAPIOptionalParams(BaseModel):
),
)
run_only_on_call_types: tuple[str, ...] | None = Field(
default=None,
description=(
"If set, the guardrail runs only for these call types and every other call type is passed "
"through without calling the guardrail endpoint. Takes precedence over skip_call_types. "
"Values are CallTypes values, the strings logged as call_type (e.g. ['acompletion', "
"'anthropic_messages', 'aresponses']). Unknown values and call_mcp_tool are rejected at "
"startup. The call type comes from the authenticated request route: /anthropic/v1/messages "
"pass-through calls are anthropic_messages, other pass-through calls are "
"pass_through_endpoint, and a batch-file record is classified by its own endpoint when Bedrock "
"and Vertex, which run record urls themselves, run it as that call type and its body agrees. "
"A call whose type cannot be resolved, including any other batch-file record, still runs "
"the guardrail."
),
)
skip_call_types: tuple[str, ...] | None = Field(
default=None,
description=(
"Call types that are passed through without calling the guardrail endpoint, e.g. "
"['aembedding', 'aimage_generation']. Same values and resolution as run_only_on_call_types. "
"Ignored when run_only_on_call_types is set."
),
)
class GenericGuardrailAPIConfigModel(
GuardrailConfigModel[GenericGuardrailAPIOptionalParams],

View file

@ -416,6 +416,22 @@ class TestMetadataExtraction:
assert request_metadata["user_api_key_hash"] == "hashed-token-value"
assert request_metadata["user_api_key_user_id"] == "test-user"
@pytest.mark.parametrize(
"request_data, expected_hash",
[
pytest.param({"metadata": {"user_api_key_token": "caller-token"}}, None, id="caller-bucket-token-ignored"),
pytest.param(
{"litellm_metadata": {"user_api_key_token": "proxy-token", "user_api_key_hash": "logged-key"}},
"logged-key",
id="hash-wins-over-token",
),
],
)
def test_token_fallback_only_from_litellm_metadata_and_only_without_hash(
self, generic_guardrail, request_data, expected_hash
):
assert generic_guardrail._extract_user_api_key_metadata(request_data).get("user_api_key_hash") == expected_hash
@pytest.mark.asyncio
async def test_metadata_extraction_empty_when_no_metadata(self, generic_guardrail):
"""Test metadata extraction returns empty dict when no metadata available"""

View file

@ -582,15 +582,20 @@ def test_ensure_litellm_metadata_populates_from_user_api_key_dict() -> None:
assert data["litellm_metadata"]["user_api_key_team_id"] == "t1"
def test_ensure_litellm_metadata_noop_when_already_present() -> None:
"""Verify _ensure_litellm_metadata does not overwrite existing litellm_metadata."""
def test_ensure_litellm_metadata_overrides_caller_identity_in_existing_bucket() -> None:
"""An existing litellm_metadata keeps its other keys, but its identity comes from the authenticated key."""
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
_ensure_litellm_metadata,
)
user_auth = UserAPIKeyAuth(user_id="should-not-appear")
data: dict = {"litellm_metadata": {"existing": "value"}}
user_auth = UserAPIKeyAuth(user_id="auth-user", key_alias="auth-alias", team_id="auth-team")
bucket: dict = {"existing": "value", "user_api_key_alias": "batch-worker", "user_api_key_team_id": "team-exempt"}
data: dict = {"litellm_metadata": bucket}
_ensure_litellm_metadata(data, user_auth)
assert data["litellm_metadata"] == {"existing": "value"}
assert data["litellm_metadata"] is bucket
assert bucket["existing"] == "value"
assert bucket["user_api_key_alias"] == "auth-alias"
assert bucket["user_api_key_team_id"] == "auth-team"
assert bucket["user_api_key_user_id"] == "auth-user"

View file

@ -1,10 +1,16 @@
"""Tests for unified guardrail."""
import copy
import json
import logging
from collections.abc import Callable
from types import SimpleNamespace
from typing import TYPE_CHECKING, Final, Literal
from unittest.mock import MagicMock
import httpx
import pytest
from fastapi import Request
import litellm
from litellm.caching import DualCache
@ -22,6 +28,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import (
openai_messages_without_tool,
)
from litellm.llms.base_llm.ocr.transformation import OCRPage, OCRResponse
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.llms.mistral.ocr.guardrail_translation.handler import OCRHandler
from litellm.llms.openai.chat.guardrail_translation.handler import (
OpenAIChatCompletionsHandler,
@ -33,12 +40,15 @@ from litellm.proxy._experimental.mcp_server.guardrail_translation.handler import
MCPGuardrailTranslationHandler,
)
from litellm.proxy._types import LiteLLMRoutes, UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail import (
unified_guardrail as unified_module,
)
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup, add_litellm_data_to_request
from litellm.proxy.utils import ProxyLogging
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.utils import CallTypes, Delta, GenericGuardrailAPIInputs, ModelResponseStream, StreamingChoices
@ -2395,3 +2405,215 @@ class TestTranslationMappingsAreReadLive:
assert not [
name for name, value in vars(unified_module).items() if isinstance(value, dict) and CallTypes.aocr in value
]
_RAW_CLI_SESSION_TOKEN: Final = "cli-session-Qm7xJ2kP9sLw4vT1nR8yAa"
def _sk_key(route: str) -> UserAPIKeyAuth:
return UserAPIKeyAuth(
api_key="sk-real-caller-key",
key_alias="prod-app",
team_id="team-prod",
metadata={"key_label": "k1"},
team_metadata={"phoenix_project_name": "team-proj", "priority": "high"},
request_route=route,
)
def _cli_session_key(route: str) -> UserAPIKeyAuth:
return UserAPIKeyAuth(
token=_RAW_CLI_SESSION_TOKEN,
key_alias="cli-session-alice",
user_id="alice",
is_session_token=True,
team_id="team-prod",
team_metadata={"phoenix_project_name": "team-proj", "priority": "high"},
request_route=route,
)
class TestGuardrailsSeeAuthenticatedIdentity:
"""A request body cannot make a guardrail vendor see another key's identity, and the real one reaches it."""
@staticmethod
def _generic_guardrail(vendor_payloads: list[dict[str, object]]) -> CustomGuardrail:
def vendor(request: httpx.Request) -> httpx.Response:
vendor_payloads.append(json.loads(request.content))
return httpx.Response(200, json={"action": "NONE"})
guardrail = GenericGuardrailAPI(api_base="https://guardrail.test", guardrail_name="generic")
guardrail.async_handler = AsyncHTTPHandler(transport=httpx.MockTransport(vendor))
return guardrail
@pytest.mark.asyncio
@pytest.mark.parametrize("bucket", ["metadata", "litellm_metadata"])
async def test_pass_through_body_cannot_forge_identity(self, monkeypatch, bucket: str) -> None:
_patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings())
vendor_payloads: list[dict[str, object]] = []
key = UserAPIKeyAuth(api_key="sk-real-caller-key", key_alias="prod-app", team_id="team-prod")
data = {
"guardrail_to_apply": self._generic_guardrail(vendor_payloads),
"prompt": "hello",
bucket: {
"user_api_key_alias": "batch-worker",
"user_api_key_team_id": "team-exempt",
"user_api_key_token": "forged-hash",
},
}
await UnifiedLLMGuardrails().async_pre_call_hook(
user_api_key_dict=key, cache=DualCache(), data=data, call_type=CallTypes.pass_through.value
)
assert len(vendor_payloads) == 1
identity = vendor_payloads[0]["request_data"]
assert identity["user_api_key_alias"] == "prod-app"
assert identity["user_api_key_team_id"] == "team-prod"
assert identity["user_api_key_hash"] == key.api_key
@pytest.mark.asyncio
async def test_mcp_tool_call_reaches_vendor_with_key_alias(self) -> None:
vendor_payloads: list[dict[str, object]] = []
key = UserAPIKeyAuth(api_key="sk-real-caller-key", key_alias="prod-app", team_id="team-prod")
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
mcp_kwargs = {
"name": "search",
"arguments": {"query": "hello"},
"server_name": "docs",
"user_api_key_auth": key,
"user_api_key_user_id": key.user_id,
"user_api_key_team_id": key.team_id,
"user_api_key_end_user_id": None,
"user_api_key_hash": key.api_key,
"headers": {},
}
data = proxy_logging._convert_mcp_to_llm_format(
proxy_logging._create_mcp_request_object_from_kwargs(mcp_kwargs), mcp_kwargs
)
data["guardrail_to_apply"] = self._generic_guardrail(vendor_payloads)
await UnifiedLLMGuardrails().async_pre_call_hook(
user_api_key_dict=key, cache=DualCache(), data=data, call_type=CallTypes.call_mcp_tool.value
)
assert len(vendor_payloads) == 1
identity = vendor_payloads[0]["request_data"]
assert identity["user_api_key_alias"] == "prod-app"
assert identity["user_api_key_team_id"] == "team-prod"
assert identity["user_api_key_hash"] == key.api_key
@pytest.mark.asyncio
async def test_pass_through_body_cannot_forge_request_route(self, monkeypatch) -> None:
"""Guardrails and call-type lookups key on user_api_key_request_route, so it must be the key's own route."""
_patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings())
key = UserAPIKeyAuth(api_key="sk-real-caller-key", request_route="/openai/v1/chat/completions")
data = {
"guardrail_to_apply": RecordingGuardrail(),
"prompt": "hello",
"litellm_metadata": {"user_api_key_request_route": "/v1/embeddings"},
}
await UnifiedLLMGuardrails().async_pre_call_hook(
user_api_key_dict=key, cache=DualCache(), data=data, call_type=CallTypes.pass_through.value
)
assert data["litellm_metadata"]["user_api_key_request_route"] == "/openai/v1/chat/completions"
@pytest.mark.asyncio
@pytest.mark.parametrize("route", ["/v1/messages", "/v1/responses", "/v1/chat/completions"])
@pytest.mark.parametrize(
"make_key", [pytest.param(_sk_key, id="sk-key"), pytest.param(_cli_session_key, id="cli-session-key")]
)
async def test_chat_path_request_keeps_proxy_metadata_and_sends_stable_hash(
self, monkeypatch, route: str, make_key: Callable[[str], UserAPIKeyAuth]
) -> None:
"""After the chat-path metadata build, the guardrail hook leaves the proxy's bucket as it was (team metadata
in user_api_key_auth_metadata included) and the vendor gets the logged key, never a raw CLI session token."""
_patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings())
request = MagicMock(spec=Request)
request.url = MagicMock()
request.url.path = route
request.url.__str__.return_value = "http://localhost" + route
request.method = "POST"
request.query_params = {}
request.headers = {"Content-Type": "application/json"}
request.client = MagicMock()
request.client.host = "127.0.0.1"
request.state = MagicMock()
key = make_key(route)
data = await add_litellm_data_to_request(
data={"model": "m", "messages": [{"role": "user", "content": "hi"}]},
request=request,
user_api_key_dict=key,
proxy_config=MagicMock(),
general_settings={},
version="v",
)
proxy_bucket = data.get("litellm_metadata")
unshared = ("litellm_parent_otel_span", "user_api_key_auth")
bucket_before = (
copy.deepcopy({k: v for k, v in proxy_bucket.items() if k not in unshared}) if proxy_bucket else None
)
vendor_payloads: list[dict[str, object]] = []
data["guardrail_to_apply"] = self._generic_guardrail(vendor_payloads)
await UnifiedLLMGuardrails().async_pre_call_hook(
user_api_key_dict=key, cache=DualCache(), data=data, call_type=CallTypes.acompletion.value
)
assert len(vendor_payloads) == 1
assert vendor_payloads[0]["request_data"]["user_api_key_hash"] == LiteLLMProxyRequestSetup.get_logged_api_key(
key
)
assert _RAW_CLI_SESSION_TOKEN not in json.dumps(vendor_payloads[0])
if bucket_before is not None:
assert data["litellm_metadata"] is proxy_bucket
assert {k: v for k, v in proxy_bucket.items() if k in bucket_before} == bucket_before
assert bucket_before["user_api_key_auth_metadata"]["priority"] == "high"
assert "user_api_key_token" not in proxy_bucket
@pytest.mark.asyncio
@pytest.mark.parametrize("bucket", ["metadata", "litellm_metadata", None])
async def test_pass_through_cli_session_key_sends_stable_hash(self, monkeypatch, bucket: str | None) -> None:
_patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings())
vendor_payloads: list[dict[str, object]] = []
key = _cli_session_key("/anthropic/v1/messages")
forged_bucket = {bucket: {"user_api_key_token": "forged-hash"}} if bucket else {}
data = {
"guardrail_to_apply": self._generic_guardrail(vendor_payloads),
"prompt": "hello",
"proxy_server_request": {"headers": {"x-tenant": "tenant-real"}},
**forged_bucket,
}
await UnifiedLLMGuardrails().async_pre_call_hook(
user_api_key_dict=key, cache=DualCache(), data=data, call_type=CallTypes.pass_through.value
)
assert len(vendor_payloads) == 1
assert vendor_payloads[0]["request_data"]["user_api_key_hash"] == "cli-session-alice"
assert _RAW_CLI_SESSION_TOKEN not in json.dumps(vendor_payloads[0])
assert vendor_payloads[0]["texts"] == ['{"prompt": "hello"}']
assert vendor_payloads[0]["request_headers"] == {"x-tenant": "[present]"}
@pytest.mark.asyncio
async def test_token_only_key_drops_forged_token_already_in_proxy_bucket(self, monkeypatch) -> None:
"""A key with no api_key logs no hash, so a user_api_key_token left in litellm_metadata would become the
vendor's hash if the hook kept it."""
_patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings())
vendor_payloads: list[dict[str, object]] = []
key = UserAPIKeyAuth(token="abc123hashed", key_alias="prod-app")
data = {
"guardrail_to_apply": self._generic_guardrail(vendor_payloads),
"prompt": "hello",
"litellm_metadata": {"user_api_key_token": "forged-hash"},
}
await UnifiedLLMGuardrails().async_pre_call_hook(
user_api_key_dict=key, cache=DualCache(), data=data, call_type=CallTypes.pass_through.value
)
assert len(vendor_payloads) == 1
assert "user_api_key_token" not in data["litellm_metadata"]
assert vendor_payloads[0]["request_data"].get("user_api_key_hash") is None

View file

@ -4,11 +4,13 @@ from datetime import datetime
from typing import Dict, List, Optional
from unittest.mock import AsyncMock
import httpx
import pytest
from fastapi import HTTPException
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_endpoints import (
CreateGuardrailRequest,
@ -33,6 +35,7 @@ from litellm.proxy.guardrails.guardrail_endpoints import (
from litellm.proxy.guardrails.guardrail_endpoints import (
test_custom_code_guardrail as run_custom_code_test_endpoint,
)
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI
MOCK_ADMIN_USER = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
from litellm.proxy.guardrails.guardrail_registry import (
@ -1487,12 +1490,12 @@ async def test_apply_guardrail_invokes_logging_pipeline(mocker):
}
def _patch_apply_guardrail_env(mocker, guardrail_result):
def _patch_apply_guardrail_env(mocker, guardrail_result, processed_data=None, guardrail=None):
mock_guardrail = mocker.Mock()
mock_guardrail.apply_guardrail = AsyncMock(return_value=guardrail_result)
mock_registry = mocker.Mock()
mock_registry.get_initialized_guardrail_callback.return_value = mock_guardrail
mock_registry.get_initialized_guardrail_callback.return_value = guardrail or mock_guardrail
mocker.patch("litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY", mock_registry)
mock_logging_obj = mocker.Mock()
@ -1500,7 +1503,7 @@ def _patch_apply_guardrail_env(mocker, guardrail_result):
mock_logging_obj.model_call_details = {}
mock_processor = mocker.Mock()
mock_processor.common_processing_pre_call_logic = AsyncMock(
return_value=({"guardrail_name": "test-guardrail"}, mock_logging_obj)
return_value=(processed_data or {"guardrail_name": "test-guardrail"}, mock_logging_obj)
)
mocker.patch(
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing",
@ -1535,11 +1538,12 @@ async def test_apply_guardrail_forwards_metadata_to_guardrail(mocker):
user_api_key_dict=UserAPIKeyAuth(),
)
mock_guardrail.apply_guardrail.assert_awaited_once_with(
inputs={"texts": ["What are tax loopholes?"]},
request_data={"metadata": {"forbidden_topics": ["tax"]}},
input_type="request",
)
mock_guardrail.apply_guardrail.assert_awaited_once()
call = mock_guardrail.apply_guardrail.await_args.kwargs
assert call["inputs"] == {"texts": ["What are tax loopholes?"]}
assert call["input_type"] == "request"
assert "messages" not in call["request_data"]
assert call["request_data"]["metadata"]["forbidden_topics"] == ["tax"]
@pytest.mark.asyncio
@ -1561,39 +1565,201 @@ async def test_apply_guardrail_forwards_metadata_and_messages_together(mocker):
user_api_key_dict=UserAPIKeyAuth(),
)
mock_guardrail.apply_guardrail.assert_awaited_once_with(
inputs={"texts": ["What are tax loopholes?"]},
request_data={
"messages": messages,
"metadata": {"forbidden_topics": ["tax"]},
},
input_type="request",
)
request_data = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]
assert request_data["messages"] == messages
assert request_data["metadata"]["forbidden_topics"] == ["tax"]
@pytest.mark.asyncio
async def test_apply_guardrail_omits_metadata_when_not_sent(mocker):
"""Without metadata, request_data stays empty (backward-compatible)."""
async def test_apply_guardrail_authenticated_identity_overrides_client_metadata(mocker):
"""A caller must not be able to claim another key's or team's identity in the body metadata."""
mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]})
caller = UserAPIKeyAuth(
api_key="sk-real-caller-key",
key_alias="real-caller",
team_id="real-team",
user_id="real-user",
)
request = ApplyGuardrailRequest(
guardrail_name="test-guardrail",
text="hello",
metadata={
"user_api_key_alias": "exempt-batch-worker",
"user_api_key_team_id": "exempt-team",
"user_api_key_user_id": "someone-else",
"user_api_key_hash": "forged-hash",
"forbidden_topics": ["tax"],
},
)
await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller)
metadata = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]["metadata"]
assert metadata["user_api_key_alias"] == "real-caller"
assert metadata["user_api_key_team_id"] == "real-team"
assert metadata["user_api_key_user_id"] == "real-user"
assert metadata["user_api_key_hash"] == caller.api_key
assert metadata["user_api_key_hash"] != "forged-hash"
assert metadata["forbidden_topics"] == ["tax"]
@pytest.mark.asyncio
async def test_apply_guardrail_drops_client_identity_fields_the_key_does_not_set(mocker):
"""Proxy-owned slots in the body never reach the guardrail, including user_api_key_token, which guardrails
map onto the key hash, and control fields the chat path also strips."""
mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]})
caller = UserAPIKeyAuth(metadata={"zguard_policy_id": "strict"}, object_permission_id="perm-real")
request = ApplyGuardrailRequest(
guardrail_name="test-guardrail",
text="hello",
metadata={
"user_api_key_alias": "exempt-batch-worker",
"user_api_key_team_id": "exempt-team",
"user_api_key_token": "forged-hash",
"user_api_key_metadata": {"zguard_policy_id": "permissive"},
"user_api_key_object_permission_id": "perm-forged",
"user_api_key": "forged-key",
"applied_guardrails": ["already-ran"],
"headers": {"x-end-user": "someone-else"},
"trace_label": "nightly",
},
)
await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller)
metadata = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]["metadata"]
assert metadata["user_api_key_alias"] is None
assert metadata["user_api_key_team_id"] is None
assert "user_api_key_token" not in metadata
assert metadata["user_api_key_metadata"] == {"zguard_policy_id": "strict"}
assert metadata["user_api_key_object_permission_id"] == "perm-real"
assert metadata["user_api_key"] is None
assert "applied_guardrails" not in metadata
assert "headers" not in metadata
assert metadata["trace_label"] == "nightly"
@pytest.mark.asyncio
async def test_apply_guardrail_forwards_real_request_headers_not_caller_supplied_ones(mocker):
"""Guardrails forward metadata headers to vendors, so they must be the proxy's view of the request."""
real_headers = {"user-agent": "real-client/1.0"}
mock_guardrail = _patch_apply_guardrail_env(
mocker,
{"texts": ["ok"]},
processed_data={"guardrail_name": "test-guardrail", "metadata": {"headers": real_headers}},
)
request = ApplyGuardrailRequest(
guardrail_name="test-guardrail",
text="hello",
metadata={"headers": {"user-agent": "forged/1.0", "x-end-user": "someone-else"}},
)
await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=UserAPIKeyAuth())
metadata = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]["metadata"]
assert metadata["headers"] == real_headers
@pytest.mark.asyncio
async def test_apply_guardrail_generic_guardrail_api_sends_authenticated_identity_to_vendor(mocker):
"""End to end through a real GenericGuardrailAPI: the vendor payload names the authenticated key even when
the body forges user_api_key_alias and user_api_key_token, which the generic guardrail maps onto the hash."""
vendor_payloads = []
def vendor(request: httpx.Request) -> httpx.Response:
vendor_payloads.append(json.loads(request.content))
return httpx.Response(200, json={"action": "NONE"})
generic_guardrail = GenericGuardrailAPI(api_base="https://guardrail.test", guardrail_name="generic")
generic_guardrail.async_handler = AsyncHTTPHandler(transport=httpx.MockTransport(vendor))
_patch_apply_guardrail_env(mocker, {"texts": ["unused"]}, guardrail=generic_guardrail)
caller = UserAPIKeyAuth(api_key="sk-real-caller-key", key_alias="real-caller", team_id="real-team")
request = ApplyGuardrailRequest(
guardrail_name="generic",
text="hello",
metadata={"user_api_key_alias": "exempt-batch-worker", "user_api_key_token": "forged-hash"},
)
response = await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller)
assert response.response_text == "hello"
assert len(vendor_payloads) == 1
identity = vendor_payloads[0]["request_data"]
assert identity["user_api_key_alias"] == "real-caller"
assert identity["user_api_key_team_id"] == "real-team"
assert identity["user_api_key_hash"] == caller.api_key
@pytest.mark.asyncio
async def test_apply_guardrail_request_route_comes_from_the_key(mocker):
"""Guardrails pick call-type behavior from user_api_key_request_route, so the body cannot choose it."""
mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]})
request = ApplyGuardrailRequest(
guardrail_name="test-guardrail",
text="hello",
metadata={"user_api_key_request_route": "/v1/embeddings"},
)
await apply_guardrail(
fastapi_request=mocker.Mock(),
request=request,
user_api_key_dict=UserAPIKeyAuth(request_route="/guardrails/apply_guardrail"),
)
metadata = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]["metadata"]
assert metadata["user_api_key_request_route"] == "/guardrails/apply_guardrail"
@pytest.mark.asyncio
async def test_apply_guardrail_cli_session_key_sends_stable_hash_to_vendor(mocker):
"""A CLI session key's raw per-login token must never reach the vendor; it gets the stable logged key."""
raw_session_token = "cli-session-Qm7xJ2kP9sLw4vT1nR8yAa"
vendor_payloads = []
def vendor(request: httpx.Request) -> httpx.Response:
vendor_payloads.append(json.loads(request.content))
return httpx.Response(200, json={"action": "NONE"})
generic_guardrail = GenericGuardrailAPI(api_base="https://guardrail.test", guardrail_name="generic")
generic_guardrail.async_handler = AsyncHTTPHandler(transport=httpx.MockTransport(vendor))
_patch_apply_guardrail_env(mocker, {"texts": ["unused"]}, guardrail=generic_guardrail)
caller = UserAPIKeyAuth(
token=raw_session_token, key_alias="cli-session-alice", user_id="alice", is_session_token=True
)
request = ApplyGuardrailRequest(
guardrail_name="generic", text="hello", metadata={"user_api_key_token": raw_session_token}
)
await apply_guardrail(fastapi_request=mocker.Mock(), request=request, user_api_key_dict=caller)
assert len(vendor_payloads) == 1
assert vendor_payloads[0]["request_data"]["user_api_key_hash"] == "cli-session-alice"
assert raw_session_token not in json.dumps(vendor_payloads[0])
@pytest.mark.asyncio
async def test_apply_guardrail_carries_authenticated_identity_when_no_metadata_sent(mocker):
"""request_data always carries the authenticated identity, even when the body has no metadata."""
mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]})
request = ApplyGuardrailRequest(guardrail_name="test-guardrail", text="hello")
await apply_guardrail(
fastapi_request=mocker.Mock(),
request=request,
user_api_key_dict=UserAPIKeyAuth(),
user_api_key_dict=UserAPIKeyAuth(key_alias="known-caller", team_id="known-team"),
)
mock_guardrail.apply_guardrail.assert_awaited_once_with(
inputs={"texts": ["hello"]},
request_data={},
input_type="request",
)
call = mock_guardrail.apply_guardrail.await_args.kwargs
assert call["inputs"] == {"texts": ["hello"]}
assert "messages" not in call["request_data"]
assert call["request_data"]["metadata"]["user_api_key_alias"] == "known-caller"
assert call["request_data"]["metadata"]["user_api_key_team_id"] == "known-team"
@pytest.mark.asyncio
async def test_apply_guardrail_forwards_explicit_empty_messages_and_metadata(mocker):
"""Explicitly-sent empty messages/metadata must be forwarded, not dropped;
only omitted fields stay out of request_data."""
"""Explicitly-sent empty messages must be forwarded, not dropped, and empty
metadata still carries the authenticated identity."""
mock_guardrail = _patch_apply_guardrail_env(mocker, {"texts": ["ok"]})
request = ApplyGuardrailRequest(
@ -1605,14 +1771,12 @@ async def test_apply_guardrail_forwards_explicit_empty_messages_and_metadata(moc
await apply_guardrail(
fastapi_request=mocker.Mock(),
request=request,
user_api_key_dict=UserAPIKeyAuth(),
user_api_key_dict=UserAPIKeyAuth(key_alias="known-caller"),
)
mock_guardrail.apply_guardrail.assert_awaited_once_with(
inputs={"texts": ["hello"]},
request_data={"messages": [], "metadata": {}},
input_type="request",
)
request_data = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]
assert request_data["messages"] == []
assert request_data["metadata"]["user_api_key_alias"] == "known-caller"
@pytest.mark.asyncio

View file

@ -1186,3 +1186,85 @@ async def test_a_dropped_record_without_a_named_guardrail_reports_none():
result = await _scan_full(_jsonl(_record("b", content="tripwire")), FakeProxyLogging(_hook))
assert result.changes == (RecordDropped(line_number=1, custom_id="b", guardrail=None),)
_CHAT_BODY = {"messages": [{"role": "user", "content": "x"}]}
@pytest.mark.asyncio
async def test_a_record_carrying_its_own_logging_obj_is_still_scanned_when_the_guardrail_fails_open(monkeypatch):
"""A caller's litellm_logging_obj used to crash the guardrail call, which fail_on_error=False then let through."""
import httpx
import litellm
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI
received = []
def _respond(request):
received.append(json.loads(request.content))
return httpx.Response(200, json={"action": "BLOCKED", "blocked_reason": "blocked by test endpoint"})
guardrail = GenericGuardrailAPI(
api_base="https://guardrail.test",
guardrail_name="fail-open-guard",
event_hook="pre_call",
default_on=True,
fail_on_error=False,
async_handler=AsyncHTTPHandler(transport=httpx.MockTransport(_respond)),
)
monkeypatch.setattr(litellm, "callbacks", [guardrail])
ProxyLogging._callback_capabilities_cache.clear()
record = _record("carrier", content="carried chat")
record["body"]["litellm_logging_obj"] = {"a": 1}
result = await _scan_full(_jsonl(record), ProxyLogging(user_api_key_cache=DualCache()))
assert len(received) == 1, "the record reached the guardrail endpoint"
assert result.changes == (RecordDropped(line_number=1, custom_id="carrier", guardrail="generic_guardrail_api"),)
ProxyLogging._callback_capabilities_cache.clear()
class RouteRecordingProxyLogging(FakeProxyLogging):
async def pre_call_hook(self, user_api_key_dict, data, call_type, guardrails_only=False):
self.seen.append((call_type, user_api_key_dict.request_route))
return data
@pytest.mark.asyncio
@pytest.mark.parametrize(
("url", "body", "expected"),
[
("/v1/chat/completions", _CHAT_BODY, ("acompletion", "/v1/chat/completions")),
("/v1/embeddings", {"input": "x"}, ("aembedding", "/v1/embeddings")),
("/custom/unmapped", _CHAT_BODY, ("acompletion", "/chat/completions")),
("/v1/rerank", _CHAT_BODY, ("acompletion", "/chat/completions")),
("", _CHAT_BODY, ("acompletion", "/chat/completions")),
("/v1/messages", _CHAT_BODY, ("anthropic_messages", None)),
("/anthropic/v1/messages", _CHAT_BODY, ("anthropic_messages", None)),
("https://api.openai.com/v1/embeddings", {"input": "x"}, ("aembedding", None)),
("/embeddings", {"input": "x"}, ("aembedding", None)),
("/v1/embeddings", _CHAT_BODY, ("aembedding", None)),
("", {"input": "x"}, ("aembedding", None)),
("/v1/responses", {"input": "x"}, ("aresponses", None)),
("/v1/completions", {"prompt": "x"}, ("atext_completion", None)),
],
)
async def test_each_record_is_scanned_under_the_route_every_provider_runs_it_as(url, body, expected):
"""
Guardrails that classify by the key's route see the record's endpoint, not the upload route, and no
route at all when some batch provider would run the record as a different call type than its url says
"""
upload_key = UserAPIKeyAuth(api_key="sk-test", request_route="/v1/files")
logging_obj = RouteRecordingProxyLogging()
await scan_batch_input_file(
file_source=_jsonl({"custom_id": "r", "method": "POST", "url": url, "body": {"model": "m", **body}}),
request_metadata={},
user_api_key_dict=upload_key,
proxy_logging_obj=logging_obj,
)
assert logging_obj.seen == [expected]
assert upload_key.request_route == "/v1/files", "the upload's own key must not be rewritten"

View file

@ -26,7 +26,12 @@ from litellm._logging import verbose_proxy_logger
from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.generic_guardrail_api import (
_extract_inbound_headers,
)
from litellm.proxy.litellm_pre_call_utils import _REDACTED_HEADER_VALUE
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS,
LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
@ -48,6 +53,7 @@ from litellm.proxy.pass_through_endpoints.success_handler import (
)
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
from litellm.types import utils as types_utils
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY,
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
@ -7622,3 +7628,166 @@ def test_passthrough_attributes_a_cli_session_to_its_alias_not_the_login_token()
metadata = kwargs["litellm_params"]["metadata"]
assert metadata["user_api_key"] == "cli-session-alice"
assert _get_spend_logs_metadata(metadata)["user_api_key"] == "cli-session-alice"
@pytest.mark.asyncio
async def test_pass_through_request_strips_caller_identity_before_guardrail_hooks():
"""
Regression: a pass-through body skips add_litellm_data_to_request, so forged user_api_key_* fields, guardrail
control fields and inbound headers reached pre_call_hook guardrails as the caller's identity. The upstream body
is unchanged because these keys never reach it.
"""
upstream_bodies = []
def transport_handler(upstream_request: httpx.Request) -> httpx.Response:
upstream_bodies.append(json.loads(upstream_request.content))
return httpx.Response(200, json={"ok": True})
real_handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.PassThroughEndpoint,
params={"timeout": resolve_pass_through_request_timeout(None)},
)
cache_dict = litellm.in_memory_llm_clients_cache.cache_dict
cache_key = next(key for key, cached in cache_dict.items() if cached is real_handler)
cache_dict[cache_key] = SimpleNamespace(client=httpx.AsyncClient(transport=httpx.MockTransport(transport_handler)))
hook_data = []
def record_hook_data(user_api_key_dict, data, call_type):
hook_data.append({key: value for key, value in data.items() if key != "litellm_logging_obj"})
return data
mock_proxy_logging = MagicMock()
mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=record_hook_data)
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={})
forged = {
"user_api_key_alias": "batch-worker",
"user_api_key_team_id": "team-exempt",
"user_api_key_token": "forged-hash",
"user_api_key_request_route": "/v1/embeddings",
"disable_global_guardrails": True,
"headers": {"x-authenticated-user": "admin@corp"},
"trace_label": "nightly",
}
forged_headers = {"x-authenticated-user": "admin@corp", "x-litellm-end-user-id": "victim", "x-tenant": "forged"}
body = {
"prompt": "hello",
"metadata": forged,
"litellm_metadata": forged,
"headers": forged_headers,
"proxy_server_request": {"headers": forged_headers},
}
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.headers = Headers(
{
"content-type": "application/json",
"x-tenant": "tenant-real",
"authorization": "Bearer sk-real-caller-key",
"x-my-key": "sk-custom-header-key",
"x-mcp-github-authorization": "Bearer mcp-upstream-token",
"cookie": "session=secret",
}
)
mock_request.query_params = QueryParams({})
mock_request.body = AsyncMock(return_value=json.dumps(body).encode())
try:
with (
patch(
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging
), # test-quality-ok: read at call time
patch("litellm.proxy.proxy_server.general_settings", {"litellm_key_header_name": "x-my-key"}),
):
response = await pass_through_request(
request=mock_request,
target="https://upstream.test/v1/generate",
custom_headers={},
user_api_key_dict=UserAPIKeyAuth(api_key="sk-real-caller-key", key_alias="prod-app"),
)
finally:
cache_dict[cache_key] = real_handler
assert response.status_code == 200
assert hook_data == [
{
"prompt": "hello",
"metadata": {"trace_label": "nightly"},
"litellm_metadata": {"trace_label": "nightly"},
"proxy_server_request": {
"headers": {
"content-type": "application/json",
"x-tenant": "tenant-real",
"cookie": _REDACTED_HEADER_VALUE,
}
},
}
]
vendor_headers = _extract_inbound_headers(request_data=hook_data[0], logging_obj=None, extra_allowlist={"x-tenant"})
assert vendor_headers is not None and vendor_headers["x-tenant"] == "tenant-real"
assert upstream_bodies == [{"prompt": "hello"}]
@pytest.mark.asyncio
async def test_pass_through_post_call_guardrails_receive_real_inbound_headers():
"""Post-call guardrails run on a copy of the body the litellm-param pop already stripped, so without an explicit
re-attach an operator's extra_headers allowlist forwarded nothing on the response side."""
def transport_handler(upstream_request: httpx.Request) -> httpx.Response:
return httpx.Response(200, json={"completion": "hi"})
real_handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.PassThroughEndpoint,
params={"timeout": resolve_pass_through_request_timeout(None)},
)
cache_dict = litellm.in_memory_llm_clients_cache.cache_dict
cache_key = next(key for key, cached in cache_dict.items() if cached is real_handler)
cache_dict[cache_key] = SimpleNamespace(client=httpx.AsyncClient(transport=httpx.MockTransport(transport_handler)))
post_call_data = []
def record_post_call(data, user_api_key_dict, response):
post_call_data.append(data)
return response
mock_proxy_logging = MagicMock()
mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=lambda user_api_key_dict, data, call_type: data)
mock_proxy_logging.post_call_success_hook = AsyncMock(side_effect=record_post_call)
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={})
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.headers = Headers(
{"content-type": "application/json", "x-tenant": "tenant-real", "authorization": "Bearer sk-real-caller-key"}
)
mock_request.query_params = QueryParams({})
mock_request.body = AsyncMock(
return_value=json.dumps({"prompt": "hello", "headers": {"x-tenant": "forged"}}).encode()
)
try:
with patch(
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging
): # test-quality-ok: read at call time
response = await pass_through_request(
request=mock_request,
target="https://upstream.test/v1/generate",
custom_headers={},
user_api_key_dict=UserAPIKeyAuth(api_key="sk-real-caller-key", key_alias="prod-app"),
guardrails_config=["gg"],
)
finally:
cache_dict[cache_key] = real_handler
assert response.status_code == 200
assert len(post_call_data) == 1, "the post-call guardrail hook did not run"
assert post_call_data[0]["proxy_server_request"] == {
"headers": {"content-type": "application/json", "x-tenant": "tenant-real"}
}
vendor_headers = _extract_inbound_headers(
request_data=post_call_data[0], logging_obj=None, extra_allowlist={"x-tenant"}
)
assert vendor_headers is not None and vendor_headers["x-tenant"] == "tenant-real"

View file

@ -61,11 +61,11 @@ async def test_apply_guardrail_endpoint_returns_correct_response(
assert response.response_text == "Redacted text: [REDACTED] and [REDACTED]"
# Verify the guardrail was called with correct parameters
mock_guardrail.apply_guardrail.assert_called_once_with(
inputs={"texts": ["Test text with PII"]},
request_data={},
input_type="request",
)
mock_guardrail.apply_guardrail.assert_called_once()
call = mock_guardrail.apply_guardrail.call_args.kwargs
assert call["inputs"] == {"texts": ["Test text with PII"]}
assert call["input_type"] == "request"
assert call["request_data"]["metadata"]["user_api_key_hash"] == user_api_key_dict.api_key
@pytest.mark.asyncio
@ -197,6 +197,8 @@ async def test_apply_guardrail_endpoint_without_optional_params(mock_proxy_loggi
assert response.response_text == "Processed text"
# Verify the guardrail was called with correct parameters
mock_guardrail.apply_guardrail.assert_called_once_with(
inputs={"texts": ["Test text"]}, request_data={}, input_type="request"
)
mock_guardrail.apply_guardrail.assert_called_once()
call = mock_guardrail.apply_guardrail.call_args.kwargs
assert call["inputs"] == {"texts": ["Test text"]}
assert call["input_type"] == "request"
assert call["request_data"]["metadata"]["user_api_key_hash"] == user_api_key_dict.api_key

View file

@ -8,8 +8,11 @@ Regression coverage for the guardrail route table bugs:
take call_types[0] to a handler-less type
"""
import pytest
from litellm.litellm_core_utils.api_route_to_call_types import (
get_call_types_for_route,
get_primary_call_type_for_route,
)
from litellm.types.utils import API_ROUTE_TO_CALL_TYPES, CallTypes
@ -110,3 +113,15 @@ class TestExistingRouteResolutionUnchanged:
def test_unknown_route_returns_none(self):
assert get_call_types_for_route("/not/a/real/route") is None
@pytest.mark.parametrize("route", ["/v1/chat/completions", "/v1/embeddings", "/v1/messages", "/v1/realtime"])
def test_primary_call_type_is_the_first_call_type_of_the_route(route):
call_types = get_call_types_for_route(route)
assert call_types is not None
assert get_primary_call_type_for_route(route) is call_types[0]
@pytest.mark.parametrize("route", [None, "", "/not/a/real/route"])
def test_primary_call_type_is_none_without_a_mapped_route(route):
assert get_primary_call_type_for_route(route) is None

View file

@ -17,6 +17,8 @@ from litellm.litellm_core_utils.realtime_streaming import (
client_sent_openai_beta_realtime_header,
)
from litellm.llms.xai.realtime.transformation import XAIRealtimeNormalizer
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.grayswan.grayswan import GraySwanGuardrail
from litellm.types.guardrails import GuardrailEventHooks
@ -3550,3 +3552,120 @@ async def test_provider_bytes_are_sent_raw_after_pacing():
assert [call.args[0] for call in backend_ws.send.await_args_list] == [b"\x00\x01", '{"type":"endStream"}']
provider_config.pace_backend_send.assert_awaited_once_with(b"\x00\x01")
@pytest.mark.asyncio
async def test_realtime_transcript_guardrail_receives_authenticated_identity(monkeypatch: pytest.MonkeyPatch):
"""Transcript guardrails get the session key's identity in litellm_metadata, as the chat path provides it."""
received_request_data = []
class IdentityRecordingGuardrail(CustomGuardrail):
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
received_request_data.append(request_data)
return inputs
monkeypatch.setattr(
litellm,
"callbacks",
[
IdentityRecordingGuardrail(
guardrail_name="identity_recorder",
event_hook=GuardrailEventHooks.realtime_input_transcription,
default_on=True,
)
],
)
key = UserAPIKeyAuth(api_key="sk-real-caller-key", key_alias="prod-app", team_id="team-prod")
streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock(), user_api_key_dict=key)
blocked = await streaming.run_realtime_guardrails("hello there")
assert blocked is False
assert len(received_request_data) == 1
identity = received_request_data[0]["litellm_metadata"]
assert identity["user_api_key_alias"] == "prod-app"
assert identity["user_api_key_team_id"] == "team-prod"
assert identity["user_api_key_hash"] == key.api_key
@pytest.mark.asyncio
async def test_realtime_grayswan_payload_carries_only_identity(monkeypatch: pytest.MonkeyPatch):
"""Gray Swan forwards litellm_metadata verbatim, so realtime must hand it identity and no key secrets."""
vendor_payloads: list[dict[str, object]] = []
class RecordingGraySwan(GraySwanGuardrail):
async def _call_grayswan_api(self, payload):
vendor_payloads.append(payload)
return {"violation": 0.0, "violated_rules": []}
monkeypatch.setattr(
litellm,
"callbacks",
[
RecordingGraySwan(
guardrail_name="grayswan",
api_key="test-key",
event_hook=GuardrailEventHooks.pre_call,
default_on=True,
)
],
)
key = UserAPIKeyAuth(
api_key="sk-real-caller-key",
key_alias="prod-app",
team_id="team-prod",
organization_metadata={"logging": [{"callback_vars": {"langfuse_secret_key": "SECRET-ORG"}}]},
jwt_claims={"email": "alice@corp.example", "name": "Alice Smith"},
team_member={"user_id": "alice", "user_email": "alice@corp.example", "role": "admin"},
)
streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock(), user_api_key_dict=key)
await streaming.run_realtime_guardrails("hello there", event_hooks=[GuardrailEventHooks.pre_call])
assert len(vendor_payloads) == 1
vendor_metadata = vendor_payloads[0]["litellm_metadata"]
assert vendor_metadata["user_api_key_alias"] == "prod-app"
assert vendor_metadata["user_api_key_team_id"] == "team-prod"
serialized = json.dumps(vendor_metadata)
for leaked in ("SECRET-ORG", "Alice Smith", "alice@corp.example"):
assert leaked not in serialized
@pytest.mark.asyncio
@pytest.mark.parametrize(
"sdk_value",
[
pytest.param({"key_alias": "forged", "team_id": "team-exempt"}, id="dict"),
pytest.param({"spend": "not-a-number"}, id="malformed-dict"),
pytest.param("sk-raw-string", id="string"),
],
)
async def test_realtime_guardrail_gets_no_identity_from_non_auth_sdk_value(
monkeypatch: pytest.MonkeyPatch, sdk_value: object
):
"""Only a proxy-authenticated UserAPIKeyAuth yields identity; an SDK-supplied value never raises or fakes one."""
received_request_data: list[dict[str, object]] = []
class IdentityRecordingGuardrail(CustomGuardrail):
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
received_request_data.append(request_data)
return inputs
monkeypatch.setattr(
litellm,
"callbacks",
[
IdentityRecordingGuardrail(
guardrail_name="identity_recorder",
event_hook=GuardrailEventHooks.realtime_input_transcription,
default_on=True,
)
],
)
streaming = RealTimeStreaming(MagicMock(), MagicMock(), MagicMock(), user_api_key_dict=sdk_value)
blocked = await streaming.run_realtime_guardrails("hello there")
assert blocked is False
assert len(received_request_data) == 1
assert not [key for key in received_request_data[0]["litellm_metadata"] if key.startswith("user_api_key")]

View file

@ -0,0 +1,53 @@
import json
from typing import Final
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.proxy._types import UserAPIKeyAuth
RAW_SESSION_TOKEN: Final = "cli-session-Qm7xJ2kP9sLw4vT1nR8yAa"
def _fully_populated_session_key() -> UserAPIKeyAuth:
return UserAPIKeyAuth(
token=RAW_SESSION_TOKEN,
is_session_token=True,
key_alias="cli-session-alice",
user_id="alice",
team_id="team-prod",
org_id="org-1",
metadata={"logging": [{"callback_name": "langfuse", "callback_vars": {"langfuse_secret_key": "SECRET-KEY"}}]},
team_metadata={"logging": [{"callback_vars": {"langfuse_secret_key": "SECRET-TEAM"}}]},
organization_metadata={"logging": [{"callback_vars": {"langfuse_secret_key": "SECRET-ORG"}}]},
project_metadata={"logging": [{"callback_vars": {"langfuse_secret_key": "SECRET-PROJECT"}}]},
jwt_claims={"sub": "alice", "email": "alice@corp.example", "name": "Alice Smith"},
team_member={"user_id": "alice", "user_email": "alice@corp.example", "role": "admin"},
config={"internal": "proxy-config"},
)
def test_transform_emits_only_identity_never_credentials_or_callback_secrets():
metadata = BaseTranslation.transform_user_api_key_dict_to_metadata(_fully_populated_session_key())
serialized = json.dumps(metadata, default=str)
assert metadata["user_api_key_alias"] == "cli-session-alice"
assert metadata["user_api_key_key_alias"] == "cli-session-alice"
assert metadata["user_api_key_hash"] == "cli-session-alice"
assert metadata["user_api_key_team_id"] == "team-prod"
assert RAW_SESSION_TOKEN not in serialized
assert "callback_vars" not in serialized
assert "SECRET" not in serialized
assert "Alice Smith" not in serialized
assert "proxy-config" not in serialized
for dropped in (
"user_api_key_token",
"user_api_key_jwt_claims",
"user_api_key_team_member",
"user_api_key_organization_metadata",
"user_api_key_project_metadata",
"user_api_key_config",
):
assert dropped not in metadata
def test_transform_of_no_key_is_empty():
assert BaseTranslation.transform_user_api_key_dict_to_metadata(None) == {}

View file

@ -0,0 +1,528 @@
import io
import json
import logging
from collections.abc import Sequence
from dataclasses import dataclass
from datetime import datetime
from typing import Final, Literal
from unittest.mock import MagicMock
import httpx
import pytest
from starlette.requests import Request
import litellm
from litellm.caching.dual_cache import DualCache
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
GenericGuardrailAPI,
initialize_guardrail,
)
from litellm.proxy.openai_files_endpoints.batch_guardrails import BatchScanResult, scan_batch_input_file
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import pass_through_request
from litellm.proxy.utils import ProxyLogging
from litellm.types.guardrails import GuardrailEventHooks, LitellmParams
from litellm.types.utils import GenericGuardrailAPIInputs
_INPUTS: Final = GenericGuardrailAPIInputs(texts=["hello"])
_CHAT_ROUTE: Final = {"metadata": {"user_api_key_request_route": "/v1/chat/completions"}}
@dataclass(frozen=True, slots=True)
class _Endpoint:
received: list[dict[str, object]]
handler: AsyncHTTPHandler
def _endpoint(action: str = "NONE") -> _Endpoint:
received: Final[list[dict[str, object]]] = [] # mutable-ok: records what the fake endpoint was sent
def _respond(request: httpx.Request) -> httpx.Response:
received.append(json.loads(request.content))
return httpx.Response(200, json={"action": action, "blocked_reason": "blocked by test endpoint"})
return _Endpoint(received=received, handler=AsyncHTTPHandler(transport=httpx.MockTransport(_respond)))
def _guardrail(
endpoint: _Endpoint,
*,
run_only_on_call_types: Sequence[str] | None = None,
skip_call_types: Sequence[str] | None = None,
guardrail_name: str = "call-type-filter",
) -> GenericGuardrailAPI:
return GenericGuardrailAPI(
api_base="https://guardrail.test",
guardrail_name=guardrail_name,
event_hook="pre_call",
default_on=True,
run_only_on_call_types=run_only_on_call_types,
skip_call_types=skip_call_types,
async_handler=endpoint.handler,
)
def _logging_obj(call_type: str) -> Logging:
return Logging(
model="gpt-test",
messages=[{"role": "user", "content": "hello"}],
stream=False,
call_type=call_type,
start_time=datetime(2026, 1, 1),
litellm_call_id="call-1",
function_id="fn-1",
)
async def _apply(
guardrail: GenericGuardrailAPI,
*,
call_type: str | None,
input_type: Literal["request", "response"] = "request",
request_data: dict[str, object] | None = None,
) -> GenericGuardrailAPIInputs:
return await guardrail.apply_guardrail(
inputs=_INPUTS,
request_data={} if request_data is None else request_data,
input_type=input_type,
logging_obj=None if call_type is None else _logging_obj(call_type),
)
def _recorded(request_data: dict[str, object]) -> list[tuple[object, object]]:
_, bucket = get_or_create_metadata_bucket(request_data)
return [
(entry["guardrail_status"], entry["guardrail_response"])
for entry in bucket.get("standard_logging_guardrail_information", [])
]
@pytest.mark.parametrize("input_type", ["request", "response"])
async def test_allowlisted_call_type_reaches_the_endpoint(input_type: Literal["request", "response"]):
endpoint: Final = _endpoint()
guardrail: Final = _guardrail(endpoint, run_only_on_call_types=["acompletion", "anthropic_messages"])
result: Final = await _apply(guardrail, call_type="anthropic_messages", input_type=input_type)
assert [(body["texts"], body["input_type"]) for body in endpoint.received] == [(["hello"], input_type)]
assert result["texts"] == ["hello"]
@pytest.mark.parametrize("input_type", ["request", "response"])
async def test_call_type_outside_the_allowlist_is_passed_through_unsent(input_type: Literal["request", "response"]):
endpoint: Final = _endpoint()
guardrail: Final = _guardrail(endpoint, run_only_on_call_types=["acompletion"])
result: Final = await _apply(guardrail, call_type="aembedding", input_type=input_type)
assert endpoint.received == []
assert result == _INPUTS
@pytest.mark.parametrize("input_type", ["request", "response"])
async def test_denylisted_call_type_is_passed_through_unsent(input_type: Literal["request", "response"]):
endpoint: Final = _endpoint()
guardrail: Final = _guardrail(endpoint, skip_call_types=["aembedding", "aspeech"])
result: Final = await _apply(guardrail, call_type="aembedding", input_type=input_type)
assert endpoint.received == []
assert result == _INPUTS
async def test_call_type_outside_the_denylist_reaches_the_endpoint():
endpoint: Final = _endpoint()
guardrail: Final = _guardrail(endpoint, skip_call_types=["aembedding"])
await _apply(guardrail, call_type="acompletion")
assert len(endpoint.received) == 1
async def test_allowlist_takes_precedence_over_denylist():
endpoint: Final = _endpoint()
guardrail: Final = _guardrail(endpoint, run_only_on_call_types=["aembedding"], skip_call_types=["aembedding"])
await _apply(guardrail, call_type="aembedding")
assert len(endpoint.received) == 1
async def test_empty_allowlist_is_treated_as_unset():
endpoint: Final = _endpoint()
guardrail: Final = _guardrail(endpoint, run_only_on_call_types=[], skip_call_types=["aembedding"])
await _apply(guardrail, call_type="acompletion")
await _apply(guardrail, call_type="aembedding")
assert len(endpoint.received) == 1
@pytest.mark.parametrize(
"request_data",
[
{},
{"call_type": "aembedding"},
{"metadata": {"user_api_key_request_route": ""}},
{"metadata": {"user_api_key_request_route": 7}},
{"user_api_key_dict": {"request_route": "/v1/embeddings"}},
],
)
async def test_unresolvable_call_type_still_runs_the_guardrail(request_data: dict[str, object]):
endpoint: Final = _endpoint()
guardrail: Final = _guardrail(endpoint, run_only_on_call_types=["acompletion"])
await _apply(guardrail, call_type=None, request_data=request_data)
assert len(endpoint.received) == 1
@pytest.mark.parametrize("input_type", ["request", "response"])
async def test_chat_route_stays_in_the_allowlist_after_the_bridge_flips_the_logging_call_type(
input_type: Literal["request", "response"],
):
endpoint: Final = _endpoint()
guardrail: Final = _guardrail(endpoint, run_only_on_call_types=["acompletion"])
await _apply(guardrail, call_type="responses", input_type=input_type, request_data=_CHAT_ROUTE)
assert [body["input_type"] for body in endpoint.received] == [input_type]
async def test_denying_responses_does_not_skip_only_the_output_side_of_a_chat_request():
endpoint: Final = _endpoint()
guardrail: Final = _guardrail(endpoint, skip_call_types=["responses"])
await _apply(guardrail, call_type="acompletion", input_type="request", request_data=_CHAT_ROUTE)
await _apply(guardrail, call_type="responses", input_type="response", request_data=_CHAT_ROUTE)
assert [body["input_type"] for body in endpoint.received] == ["request", "response"]
@pytest.mark.parametrize("metadata_field", ["metadata", "litellm_metadata"])
async def test_route_call_type_beats_the_logging_obj(metadata_field: str):
endpoint: Final = _endpoint()
guardrail: Final = _guardrail(endpoint, run_only_on_call_types=["acompletion"])
await _apply(
guardrail,
call_type="acompletion",
request_data={metadata_field: {"user_api_key_request_route": "/v1/embeddings"}},
)
assert endpoint.received == []
async def test_litellm_metadata_route_beats_metadata_route():
endpoint: Final = _endpoint()
guardrail: Final = _guardrail(endpoint, run_only_on_call_types=["acompletion"])
await _apply(
guardrail,
call_type=None,
request_data={
"metadata": {"user_api_key_request_route": "/v1/chat/completions"},
"litellm_metadata": {"user_api_key_request_route": "/v1/embeddings"},
},
)
assert endpoint.received == []
@pytest.mark.parametrize(
("run_only_on_call_types", "skip_call_types", "expected_calls"),
[(None, ["_arealtime"], 0), (["acompletion"], None, 0), (["_arealtime"], None, 1), (None, None, 1)],
)
async def test_realtime_transcripts_resolve_as_realtime(
run_only_on_call_types: list[str] | None, skip_call_types: list[str] | None, expected_calls: int
):
endpoint: Final = _endpoint()
guardrail: Final = _guardrail(
endpoint,
run_only_on_call_types=run_only_on_call_types,
skip_call_types=skip_call_types,
)
litellm.logging_callback_manager.add_litellm_callback(guardrail)
streaming: Final = RealTimeStreaming(
websocket=MagicMock(),
backend_ws=MagicMock(),
logging_obj=_logging_obj("_arealtime"),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed", request_route="/v1/realtime"),
)
blocked: Final = await streaming.run_realtime_guardrails("hello there", event_hooks=[GuardrailEventHooks.pre_call])
assert blocked is False
assert len(endpoint.received) == expected_calls
_BATCH_RECORDS: Final = (
{
"custom_id": "chat",
"method": "POST",
"url": "/v1/chat/completions",
"body": {"model": "m", "messages": [{"role": "user", "content": "chat text"}]},
},
{"custom_id": "embed", "method": "POST", "url": "/v1/embeddings", "body": {"model": "m", "input": "embed text"}},
)
def _batch_file(records: Sequence[dict[str, object]]) -> io.BytesIO:
return io.BytesIO("\n".join(json.dumps(record) for record in records).encode())
async def _scan_batch(records: Sequence[dict[str, object]], upload_route: str = "/v1/files") -> None:
result: Final = await scan_batch_input_file(
file_source=_batch_file(records),
request_metadata={},
user_api_key_dict=UserAPIKeyAuth(api_key="hashed", request_route=upload_route),
proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()),
)
assert isinstance(result, BatchScanResult)
@pytest.mark.parametrize(
("guardrail_options", "expected_texts"),
[
({"run_only_on_call_types": ["acompletion"]}, [["chat text"]]),
({"skip_call_types": ["aembedding"]}, [["chat text"]]),
({"run_only_on_call_types": ["aembedding"]}, [["embed text"]]),
({}, [["chat text"], ["embed text"]]),
],
)
@pytest.mark.parametrize("upload_route", ["/v1/files", "/files"])
async def test_batch_file_records_are_filtered_by_their_own_call_type(
guardrail_options: dict[str, list[str]], expected_texts: list[list[str]], upload_route: str
):
endpoint: Final = _endpoint()
litellm.logging_callback_manager.add_litellm_callback(_guardrail(endpoint, **guardrail_options))
await _scan_batch(_BATCH_RECORDS, upload_route)
assert sorted(body["texts"] for body in endpoint.received) == expected_texts
@pytest.mark.parametrize(
"guardrail_options",
[{"run_only_on_call_types": ["acompletion"]}, {"skip_call_types": ["anthropic_messages"]}],
)
@pytest.mark.parametrize("url", ["/v1/messages", "/anthropic/v1/messages"])
async def test_chat_record_relabeled_as_messages_is_still_scanned(guardrail_options: dict[str, list[str]], url: str):
endpoint: Final = _endpoint()
litellm.logging_callback_manager.add_litellm_callback(_guardrail(endpoint, **guardrail_options))
await _scan_batch(
[
{
"custom_id": "relabeled",
"method": "POST",
"url": url,
"body": {"model": "m", "messages": [{"role": "user", "content": "relabeled chat"}]},
}
]
)
assert [body["texts"] for body in endpoint.received] == [["relabeled chat"]], (
"Bedrock and Vertex run this record as chat, so a filter must not skip it on its url"
)
@pytest.mark.parametrize("caller_logging_obj", [{}, "x", {"call_type": "aembedding"}])
async def test_batch_record_carrying_its_own_logging_obj_is_scanned_under_a_filter(caller_logging_obj: object):
endpoint: Final = _endpoint()
litellm.logging_callback_manager.add_litellm_callback(_guardrail(endpoint, run_only_on_call_types=["acompletion"]))
await _scan_batch(
[
{
"custom_id": "carrier",
"method": "POST",
"url": "/v1/messages",
"body": {
"model": "m",
"messages": [{"role": "user", "content": "carried chat"}],
"litellm_logging_obj": caller_logging_obj,
},
}
]
)
assert [body["texts"] for body in endpoint.received] == [["carried chat"]]
@pytest.mark.parametrize(
"request_data",
[{}, {"metadata": {"user_api_key_request_route": "/not/a/mapped/route"}}],
)
async def test_logging_obj_call_type_is_used_when_the_route_does_not_resolve(request_data: dict[str, object]):
endpoint: Final = _endpoint()
guardrail: Final = _guardrail(endpoint, run_only_on_call_types=["acompletion"])
await _apply(guardrail, call_type="aembedding", request_data=request_data)
assert endpoint.received == []
async def test_empty_logging_call_type_runs_the_guardrail():
endpoint: Final = _endpoint()
guardrail: Final = _guardrail(endpoint, run_only_on_call_types=["acompletion"])
await _apply(guardrail, call_type="")
assert len(endpoint.received) == 1
@pytest.mark.parametrize(
("guardrail_options", "call_type", "reason"),
[
({"run_only_on_call_types": ["acompletion"]}, "aembedding", "not in run_only_on_call_types"),
({"skip_call_types": ["aembedding"]}, "aembedding", "in skip_call_types"),
],
)
async def test_skipped_call_records_one_not_run_entry(
guardrail_options: dict[str, list[str]], call_type: str, reason: str
):
endpoint: Final = _endpoint()
guardrail: Final = _guardrail(endpoint, **guardrail_options)
request_data: Final[dict[str, object]] = {"model": "gpt-test"}
await _apply(guardrail, call_type=call_type, request_data=request_data)
assert _recorded(request_data) == [("not_run", f"skipped: call type {call_type} {reason}")]
async def test_scanned_call_still_records_success():
endpoint: Final = _endpoint()
guardrail: Final = _guardrail(endpoint, skip_call_types=["aembedding"])
request_data: Final[dict[str, object]] = {"model": "gpt-test"}
await _apply(guardrail, call_type="acompletion", request_data=request_data)
assert [status for status, _ in _recorded(request_data)] == ["success"]
@pytest.mark.parametrize("option", ["run_only_on_call_types", "skip_call_types"])
def test_single_string_config_is_rejected(option: str):
with pytest.raises(ValueError, match=f"{option} must be a list of strings"):
_guardrail(_endpoint(), **{option: "aembedding"})
@pytest.mark.parametrize("option", ["run_only_on_call_types", "skip_call_types"])
def test_unknown_call_type_is_rejected(option: str):
with pytest.raises(ValueError, match=f"{option} contains unknown call type"):
_guardrail(_endpoint(), **{option: ["acompletion", "chat_completion"]})
def test_call_types_member_name_is_rejected_with_its_value():
with pytest.raises(ValueError, match="'pass_through' -> 'pass_through_endpoint'"):
_guardrail(_endpoint(), skip_call_types=["pass_through"])
@pytest.mark.parametrize("option", ["run_only_on_call_types", "skip_call_types"])
def test_mcp_tool_calls_are_rejected_because_they_never_resolve(option: str):
with pytest.raises(ValueError, match="call_mcp_tool"):
_guardrail(_endpoint(), **{option: ["call_mcp_tool"]})
async def test_call_types_value_that_differs_from_its_name_matches():
endpoint: Final = _endpoint()
guardrail: Final = _guardrail(endpoint, skip_call_types=["pass_through_endpoint", "_arealtime"])
await _apply(guardrail, call_type="pass_through_endpoint")
assert endpoint.received == []
def test_one_list_does_not_warn(caplog: pytest.LogCaptureFixture):
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
_guardrail(_endpoint(), run_only_on_call_types=["acompletion"], skip_call_types=None)
assert caplog.records == []
def test_setting_both_lists_warns_that_the_denylist_is_ignored(caplog: pytest.LogCaptureFixture):
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
_guardrail(_endpoint(), run_only_on_call_types=["acompletion"], skip_call_types=["aembedding"])
assert any("skip_call_types" in record.getMessage() for record in caplog.records)
@pytest.mark.parametrize("call_type", ["acompletion", "aembedding", "anthropic_messages", None])
async def test_no_filter_config_keeps_every_call_type_scanned(call_type: str | None):
endpoint: Final = _endpoint()
guardrail: Final = _guardrail(endpoint)
await _apply(guardrail, call_type=call_type, input_type="response")
assert len(endpoint.received) == 1
@pytest.mark.parametrize(
"config",
[
{"run_only_on_call_types": ["acompletion"]},
{"optional_params": {"run_only_on_call_types": ["acompletion"]}},
{"skip_call_types": ["aembedding"]},
{"optional_params": {"skip_call_types": ["aembedding"]}},
],
)
def test_initialize_guardrail_forwards_call_type_filters(config: dict[str, object]):
litellm_params: Final = LitellmParams.model_validate(
{
"guardrail": "generic_guardrail_api",
"mode": "pre_call",
"api_base": "https://guardrail.test",
**config,
}
)
guardrail: Final = initialize_guardrail(litellm_params, {"guardrail_name": "from-config"})
assert guardrail.call_type_filter.allows("acompletion")
assert not guardrail.call_type_filter.allows("aembedding")
def _pass_through_http_request(body: dict[str, object]) -> Request:
payload: Final = json.dumps(body).encode()
async def _receive() -> dict[str, object]:
return {"type": "http.request", "body": payload, "more_body": False}
return Request(
{
"type": "http",
"method": "POST",
"path": "/custom/pass-through",
"query_string": b"",
"headers": [(b"content-type", b"application/json")],
},
_receive,
)
@pytest.mark.parametrize(
"forged",
[
{},
{"litellm_metadata": {"user_api_key_request_route": "/v1/embeddings"}},
],
)
async def test_pass_through_caller_cannot_forge_a_skipped_route(forged: dict[str, object]):
endpoint: Final = _endpoint(action="BLOCKED")
guardrail: Final = _guardrail(endpoint, skip_call_types=["aembedding"], guardrail_name="pass-through-guard")
litellm.logging_callback_manager.add_litellm_callback(guardrail)
with pytest.raises(ProxyException, match="blocked by test endpoint"):
await pass_through_request(
request=_pass_through_http_request({"input": "hello", **forged}),
target="http://127.0.0.1:9/custom",
custom_headers={},
user_api_key_dict=UserAPIKeyAuth(api_key="hashed", request_route="/custom/pass-through"),
guardrails_config={"pass-through-guard": {}},
)
assert len(endpoint.received) == 1