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