From c501874bef052116ce8546b94a7307782fd21a5a Mon Sep 17 00:00:00 2001 From: Caduri Katzav Date: Mon, 28 Sep 2026 18:14:45 +0300 Subject: [PATCH] feat(guardrails): add call-type filters to generic_guardrail_api Adds run_only_on_call_types (allowlist) and skip_call_types (denylist) so an operator can keep calls such as embeddings away from the guardrail endpoint. A filtered call is passed through untouched on both the request and response hooks, nothing is sent, and a not_run guardrail entry records why. The allowlist wins when both are set. Values are CallTypes values; unknown values, a bare string instead of a list, and call_mcp_tool (MCP tool calls reach the guardrail without a call type) are rejected at startup. A call whose type cannot be resolved still runs the guardrail. Both options default to off, so behavior is unchanged unless configured The call type comes from the authenticated user_api_key_request_route, then the logging object. The route is stable across both hooks, while the logging object's call type is rewritten to "responses" when a chat or messages call is bridged to the Responses API Batch-file scans now run each record under its own endpoint route instead of the upload route, so a filter, and any guardrail that reads the key's route, sees the record as the chat or embeddings call it is. A record only gets that route when its body shape agrees with its url and the Bedrock and Vertex record classifiers would run it as the same call type. Any other record, such as a chat body labeled /v1/messages, gets no route and is always scanned The batch scan also strips a litellm_logging_obj key from each record body before the guardrails run. On main a truthy caller value there made the generic guardrail fail, and fail_on_error: false then shipped the record unscanned Adds get_primary_call_type_for_route and uses it in the unified guardrail so both resolve a route to a call type the same way. GenericGuardrailAPI also accepts an injected async_handler, used by the tests --- .../api_route_to_call_types.py | 9 +- .../generic_guardrail_api/__init__.py | 2 + .../generic_guardrail_api/call_type_filter.py | 114 ++++ .../generic_guardrail_api/config_parsing.py | 15 + .../generic_guardrail_api.py | 30 +- .../unified_guardrail/unified_guardrail.py | 41 +- .../batch_guardrails.py | 80 ++- .../guardrail_hooks/generic_guardrail_api.py | 25 + .../test_batch_guardrails.py | 82 +++ .../test_api_route_to_call_types.py | 15 + tests/unit/proxy/guardrails/__init__.py | 0 .../guardrails/guardrail_hooks/__init__.py | 0 .../generic_guardrail_api/__init__.py | 0 .../test_call_type_filter.py | 528 ++++++++++++++++++ 14 files changed, 901 insertions(+), 40 deletions(-) create mode 100644 litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/call_type_filter.py create mode 100644 litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/config_parsing.py create mode 100644 tests/unit/proxy/guardrails/__init__.py create mode 100644 tests/unit/proxy/guardrails/guardrail_hooks/__init__.py create mode 100644 tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py create mode 100644 tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_call_type_filter.py diff --git a/litellm/litellm_core_utils/api_route_to_call_types.py b/litellm/litellm_core_utils/api_route_to_call_types.py index 428d7563d4a..8212f7cee20 100644 --- a/litellm/litellm_core_utils/api_route_to_call_types.py +++ b/litellm/litellm_core_utils/api_route_to_call_types.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py index e3511d46544..2ec96fd8791 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py @@ -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) diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/call_type_filter.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/call_type_filter.py new file mode 100644 index 00000000000..93341f533a5 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/call_type_filter.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/config_parsing.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/config_parsing.py new file mode 100644 index 00000000000..10934c8e244 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/config_parsing.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py index 320682c89b7..ac0bf28dcac 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py @@ -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())) @@ -431,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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index 0743af68e18..71cfd6aaecf 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -18,7 +18,7 @@ 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 @@ -76,12 +76,10 @@ def resolve_endpoint_translation( 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 @@ -313,11 +311,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) @@ -421,20 +415,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, @@ -1081,8 +1068,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]] = [] @@ -1099,12 +1086,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) diff --git a/litellm/proxy/openai_files_endpoints/batch_guardrails.py b/litellm/proxy/openai_files_endpoints/batch_guardrails.py index 1db4474fc40..625fb661df5 100644 --- a/litellm/proxy/openai_files_endpoints/batch_guardrails.py +++ b/litellm/proxy/openai_files_endpoints/batch_guardrails.py @@ -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, diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py index 44e2cc2404f..3194be8924e 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py @@ -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], diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_batch_guardrails.py b/tests/test_litellm/proxy/openai_files_endpoint/test_batch_guardrails.py index 8c8dc5d799f..04e559cb3d3 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_batch_guardrails.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_batch_guardrails.py @@ -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" diff --git a/tests/unit/litellm_core_utils/test_api_route_to_call_types.py b/tests/unit/litellm_core_utils/test_api_route_to_call_types.py index 42ca91bfd8f..c8f632040bf 100644 --- a/tests/unit/litellm_core_utils/test_api_route_to_call_types.py +++ b/tests/unit/litellm_core_utils/test_api_route_to_call_types.py @@ -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 diff --git a/tests/unit/proxy/guardrails/__init__.py b/tests/unit/proxy/guardrails/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/__init__.py b/tests/unit/proxy/guardrails/guardrail_hooks/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_call_type_filter.py b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_call_type_filter.py new file mode 100644 index 00000000000..7bf1ba66413 --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_call_type_filter.py @@ -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