mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge c501874bef into 6bc17f98d7
This commit is contained in:
commit
15d9e6fd13
28 changed files with 1847 additions and 191 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")]
|
||||
|
|
|
|||
|
|
@ -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) == {}
|
||||
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue