diff --git a/litellm/constants.py b/litellm/constants.py index c4adc0de22f..b8048ad30b4 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1,10 +1,20 @@ import os import sys +from enum import Enum from types import MappingProxyType from typing import Final, Literal from litellm.litellm_core_utils.env_utils import get_env_int, get_env_int_in_range, get_env_int_or_none +SERVER_STREAMING_CLASSIFICATION_KEY: Final = "litellm_server_streaming_classification" + + +class ServerStreamingClassification(str, Enum): + MARKER = "litellm-server-streaming" + + +SERVER_STREAMING_CLASSIFICATION_MARKER: Final = ServerStreamingClassification.MARKER + DEFER_PYDANTIC_BUILD: Final = os.getenv("DEFER_PYDANTIC_BUILD", "true") in ("true", "1", "on") DEFAULT_HEALTH_CHECK_PROMPT: Final = str(os.getenv("DEFAULT_HEALTH_CHECK_PROMPT", "test from litellm")) AZURE_DEFAULT_RESPONSES_API_VERSION: Final = str(os.getenv("AZURE_DEFAULT_RESPONSES_API_VERSION", "preview")) diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 94f1207b9ae..383a09f0931 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -6,7 +6,7 @@ import secrets from collections.abc import Mapping, Sequence from datetime import datetime from types import MappingProxyType -from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, Optional, get_args +from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, Optional, cast, get_args import httpx @@ -21,11 +21,14 @@ from litellm.litellm_core_utils.core_helpers import ( ) from litellm.secret_managers.main import str_to_bool from litellm.types.guardrails import ( + DEFAULT_GUARDRAIL_STREAM_SCOPE, DynamicGuardrailParams, GuardrailEventHooks, + GuardrailStreamScope, LitellmParams, LoggingOnlyScope, Mode, + runtime_stream_scope, ) from litellm.types.llms.openai import AllMessageValues from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel @@ -49,6 +52,8 @@ from litellm.constants import ( GUARDRAIL_SCANNED_MESSAGES_CACHE_TTL_SECONDS, LOGS_GUARDRAIL_INFORMATION_MARKER, PRE_CALL_EXECUTED_GUARDRAILS_KEY, + SERVER_STREAMING_CLASSIFICATION_KEY, + SERVER_STREAMING_CLASSIFICATION_MARKER, ) from litellm.exceptions import ( BlockedPiiEntityError, @@ -173,6 +178,42 @@ def get_session_id_from_request_data(request_data: dict[str, Any]) -> str | None return None +_REALTIME_STREAMING_HOOKS: Final = frozenset({GuardrailEventHooks.realtime_input_transcription}) + + +def without_server_streaming_classification(data: Mapping[str, object]) -> dict[str, object]: + return { + key: value + for key, value in data.items() + if key != SERVER_STREAMING_CLASSIFICATION_KEY or value != SERVER_STREAMING_CLASSIFICATION_MARKER + } + + +def guardrail_request_data_with_streaming( + data: Mapping[str, object], + *, + is_streaming: bool, +) -> dict[str, object]: + data_without_server_classification: Final = without_server_streaming_classification(data) + if not is_streaming: + return data_without_server_classification + return { + **data_without_server_classification, + SERVER_STREAMING_CLASSIFICATION_KEY: SERVER_STREAMING_CLASSIFICATION_MARKER, + } + + +def _request_is_streaming(data: object, event_type: GuardrailEventHooks | None = None) -> bool: + if event_type in _REALTIME_STREAMING_HOOKS: + return True + if not isinstance(data, Mapping): + return False + return ( + data.get("stream") is True + or data.get(SERVER_STREAMING_CLASSIFICATION_KEY) is SERVER_STREAMING_CLASSIFICATION_MARKER + ) + + class CustomGuardrail(CustomLogger): # If True, during_call runs async_moderation_hook instead of the unified apply_guardrail path. use_native_during_call_hook: ClassVar[bool] = False @@ -183,6 +224,9 @@ class CustomGuardrail(CustomLogger): records_own_guardrail_information: ClassVar[bool] = False logging_only_scope: LoggingOnlyScope | None + stream_scope_default: GuardrailStreamScope = DEFAULT_GUARDRAIL_STREAM_SCOPE + stream_scope_by_hook: tuple[tuple[str, GuardrailStreamScope], ...] = () + timeout: float | httpx.Timeout | None = None def __init_subclass__(cls, **kwargs: object) -> None: # kwargs-ok: forwarded to cooperative __init_subclass__ hooks @@ -258,6 +302,8 @@ class CustomGuardrail(CustomLogger): self.run_in_parallel: bool = run_in_parallel self.scan_raw_request: bool = scan_raw_request self.only_scan_new_messages: bool = only_scan_new_messages + stream_scope_arg: Final[object] = cast(object, kwargs.pop("stream_scope", None)) # cast-ok: config + self.apply_stream_scope(stream_scope_arg) self.logging_only_scope = None if timeout is not None: self.timeout = timeout @@ -1099,6 +1145,23 @@ class CustomGuardrail(CustomLogger): return name in suppressed_compression_guardrails() + def apply_stream_scope(self, stream_scope: object) -> None: + default, by_hook = runtime_stream_scope(stream_scope) + self.stream_scope_default = default + self.stream_scope_by_hook = tuple(by_hook.items()) + + def stream_scope_allows(self, data: object, event_type: GuardrailEventHooks) -> bool: + scope: Final = next( + (scope for hook, scope in self.stream_scope_by_hook if hook == event_type.value), + self.stream_scope_default, + ) + if scope == "both": + return True + is_streaming: Final = _request_is_streaming(data, event_type) + if scope == "streaming": + return is_streaming + return not is_streaming + def should_run_guardrail( self, data, @@ -1142,8 +1205,10 @@ class CustomGuardrail(CustomLogger): data, self.event_hook, event_type ) if result is not None: - return result - return True + tagged_result: Final[bool] = bool(cast(object, result)) # cast-ok: helper return + data_obj: Final[object] = cast(object, data) # cast-ok: data param + return tagged_result and self.stream_scope_allows(data_obj, event_type) + return self.stream_scope_allows(cast(object, data), event_type) # cast-ok: data param return False if ( @@ -1167,8 +1232,9 @@ class CustomGuardrail(CustomLogger): ) result = EnterpriseCustomGuardrailHelper._should_run_if_mode_by_tag(data, self.event_hook, event_type) if result is not None: - return result - return True + mode_tag_result: Final[bool] = bool(cast(object, result)) # cast-ok: helper return + return mode_tag_result and self.stream_scope_allows(cast(object, data), event_type) # cast-ok: data + return self.stream_scope_allows(cast(object, data), event_type) # cast-ok: data param def _event_hook_is_event_type(self, event_type: GuardrailEventHooks) -> bool: """ diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index a781be610a6..473be107cbb 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -782,7 +782,7 @@ class RealTimeStreaming: isinstance(cb, CustomGuardrail) and any( cb.should_run_guardrail( - data=self.request_data, + data={**self.request_data, "stream": True}, event_type=et, ) for et in event_hooks @@ -847,7 +847,7 @@ class RealTimeStreaming: if event_hooks is None: event_hooks = [GuardrailEventHooks.realtime_input_transcription] _realtime_event_types: Final = event_hooks - _check_data: Final = {**self.request_data, "transcript": transcript} + _check_data: Final = {**self.request_data, "transcript": transcript, "stream": True} _already_run: Final[set] = set() for callback in litellm.callbacks: diff --git a/litellm/llms/bedrock/passthrough/transformation.py b/litellm/llms/bedrock/passthrough/transformation.py index c61576fa7b2..0d7f7cd3562 100644 --- a/litellm/llms/bedrock/passthrough/transformation.py +++ b/litellm/llms/bedrock/passthrough/transformation.py @@ -22,6 +22,13 @@ if TYPE_CHECKING: from litellm.types.utils import CostResponseTypes +BEDROCK_STREAMING_ACTIONS: Final = frozenset({"invoke-with-response-stream", "converse-stream"}) + + +def is_bedrock_streaming_endpoint(endpoint: str) -> bool: + return endpoint.partition("?")[0].rstrip("/").rsplit("/", 1)[-1] in BEDROCK_STREAMING_ACTIONS + + _TEXT_ONLY_DELTA_FIELDS: Final = frozenset({"content", "role"}) diff --git a/litellm/llms/pass_through/guardrail_translation/handler.py b/litellm/llms/pass_through/guardrail_translation/handler.py index 1f295a6e656..fd58ca2ddbb 100644 --- a/litellm/llms/pass_through/guardrail_translation/handler.py +++ b/litellm/llms/pass_through/guardrail_translation/handler.py @@ -10,6 +10,7 @@ from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final, Optional from litellm._logging import verbose_proxy_logger +from litellm.integrations.custom_guardrail import without_server_streaming_classification from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation from litellm.proxy._types import PassThroughGuardrailSettings from litellm.types.utils import GenericGuardrailAPIInputs @@ -80,7 +81,9 @@ class PassThroughEndpointHandler(BaseTranslation): from litellm.litellm_core_utils.safe_json_dumps import safe_dumps payload_to_check: Final = { - k: v for k, v in data.items() if not k.startswith("_") and k not in ("metadata", "litellm_logging_obj") + k: v + for k, v in without_server_streaming_classification(data).items() + if not k.startswith("_") and k not in ("metadata", "litellm_logging_obj") } verbose_proxy_logger.debug("PassThroughEndpointHandler: Using full payload for guardrail") return safe_dumps(payload_to_check) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 9322fb77615..7dcca8f9ec0 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -11906,6 +11906,34 @@ "description": "When True (default), after sensitive data is detected and routed, all subsequent requests in the same session will continue routing to the same model.", "title": "Sticky Session Routing" }, + "stream_scope": { + "anyOf": [ + { + "enum": [ + "streaming", + "non_streaming", + "both" + ], + "type": "string" + }, + { + "additionalProperties": { + "enum": [ + "streaming", + "non_streaming", + "both" + ], + "type": "string" + }, + "type": "object" + }, + { + "type": "null" + } + ], + "description": "Whether this guardrail runs on streaming requests, non-streaming requests, or both. A string applies to every configured mode. A map overrides named modes (pre_call, during_call, post_call, ...); omitted keys default to both. Unset means both, matching historical behavior.", + "title": "Stream Scope" + }, "template_id": { "anyOf": [ { @@ -14831,6 +14859,34 @@ "description": "When True (default), after sensitive data is detected and routed, all subsequent requests in the same session will continue routing to the same model.", "title": "Sticky Session Routing" }, + "stream_scope": { + "anyOf": [ + { + "enum": [ + "streaming", + "non_streaming", + "both" + ], + "type": "string" + }, + { + "additionalProperties": { + "enum": [ + "streaming", + "non_streaming", + "both" + ], + "type": "string" + }, + "type": "object" + }, + { + "type": "null" + } + ], + "description": "Whether this guardrail runs on streaming requests, non-streaming requests, or both. A string applies to every configured mode. A map overrides named modes (pre_call, during_call, post_call, ...); omitted keys default to both. Unset means both, matching historical behavior.", + "title": "Stream Scope" + }, "template_id": { "anyOf": [ { diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index 8e1209401ab..6072d8bf22a 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -69,6 +69,7 @@ from litellm.types.guardrails import ( PresidioPresidioConfigModelUserInterface, SupportedGuardrailIntegrations, ToolPermissionGuardrailConfigModel, + with_tolerated_stream_scope, ) from litellm.types.llms.base import LiteLLMBaseModel from litellm.types.proxy.guardrails.guardrail_hooks.hide_secrets import ( @@ -146,7 +147,7 @@ def _get_guardrails_list_response( GuardrailInfoResponse( guardrail_id=guardrail.get("guardrail_id"), guardrail_name=guardrail.get("guardrail_name"), - litellm_params=masked_params, + litellm_params=with_tolerated_stream_scope(masked_params), guardrail_info=guardrail.get("guardrail_info"), ) ) @@ -289,7 +290,7 @@ async def list_guardrails_v2( ) masked_litellm_params = ( parse_tolerant_litellm_params( - masked_litellm_params_dict, + with_tolerated_stream_scope(masked_litellm_params_dict), guardrail.get("guardrail_name") or "Unknown", params_model=BaseLitellmParams, ) @@ -336,7 +337,7 @@ async def list_guardrails_v2( ) masked_in_memory_litellm_params_typed = ( parse_tolerant_litellm_params( - masked_in_memory_litellm_params, + with_tolerated_stream_scope(masked_in_memory_litellm_params), guardrail.get("guardrail_name") or "Unknown", params_model=BaseLitellmParams, ) @@ -1263,7 +1264,7 @@ async def patch_guardrail( # Update litellm_params if default_on is provided or pii_entities_config is provided existing_litellm_params: Final = _as_str_object_mapping(dict(existing_guardrail.get("litellm_params", {}))) current_litellm_params: Final = parse_tolerant_litellm_params( - existing_litellm_params, + with_tolerated_stream_scope(existing_litellm_params), existing_guardrail.get("guardrail_name") or "Unknown", ) requested_litellm_params: Final[Mapping[str, object]] = ( @@ -1275,7 +1276,7 @@ async def patch_guardrail( MappingProxyType({**current_litellm_params.model_dump(exclude_unset=True), **requested_litellm_params}) ) try: - parsed_litellm_params: Final = LitellmParams(**merged_litellm_params) + parsed_litellm_params: Final = LitellmParams(**with_tolerated_stream_scope(merged_litellm_params)) except ValidationError as validation_error: raise HTTPException( status_code=422, @@ -1436,7 +1437,7 @@ async def get_guardrail_info(guardrail_id: str): ) masked_litellm_params = ( parse_tolerant_litellm_params( - masked_litellm_params_dict, + with_tolerated_stream_scope(masked_litellm_params_dict), result.get("guardrail_name") or "Unknown", params_model=BaseLitellmParams, ) diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index d14cc940885..439d17b8644 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -54,6 +54,7 @@ from litellm.types.guardrails import ( LakeraCategoryThresholds, LitellmParams, SupportedGuardrailIntegrations, + with_tolerated_stream_scope, ) from .guardrail_hooks.llm_as_a_judge import ( @@ -634,6 +635,8 @@ def _configure_callback_scoping( "skip_tool_message_in_guardrail are enabled together, which excludes every message from " "scanning, so no request content would ever be scanned. Remove one of the two." ) + if isinstance(custom_guardrail_callback, CustomGuardrail): # pyright: ignore[reportUnnecessaryIsInstance] # module-path classes may only subclass CustomLogger + custom_guardrail_callback.apply_stream_scope(litellm_params.stream_scope) _apply_configured_bool_overrides(custom_guardrail_callback, litellm_params) @@ -718,9 +721,11 @@ class InMemoryGuardrailHandler: if isinstance(litellm_params_data, dict): if reject_invalid_logging_only_scope: - litellm_params = LitellmParams(**litellm_params_data) + litellm_params = LitellmParams(**with_tolerated_stream_scope(litellm_params_data)) else: - litellm_params = parse_tolerant_litellm_params(litellm_params_data, guardrail["guardrail_name"]) + litellm_params = parse_tolerant_litellm_params( + with_tolerated_stream_scope(litellm_params_data), guardrail["guardrail_name"] + ) else: litellm_params = litellm_params_data @@ -863,14 +868,17 @@ class InMemoryGuardrailHandler: # Extract additional params from litellm_params to pass to custom guardrail # This matches the behavior of other guardrail initializers (e.g., initialize_lakera) # and aligns with the documented behavior for custom guardrails - if hasattr(litellm_params, "model_dump"): - extra_params = litellm_params.model_dump(exclude_none=True) - else: - extra_params = dict(litellm_params) if litellm_params else {} - - # Remove params that are handled explicitly or are internal - for key in ["guardrail", "mode", "default_on"]: - extra_params.pop(key, None) + excluded_extra_param_keys: Final = frozenset(("guardrail", "mode", "default_on", "stream_scope")) + extra_params_items: Final = ( + litellm_params.model_dump(exclude_none=True).items() + if hasattr(litellm_params, "model_dump") + else iter(litellm_params) + if litellm_params + else () + ) + extra_params: Final = MappingProxyType( + {key: value for key, value in extra_params_items if key not in excluded_extra_param_keys} + ) _guardrail_callback: Final = _guardrail_class( guardrail_name=guardrail["guardrail_name"], @@ -1007,7 +1015,7 @@ class InMemoryGuardrailHandler: return params.model_dump() if isinstance(params, dict): try: - return parse_tolerant_litellm_params(params, guardrail_name).model_dump() + return parse_tolerant_litellm_params(with_tolerated_stream_scope(params), guardrail_name).model_dump() except ValidationError as e: verbose_proxy_logger.warning( "Could not normalize guardrail litellm_params for comparison; treating the guardrail as changed. Error: %s", diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 46cdf328af6..66dca57e782 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -32,6 +32,7 @@ from litellm.constants import ( SESSION_ID_OMITTED_METADATA_KEY, X_LITELLM_DISABLE_CALLBACKS, ) +from litellm.integrations.custom_guardrail import without_server_streaming_classification from litellm.litellm_core_utils.core_helpers import is_codex_user_agent from litellm.litellm_core_utils.credential_accessor import CredentialAccessor from litellm.litellm_core_utils.initialize_dynamic_callback_params import ( @@ -2094,7 +2095,9 @@ def refresh_proxy_server_request_body_snapshot( | _TRANSPORT_ONLY_CREDENTIAL_KEYS | _CALLBACK_CREDENTIAL_KEYS ) - body: Final = {k: v for k, v in data.items() if k not in _body_snapshot_exclude} + body: Final = { + k: v for k, v in without_server_streaming_classification(data).items() if k not in _body_snapshot_exclude + } proxy_server_request["body"] = body if guardrails_applied and isinstance(logging_obj, Logging): metadata: Final = data.get(get_metadata_variable_name_from_kwargs(data)) diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index ca142088383..1eefb58d7fa 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -44,12 +44,14 @@ from litellm.constants import ( AZURE_SPEECH_SUBSCRIPTION_KEY_HEADER, BEDROCK_AGENT_RUNTIME_PASS_THROUGH_ROUTES, ) +from litellm.integrations.custom_guardrail import guardrail_request_data_with_streaming from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix from litellm.llms.anthropic.common_utils import AnthropicModelInfo, merge_anthropic_beta_headers from litellm.llms.azure.passthrough.transformation import ( foreign_azure_deployment, is_azure_body_model_inference_endpoint, ) +from litellm.llms.bedrock.passthrough.transformation import is_bedrock_streaming_endpoint from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.llms.deepgram.common_utils import ( deepgram_listen_callback_params, @@ -946,8 +948,6 @@ BEDROCK_ENDPOINT_ACTIONS: Final = { "count-tokens", } -BEDROCK_STREAMING_ACTIONS: Final = {"invoke-with-response-stream", "converse-stream"} - def is_bedrock_count_tokens_endpoint(endpoint: str) -> bool: return "count_tokens" in endpoint or "count-tokens" in endpoint @@ -1077,7 +1077,7 @@ async def handle_bedrock_passthrough_router_model( from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing # Detect streaming based on endpoint - is_streaming: Final = any(action in endpoint for action in BEDROCK_STREAMING_ACTIONS) + is_streaming: Final = is_bedrock_streaming_endpoint(endpoint) verbose_proxy_logger.debug( "Bedrock router passthrough: model='%s', endpoint='%s', streaming=%s", model, endpoint, is_streaming @@ -1085,15 +1085,19 @@ async def handle_bedrock_passthrough_router_model( # Use the common processing path (same as non-router models) # This ensures all metadata, hooks, and logging are properly initialized - data: Final[dict[str, object]] = {} + bedrock_payload: Final[dict[str, object]] = { + "model": model, + "method": request.method, + "endpoint": endpoint, + "data": request_body, + "custom_llm_provider": "bedrock", + } + data: Final[dict[str, object]] = guardrail_request_data_with_streaming( + MappingProxyType(bedrock_payload), + is_streaming=is_streaming, + ) base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data) - data["model"] = model - data["method"] = request.method - data["endpoint"] = endpoint - data["data"] = request_body - data["custom_llm_provider"] = "bedrock" - # Use the common passthrough processing to handle metadata and hooks # This also handles all response formatting (streaming/non-streaming) and exceptions try: @@ -1285,14 +1289,19 @@ async def bedrock_llm_proxy_route( "Bedrock passthrough: Using direct Bedrock model '%s' for endpoint '%s'", model, endpoint ) - data: Final[dict[str, object]] = {} + is_streaming: Final = is_bedrock_streaming_endpoint(endpoint) + passthrough_payload: Final[dict[str, object]] = { + "method": request.method, + "endpoint": endpoint, + "data": request_body, + "custom_llm_provider": "bedrock", + } + data: Final[dict[str, object]] = guardrail_request_data_with_streaming( + MappingProxyType(passthrough_payload), + is_streaming=is_streaming, + ) base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data) - data["method"] = request.method - data["endpoint"] = endpoint - data["data"] = request_body - data["custom_llm_provider"] = "bedrock" - try: result: Final = await base_llm_response_processor.base_passthrough_process_llm_request( request=request, diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 07053ad0f32..ee958679c1f 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -48,7 +48,11 @@ from litellm.constants import ( SESSION_ID_OMITTED_METADATA_KEY, WEBSOCKET_CLOSE_REASON_MAX_BYTES, ) -from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + guardrail_request_data_with_streaming, + without_server_streaming_classification, +) from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.core_helpers import ( bind_budget_reservation_to_callbacks, @@ -603,6 +607,12 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): from litellm.proxy.proxy_server import llm_router _parsed_body = _parsed_body or {} + parsed_body_typed: Final[Mapping[str, object]] = cast(Mapping[str, object], _parsed_body) # cast-ok: json + server_marker_free_body: Final = without_server_streaming_classification(parsed_body_typed) + # The marker-free body must propagate through the caller's request dict, so + # downstream guardrail scans and snapshots never observe the server streaming marker. + _parsed_body.clear() + _parsed_body.update(server_marker_free_body) managed_model: Final = get_model_from_request( request_data=_parsed_body, route=get_request_route(request), @@ -828,7 +838,7 @@ def _build_passthrough_failure_request_payload( error response. Spend tracking only attributes a recovered cost when it comes paired with a usage object, so both keys are written together. """ - request_payload: Final[dict] = dict(parsed_body or {}) + request_payload: Final[dict] = dict(cast(Mapping[str, object], parsed_body or {})) # cast-ok: json body if kwargs: request_payload.update(kwargs) if logging_obj is not None: @@ -1163,7 +1173,7 @@ async def pass_through_request( is_multipart: Final = HttpPassThroughEndpointHelpers.is_multipart(request) and not custom_body if custom_body: - _parsed_body = custom_body + _parsed_body = dict(custom_body) elif is_multipart: # Don't parse multipart body here - it will be handled by make_multipart_http_request _parsed_body = {} @@ -1232,6 +1242,14 @@ async def pass_through_request( if _parsed_body is None: _parsed_body = {} _parsed_body["litellm_logging_obj"] = logging_obj + is_streaming_pass_through: Final = bool( + HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body( + parsed_body=_parsed_body, + stream=stream, + ) + ) + typed_body: Final[Mapping[str, object]] = cast(Mapping[str, object], _parsed_body) # cast-ok: json + _parsed_body = guardrail_request_data_with_streaming(typed_body, is_streaming=is_streaming_pass_through) ### CALL HOOKS ### - modify incoming data / reject request before calling the model _parsed_body = await proxy_logging_obj.pre_call_hook( @@ -2518,10 +2536,10 @@ async def websocket_passthrough_request( ) ### CALL HOOKS ### - modify incoming data / reject request before calling the model - websocket_data: dict[str, object] = {} - websocket_data = await proxy_logging_obj.pre_call_hook( + websocket_hook_data: Final = guardrail_request_data_with_streaming(MappingProxyType({}), is_streaming=True) + await proxy_logging_obj.pre_call_hook( user_api_key_dict=user_api_key_dict, - data=websocket_data, + data=websocket_hook_data, call_type="pass_through_endpoint", ) diff --git a/litellm/proxy/policy_engine/pipeline_executor.py b/litellm/proxy/policy_engine/pipeline_executor.py index d80eb56d8e3..5ccf2920030 100644 --- a/litellm/proxy/policy_engine/pipeline_executor.py +++ b/litellm/proxy/policy_engine/pipeline_executor.py @@ -8,7 +8,8 @@ pass/fail actions (allow, block, next, modify_response) and data forwarding. import copy import time from collections.abc import Callable, Mapping, Sequence -from typing import TYPE_CHECKING, Final, Literal, TypeVar +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, Literal, TypeVar, cast from pydantic import BaseModel @@ -29,6 +30,7 @@ from litellm.proxy.common_utils.callback_utils import add_guardrail_to_applied_g from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( UnifiedLLMGuardrails, ) +from litellm.types.guardrails import GuardrailEventHooks from litellm.types.proxy.policy_engine.pipeline_types import ( PipelineExecutionResult, PipelineStep, @@ -265,6 +267,30 @@ class _LegacyHookStreamAdapter(CustomGuardrail): return recorder.inputs +_PIPELINE_EVENT_HOOKS: Final = MappingProxyType( + { + "pre_call": GuardrailEventHooks.pre_call, + "post_call": GuardrailEventHooks.post_call, + "during_call": GuardrailEventHooks.during_call, + } +) + + +def _pipeline_stream_scope_allows( + callback: CustomGuardrail, + hook_input: Mapping[str, object], + mode: str, + streaming_chunks: list[object] | None, +) -> bool: + event_type: Final = _PIPELINE_EVENT_HOOKS.get(mode) + if event_type is None: + return True + return callback.stream_scope_allows( + hook_input if streaming_chunks is None else {**hook_input, "stream": True}, + event_type, + ) + + def _prepare_hook_input( step: PipelineStep, callback: CustomGuardrail, @@ -494,7 +520,7 @@ class PipelineExecutor: streaming_chunks: list[object] | None = None, # mutable-ok: shared buffered-stream chunks, read per step endpoint_translation: "BaseTranslation | None" = None, ) -> tuple[ - Literal["pass", "fail", "error"], + Literal["pass", "fail", "error", "skip"], dict | None, str | None, Exception | None, @@ -504,7 +530,7 @@ class PipelineExecutor: Returns: Tuple of (outcome, modified_data, error_detail, original_exception): - - outcome: "pass", "fail", or "error" + - outcome: "pass", "fail", "error", or "skip" - modified_data: dict if guardrail returned modified data, else None - error_detail: error message string if fail/error, else None - original_exception: the exception the guardrail raised, so the @@ -516,6 +542,10 @@ class PipelineExecutor: verbose_proxy_logger.warning("Pipeline: guardrail '%s' not found in callbacks", step.guardrail) return ("error", None, f"Guardrail '{step.guardrail}' not found", None) + hook_data: Final[Mapping[str, object]] = cast(Mapping[str, object], data) # cast-ok: payload + if not _pipeline_stream_scope_allows(callback, hook_data, mode, streaming_chunks): + return ("skip", None, None, None) + hook_input, scans_raw_request = _prepare_hook_input(step, callback, data, raw_request_snapshot) snapshot_entries_before: Final = len(_recorded_guardrail_information(hook_input)) @@ -702,10 +732,13 @@ def _pipeline_action_for_outcome(step: PipelineStep, outcome: str) -> str: """ Map pipeline step outcome to the configured action. + - skip -> next (stream_scope mismatch; do not apply on_pass/on_fail) - pass -> on_pass - fail -> on_fail (content/policy intervention) - error -> on_error if set, else on_fail (backward compatible) """ + if outcome == "skip": + return "next" if outcome == "pass": return step.on_pass if outcome == "fail": diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 201b92f8107..ecc3d0d747c 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -2,11 +2,12 @@ from collections.abc import Mapping from datetime import datetime from enum import Enum from types import MappingProxyType -from typing import Final, Literal +from typing import Final, Literal, cast from pydantic import ConfigDict, Field, field_validator, model_validator from typing_extensions import ReadOnly, Required, TypedDict +from litellm._logging import verbose_logger from litellm.constants import BEDROCK_APPLY_GUARDRAIL_CHUNK_BUDGET_CHARS from litellm.types.llms.base import LiteLLMBaseModel from litellm.types.proxy.guardrails.guardrail_hooks.agent_365 import ( @@ -901,6 +902,93 @@ class ContentFilterConfigModel(LiteLLMBaseModel): MCP_SECURITY_ON_VIOLATION: Final = frozenset({"block", "alert"}) +GuardrailStreamScope = Literal["streaming", "non_streaming", "both"] +DEFAULT_GUARDRAIL_STREAM_SCOPE: Final[GuardrailStreamScope] = "both" + + +class GuardrailEventHooks(str, Enum): + pre_call = "pre_call" + post_call = "post_call" + during_call = "during_call" + logging_only = "logging_only" + pre_mcp_call = "pre_mcp_call" + during_mcp_call = "during_mcp_call" + post_mcp_call = "post_mcp_call" + realtime_input_transcription = "realtime_input_transcription" + + +GUARDRAIL_EVENT_HOOK_VALUES: Final = frozenset(member.value for member in GuardrailEventHooks) + +_GUARDRAIL_STREAM_SCOPES: Final[Mapping[str, GuardrailStreamScope]] = MappingProxyType( + { + "streaming": "streaming", + "non_streaming": "non_streaming", + "both": "both", + } +) + + +def _as_guardrail_stream_scope(value: object) -> GuardrailStreamScope: + if not isinstance(value, str): + raise ValueError(f"stream_scope values must be strings, got {type(value).__name__}") + scope: Final = _GUARDRAIL_STREAM_SCOPES.get(value.lower()) + if scope is None: + raise ValueError(f"stream_scope must be one of both, streaming, non_streaming, got {value!r}") + return scope + + +def _validated_stream_scope_hook(key: object) -> str: + if not isinstance(key, str): + raise ValueError(f"stream_scope keys must be strings, got {type(key).__name__}") + hook: Final = key.lower() + if hook not in GUARDRAIL_EVENT_HOOK_VALUES: + raise ValueError( + f"stream_scope keys must be guardrail modes ({sorted(GUARDRAIL_EVENT_HOOK_VALUES)}), got {key!r}" + ) + return hook + + +def coerce_stream_scope(value: object) -> GuardrailStreamScope | dict[str, GuardrailStreamScope] | None: + if value is None: + return None + if isinstance(value, str): + return _as_guardrail_stream_scope(value) + if isinstance(value, Mapping): + scope_map: Final[Mapping[str, object]] = cast(Mapping[str, object], value) # cast-ok: keys validated below + return { + _validated_stream_scope_hook(key): _as_guardrail_stream_scope(scope) for key, scope in scope_map.items() + } + raise ValueError(f"stream_scope must be a string or mapping, got {type(value).__name__}") + + +def stored_stream_scope(value: object) -> GuardrailStreamScope | dict[str, GuardrailStreamScope] | None: + try: + return coerce_stream_scope(value) + except ValueError: + verbose_logger.warning("Ignoring invalid stored stream_scope value of type %s", type(value).__name__) + return None + + +def with_tolerated_stream_scope(params: Mapping[str, object]) -> dict[str, object]: + if "stream_scope" not in params: + return dict(params) + return { + **params, + "stream_scope": stored_stream_scope(params["stream_scope"]), + } + + +def runtime_stream_scope( + stream_scope: object, +) -> tuple[GuardrailStreamScope, MappingProxyType[str, GuardrailStreamScope]]: + coerced: Final = coerce_stream_scope(stream_scope) + if coerced is None: + return DEFAULT_GUARDRAIL_STREAM_SCOPE, MappingProxyType({}) + if isinstance(coerced, str): + return coerced, MappingProxyType({}) + return DEFAULT_GUARDRAIL_STREAM_SCOPE, MappingProxyType(coerced) + + LoggingOnlyScope = Literal["input", "output", "both"] @@ -1144,6 +1232,21 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up ), ) + stream_scope: GuardrailStreamScope | dict[str, GuardrailStreamScope] | None = Field( + default=None, + description=( + "Whether this guardrail runs on streaming requests, non-streaming requests, or both. " + "A string applies to every configured mode. A map overrides named modes " + "(pre_call, during_call, post_call, ...); omitted keys default to both. " + "Unset means both, matching historical behavior." + ), + ) + + @field_validator("stream_scope", mode="before") + @classmethod + def normalize_stream_scope(cls, v: object) -> GuardrailStreamScope | dict[str, GuardrailStreamScope] | None: + return coerce_stream_scope(v) + logging_only_scope: LoggingOnlyScope | None = Field( default=None, description=( @@ -1278,17 +1381,6 @@ class guardrailConfig(TypedDict): guardrails: list[Guardrail] -class GuardrailEventHooks(str, Enum): - pre_call = "pre_call" - post_call = "post_call" - during_call = "during_call" - logging_only = "logging_only" - pre_mcp_call = "pre_mcp_call" - during_mcp_call = "during_mcp_call" - post_mcp_call = "post_mcp_call" - realtime_input_transcription = "realtime_input_transcription" - - class DynamicGuardrailParams(TypedDict): extra_body: ReadOnly[dict[str, object]] diff --git a/litellm/types/proxy/policy_engine/pipeline_types.py b/litellm/types/proxy/policy_engine/pipeline_types.py index 5278754701d..0a089d2fe9b 100644 --- a/litellm/types/proxy/policy_engine/pipeline_types.py +++ b/litellm/types/proxy/policy_engine/pipeline_types.py @@ -87,7 +87,7 @@ class PipelineStepResult(LiteLLMBaseModel): """Result of executing a single pipeline step.""" guardrail_name: str - outcome: Literal["pass", "fail", "error"] + outcome: Literal["pass", "fail", "error", "skip"] action_taken: str modified_data: dict[str, Any] | None = None error_detail: str | None = None diff --git a/tests/integration/observability/test_guardrail_stream_scope.py b/tests/integration/observability/test_guardrail_stream_scope.py new file mode 100644 index 00000000000..8b832f07aa4 --- /dev/null +++ b/tests/integration/observability/test_guardrail_stream_scope.py @@ -0,0 +1,856 @@ +from __future__ import annotations + +import base64 +import json +import os +import uuid +from collections.abc import Iterator, Mapping, Sequence +from dataclasses import dataclass +from itertools import chain +from pathlib import Path +from types import MappingProxyType +from typing import Final, Literal, TypeAlias + +import httpx +import pytest +import yaml +from integration._support.client import Gateway, JsonValue, Scenario, eventually +from integration._support.database import read_rows, write_rows +from integration._support.process import owned_proxy_process +from integration._support.upstream import ( + _aws_event_frame, # pyright: ignore[reportPrivateUsage] # project Bedrock event-stream encoder +) +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import TypeAdapter + +import litellm + +JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +EndpointScope: TypeAlias = Literal["streaming", "non_streaming"] +ModelKind: TypeAlias = Literal["router", "direct"] +StreamingAction: TypeAlias = Literal["converse-stream", "invoke-with-response-stream"] +NonStreamingAction: TypeAlias = Literal["converse", "invoke"] +BedrockAction: TypeAlias = StreamingAction | NonStreamingAction +BEDROCK_MODEL_ID: Final = "anthropic.claude-sonnet-5-v1:0" +BEDROCK_FALSE_POSITIVE_MODEL_ID: Final = "anthropic.claude-converse-stream-test-v1:0" +BEDROCK_EVENT_STREAM: Final = "application/vnd.amazon.eventstream" +GUARDRAIL_PATH: Final = "/beta/litellm_basic_guardrail_api" +SERVER_STREAMING_CLASSIFICATION_KEY: Final = "litellm_server_streaming_classification" +STREAMING_ACTIONS: Final[tuple[StreamingAction, ...]] = ( + "converse-stream", + "invoke-with-response-stream", +) +NON_STREAMING_ACTIONS: Final[tuple[NonStreamingAction, ...]] = ("converse", "invoke") +SCOPES: Final[tuple[EndpointScope, ...]] = ("streaming", "non_streaming") +HOSTILE_CLASSIFICATION_CASES: Final = ( + pytest.param("is_streaming_request", True, id="boolean-marker"), + pytest.param("is_streaming_request", "litellm-server-streaming", id="server-marker-string"), + pytest.param("litellm_server_streaming_classification", True, id="classification-field"), +) +WORKTREE: Final = Path(__file__).resolve().parents[3] +LITELLM_PATH: Final = Path(litellm.__file__).resolve() +assert LITELLM_PATH.is_relative_to(WORKTREE), (LITELLM_PATH, WORKTREE) +print(f"stream_scope repro litellm import: {LITELLM_PATH}") # noqa: T201 # required worktree evidence + + +def _json(value: object) -> bytes: + return json.dumps(value, separators=(",", ":")).encode() + + +def _strings(value: JsonValue) -> tuple[str, ...]: + if isinstance(value, str): + return (value,) + if isinstance(value, list): + return tuple(chain.from_iterable(_strings(item) for item in value)) + if isinstance(value, dict): + return tuple(chain.from_iterable(_strings(item) for item in value.values())) + return () + + +def _key_names(value: JsonValue) -> tuple[str, ...]: + if isinstance(value, dict): + return tuple(value) + tuple(chain.from_iterable(_key_names(item) for item in value.values())) + if isinstance(value, list): + return tuple(chain.from_iterable(_key_names(item) for item in value)) + return () + + +def _marker(body: JsonValue) -> str: + return next(value for value in _strings(body) if value.startswith("scope-")) + + +def _bedrock_converse_stream(marker: str) -> bytes: + return b"".join( + _aws_event_frame(event_type, payload, marker, marker) + for event_type, payload in ( + ("messageStart", {"role": "assistant"}), + ( + "contentBlockDelta", + {"delta": {"text": f"scripted Bedrock reply {marker}"}, "contentBlockIndex": 0}, + ), + ("messageStop", {"stopReason": "end_turn"}), + ("metadata", {"usage": {"inputTokens": 11, "outputTokens": 4, "totalTokens": 15}}), + ) + ) + + +def _invoke_chunk(payload: Mapping[str, JsonValue], marker: str) -> bytes: + encoded: Final = base64.b64encode(_json(payload)).decode() + return _aws_event_frame("chunk", {"bytes": encoded}, marker, marker) + + +def _bedrock_invoke_stream(marker: str) -> bytes: + events: Final = ( + { + "type": "message_start", + "message": { + "id": f"msg-{marker}", + "type": "message", + "role": "assistant", + "model": BEDROCK_MODEL_ID, + "content": [], + "stop_reason": None, + "stop_sequence": None, + "usage": {"input_tokens": 11, "output_tokens": 0}, + }, + }, + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "text_delta", "text": f"scripted Bedrock reply {marker}"}, + }, + {"type": "content_block_stop", "index": 0}, + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn"}, + "usage": {"input_tokens": 11, "output_tokens": 4}, + }, + {"type": "message_stop"}, + ) + return b"".join(_invoke_chunk(event, marker) for event in events) + + +def _provider(request: Request) -> Reply: + if not request.body: + return Reply(status=400, body=_json({"error": "empty request body"})) + body: Final = JSON_OBJECT.validate_json(request.body) + marker: Final = _marker(body) + target: Final = request.target.split("?", 1)[0] + if target.startswith("/passthrough"): + if body.get("stream") is True: + streamed_response: Final = _json({"received": body}) + return Reply( + content_type="text/event-stream", + chunks=(b"data: " + streamed_response + b"\n\n", b"data: [DONE]\n\n"), + ) + return Reply(body=_json({"received": body})) + if target.endswith("/converse-stream"): + return Reply(body=_bedrock_converse_stream(marker), content_type=BEDROCK_EVENT_STREAM) + if target.endswith("/invoke-with-response-stream"): + return Reply(body=_bedrock_invoke_stream(marker), content_type=BEDROCK_EVENT_STREAM) + if target.endswith("/converse") or target.endswith("/invoke"): + return Reply(body=_json({"output": f"scripted Bedrock reply {marker}"})) + if target == "/v1/chat/completions": + if body.get("stream") is True: + chunk: Final = { + "id": f"chatcmpl-{marker}", + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "delta": {"role": "assistant", "content": f"scripted chat reply {marker}"}, + "finish_reason": None, + } + ], + } + final_chunk: Final = { + **chunk, + "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}], + } + return Reply( + content_type="text/event-stream", + chunks=( + b"data: " + _json(chunk) + b"\n\n", + b"data: " + _json(final_chunk) + b"\n\n", + b"data: [DONE]\n\n", + ), + ) + return Reply( + body=_json( + { + "id": f"chatcmpl-{marker}", + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": f"scripted chat reply {marker}"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + } + ) + ) + return Reply(status=404, body=_json({"error": f"unexpected upstream path: {target}"})) + + +def _sink(request: Request) -> Reply: + assert request.target.endswith(GUARDRAIL_PATH), request.target + assert b"scope-" in request.body, request.body.decode() + return Reply(body=_json({"action": "NONE"})) + + +def _rail(name: str, sink: Wire, scope: EndpointScope) -> dict[str, JsonValue]: + return { + "guardrail_name": name, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": False, + "stream_scope": scope, + "api_base": f"{sink.url}/{name}", + "api_key": "synthetic-guardrail-key", + }, + } + + +def _chat_proxy_config(provider_url: str, guardrails: list[dict[str, JsonValue]]) -> dict[str, object]: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + return { + **config, + "guardrails": guardrails, + "model_list": [ + { + "model_name": "scope-invalid-config-chat", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_base": f"{provider_url}/v1", + "api_key": "synthetic-provider-key", + }, + } + ], + } + + +@dataclass(frozen=True, slots=True) +class ReproRig: + candidate: Gateway + direct_candidate: Gateway + scenario: Scenario + models: Mapping[str, str] + rails: Mapping[str, str] + provider: Wire + sink: Wire + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[ReproRig]: + with httpx.Client( + base_url=os.environ["INTEGRATION_PROXY_URL"], + timeout=30, + trust_env=False, + ) as root_client: + root_gateway: Final = Gateway( + root_client, + os.environ.get("INTEGRATION_MASTER_KEY", "sk-integration-master"), + os.environ["INTEGRATION_UPSTREAM_URL"], + ) + directory: Final = tmp_path_factory.mktemp("guardrail-stream-scope-repro") + with wire_server(_provider) as provider, wire_server(_sink) as sink: + rails: Final = MappingProxyType( + { + "bedrock_streaming": "bedrock_streaming", + "bedrock_non_streaming": "bedrock_non_streaming", + "chat_streaming": "chat_streaming", + "passthrough_streaming": "passthrough_streaming", + "passthrough_non_streaming": "passthrough_non_streaming", + "passthrough_spoof_streaming": "passthrough_spoof_streaming", + } + ) + models: Final = MappingProxyType( + { + "chat": "scope-chat", + "bedrock_router": "scope-bedrock-router", + "bedrock_false_positive_router": "scope-bedrock-converse-stream-model", + } + ) + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [ + _rail(rails["bedrock_streaming"], sink, "streaming"), + _rail(rails["bedrock_non_streaming"], sink, "non_streaming"), + _rail(rails["chat_streaming"], sink, "streaming"), + _rail(rails["passthrough_streaming"], sink, "streaming"), + _rail(rails["passthrough_non_streaming"], sink, "non_streaming"), + _rail(rails["passthrough_spoof_streaming"], sink, "streaming"), + ] + config["model_list"] = [ + { + "model_name": models["chat"], + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_base": f"{provider.url}/v1", + "api_key": "synthetic-provider-key", + }, + }, + { + "model_name": models["bedrock_router"], + "litellm_params": { + "model": f"bedrock/{BEDROCK_MODEL_ID}", + "api_base": provider.url, + "aws_access_key_id": "AKIASYNTHETICSTREAMSCOPE", + "aws_secret_access_key": "synthetic-bedrock-secret", + "aws_region_name": "us-east-1", + }, + }, + { + "model_name": models["bedrock_false_positive_router"], + "litellm_params": { + "model": f"bedrock/{BEDROCK_FALSE_POSITIVE_MODEL_ID}", + "api_base": provider.url, + "aws_access_key_id": "AKIASYNTHETICSTREAMSCOPE", + "aws_secret_access_key": "synthetic-bedrock-secret", + "aws_region_name": "us-east-1", + }, + }, + ] + config["environment_variables"] = { + "AWS_BEDROCK_RUNTIME_ENDPOINT": provider.url, + "AWS_ACCESS_KEY_ID": "AKIASYNTHETICSTREAMSCOPE", + "AWS_SECRET_ACCESS_KEY": "synthetic-bedrock-secret", + "AWS_REGION": "us-east-1", + "AWS_REGION_NAME": "us-east-1", + } + config["general_settings"]["pass_through_endpoints"] = [ + { + "path": "/pt-forward", + "target": f"{provider.url}/passthrough", + "include_subpath": True, + }, + { + "path": "/pt-spoof", + "target": f"{provider.url}/passthrough", + "include_subpath": True, + "guardrails": {rails["passthrough_spoof_streaming"]: None}, + }, + { + "path": "/pt-scope", + "target": f"{provider.url}/passthrough", + "include_subpath": True, + "guardrails": { + rails["passthrough_streaming"]: None, + rails["passthrough_non_streaming"]: None, + }, + }, + ] + config_path: Final = directory / "stream-scope-repro.yaml" + config_path.write_text(yaml.safe_dump(config)) + with owned_proxy_process(root_gateway, directory, {}, config=config_path, workers=1) as owned: + with owned.gateway.scenario() as scenario: + direct_config: Final = { + **config, + "model_list": [ + *config["model_list"], + { + "model_name": f"scope-unused-{uuid.uuid4().hex}*", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_base": f"{provider.url}/v1", + "api_key": "synthetic-provider-key", + }, + }, + ], + } + direct_config_path: Final = directory / "stream-scope-direct.yaml" + direct_config_path.write_text(yaml.safe_dump(direct_config)) + with owned_proxy_process( + root_gateway, + directory, + {}, + config=direct_config_path, + workers=1, + ) as direct: + yield ReproRig( + owned.gateway, + direct.gateway, + scenario, + models, + rails, + provider, + sink, + ) + + +def _bedrock_request( + action: BedrockAction, + model_path: str, + marker: str, +) -> tuple[str, dict[str, JsonValue]]: + if action in ("converse", "converse-stream"): + body: Final = { + "messages": [{"role": "user", "content": [{"text": marker}]}], + "inferenceConfig": {"maxTokens": 16}, + } + else: + body = { + "anthropic_version": "bedrock-2023-05-31", + "max_tokens": 16, + "messages": [{"role": "user", "content": [{"type": "text", "text": marker}]}], + } + return f"/bedrock/model/{model_path}/{action}", body + + +def _matching_requests(wire: Wire, marker: str) -> tuple[Request, ...]: + return tuple(request for request in wire.drain() if marker.encode() in request.body) + + +def _rail_scans(rows: Sequence[Request], rail_name: str, marker: str) -> tuple[Request, ...]: + return tuple( + request for request in rows if request.target.startswith(f"/{rail_name}/") and marker.encode() in request.body + ) + + +def _chat_request_with_scans( + gateway: Gateway, + sink: Wire, + model: str, + marker: str, + streamed: bool, + rail_name: str, +) -> tuple[httpx.Response, tuple[Request, ...]]: + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": marker}], + "stream": streamed, + "guardrails": [rail_name], + }, + ) + return response, _rail_scans(sink.drain(), rail_name, marker) + + +def _spend_row_for_call(call_id: str, content: bytes) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + "SELECT request_id, litellm_call_id, spend, prompt_tokens, completion_tokens, metadata " + 'FROM "LiteLLM_SpendLogs" WHERE request_id=%s OR litellm_call_id=%s', + (call_id, call_id), + ), + lambda values: len(values) >= 1, + seconds=70, + ) + assert len(rows) == 1, (call_id, rows, content) + assert call_id in (rows[0]["request_id"], rows[0]["litellm_call_id"]), content + return rows[0] + + +@pytest.mark.parametrize("model_kind", ("router", "direct")) +@pytest.mark.parametrize("action", STREAMING_ACTIONS) +@pytest.mark.parametrize("scope", SCOPES) +def test_bedrock_streaming_actions_run_streaming_scoped_rails( + rig: ReproRig, + model_kind: ModelKind, + action: StreamingAction, + scope: EndpointScope, +) -> None: + marker: Final = f"scope-bedrock-stream-{uuid.uuid4().hex}" + model_path: Final = rig.models["bedrock_router"] if model_kind == "router" else BEDROCK_MODEL_ID + path, body = _bedrock_request(action, model_path, marker) + expected_body: Final = ( + _bedrock_converse_stream(marker) if action == "converse-stream" else _bedrock_invoke_stream(marker) + ) + call_id: Final = f"stream-scope-{uuid.uuid4().hex}" + key: Final = rig.scenario.key(guardrails=[rig.rails[f"bedrock_{scope}"]]) + candidate: Final = rig.candidate if model_kind == "router" else rig.direct_candidate + response: Final = candidate.request( + "POST", + path, + body, + key=key, + headers={"x-litellm-call-id": call_id}, + ) + assert response.status_code == 200, response.content + assert response.headers.get("content-type") == BEDROCK_EVENT_STREAM, dict(response.headers) + assert response.content == expected_body, response.content + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows, response.content) + provider_body: Final = JSON_OBJECT.validate_json(provider_rows[0].body) + key_names: Final = _key_names(provider_body) + assert "is_streaming_request" not in key_names, provider_body + assert not tuple(name for name in key_names if name.startswith("litellm_")), provider_body + _spend_row_for_call(call_id, response.content) + sink_rows: Final = _rail_scans(rig.sink.drain(), rig.rails[f"bedrock_{scope}"], marker) + expected_scans: Final = int(scope == "streaming") + assert len(sink_rows) == expected_scans, (marker, model_kind, action, scope, sink_rows, response.content) + + +@pytest.mark.parametrize("model_kind", ("router", "direct")) +@pytest.mark.parametrize("action", NON_STREAMING_ACTIONS) +@pytest.mark.parametrize("scope", SCOPES) +def test_bedrock_non_streaming_actions_run_non_streaming_scoped_rails( + rig: ReproRig, + model_kind: ModelKind, + action: NonStreamingAction, + scope: EndpointScope, +) -> None: + marker: Final = f"scope-bedrock-nonstream-{uuid.uuid4().hex}" + model_path: Final = rig.models["bedrock_router"] if model_kind == "router" else BEDROCK_MODEL_ID + path, body = _bedrock_request(action, model_path, marker) + key: Final = rig.scenario.key(guardrails=[rig.rails[f"bedrock_{scope}"]]) + candidate: Final = rig.candidate if model_kind == "router" else rig.direct_candidate + response: Final = candidate.request("POST", path, body, key=key) + assert response.status_code == 200, response.text + assert marker in response.text, response.text + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows, response.text) + provider_body: Final = JSON_OBJECT.validate_json(provider_rows[0].body) + key_names: Final = _key_names(provider_body) + assert "is_streaming_request" not in key_names, provider_body + assert not tuple(name for name in key_names if name.startswith("litellm_")), provider_body + sink_rows: Final = _rail_scans(rig.sink.drain(), rig.rails[f"bedrock_{scope}"], marker) + expected_scans: Final = int(scope == "non_streaming") + assert len(sink_rows) == expected_scans, (marker, model_kind, action, scope, sink_rows, response.text) + + +@pytest.mark.parametrize("model_kind", ("router", "direct")) +def test_bedrock_model_id_streaming_action_text_on_converse_is_non_streaming( + rig: ReproRig, + model_kind: ModelKind, +) -> None: + marker: Final = f"scope-bedrock-converse-model-{uuid.uuid4().hex}" + model_path: Final = ( + rig.models["bedrock_false_positive_router"] if model_kind == "router" else BEDROCK_FALSE_POSITIVE_MODEL_ID + ) + path, body = _bedrock_request("converse", model_path, marker) + call_id: Final = f"stream-scope-bedrock-{uuid.uuid4().hex}" + key: Final = rig.scenario.key( + guardrails=[ + rig.rails["bedrock_streaming"], + rig.rails["bedrock_non_streaming"], + ] + ) + candidate: Final = rig.candidate if model_kind == "router" else rig.direct_candidate + response: Final = candidate.request( + "POST", + path, + body, + key=key, + headers={"x-litellm-call-id": call_id}, + ) + assert response.status_code == 200, response.text + assert marker in response.text, response.text + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, model_kind, provider_rows, response.text) + sink_rows: Final = rig.sink.drain() + streaming_rows: Final = _rail_scans(sink_rows, rig.rails["bedrock_streaming"], marker) + non_streaming_rows: Final = _rail_scans(sink_rows, rig.rails["bedrock_non_streaming"], marker) + assert streaming_rows == (), (marker, model_kind, streaming_rows, response.text) + assert len(non_streaming_rows) == 1, (marker, model_kind, non_streaming_rows, response.text) + + +@pytest.mark.parametrize("streamed", (False, True), ids=("stream-absent", "stream-true")) +def test_configured_passthrough_forwards_caller_is_streaming_request_field( + rig: ReproRig, + streamed: bool, +) -> None: + marker: Final = f"scope-passthrough-forward-{uuid.uuid4().hex}" + caller_value: Final = f"caller-{uuid.uuid4().hex}" + body: Final = { + "marker": marker, + "is_streaming_request": caller_value, + **({"stream": True} if streamed else {}), + } + response: Final = rig.candidate.request("POST", "/pt-forward", body) + assert response.status_code == 200, response.text + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows, response.text) + upstream_body: Final = JSON_OBJECT.validate_json(provider_rows[0].body) + assert upstream_body.get("is_streaming_request") == caller_value, ( + marker, + caller_value, + upstream_body, + response.text, + ) + assert upstream_body == body, (marker, body, upstream_body, response.text) + if streamed: + assert response.headers.get("content-type", "").lower().startswith("text/event-stream"), dict(response.headers) + event_body: Final = response.text.removeprefix("data: ").split("\n", maxsplit=1)[0] + response_body: Final = JSON_OBJECT.validate_json(event_body) + else: + response_body = JSON_OBJECT.validate_json(response.content) + assert response_body == {"received": body}, response.text + + +@pytest.mark.parametrize(("hostile_field", "hostile_value"), HOSTILE_CLASSIFICATION_CASES) +def test_client_cannot_spoof_server_stream_classification( + rig: ReproRig, + hostile_field: str, + hostile_value: JsonValue, +) -> None: + chat_marker: Final = f"scope-chat-spoof-{uuid.uuid4().hex}" + chat_body: Final = { + "model": rig.models["chat"], + "messages": [{"role": "user", "content": chat_marker}], + "stream": False, + hostile_field: hostile_value, + } + chat_key: Final = rig.scenario.key(guardrails=[rig.rails["chat_streaming"]]) + chat_response: Final = rig.candidate.request( + "POST", + "/v1/chat/completions", + chat_body, + key=chat_key, + ) + chat_provider_rows: Final = _matching_requests(rig.provider, chat_marker) + chat_upstream_body: Final = chat_provider_rows[0].body.decode() if chat_provider_rows else "" + print( # noqa: T201 # required chat hostile-body observation + f"chat hostile {hostile_field}={hostile_value!r}: " + f"status={chat_response.status_code}, response={chat_response.text!r}, upstream={chat_upstream_body}" + ) + assert chat_response.status_code == 200, chat_response.text + assert len(chat_provider_rows) == 1, (chat_marker, chat_provider_rows, chat_response.text) + assert JSON_OBJECT.validate_json(chat_provider_rows[0].body) == { + "messages": [{"role": "user", "content": chat_marker}], + "model": "gpt-4o-mini", + hostile_field: hostile_value, + }, (hostile_field, hostile_value, chat_upstream_body) + assert _rail_scans(rig.sink.drain(), rig.rails["chat_streaming"], chat_marker) == (), ( + chat_marker, + hostile_field, + hostile_value, + chat_response.text, + ) + + +@pytest.mark.parametrize(("hostile_field", "hostile_value"), HOSTILE_CLASSIFICATION_CASES) +def test_configured_passthrough_cannot_spoof_server_stream_classification( + rig: ReproRig, + hostile_field: str, + hostile_value: JsonValue, +) -> None: + passthrough_marker: Final = f"scope-passthrough-spoof-{uuid.uuid4().hex}" + passthrough_body: Final = { + "marker": passthrough_marker, + "stream": False, + hostile_field: hostile_value, + } + passthrough_response: Final = rig.candidate.request("POST", "/pt-spoof", passthrough_body) + assert passthrough_response.status_code == 200, passthrough_response.text + passthrough_provider_rows: Final = _matching_requests(rig.provider, passthrough_marker) + assert len(passthrough_provider_rows) == 1, ( + passthrough_marker, + passthrough_provider_rows, + passthrough_response.text, + ) + assert _rail_scans(rig.sink.drain(), rig.rails["passthrough_spoof_streaming"], passthrough_marker) == (), ( + passthrough_marker, + hostile_field, + hostile_value, + passthrough_response.text, + ) + + +@pytest.mark.parametrize( + ("streamed", "request_fields"), + ((True, {"stream": True}), (False, {"stream": False}), (False, {})), + ids=("stream-true", "stream-false", "stream-absent"), +) +def test_passthrough_scope_follows_proxy_stream_decision( + rig: ReproRig, + streamed: bool, + request_fields: dict[str, JsonValue], +) -> None: + marker: Final = f"scope-passthrough-scope-{uuid.uuid4().hex}" + body: Final = {"marker": marker, **request_fields} + response: Final = rig.candidate.request("POST", "/pt-scope", body) + assert response.status_code == 200, response.text + if streamed: + expected_frame: Final = f"data: {_json({'received': body}).decode()}\n\ndata: [DONE]\n\n" + assert response.headers.get("content-type", "").lower().startswith("text/event-stream"), dict(response.headers) + assert response.text == expected_frame, response.text + assert response.headers.get("transfer-encoding", "").lower() == "chunked", dict(response.headers) + else: + response_body: Final = JSON_OBJECT.validate_json(response.content) + assert response_body == {"received": body}, response.text + assert "content-length" in response.headers, dict(response.headers) + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows, response.text) + upstream_body: Final = JSON_OBJECT.validate_json(provider_rows[0].body) + assert upstream_body == body, (marker, body, upstream_body, response.text) + assert SERVER_STREAMING_CLASSIFICATION_KEY not in upstream_body, upstream_body + sink_rows: Final = rig.sink.drain() + streaming_rows: Final = _rail_scans(sink_rows, rig.rails["passthrough_streaming"], marker) + non_streaming_rows: Final = _rail_scans(sink_rows, rig.rails["passthrough_non_streaming"], marker) + assert len(streaming_rows) == int(streamed), (marker, streamed, streaming_rows, response.text) + assert len(non_streaming_rows) == int(not streamed), (marker, streamed, non_streaming_rows, response.text) + + +def test_invalid_yaml_stream_scope_keeps_rail_running_on_both_shapes(rig: ReproRig, tmp_path: Path) -> None: + name: Final = f"scope-invalid-yaml-{uuid.uuid4().hex}" + invalid_rail: Final = { + "guardrail_name": name, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": False, + "stream_scope": "sometimes", + "api_base": f"{rig.sink.url}/{name}", + "api_key": "synthetic-guardrail-key", + }, + } + config: Final = _chat_proxy_config(rig.provider.url, [invalid_rail]) + config_path: Final = tmp_path / "invalid-stream-scope.yaml" + config_path.write_text(yaml.safe_dump(config)) + with owned_proxy_process(rig.candidate, tmp_path, {}, config=config_path, workers=1) as owned: + markers: Final = tuple(f"scope-invalid-yaml-{int(streamed)}-{uuid.uuid4().hex}" for streamed in (False, True)) + observations: Final = tuple( + _chat_request_with_scans( + owned.gateway, + rig.sink, + "scope-invalid-config-chat", + marker, + streamed, + name, + ) + for streamed, marker in zip((False, True), markers) + ) + assert tuple(response.status_code for response, _ in observations) == (200, 200), tuple( + response.text for response, _ in observations + ) + assert tuple(len(sink_rows) for _, sink_rows in observations) == (1, 1), (markers, observations) + listed: Final = owned.gateway.request("GET", "/guardrails/list") + assert listed.status_code == 200, listed.text + list_payload: Final = JSON_OBJECT.validate_json(listed.content) + listed_guardrails: Final = list_payload.get("guardrails") + assert isinstance(listed_guardrails, list), list_payload + listed_rail: Final = next( + (row for row in listed_guardrails if isinstance(row, dict) and row.get("guardrail_name") == name), + None, + ) + assert isinstance(listed_rail, dict), list_payload + listed_params: Final = listed_rail.get("litellm_params") + assert isinstance(listed_params, dict), listed_rail + assert listed_params.get("stream_scope") is None, listed_rail + + +def test_persisted_invalid_stream_scope_row_stays_readable_and_enforced( + rig: ReproRig, + tmp_path: Path, +) -> None: + guardrail_id: Final = str(uuid.uuid4()) + guardrail_name: Final = f"scope-invalid-persisted-{uuid.uuid4().hex}" + params: Final = { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": False, + "api_base": f"{rig.sink.url}/{guardrail_name}", + "api_key": "synthetic-guardrail-key", + "stream_scope": "sometimes", + } + database_url: Final = os.environ.get("INTEGRATION_PROXY_DATABASE_URL") or os.environ["DATABASE_URL"] + write_rows( + 'INSERT INTO "LiteLLM_GuardrailsTable" ' + '("guardrail_id", "guardrail_name", "litellm_params", "created_at", "updated_at") ' + "VALUES (%s, %s, %s::jsonb, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)", + (guardrail_id, guardrail_name, json.dumps(params)), + database_url=database_url, + ) + try: + config: Final = _chat_proxy_config(rig.provider.url, []) + config_path: Final = tmp_path / "persisted-invalid-stream-scope.yaml" + config_path.write_text(yaml.safe_dump(config)) + with owned_proxy_process(rig.candidate, tmp_path, {}, config=config_path, workers=1) as owned: + info: Final = eventually( + lambda: owned.gateway.request("GET", f"/guardrails/{guardrail_id}/info"), + lambda response: response.status_code != 404, + seconds=20, + ) + listed: Final = owned.gateway.request("GET", "/v2/guardrails/list") + list_payload: Final = JSON_OBJECT.validate_json(listed.content) if listed.status_code == 200 else {} + listed_guardrails: Final = list_payload.get("guardrails") + includes_row: Final = isinstance(listed_guardrails, list) and any( + isinstance(row, dict) and row.get("guardrail_id") == guardrail_id for row in listed_guardrails + ) + markers: Final = ( + f"scope-invalid-persisted-0-{uuid.uuid4().hex}", + f"scope-invalid-persisted-1-{uuid.uuid4().hex}", + ) + observations: Final = tuple( + _chat_request_with_scans( + owned.gateway, + rig.sink, + "scope-invalid-config-chat", + marker, + streamed, + guardrail_name, + ) + for streamed, marker in zip((False, True), markers) + ) + assert ( + info.status_code == 200 + and listed.status_code == 200 + and includes_row + and tuple(response.status_code for response, _ in observations) == (200, 200) + and tuple(len(sink_rows) for _, sink_rows in observations) == (1, 1) + ), { + "info": (info.status_code, info.text), + "list": (listed.status_code, listed.text), + "includes_row": includes_row, + "responses": tuple((response.status_code, response.text) for response, _ in observations), + "scan_counts": tuple(len(sink_rows) for _, sink_rows in observations), + } + finally: + write_rows( + 'DELETE FROM "LiteLLM_GuardrailsTable" WHERE guardrail_id=%s', + (guardrail_id,), + database_url=database_url, + ) + + +def test_management_rejects_invalid_stream_scope(rig: ReproRig) -> None: + name: Final = f"scope-invalid-management-{uuid.uuid4().hex}" + invalid_params: Final = { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": False, + "api_base": f"{rig.sink.url}/{name}", + "api_key": "synthetic-guardrail-key", + "stream_scope": "sometimes", + } + created_invalid: Final = rig.candidate.request( + "POST", + "/guardrails", + {"guardrail": {"guardrail_name": name, "litellm_params": invalid_params}}, + ) + assert created_invalid.status_code == 422, created_invalid.text + + valid_params: Final = {**invalid_params, "stream_scope": "both"} + created: Final = rig.candidate.request( + "POST", + "/guardrails", + {"guardrail": {"guardrail_name": name, "litellm_params": valid_params}}, + ) + assert created.status_code == 200, created.text + guardrail_id: Final = str(created.json()["guardrail_id"]) + try: + put_response: Final = rig.candidate.request( + "PUT", + f"/guardrails/{guardrail_id}", + {"guardrail": {"guardrail_name": name, "litellm_params": invalid_params}}, + ) + patch_response: Final = rig.candidate.request( + "PATCH", + f"/guardrails/{guardrail_id}", + {"litellm_params": {"stream_scope": "sometimes"}}, + ) + assert put_response.status_code == 422, put_response.text + assert patch_response.status_code == 422, patch_response.text + finally: + deleted: Final = rig.candidate.request("DELETE", f"/guardrails/{guardrail_id}") + assert deleted.status_code == 200, deleted.text diff --git a/tests/integration/observability/test_guardrail_stream_scope_chaos.py b/tests/integration/observability/test_guardrail_stream_scope_chaos.py new file mode 100644 index 00000000000..f8a9d0a4bba --- /dev/null +++ b/tests/integration/observability/test_guardrail_stream_scope_chaos.py @@ -0,0 +1,1079 @@ +from __future__ import annotations + +import json +import os +import shutil +import signal +import socket +import subprocess +import uuid +from collections.abc import Iterator, Mapping, Sequence +from concurrent.futures import Future, ThreadPoolExecutor +from contextlib import ExitStack, contextmanager +from dataclasses import dataclass +from itertools import chain +from pathlib import Path +from threading import Barrier, Event +from types import MappingProxyType +from typing import Final, Literal, TypeAlias, cast + +import httpx +import psutil +import pytest +import yaml +from integration._support.client import Gateway, JsonValue, eventually +from integration._support.database import read_rows +from integration._support.process import OwnedProxy, owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import TypeAdapter + +ChaosEndpoint: TypeAlias = Literal["chat", "messages", "responses"] +CHAOS_ENDPOINTS: Final[tuple[ChaosEndpoint, ...]] = ("chat", "messages", "responses") +CHAOS_MODELS: Final = MappingProxyType( + {"chat": "chaos-chat", "messages": "chaos-messages", "responses": "chaos-responses"} +) +GUARDRAIL_PATH: Final = "/beta/litellm_basic_guardrail_api" +JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +POSTGRES_IMAGE: Final = "postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5" + + +def _json(value: object) -> bytes: + return json.dumps(value, separators=(",", ":")).encode() + + +def _texts(value: JsonValue) -> tuple[str, ...]: + if isinstance(value, str): + return (value,) + if isinstance(value, list): + return tuple(chain.from_iterable(_texts(item) for item in value)) + if isinstance(value, dict): + return tuple(chain.from_iterable(_texts(item) for item in value.values())) + return () + + +def _marker(body: Mapping[str, JsonValue]) -> str: + return next((text for text in _texts(dict(body)) if text.startswith("audit-")), "audit-chaos") + + +def _sse(events: Sequence[Mapping[str, JsonValue]]) -> tuple[bytes, ...]: + return tuple(f"data: {json.dumps(event, separators=(',', ':'))}\n\n".encode() for event in events) + ( + b"data: [DONE]\n\n", + ) + + +def _messages_stream(message: Mapping[str, JsonValue]) -> tuple[bytes, ...]: + content: Final = cast(list[JsonValue], message["content"]) + text: Final = cast(dict[str, JsonValue], content[0])["text"] + assert isinstance(text, str) + return ( + f"event: message_start\ndata: {json.dumps({**message, 'content': [], 'stop_reason': None, 'usage': {'input_tokens': 2, 'output_tokens': 0}})}\n\n".encode(), + b'event: content_block_start\ndata: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}\n\n', + f"event: content_block_delta\ndata: {json.dumps({'type': 'content_block_delta', 'index': 0, 'delta': {'type': 'text_delta', 'text': text}})}\n\n".encode(), + b'event: content_block_stop\ndata: {"type":"content_block_stop","index":0}\n\n', + f"event: message_delta\ndata: {json.dumps({'type': 'message_delta', 'delta': {'stop_reason': 'end_turn', 'stop_sequence': None}, 'usage': {'output_tokens': 2}})}\n\n".encode(), + b'event: message_stop\ndata: {"type":"message_stop"}\n\n', + ) + + +def _responses_stream( + response: Mapping[str, JsonValue], + output: Mapping[str, JsonValue], + marker: str, +) -> tuple[bytes, ...]: + events: Final[tuple[dict[str, JsonValue], ...]] = ( + {"type": "response.created", "response": {**response, "status": "in_progress", "output": []}}, + {"type": "response.in_progress", "response": {**response, "status": "in_progress", "output": []}}, + {"type": "response.output_item.added", "item": dict(output), "output_index": 0}, + { + "type": "response.content_part.added", + "item_id": f"msg-{marker}", + "output_index": 0, + "content_index": 0, + "part": {"type": "output_text", "text": "", "annotations": []}, + }, + { + "type": "response.output_text.delta", + "item_id": f"msg-{marker}", + "output_index": 0, + "content_index": 0, + "delta": marker, + }, + { + "type": "response.output_text.done", + "item_id": f"msg-{marker}", + "output_index": 0, + "content_index": 0, + "text": marker, + }, + { + "type": "response.content_part.done", + "item_id": f"msg-{marker}", + "output_index": 0, + "content_index": 0, + "part": cast(list[JsonValue], output["content"])[0], + }, + {"type": "response.output_item.done", "item": dict(output), "output_index": 0}, + {"type": "response.completed", "response": dict(response)}, + ) + return tuple( + f"event: {event['type']}\ndata: {json.dumps({**event, 'sequence_number': index}, separators=(',', ':'))}\n\n".encode() + for index, event in enumerate(events) + ) + + +def _provider(request: Request) -> Reply: + if request.method == "GET" and request.target.partition("?")[0] == "/v1/models": + return Reply(body=_json({"object": "list", "data": [{"id": "gpt-4o-mini", "object": "model"}]})) + if not request.body: + return Reply(status=400, body=_json({"error": "request body is required"})) + body: Final = JSON_OBJECT.validate_json(request.body) + marker: Final = _marker(body) + streamed: Final = bool(body.get("stream")) + if request.target == "/v1/chat/completions": + if streamed: + return Reply( + content_type="text/event-stream", + chunks=_sse( + ( + { + "id": f"chatcmpl-{marker}", + "object": "chat.completion.chunk", + "choices": [{"index": 0, "delta": {"content": marker}, "finish_reason": None}], + }, + ) + ), + ) + return Reply( + body=_json( + { + "id": f"chatcmpl-{marker}", + "object": "chat.completion", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": marker}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 2, "completion_tokens": 2, "total_tokens": 4}, + } + ) + ) + if request.target == "/v1/messages": + message: Final = { + "id": f"msg-{marker}", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5-20250929", + "content": [{"type": "text", "text": marker}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 2, "output_tokens": 2}, + } + return ( + Reply( + content_type="text/event-stream", + chunks=_messages_stream(message), + ) + if streamed + else Reply(body=_json(message)) + ) + if request.target == "/v1/responses": + response: Final = { + "id": f"resp-{marker}", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [ + { + "type": "message", + "id": f"msg-{marker}", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": marker, "annotations": []}], + } + ], + "usage": {"input_tokens": 2, "output_tokens": 2, "total_tokens": 4}, + } + return ( + Reply( + content_type="text/event-stream", + chunks=_responses_stream(response, cast(dict[str, JsonValue], response["output"][0]), marker), + ) + if streamed + else Reply(body=_json(response)) + ) + return Reply(status=404, body=_json({"error": f"unexpected provider target {request.target}"})) + + +def _sink(request: Request) -> Reply: + assert request.target.endswith(GUARDRAIL_PATH), request.target + body: Final = JSON_OBJECT.validate_json(request.body) + assert body.get("litellm_call_id") is not None or any("audit-" in text for text in _texts(body)), ( + request.body.decode() + ) + return Reply(body=_json({"action": "NONE"})) + + +def _rail( + name: str, + sink_url: str, + scope: Literal["streaming", "non_streaming"], + *, + default_on: bool = False, +) -> dict[str, JsonValue]: + return { + "guardrail_name": name, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": default_on, + "stream_scope": scope, + "api_base": f"{sink_url}/{name}", + "api_key": "synthetic-chaos-key", + }, + } + + +def _config(provider_url: str, rails: Sequence[dict[str, JsonValue]]) -> dict[str, JsonValue]: + base: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + return cast( + dict[str, JsonValue], + { + **base, + "guardrails": list(rails), + "model_list": [ + { + "model_name": CHAOS_MODELS["chat"], + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_base": f"{provider_url}/v1", + "api_key": "synthetic-provider-key", + "num_retries": 0, + }, + }, + { + "model_name": CHAOS_MODELS["messages"], + "litellm_params": { + "model": "anthropic/claude-sonnet-4-5-20250929", + "api_base": provider_url, + "api_key": "synthetic-provider-key", + "num_retries": 0, + }, + }, + { + "model_name": CHAOS_MODELS["responses"], + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_base": f"{provider_url}/v1", + "api_key": "synthetic-provider-key", + "num_retries": 0, + }, + }, + ], + "environment_variables": { + **base.get("environment_variables", {}), + "OPENAI_API_BASE": provider_url, + "OPENAI_API_KEY": "synthetic-provider-key", + "ANTHROPIC_API_BASE": provider_url, + "ANTHROPIC_API_KEY": "synthetic-provider-key", + }, + }, + ) + + +@dataclass(frozen=True, slots=True) +class ChaosRig: + gateway: Gateway + provider: Wire + directory: Path + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[ChaosRig]: + with ( + httpx.Client( + base_url=os.environ["INTEGRATION_PROXY_URL"], + timeout=30, + trust_env=False, + ) as root_client, + wire_server(_provider) as provider, + ): + root_gateway: Final = Gateway( + root_client, + os.environ.get("INTEGRATION_MASTER_KEY", "sk-integration-master"), + os.environ["INTEGRATION_UPSTREAM_URL"], + ) + yield ChaosRig(root_gateway, provider, tmp_path_factory.mktemp("stream-scope-chaos")) + + +@dataclass(frozen=True, slots=True) +class CallPlan: + marker: str + call_id: str + endpoint: ChaosEndpoint + streamed: bool + + +def _plans(prefix: str, count: int) -> tuple[CallPlan, ...]: + return tuple( + CallPlan( + f"audit-{prefix}-{index}-{uuid.uuid4().hex}", + f"{prefix}-{uuid.uuid4().hex}", + CHAOS_ENDPOINTS[index % len(CHAOS_ENDPOINTS)], + index % 2 == 1, + ) + for index in range(count) + ) + + +def _request(gateway: Gateway, plan: CallPlan, rails: Sequence[str]) -> httpx.Response: + guardrail_field: Final[dict[str, JsonValue]] = {"guardrails": list(rails)} if rails else {} + match plan.endpoint: + case "chat": + return gateway.request( + "POST", + "/v1/chat/completions", + { + "model": CHAOS_MODELS["chat"], + "messages": [{"role": "user", "content": plan.marker}], + **guardrail_field, + **({"stream": True} if plan.streamed else {}), + }, + headers={"x-litellm-call-id": plan.call_id}, + ) + case "messages": + return gateway.request( + "POST", + "/v1/messages", + { + "model": CHAOS_MODELS["messages"], + "max_tokens": 32, + "messages": [{"role": "user", "content": plan.marker}], + **guardrail_field, + **({"stream": True} if plan.streamed else {}), + }, + headers={"x-litellm-call-id": plan.call_id}, + ) + case "responses": + return gateway.request( + "POST", + "/v1/responses", + { + "model": CHAOS_MODELS["responses"], + "input": plan.marker, + **guardrail_field, + **({"stream": True} if plan.streamed else {}), + }, + headers={"x-litellm-call-id": plan.call_id}, + ) + raise AssertionError(plan.endpoint) + + +def _rows_for_marker(rows: Sequence[Request], marker: str) -> tuple[Request, ...]: + return tuple(request for request in rows if marker.encode() in request.body) + + +def _rail_scans(rows: Sequence[Request], rail_name: str, marker: str) -> tuple[Request, ...]: + return tuple( + request for request in rows if request.target.startswith(f"/{rail_name}/") and marker.encode() in request.body + ) + + +def _spend_rows(call_id: str, database_url: str | None = None) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT request_id, litellm_call_id FROM "LiteLLM_SpendLogs" WHERE request_id=%s OR litellm_call_id=%s', + (call_id, call_id), + database_url=database_url, + ) + + +def _one_spend_row( + call_id: str, + database_url: str | None = None, + *, + seconds: float = 70, +) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: _spend_rows(call_id, database_url), + lambda values: len(values) >= 1, + seconds=seconds, + ) + assert len(rows) == 1, (call_id, rows) + return rows[0] + + +def _expected_in_scope(plan: CallPlan, sink_a_name: str, sink_b_name: str) -> tuple[str, str]: + return (sink_a_name, sink_b_name) if plan.streamed else (sink_b_name, sink_a_name) + + +def _assert_successful_calls( + plans: Sequence[CallPlan], + responses: Sequence[httpx.Response], + provider_rows: Sequence[Request], + sink_a_rows: Sequence[Request], + sink_b_rows: Sequence[Request], + sink_a_name: str, + sink_b_name: str, +) -> None: + for plan, response in zip(plans, responses): + if response.status_code != 200: + continue + assert plan.marker in response.text, (plan, response.text) + provider_match: Final = _rows_for_marker(provider_rows, plan.marker) + assert len(provider_match) == 1, (plan, provider_match) + in_sink, out_sink = _expected_in_scope(plan, sink_a_name, sink_b_name) + in_rows: Final = _rail_scans( + sink_a_rows if in_sink == sink_a_name else sink_b_rows, + in_sink, + plan.marker, + ) + out_rows: Final = _rail_scans( + sink_a_rows if out_sink == sink_a_name else sink_b_rows, + out_sink, + plan.marker, + ) + assert len(in_rows) == 1 and len(out_rows) == 0, (plan, in_rows, out_rows) + spend: Final = _one_spend_row(plan.call_id) + assert plan.call_id in (spend.get("request_id"), spend.get("litellm_call_id")), (plan, spend) + + +def _call_wave(gateway: Gateway, plans: Sequence[CallPlan], rails: Sequence[str]) -> tuple[httpx.Response, ...]: + with ThreadPoolExecutor(max_workers=20) as pool: + futures: Final[tuple[Future[httpx.Response], ...]] = tuple( + pool.submit(_request, gateway, plan, rails) for plan in plans + ) + return tuple(future.result() for future in futures) + + +@contextmanager +def _owned_proxy( + rig: ChaosRig, + directory: Path, + rails: Sequence[dict[str, JsonValue]], + *, + workers: int = 1, +) -> Iterator[OwnedProxy]: + config_path: Final = directory / f"chaos-{uuid.uuid4().hex}.yaml" + config_path.write_text(yaml.safe_dump(_config(rig.provider.url, rails))) + with owned_proxy_process(rig.gateway, directory, {}, config=config_path, workers=workers) as owned: + yield owned + + +def _free_port() -> int: + with socket.socket() as reserve: + reserve.bind(("127.0.0.1", 0)) + return int(reserve.getsockname()[1]) + + +def _assert_down_wave( + plans: Sequence[CallPlan], + responses: Sequence[httpx.Response], + provider_rows: Sequence[Request], + sink_a_rows: Sequence[Request], + sink_b_rows: Sequence[Request], + sink_a_name: str, + sink_b_name: str, +) -> None: + for plan, response in zip(plans, responses): + if plan.streamed: + assert response.status_code >= 500 and response.content, (plan, response.status_code, response.text) + assert _rows_for_marker(provider_rows, plan.marker) == (), (plan, provider_rows) + assert _rail_scans(sink_b_rows, sink_b_name, plan.marker) == (), (plan, sink_b_name) + continue + assert response.status_code == 200 and plan.marker in response.text, (plan, response.status_code, response.text) + assert len(_rows_for_marker(provider_rows, plan.marker)) == 1, (plan, provider_rows) + assert _rail_scans(sink_a_rows, sink_a_name, plan.marker) == (), (plan, sink_a_name) + assert len(_rail_scans(sink_b_rows, sink_b_name, plan.marker)) == 1, (plan, sink_b_name) + spend: Final = _one_spend_row(plan.call_id) + assert plan.call_id in (spend.get("request_id"), spend.get("litellm_call_id")), (plan, spend) + + +def test_h1_sink_outage_keeps_scope_isolated_through_recovery(rig: ChaosRig, tmp_path: Path) -> None: + port_a: Final = _free_port() + started: Final = Event() + unavailable: Final = Event() + release: Final = Event() + + def gated_sink(request: Request) -> Reply: + started.set() + assert release.wait(timeout=45), "sink outage gate was not released" + if unavailable.is_set(): + return Reply(status=503, body=_json({"error": "synthetic sink outage"})) + return _sink(request) + + with wire_server(_sink) as sink_b, ExitStack() as sink_a_stack: + sink_a: Final = sink_a_stack.enter_context(wire_server(gated_sink, port=port_a)) + name_a: Final = f"h1-stream-{uuid.uuid4().hex}" + name_b: Final = f"h1-non-stream-{uuid.uuid4().hex}" + rails: Final = ( + _rail(name_a, sink_a.url, "streaming"), + _rail(name_b, sink_b.url, "non_streaming"), + ) + with _owned_proxy(rig, tmp_path, rails) as owned, ThreadPoolExecutor(max_workers=30) as pool: + outage_plans: Final = _plans("h1-burst", 30) + futures: Final[tuple[Future[httpx.Response], ...]] = tuple( + pool.submit(_request, owned.gateway, plan, (name_a, name_b)) for plan in outage_plans + ) + try: + assert started.wait(timeout=30), "streaming rail did not reach sink A" + unavailable.set() + finally: + release.set() + sink_a_stack.close() + outage_responses: Final = tuple(future.result(timeout=70) for future in futures) + outage_provider: Final = rig.provider.drain() + outage_a: Final = sink_a.drain() + outage_b: Final = sink_b.drain() + _assert_down_wave(outage_plans, outage_responses, outage_provider, outage_a, outage_b, name_a, name_b) + + with wire_server(_sink, port=port_a) as recovered_sink_a: + recovery_plans: Final = _plans("h1-recovery", 20) + recovery_responses: Final = _call_wave(owned.gateway, recovery_plans, (name_a, name_b)) + recovery_provider: Final = rig.provider.drain() + recovery_a: Final = recovered_sink_a.drain() + recovery_b: Final = sink_b.drain() + assert tuple(response.status_code for response in recovery_responses) == (200,) * 20, recovery_responses + _assert_successful_calls( + recovery_plans, + recovery_responses, + recovery_provider, + recovery_a, + recovery_b, + name_a, + name_b, + ) + + +def test_h2_stream_sink_gate_does_not_block_out_of_scope_calls(rig: ChaosRig, tmp_path: Path) -> None: + started: Final = Event() + release: Final = Event() + blocked_marker: Final = f"audit-h2-stream-{uuid.uuid4().hex}" + + def gated_sink(request: Request) -> Reply: + if blocked_marker.encode() in request.body: + started.set() + assert release.wait(timeout=45), "stream sink gate was not released" + return _sink(request) + + with wire_server(gated_sink) as sink_a, wire_server(_sink) as sink_b: + name_a: Final = f"h2-stream-{uuid.uuid4().hex}" + name_b: Final = f"h2-non-stream-{uuid.uuid4().hex}" + rails: Final = (_rail(name_a, sink_a.url, "streaming"), _rail(name_b, sink_b.url, "non_streaming")) + with _owned_proxy(rig, tmp_path, rails) as owned, ThreadPoolExecutor(max_workers=2) as pool: + streaming_plan: Final = CallPlan(blocked_marker, f"h2-stream-{uuid.uuid4().hex}", "chat", True) + non_streaming_plan: Final = CallPlan( + f"audit-h2-non-stream-{uuid.uuid4().hex}", + f"h2-non-stream-{uuid.uuid4().hex}", + "messages", + False, + ) + streaming_future: Final = pool.submit(_request, owned.gateway, streaming_plan, (name_a, name_b)) + assert started.wait(timeout=30), "stream request did not reach the gated sink" + try: + non_streaming_future: Final = pool.submit( + _request, + owned.gateway, + non_streaming_plan, + (name_a, name_b), + ) + non_streaming_response: Final = non_streaming_future.result(timeout=15) + assert non_streaming_response.status_code == 200, non_streaming_response.text + assert non_streaming_plan.marker in non_streaming_response.text, non_streaming_response.text + finally: + release.set() + streaming_response: Final = streaming_future.result(timeout=30) + assert streaming_response.status_code == 200, streaming_response.text + assert streaming_plan.marker in streaming_response.text, streaming_response.text + provider_rows: Final = rig.provider.drain() + sink_a_rows: Final = sink_a.drain() + sink_b_rows: Final = sink_b.drain() + _assert_successful_calls( + (streaming_plan, non_streaming_plan), + (streaming_response, non_streaming_response), + provider_rows, + sink_a_rows, + sink_b_rows, + name_a, + name_b, + ) + + +def _worker_processes(owned: OwnedProxy) -> tuple[psutil.Process, ...]: + return tuple( + child for child in psutil.Process(owned.process.pid).children() if "spawn_main" in _process_command(child) + ) + + +def _process_command(process: psutil.Process) -> str: + try: + return " ".join(process.cmdline()) + except psutil.Error: + return "" + + +def _safe_response(future: Future[httpx.Response]) -> httpx.Response | None: + try: + return future.result(timeout=70) + except (httpx.HTTPError, TimeoutError): + return None + + +def test_h3_worker_kill_mid_burst_keeps_remaining_worker_serving(rig: ChaosRig, tmp_path: Path) -> None: + plans: Final = _plans("h3-burst", 30) + gate_markers: Final = frozenset(plan.marker for plan in plans if plan.streamed) + started: Final = Event() + release: Final = Event() + + def gated_sink(request: Request) -> Reply: + if any(marker.encode() in request.body for marker in gate_markers): + started.set() + assert release.wait(timeout=60), "worker-kill sink gate was not released" + return _sink(request) + + with wire_server(gated_sink) as sink_a, wire_server(_sink) as sink_b: + name_a: Final = f"h3-stream-{uuid.uuid4().hex}" + name_b: Final = f"h3-non-stream-{uuid.uuid4().hex}" + rails: Final = ( + _rail(name_a, sink_a.url, "streaming", default_on=True), + _rail(name_b, sink_b.url, "non_streaming", default_on=True), + ) + with _owned_proxy(rig, tmp_path, rails, workers=2) as owned, ThreadPoolExecutor(max_workers=30) as pool: + workers: Final = _worker_processes(owned) + assert len(workers) == 2, tuple(worker.pid for worker in workers) + futures: Final[tuple[Future[httpx.Response], ...]] = tuple( + pool.submit(_request, owned.gateway, plan, ()) for plan in plans + ) + assert started.wait(timeout=30), "stream requests did not reach the owned sink" + victim: Final = workers[0] + survivor: Final = workers[1] + survivor_plan: Final = CallPlan( + f"audit-h3-survivor-{uuid.uuid4().hex}", + f"h3-survivor-{uuid.uuid4().hex}", + "messages", + False, + ) + try: + victim.send_signal(signal.SIGKILL) + survivor_response: Final = _request(owned.gateway, survivor_plan, ()) + assert survivor_response.status_code == 200 and survivor_plan.marker in survivor_response.text, ( + survivor_response.status_code, + survivor_response.text, + ) + assert survivor.is_running(), survivor.pid + finally: + release.set() + responses: Final = tuple(_safe_response(future) for future in futures) + assert owned.gateway.request("GET", "/health/liveliness").status_code == 200 + successful: Final = tuple( + (plan, response) + for plan, response in zip(plans, responses) + if response is not None and response.status_code == 200 + ) + provider_rows: Final = rig.provider.drain() + sink_a_rows: Final = sink_a.drain() + sink_b_rows: Final = sink_b.drain() + successful_plans: Final = (*tuple(plan for plan, _ in successful), survivor_plan) + successful_responses: Final = (*tuple(response for _, response in successful), survivor_response) + _assert_successful_calls( + successful_plans, + successful_responses, + provider_rows, + sink_a_rows, + sink_b_rows, + name_a, + name_b, + ) + + +def _create_stored_rail( + gateway: Gateway, + name: str, + sink_url: str, + stream_scope: Literal["streaming"] | None = "streaming", +) -> str: + response: Final = gateway.request( + "POST", + "/guardrails", + { + "guardrail": { + "guardrail_name": name, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": False, + "api_base": f"{sink_url}/{name}", + "api_key": "synthetic-chaos-key", + **({"stream_scope": stream_scope} if stream_scope is not None else {}), + }, + } + }, + ) + assert response.status_code == 200, response.text + identity: Final = JSON_OBJECT.validate_json(response.content).get("guardrail_id") + assert isinstance(identity, str), response.text + return identity + + +def _assert_one_scope_wave( + plans: Sequence[CallPlan], + responses: Sequence[httpx.Response], + provider_rows: Sequence[Request], + sink_rows: Sequence[Request], + sink_name: str, + *, + database_url: str | None = None, +) -> None: + for plan, response in zip(plans, responses): + assert response.status_code == 200 and plan.marker in response.text, (plan, response.status_code, response.text) + provider_match: Final = _rows_for_marker(provider_rows, plan.marker) + sink_match: Final = _rail_scans(sink_rows, sink_name, plan.marker) + assert len(provider_match) == 1, (plan, provider_match) + assert len(sink_match) == int(plan.streamed), (plan, sink_match) + spend: Final = _one_spend_row(plan.call_id, database_url) + assert plan.call_id in (spend.get("request_id"), spend.get("litellm_call_id")), (plan, spend) + + +def test_h4_stored_scope_survives_owned_proxy_restart(rig: ChaosRig, tmp_path: Path) -> None: + with wire_server(_sink) as sink: + name: Final = f"h4-stored-{uuid.uuid4().hex}" + identity: Final = _create_stored_rail(rig.gateway, name, sink.url) + try: + with ExitStack() as first_stack: + first: Final = first_stack.enter_context(_owned_proxy(rig, tmp_path, ())) + first_plans: Final = _plans("h4-before", 20) + first_responses: Final = _call_wave(first.gateway, first_plans, (name,)) + first_provider: Final = rig.provider.drain() + first_sink: Final = sink.drain() + assert tuple(response.status_code for response in first_responses) == (200,) * 20, first_responses + _assert_one_scope_wave(first_plans, first_responses, first_provider, first_sink, name) + first_stack.close() + with _owned_proxy(rig, tmp_path, ()) as restarted: + recovery_plans: Final = _plans("h4-after", 20) + recovery_responses: Final = _call_wave(restarted.gateway, recovery_plans, (name,)) + recovery_provider: Final = rig.provider.drain() + recovery_sink: Final = sink.drain() + assert tuple(response.status_code for response in recovery_responses) == (200,) * 20, ( + recovery_responses, + ) + _assert_one_scope_wave( + recovery_plans, + recovery_responses, + recovery_provider, + recovery_sink, + name, + ) + finally: + deleted: Final = rig.gateway.request("DELETE", f"/guardrails/{identity}") + assert deleted.status_code == 200, deleted.text + + +@contextmanager +def _owned_postgres(directory: Path) -> Iterator[PostgresCluster]: + docker: Final = shutil.which("docker") + assert docker is not None, "Docker CLI is required for the H5 PostgreSQL outage test" + docker_info: Final = subprocess.run( + [docker, "info", "--format", "{{.ServerVersion}}"], + capture_output=True, + text=True, + check=False, + ) + assert docker_info.returncode == 0, docker_info.stderr + directory.mkdir(parents=True, exist_ok=True) + port: Final = _free_port() + container_name: Final = f"litellm-stream-scope-h5-{uuid.uuid4().hex}" + password: Final = uuid.uuid4().hex + created: Final = subprocess.run( + [ + docker, + "create", + "--name", + container_name, + "--env", + "POSTGRES_USER=postgres", + "--env", + "POSTGRES_PASSWORD", + "--env", + "POSTGRES_DB=postgres", + "--publish", + f"127.0.0.1:{port}:5432/tcp", + POSTGRES_IMAGE, + ], + env=os.environ | {"POSTGRES_PASSWORD": password}, + capture_output=True, + text=True, + check=False, + ) + assert created.returncode == 0, created.stderr + cluster: Final = PostgresCluster( + f"postgresql://postgres:{password}@127.0.0.1:{port}/postgres?sslmode=disable", + container_name, + docker, + directory / "postgres.log", + ) + try: + started: Final = _start_postgres(cluster) + assert started.returncode == 0, started.stderr + assert eventually(lambda: _postgres_is_ready(cluster), bool, seconds=70) + yield cluster + finally: + logs: Final = subprocess.run( + [docker, "logs", container_name], + capture_output=True, + text=True, + check=False, + ) + cluster.log_path.write_text(logs.stdout + logs.stderr) + removed: Final = subprocess.run( + [docker, "rm", "-f", container_name], + capture_output=True, + text=True, + check=False, + ) + assert removed.returncode == 0, removed.stderr + remaining: Final = subprocess.run( + [docker, "ps", "--all", "--quiet", "--filter", f"name={container_name}"], + capture_output=True, + text=True, + check=False, + ) + assert remaining.returncode == 0, remaining.stderr + assert not remaining.stdout.strip(), remaining.stdout + + +@dataclass(frozen=True, slots=True) +class PostgresCluster: + database_url: str + container_name: str + docker: str + log_path: Path + + +def _start_postgres(cluster: PostgresCluster) -> subprocess.CompletedProcess[str]: + return subprocess.run( + [cluster.docker, "start", cluster.container_name], + capture_output=True, + text=True, + check=False, + ) + + +def _postgres_is_ready(cluster: PostgresCluster) -> bool: + readiness: Final = subprocess.run( + [ + cluster.docker, + "exec", + cluster.container_name, + "pg_isready", + "-h", + "127.0.0.1", + "-p", + "5432", + "-d", + "postgres", + "-U", + "postgres", + ], + capture_output=True, + text=True, + check=False, + ) + return readiness.returncode == 0 + + +def _postgres_is_running(cluster: PostgresCluster) -> bool: + state: Final = subprocess.run( + [cluster.docker, "inspect", "--format", "{{.State.Running}}", cluster.container_name], + capture_output=True, + text=True, + check=False, + ) + return state.returncode == 0 and state.stdout.strip() == "true" + + +@pytest.mark.timeout(180) +def test_h5_stored_scope_survives_owned_postgres_outage_and_recovers_once( + rig: ChaosRig, + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + postgres_directory: Final = Path(os.environ["INTEGRATION_RESULTS_DIR"]) / f"owned-postgres-{uuid.uuid4().hex}" + with _owned_postgres(postgres_directory) as database, wire_server(_sink) as sink: + monkeypatch.setenv("INTEGRATION_PROXY_DATABASE_URL", database.database_url) + name: Final = f"h5-stored-{uuid.uuid4().hex}" + with _owned_proxy(rig, tmp_path, ()) as registrar: + _create_stored_rail(registrar.gateway, name, sink.url) + with _owned_proxy(rig, tmp_path, ()) as owned: + try: + preflight_plans: Final = ( + CallPlan( + f"audit-h5-preflight-{uuid.uuid4().hex}", + f"h5-preflight-{uuid.uuid4().hex}", + "chat", + True, + ), + ) + preflight_responses: Final = _call_wave(owned.gateway, preflight_plans, (name,)) + preflight_provider: Final = rig.provider.drain() + preflight_sink: Final = sink.drain() + assert tuple(response.status_code for response in preflight_responses) == (200,), (preflight_responses,) + _assert_one_scope_wave( + preflight_plans, + preflight_responses, + preflight_provider, + preflight_sink, + name, + database_url=database.database_url, + ) + outage_result: Final = subprocess.run( + [database.docker, "stop", database.container_name], + capture_output=True, + text=True, + check=False, + ) + assert outage_result.returncode == 0, outage_result.stderr + outage_plans: Final = _plans("h5-outage", 20) + outage_responses: Final = _call_wave(owned.gateway, outage_plans, (name,)) + outage_provider: Final = rig.provider.drain() + outage_sink: Final = sink.drain() + for plan, response in zip(outage_plans, outage_responses): + assert response.status_code == 200 and plan.marker in response.text, ( + plan, + response.status_code, + response.text, + ) + assert len(_rows_for_marker(outage_provider, plan.marker)) == 1, (plan, outage_provider) + sink_match: Final = _rail_scans(outage_sink, name, plan.marker) + assert len(sink_match) == int(plan.streamed), (plan, sink_match) + + recovered_database: Final = _start_postgres(database) + assert recovered_database.returncode == 0, recovered_database.stderr + postgres_ready: Final = eventually(lambda: _postgres_is_ready(database), bool, seconds=70) + assert postgres_ready + readiness: Final = eventually( + lambda: owned.gateway.client.get("/health/readiness"), + lambda response: response.status_code == 200 and response.json().get("db") == "connected", + seconds=70, + ) + assert readiness.json().get("db") == "connected", readiness.text + recovery_plans: Final = _plans("h5-recovery", 20) + recovery_responses: Final = _call_wave(owned.gateway, recovery_plans, (name,)) + recovery_provider: Final = rig.provider.drain() + recovery_sink: Final = sink.drain() + assert tuple(response.status_code for response in recovery_responses) == (200,) * 20, recovery_responses + _assert_one_scope_wave( + recovery_plans, + recovery_responses, + recovery_provider, + recovery_sink, + name, + database_url=database.database_url, + ) + finally: + if not _postgres_is_running(database): + restarted: Final = _start_postgres(database) + assert restarted.returncode == 0, restarted.stderr + assert eventually(lambda: _postgres_is_ready(database), bool, seconds=70) + + +@pytest.mark.timeout(360) +def test_h6_spend_rows_for_requests_served_during_postgres_restart_land_once( + rig: ChaosRig, + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + pytest.skip("BUG: LIT-9050 spend rows for requests served during a Postgres restart are dropped") + postgres_directory: Final = Path(os.environ["INTEGRATION_RESULTS_DIR"]) / f"owned-postgres-{uuid.uuid4().hex}" + with _owned_postgres(postgres_directory) as database, wire_server(_sink) as sink: + monkeypatch.setenv("INTEGRATION_PROXY_DATABASE_URL", database.database_url) + name: Final = f"h6-stored-{uuid.uuid4().hex}" + plans: Final = _plans("h6-postgres-restart", 20) + outage_markers: Final = frozenset(plan.marker for plan in plans[:10]) + arrivals: Final = Barrier(len(plans) + 1) + during_outage: Final = Event() + after_restart: Final = Event() + + def _gated_provider(request: Request) -> Reply: + if request.method == "GET" and request.target.partition("?")[0] == "/v1/models": + return _provider(request) + body: Final = JSON_OBJECT.validate_json(request.body) + marker: Final = _marker(body) + arrivals.wait(timeout=70) + gate: Final = during_outage if marker in outage_markers else after_restart + assert gate.wait(timeout=70), marker + return _provider(request) + + with wire_server(_gated_provider) as provider: + h6_rig: Final = ChaosRig(rig.gateway, provider, rig.directory) + with _owned_proxy(h6_rig, tmp_path, ()) as registrar: + _create_stored_rail(registrar.gateway, name, sink.url, stream_scope=None) + with _owned_proxy(h6_rig, tmp_path, ()) as owned, ThreadPoolExecutor(max_workers=len(plans)) as pool: + futures: Final = tuple( + pool.submit(_request, owned.gateway, plan, (name,)) for plan in plans + ) + try: + arrivals.wait(timeout=70) + stopped: Final = subprocess.run( + [database.docker, "stop", database.container_name], + capture_output=True, + text=True, + check=False, + ) + assert stopped.returncode == 0, stopped.stderr + during_outage.set() + outage_responses: Final = tuple(future.result(timeout=70) for future in futures[:10]) + for plan, response in zip(plans[:10], outage_responses): + assert response.status_code == 200 and plan.marker in response.text, ( + plan, + response.status_code, + response.text, + ) + + restarted: Final = _start_postgres(database) + assert restarted.returncode == 0, restarted.stderr + postgres_ready: Final = eventually(lambda: _postgres_is_ready(database), bool, seconds=70) + assert postgres_ready + readiness: Final = eventually( + lambda: owned.gateway.client.get("/health/readiness"), + lambda response: response.status_code == 200 + and JSON_OBJECT.validate_python(cast(object, response.json())).get("db") == "connected", + seconds=70, + ) + readiness_body: Final = JSON_OBJECT.validate_python(cast(object, readiness.json())) + assert readiness_body.get("db") == "connected", readiness.text + after_restart.set() + responses: Final = tuple(future.result(timeout=70) for future in futures) + finally: + during_outage.set() + after_restart.set() + if not _postgres_is_running(database): + recovered: Final = _start_postgres(database) + assert recovered.returncode == 0, recovered.stderr + assert eventually(lambda: _postgres_is_ready(database), bool, seconds=70) + + provider_rows: Final = provider.drain() + sink_rows: Final = sink.drain() + for plan, response in zip(plans, responses): + assert response.status_code == 200 and plan.marker in response.text, ( + plan, + response.status_code, + response.text, + ) + assert len(_rows_for_marker(provider_rows, plan.marker)) == 1, (plan, provider_rows) + assert len(_rail_scans(sink_rows, name, plan.marker)) == 1, (plan, sink_rows) + + spend_rows: Final = eventually( + lambda: tuple(_spend_rows(plan.call_id, database.database_url) for plan in plans), + lambda values: all(len(rows) == 1 for rows in values), + seconds=70, + return_last_on_timeout=True, + ) + counts: Final = tuple(len(rows) for rows in spend_rows) + missing_ids: Final = tuple( + plan.call_id for plan, rows in zip(plans, spend_rows) if not rows + ) + duplicate_ids: Final = tuple( + plan.call_id for plan, rows in zip(plans, spend_rows) if len(rows) > 1 + ) + assert counts == (1,) * len(plans), { + "missing_ids": missing_ids, + "duplicate_ids": duplicate_ids, + "counts": counts, + } diff --git a/tests/integration/observability/test_guardrail_stream_scope_matrix.py b/tests/integration/observability/test_guardrail_stream_scope_matrix.py new file mode 100644 index 00000000000..b5c66b202cc --- /dev/null +++ b/tests/integration/observability/test_guardrail_stream_scope_matrix.py @@ -0,0 +1,2076 @@ +from __future__ import annotations + +import asyncio +import contextlib +import json +import os +import socket +import socketserver +import ssl +import threading +import uuid +from asyncio import run, wait_for +from collections.abc import Generator, Iterable, Iterator, Mapping, Sequence +from contextlib import contextmanager +from dataclasses import dataclass +from itertools import chain +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final, Literal, TypeAlias, cast + +import anthropic +import httpx +import openai +import pytest +import websockets +import yaml +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import rsa +from integration._support.client import Gateway, JsonValue, Scenario, eventually +from integration._support.database import read_rows +from integration._support.mcp import McpCaller, echo_tool, register_mcp, scripted_peer +from integration._support.process import owned_proxy_process +from integration._support.tls import server_context, write_self_signed_cert +from integration._support.upstream import delete_scenario, register_scenario +from integration._support.wire import Reply, Request, Wire, wire_server +from integration.cost_calculation.cost_tracking_case import RealtimeResponse +from pydantic import TypeAdapter +from websockets.asyncio.server import ServerConnection, serve + +Endpoint: TypeAlias = Literal["chat", "messages", "responses"] +Mode: TypeAlias = Literal["pre_call", "during_call", "post_call", "logging_only"] +Scope: TypeAlias = Literal["streaming", "non_streaming"] +MatrixScope: TypeAlias = Scope | None +LoggingPhase: TypeAlias = Literal["request", "response"] +Endpoints: Final[tuple[Endpoint, ...]] = ("chat", "messages", "responses") +Modes: Final[tuple[Mode, ...]] = ("pre_call", "during_call", "post_call", "logging_only") +Scopes: Final[tuple[MatrixScope, ...]] = ("streaming", "non_streaming", None) +PROVIDER_TEXT: Final = "provider stream_scope control" +GUARDRAIL_PATH: Final = "/beta/litellm_basic_guardrail_api" +F1_BLOCKED_WORD: Final = "streamscope-mcp-blocked-word" +F3_STREAM_BLOCKED_WORD: Final = "matrix-f3-stream-block" +F3_NON_STREAM_BLOCKED_WORD: Final = "matrix-f3-non-stream-block" +JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +LOGGING_PHASE: Final = TypeAdapter(LoggingPhase) + + +def _json(value: object) -> bytes: + return json.dumps(value, separators=(",", ":")).encode() + + +def _strings(value: JsonValue) -> tuple[str, ...]: + if isinstance(value, str): + return (value,) + if isinstance(value, list): + return tuple(chain.from_iterable(_strings(item) for item in value)) + if isinstance(value, dict): + return tuple(chain.from_iterable(_strings(item) for item in value.values())) + return () + + +def _marker(body: Mapping[str, JsonValue]) -> str: + return next((text for text in _strings(dict(body)) if text.startswith("audit-")), "audit-provider") + + +def _sse(events: Iterable[Mapping[str, JsonValue]]) -> tuple[bytes, ...]: + return tuple(f"data: {json.dumps(event, separators=(',', ':'))}\n\n".encode() for event in events) + ( + b"data: [DONE]\n\n", + ) + + +def _messages_stream(message: Mapping[str, JsonValue]) -> tuple[bytes, ...]: + content: Final = cast(list[JsonValue], message["content"]) + text: Final = cast(dict[str, JsonValue], content[0])["text"] + return ( + f"event: message_start\ndata: {json.dumps({**message, 'content': [], 'stop_reason': None, 'usage': {'input_tokens': 11, 'output_tokens': 0}})}\n\n".encode(), + b'event: content_block_start\ndata: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}\n\n', + f"event: content_block_delta\ndata: {json.dumps({'type': 'content_block_delta', 'index': 0, 'delta': {'type': 'text_delta', 'text': text}})}\n\n".encode(), + b'event: content_block_stop\ndata: {"type":"content_block_stop","index":0}\n\n', + f"event: message_delta\ndata: {json.dumps({'type': 'message_delta', 'delta': {'stop_reason': 'end_turn', 'stop_sequence': None}, 'usage': {'output_tokens': 4}})}\n\n".encode(), + b'event: message_stop\ndata: {"type":"message_stop"}\n\n', + ) + + +def _responses_stream( + response: Mapping[str, JsonValue], output: Mapping[str, JsonValue], marker: str +) -> tuple[bytes, ...]: + events: Final[tuple[dict[str, JsonValue], ...]] = ( + {"type": "response.created", "response": {**response, "status": "in_progress", "output": []}}, + {"type": "response.in_progress", "response": {**response, "status": "in_progress", "output": []}}, + {"type": "response.output_item.added", "item": output, "output_index": 0}, + { + "type": "response.content_part.added", + "item_id": f"msg-{marker}", + "output_index": 0, + "content_index": 0, + "part": {"type": "output_text", "text": "", "annotations": []}, + }, + { + "type": "response.output_text.delta", + "item_id": f"msg-{marker}", + "output_index": 0, + "content_index": 0, + "delta": f"{PROVIDER_TEXT} {marker}", + }, + { + "type": "response.output_text.done", + "item_id": f"msg-{marker}", + "output_index": 0, + "content_index": 0, + "text": f"{PROVIDER_TEXT} {marker}", + }, + { + "type": "response.content_part.done", + "item_id": f"msg-{marker}", + "output_index": 0, + "content_index": 0, + "part": cast(list[JsonValue], output["content"])[0], + }, + {"type": "response.output_item.done", "item": output, "output_index": 0}, + {"type": "response.completed", "response": response}, + ) + return tuple( + f"event: {event['type']}\ndata: {json.dumps({**event, 'sequence_number': index}, separators=(',', ':'))}\n\n".encode() + for index, event in enumerate(events) + ) + + +def _provider_reply(target: str, body: Mapping[str, JsonValue]) -> Reply: + marker: Final = _marker(body) + streamed: Final = bool(body.get("stream")) + if target == "/v1/chat/completions": + if streamed: + common: Final = {"id": f"chatcmpl-{marker}", "object": "chat.completion.chunk", "created": 1} + return Reply( + content_type="text/event-stream", + chunks=_sse( + ( + { + **common, + "choices": [{"index": 0, "delta": {"role": "assistant"}, "finish_reason": None}], + }, + { + **common, + "choices": [ + {"index": 0, "delta": {"content": f"{PROVIDER_TEXT} {marker}"}, "finish_reason": None} + ], + }, + {**common, "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]}, + ) + ), + ) + return Reply( + body=_json( + { + "id": f"chatcmpl-{marker}", + "object": "chat.completion", + "created": 1, + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": f"{PROVIDER_TEXT} {marker}"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15}, + } + ) + ) + if target == "/v1/messages": + message: Final = { + "id": f"msg-{marker}", + "type": "message", + "role": "assistant", + "model": "synthetic-anthropic-model", + "content": [{"type": "text", "text": f"{PROVIDER_TEXT} {marker}"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 11, "output_tokens": 4}, + } + return ( + Reply(content_type="text/event-stream", chunks=_messages_stream(message)) + if streamed + else Reply(body=_json(message)) + ) + if target == "/v1/responses": + output: Final = { + "type": "message", + "id": f"msg-{marker}", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": f"{PROVIDER_TEXT} {marker}", "annotations": []}], + } + response: Final = { + "id": f"resp-{marker}", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [output], + "usage": {"input_tokens": 11, "output_tokens": 4, "total_tokens": 15}, + } + return ( + Reply(content_type="text/event-stream", chunks=_responses_stream(response, output, marker)) + if streamed + else Reply(body=_json(response)) + ) + if target.endswith(":generateContent"): + return Reply( + body=_json( + {"candidates": [{"content": {"role": "model", "parts": [{"text": f"{PROVIDER_TEXT} {marker}"}]}}]} + ) + ) + if target.endswith(":streamGenerateContent"): + return Reply( + content_type="text/event-stream", + chunks=( + f"data: {_json({'candidates': [{'content': {'role': 'model', 'parts': [{'text': f'{PROVIDER_TEXT} {marker}'}]}}]}).decode()}\n\n".encode(), + ), + ) + return Reply(status=404, body=_json({"error": {"message": f"unexpected target {target}"}})) + + +def _provider(request: Request) -> Reply: + if not request.body: + return Reply(status=400, body=_json({"error": "empty request body"})) + body: Final = JSON_OBJECT.validate_json(request.body) + if request.target.startswith("/passthrough"): + marker: Final = _marker(body) + if bool(body.get("stream")): + return Reply( + content_type="text/event-stream", + chunks=_sse(({"text": f"{PROVIDER_TEXT} {marker}"},)), + ) + return Reply(body=_json({"received": body})) + if "audit-g3-401-" in request.body.decode(): + return Reply(status=401, body=_json({"error": {"message": "synthetic provider unauthorized"}})) + target: Final = request.target.split("?", 1)[0] + provider_target: Final = "/v1/messages" if target.endswith("/anthropic/v1/messages") else target + return _provider_reply(provider_target, body) + + +def _sink(request: Request) -> Reply: + assert request.target.endswith(GUARDRAIL_PATH), request.target + JSON_OBJECT.validate_json(request.body) + if "/g1-" in request.target: + return Reply(status=500, body=_json({"error": "synthetic sink failure"})) + if "/g2-" in request.target or "/f2-block-" in request.target: + return Reply(body=_json({"action": "BLOCKED", "blocked_reason": "synthetic policy block"})) + return Reply(body=_json({"action": "NONE"})) + + +def _rail( + name: str, + sink: Wire, + *, + mode: str | list[str], + scope: str | Mapping[str, str] | None = None, + default_on: bool = False, +) -> dict[str, JsonValue]: + return { + "guardrail_name": name, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": mode, + "default_on": default_on, + **({"stream_scope": dict(scope)} if isinstance(scope, Mapping) else {}), + **({"stream_scope": scope} if isinstance(scope, str) else {}), + "api_base": f"{sink.url}/{name}", + "api_key": "synthetic-guardrail-key", + }, + } + + +def _realtime_filter_rail(name: str, scope: Scope, keyword: str) -> dict[str, JsonValue]: + return { + "guardrail_name": name, + "litellm_params": { + "guardrail": "litellm_content_filter", + "mode": "realtime_input_transcription", + "default_on": True, + "stream_scope": scope, + "blocked_words": [{"keyword": keyword, "action": "BLOCK"}], + }, + } + + +def _mcp_filter_rail(name: str, scope: MatrixScope) -> dict[str, JsonValue]: + return { + "guardrail_name": name, + "litellm_params": { + "guardrail": "litellm_content_filter", + "mode": "pre_mcp_call", + "default_on": False, + **({"stream_scope": scope} if scope is not None else {}), + "blocked_words": [{"keyword": F1_BLOCKED_WORD, "action": "BLOCK"}], + }, + } + + +def _mode_scope_pairs() -> tuple[tuple[Mode, MatrixScope], ...]: + return tuple(chain.from_iterable(tuple((mode, scope) for scope in Scopes) for mode in Modes)) + + +def _rail_names() -> Mapping[str, str]: + return MappingProxyType( + { + **{f"a_{mode}_{scope or 'unset'}": f"a_{mode}_{scope or 'unset'}" for mode, scope in _mode_scope_pairs()}, + "c1_default": "c1_default", + "c2_key": "c2_key", + "c3_team": "c3_team", + "c4_modes": "c4_modes", + "c5_omitted": "c5_omitted", + "c6_both": "c6_both", + "c6_unset": "c6_unset", + "e_stream": "e_stream", + "e_non_stream": "e_non_stream", + "f1_stream": "f1_stream", + "f1_non_stream": "f1_non_stream", + "f1_unset": "f1_unset", + "f2_stream": "f2_stream", + "f2_non_stream": "f2-block-non-stream", + "f3_stream": "f3_stream", + "f3_non_stream": "f3_non_stream", + "g1_failure": "g1-failure", + "g2_block": "g2-block", + "g3_provider": "g3-provider", + "z_logging_only_barrier": "z_logging_only_barrier", + } + ) + + +def _configured_rails(names: Mapping[str, str], sink: Wire) -> tuple[dict[str, JsonValue], ...]: + return tuple( + _rail(names[f"a_{mode}_{scope or 'unset'}"], sink, mode=mode, scope=scope) + for mode, scope in _mode_scope_pairs() + ) + ( + _rail(names["c1_default"], sink, mode="post_call", scope="streaming"), + _rail(names["c2_key"], sink, mode="post_call", scope="streaming"), + _rail(names["c3_team"], sink, mode="post_call", scope="streaming"), + _rail( + names["c4_modes"], + sink, + mode=["pre_call", "post_call"], + scope={"pre_call": "non_streaming", "post_call": "streaming"}, + ), + _rail( + names["c5_omitted"], + sink, + mode=["pre_call", "post_call"], + scope={"post_call": "streaming"}, + ), + _rail(names["c6_both"], sink, mode="post_call", scope="both"), + _rail(names["c6_unset"], sink, mode="post_call"), + _rail(names["e_stream"], sink, mode="pre_call", scope="streaming"), + _rail(names["e_non_stream"], sink, mode="pre_call", scope="non_streaming"), + _mcp_filter_rail(names["f1_stream"], "streaming"), + _mcp_filter_rail(names["f1_non_stream"], "non_streaming"), + _mcp_filter_rail(names["f1_unset"], None), + _rail(names["f2_stream"], sink, mode="post_call", scope="streaming"), + _rail(names["f2_non_stream"], sink, mode="post_call", scope="non_streaming"), + _realtime_filter_rail(names["f3_stream"], "streaming", F3_STREAM_BLOCKED_WORD), + _realtime_filter_rail(names["f3_non_stream"], "non_streaming", F3_NON_STREAM_BLOCKED_WORD), + _rail(names["g1_failure"], sink, mode="pre_call", scope="streaming"), + _rail(names["g2_block"], sink, mode="pre_call", scope="streaming"), + _rail(names["g3_provider"], sink, mode="pre_call", scope="both"), + _rail(names["z_logging_only_barrier"], sink, mode="logging_only"), + ) + + +def _realtime_transcription_response(transcript: str) -> RealtimeResponse: + return RealtimeResponse( + content_type="application/x-realtime", + events=( + { + "type": "conversation.item.input_audio_transcription.completed", + "event_id": "evt_$REQUEST_ID", + "item_id": "item_$REQUEST_ID", + "content_index": 0, + "transcript": transcript, + }, + { + "type": "response.done", + "event_id": "evt_$REQUEST_ID", + "response": { + "id": "resp_$REQUEST_ID", + "object": "realtime.response", + "status": "completed", + "output": [], + "usage": { + "total_tokens": 0, + "input_tokens": 0, + "output_tokens": 0, + "input_token_details": { + "text_tokens": 0, + "audio_tokens": 0, + "cached_tokens": 0, + "cached_tokens_details": {"text_tokens": 0, "audio_tokens": 0}, + }, + "output_token_details": {"text_tokens": 0, "audio_tokens": 0}, + }, + }, + }, + ), + ) + + +async def _collect_realtime_events(websocket: websockets.ClientConnection) -> tuple[dict[str, JsonValue], ...]: + event: Final = JSON_OBJECT.validate_json(await websocket.recv()) + if event.get("type") == "response.done": + return (event,) + return (event, *await _collect_realtime_events(websocket)) + + +async def _realtime_transcription_events( + url: str, + key: str, + model: str, +) -> tuple[dict[str, JsonValue], ...]: + websocket_url: Final = ( + f"{url.replace('http://', 'ws://').replace('https://', 'wss://').rstrip('/')}/v1/realtime?model={model}" + ) + async with websockets.connect(websocket_url, additional_headers={"Authorization": f"Bearer {key}"}) as websocket: + session: Final = JSON_OBJECT.validate_json(await websocket.recv()) + assert session.get("type") == "session.created", session + await websocket.send(json.dumps({"type": "response.create"})) + return await _collect_realtime_events(websocket) + + +@dataclass(frozen=True, slots=True) +class MatrixRig: + candidate: Gateway + scenario: Scenario + models: Mapping[str, str] + rails: Mapping[str, str] + key: str + provider: Wire + sink: Wire + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[MatrixRig]: + with httpx.Client( + base_url=os.environ["INTEGRATION_PROXY_URL"], + timeout=30, + trust_env=False, + ) as root_client: + root_gateway: Final = Gateway( + root_client, + os.environ.get("INTEGRATION_MASTER_KEY", "sk-integration-master"), + os.environ["INTEGRATION_UPSTREAM_URL"], + ) + directory: Final = tmp_path_factory.mktemp("guardrail-stream-scope-matrix") + with wire_server(_provider) as provider, wire_server(_sink) as sink: + names: Final = _rail_names() + config_base: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + pass_through_paths: Final = ( + { + "path": "/pt", + "target": f"{provider.url}/passthrough", + "include_subpath": True, + "guardrails": {names["e_stream"]: None, names["e_non_stream"]: None}, + }, + { + "path": "/anthropic/v1/messages", + "target": f"{provider.url}/v1/messages", + "guardrails": {names["e_stream"]: None, names["e_non_stream"]: None}, + }, + { + "path": "/gemini/v1beta/models", + "target": "", + "include_subpath": True, + "guardrails": {names["e_stream"]: None, names["e_non_stream"]: None}, + }, + ) + config: Final = { + **config_base, + "guardrails": _configured_rails(names, sink), + "model_list": [ + { + "model_name": "matrix-pipeline-model", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_base": f"{provider.url}/v1", + "api_key": "synthetic-provider-key", + }, + } + ], + "environment_variables": { + **config_base.get("environment_variables", {}), + "ANTHROPIC_API_BASE": provider.url, + "ANTHROPIC_API_KEY": "synthetic-anthropic-key", + "GEMINI_API_BASE": provider.url, + "GEMINI_API_KEY": "synthetic-gemini-key", + }, + "general_settings": { + **config_base["general_settings"], + "pass_through_endpoints": pass_through_paths, + }, + "policies": { + "matrix-f2": { + "guardrails": {"add": [names["f2_stream"], names["f2_non_stream"]]}, + "pipeline": { + "mode": "post_call", + "steps": [ + {"guardrail": names["f2_stream"], "on_pass": "next", "on_fail": "block"}, + {"guardrail": names["f2_non_stream"], "on_pass": "next", "on_fail": "block"}, + ], + }, + } + }, + "policy_attachments": [{"policy": "matrix-f2", "models": ["matrix-pipeline-model"]}], + } + config_path: Final = directory / "stream-scope-matrix.yaml" + config_path.write_text(yaml.safe_dump(config)) + with owned_proxy_process(root_gateway, directory, {}, config=config_path, workers=1) as owned: + with owned.gateway.scenario() as scenario: + models: Final = MappingProxyType( + { + "chat": scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{provider.url}/v1", + api_key="synthetic-provider-key", + ), + "messages": scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", + api_base=provider.url, + api_key="synthetic-provider-key", + ), + "responses": scenario.model( + model="openai/gpt-4o-mini", + api_base=f"{provider.url}/v1", + api_key="synthetic-provider-key", + ), + } + ) + key: Final = scenario.key(guardrails=[names["c2_key"]]) + yield MatrixRig(owned.gateway, scenario, models, names, key, provider, sink) + + +def _cell_body( + rig: MatrixRig, + endpoint: Endpoint, + marker: str, + streamed: bool, + rail: str | None, +) -> tuple[str, dict[str, JsonValue]]: + selected: Final = [] if rail is None else [rail] + if endpoint == "chat": + body: Final = { + "model": rig.models["chat"], + "messages": [{"role": "user", "content": marker}], + "guardrails": selected, + **({"stream": True} if streamed else {}), + } + return "/v1/chat/completions", body + if endpoint == "messages": + body = { + "model": rig.models["messages"], + "max_tokens": 32, + "messages": [{"role": "user", "content": marker}], + "guardrails": selected, + **({"stream": True} if streamed else {}), + } + return "/v1/messages", body + body = { + "model": rig.models["responses"], + "input": marker, + "guardrails": selected, + **({"stream": True} if streamed else {}), + } + return "/v1/responses", body + + +def _matching_requests(wire: Wire, marker: str) -> tuple[Request, ...]: + return tuple(request for request in wire.drain() if marker.encode() in request.body) + + +def _rail_scans(rows: Sequence[Request], rail_name: str, marker: str) -> tuple[Request, ...]: + return tuple( + request for request in rows if request.target.startswith(f"/{rail_name}/") and marker.encode() in request.body + ) + + +def _logging_phases(rows: Sequence[Request]) -> tuple[str, ...]: + return tuple( + sorted(LOGGING_PHASE.validate_python(JSON_OBJECT.validate_json(request.body)["input_type"]) for request in rows) + ) + + +@dataclass(slots=True) +class _SinkRowsAccumulator: + sink: Wire + rows: tuple[Request, ...] = () + + def drain(self) -> tuple[Request, ...]: + self.rows = (*self.rows, *self.sink.drain()) + return self.rows + + +def _logging_only_scans( + sink: Wire, + marker: str, + rail: str, + barrier: str, +) -> tuple[Request, ...]: + accumulator: Final = _SinkRowsAccumulator(sink) + eventually( + accumulator.drain, + lambda rows: bool(_rail_scans(rows, barrier, marker)), + seconds=70, + ) + return _rail_scans(accumulator.rows, rail, marker) + + +def _scope_scan_count(scope: MatrixScope, streamed: bool) -> int: + if scope is None or scope == "both": + return 1 + return int((scope == "streaming") == streamed) + + +def _expected_logging_phases(scope: MatrixScope, streamed: bool) -> tuple[str, ...]: + return ("request", "response") if _scope_scan_count(scope, streamed) else () + + +def _spend_row_for_call(call_id: str, content: bytes) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + "SELECT request_id, litellm_call_id, spend, prompt_tokens, completion_tokens, metadata " + 'FROM "LiteLLM_SpendLogs" WHERE request_id=%s OR litellm_call_id=%s', + (call_id, call_id), + ), + lambda values: len(values) >= 1, + seconds=70, + ) + assert len(rows) == 1, (call_id, rows, content) + assert call_id in (rows[0]["request_id"], rows[0]["litellm_call_id"]), content + return rows[0] + + +def _cell_rows( + rig: MatrixRig, + marker: str, + response: httpx.Response, + mode: Mode, + streamed: bool, + scope: MatrixScope, + call_id: str, + rail: str, +) -> tuple[tuple[Request, ...], tuple[Request, ...]]: + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows, response.text) + _spend_row_for_call(call_id, response.content) + sink_rows: Final = ( + _logging_only_scans( + rig.sink, + marker, + rail, + rig.rails["z_logging_only_barrier"], + ) + if mode == "logging_only" + else _rail_scans(rig.sink.drain(), rail, marker) + ) + return provider_rows, sink_rows + + +def _run_matrix_cell( + rig: MatrixRig, + endpoint: Endpoint, + streamed: bool, + mode: Mode, + scope: MatrixScope, +) -> None: + marker: Final = f"audit-a-{endpoint}-{int(streamed)}-{mode}-{scope or 'unset'}-{uuid.uuid4().hex}" + call_id: Final = f"matrix-a-{uuid.uuid4().hex}" + rail: Final = rig.rails[f"a_{mode}_{scope or 'unset'}"] + path, body = _cell_body(rig, endpoint, marker, streamed, rail) + response: Final = rig.candidate.request( + "POST", + path, + body, + headers={"x-litellm-call-id": call_id}, + ) + assert response.status_code == 200, response.text + assert PROVIDER_TEXT in response.text and marker in response.text, response.text + provider_rows, sink_rows = _cell_rows(rig, marker, response, mode, streamed, scope, call_id, rail) + assert len(provider_rows) == 1, (marker, provider_rows, response.text) + if mode == "logging_only": + assert _logging_phases(sink_rows) == _expected_logging_phases(scope, streamed), ( + marker, + sink_rows, + response.text, + ) + else: + expected_scans: Final = _scope_scan_count(scope, streamed) + assert len(sink_rows) == expected_scans, (marker, sink_rows, response.text) + + +@pytest.mark.parametrize("endpoint", Endpoints) +@pytest.mark.parametrize("streamed", (False, True), ids=("S0", "S1")) +@pytest.mark.parametrize("mode", Modes) +@pytest.mark.parametrize("scope", Scopes, ids=("streaming", "non_streaming", "unset")) +def test_a_stream_scope_matrix( + rig: MatrixRig, + endpoint: Endpoint, + streamed: bool, + mode: Mode, + scope: MatrixScope, +) -> None: + _run_matrix_cell(rig, endpoint, streamed, mode, scope) + + +SDK_CASES: Final[tuple[str, ...]] = ( + "openai-chat-sync", + "openai-chat-async", + "openai-responses-sync", + "openai-responses-async", + "anthropic-messages-sync", + "anthropic-messages-async", +) + + +def _openai_sdk_base_url(rig: MatrixRig) -> str: + return f"{str(rig.candidate.client.base_url).rstrip('/')}/v1" + + +def _anthropic_sdk_base_url(rig: MatrixRig) -> str: + return str(rig.candidate.client.base_url).rstrip("/") + + +def _verify_sdk_call(rig: MatrixRig, marker: str, call_id: str, text: str, streamed: bool) -> None: + assert marker in text and PROVIDER_TEXT in text, text + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows) + _spend_row_for_call(call_id, text.encode()) + sink_rows: Final = _rail_scans(rig.sink.drain(), rig.rails["c2_key"], marker) + assert len(sink_rows) == int(streamed), (marker, streamed, sink_rows) + + +def _run_sync_sdk(rig: MatrixRig, sdk: str, marker: str, call_id: str, streamed: bool) -> str: + if sdk == "openai-chat-sync": + with openai.OpenAI( + api_key=rig.key, + base_url=_openai_sdk_base_url(rig), + max_retries=0, + http_client=httpx.Client(trust_env=False, timeout=30), + ) as client: + response: Final = client.chat.completions.create( + model=rig.models["chat"], + messages=[{"role": "user", "content": marker}], + stream=streamed, + extra_headers={"x-litellm-call-id": call_id}, + ) + if streamed: + return "".join( + chunk.choices[0].delta.content or "" for chunk in response if chunk.choices[0].delta.content + ) + return cast(str, response.choices[0].message.content) + if sdk == "openai-responses-sync": + with openai.OpenAI( + api_key=rig.key, + base_url=_openai_sdk_base_url(rig), + max_retries=0, + http_client=httpx.Client(trust_env=False, timeout=30), + ) as client: + response = client.responses.create( + model=rig.models["responses"], + input=marker, + stream=streamed, + extra_headers={"x-litellm-call-id": call_id}, + ) + if streamed: + return "".join(event.delta for event in response if event.type == "response.output_text.delta") + return response.output_text + if sdk == "anthropic-messages-sync": + with anthropic.Anthropic( + api_key=rig.key, + base_url=_anthropic_sdk_base_url(rig), + max_retries=0, + http_client=httpx.Client(trust_env=False, timeout=30), + ) as client: + response = client.messages.create( + model=rig.models["messages"], + max_tokens=32, + messages=[{"role": "user", "content": marker}], + stream=streamed, + extra_headers={"x-litellm-call-id": call_id}, + ) + if streamed: + return "".join( + event.delta.text + for event in response + if event.type == "content_block_delta" and event.delta.type == "text_delta" + ) + return "".join(block.text for block in response.content if block.type == "text") + raise AssertionError(sdk) + + +async def _run_async_sdk(rig: MatrixRig, sdk: str, marker: str, call_id: str, streamed: bool) -> str: + if sdk == "openai-chat-async": + async with openai.AsyncOpenAI( + api_key=rig.key, + base_url=_openai_sdk_base_url(rig), + max_retries=0, + http_client=httpx.AsyncClient(trust_env=False, timeout=30), + ) as client: + response = await client.chat.completions.create( + model=rig.models["chat"], + messages=[{"role": "user", "content": marker}], + stream=streamed, + extra_headers={"x-litellm-call-id": call_id}, + ) + if streamed: + chunks: Final = [ + chunk.choices[0].delta.content or "" async for chunk in response if chunk.choices[0].delta.content + ] + return "".join(chunks) + return cast(str, response.choices[0].message.content) + if sdk == "openai-responses-async": + async with openai.AsyncOpenAI( + api_key=rig.key, + base_url=_openai_sdk_base_url(rig), + max_retries=0, + http_client=httpx.AsyncClient(trust_env=False, timeout=30), + ) as client: + response = await client.responses.create( + model=rig.models["responses"], + input=marker, + stream=streamed, + extra_headers={"x-litellm-call-id": call_id}, + ) + if streamed: + events: Final = [event.delta async for event in response if event.type == "response.output_text.delta"] + return "".join(events) + return response.output_text + if sdk == "anthropic-messages-async": + async with anthropic.AsyncAnthropic( + api_key=rig.key, + base_url=_anthropic_sdk_base_url(rig), + max_retries=0, + http_client=httpx.AsyncClient(trust_env=False, timeout=30), + ) as client: + response = await client.messages.create( + model=rig.models["messages"], + max_tokens=32, + messages=[{"role": "user", "content": marker}], + stream=streamed, + extra_headers={"x-litellm-call-id": call_id}, + ) + if streamed: + events: Final = [ + event.delta.text + async for event in response + if event.type == "content_block_delta" and event.delta.type == "text_delta" + ] + return "".join(events) + return "".join(block.text for block in response.content if block.type == "text") + raise AssertionError(sdk) + + +@pytest.mark.parametrize("sdk", SDK_CASES) +@pytest.mark.parametrize("streamed", (False, True), ids=("S0", "S1")) +def test_b_streaming_scope_classifies_sdk_streams(rig: MatrixRig, sdk: str, streamed: bool) -> None: + marker: Final = f"audit-b-{sdk}-{int(streamed)}-{uuid.uuid4().hex}" + call_id: Final = f"matrix-b-{uuid.uuid4().hex}" + if sdk.endswith("-async"): + text: Final = run(_run_async_sdk(rig, sdk, marker, call_id, streamed)) + else: + text = _run_sync_sdk(rig, sdk, marker, call_id, streamed) + _verify_sdk_call(rig, marker, call_id, text, streamed) + + +def _yaml_proxy_config( + rig: MatrixRig, + guardrail_name: str, + parameters: Mapping[str, JsonValue], +) -> dict[str, JsonValue]: + config_base: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + return cast( + dict[str, JsonValue], + { + **config_base, + "guardrails": [{"guardrail_name": guardrail_name, "litellm_params": dict(parameters)}], + "model_list": [ + { + "model_name": "scope-yaml-chat", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_base": f"{rig.provider.url}/v1", + "api_key": "synthetic-provider-key", + }, + } + ], + "environment_variables": { + **config_base.get("environment_variables", {}), + "OPENAI_API_BASE": rig.provider.url, + "OPENAI_API_KEY": "synthetic-provider-key", + }, + }, + ) + + +def _management_params( + name: str, + rig: MatrixRig, + *, + mode: str | list[str] = "pre_call", + scope: JsonValue = "streaming", + default_on: bool = False, +) -> dict[str, JsonValue]: + return { + "guardrail": "generic_guardrail_api", + "mode": mode, + "default_on": default_on, + "api_base": f"{rig.sink.url}/{name}", + "api_key": "synthetic-guardrail-key", + "stream_scope": scope, + } + + +def _raw_chat( + gateway: Gateway, + marker: str, + streamed: bool, + model: str, + *, + rails: Sequence[str] = (), + key: str | None = None, + call_id: str | None = None, +) -> httpx.Response: + return gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": marker}], + "guardrails": list(rails), + **({"stream": True} if streamed else {}), + }, + key=key, + headers={} if call_id is None else {"x-litellm-call-id": call_id}, + ) + + +def _assert_raw_call( + rig: MatrixRig, + marker: str, + response: httpx.Response, + streamed: bool, + expected_scans: int, + rail: str, + call_id: str | None = None, +) -> tuple[Request, ...]: + assert response.status_code == 200, response.text + assert PROVIDER_TEXT in response.text and marker in response.text, response.text + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows) + if call_id is not None: + _spend_row_for_call(call_id, response.content) + sink_rows: Final = _rail_scans(rig.sink.drain(), rail, marker) + assert len(sink_rows) == expected_scans, (marker, streamed, expected_scans, sink_rows) + return sink_rows + + +def test_c1_default_on_rail_respects_stream_scope(rig: MatrixRig, tmp_path: Path) -> None: + name: Final = f"c1-default-{uuid.uuid4().hex}" + parameters: Final = { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": True, + "stream_scope": "streaming", + "api_base": f"{rig.sink.url}/{name}", + "api_key": "synthetic-guardrail-key", + } + config_path: Final = tmp_path / "default-on-stream-scope.yaml" + config_path.write_text(yaml.safe_dump(_yaml_proxy_config(rig, name, parameters))) + with owned_proxy_process(rig.candidate, tmp_path, {}, config=config_path, workers=1) as owned: + for streamed in (False, True): + marker: Final = f"audit-c1-{int(streamed)}-{uuid.uuid4().hex}" + response: Final = _raw_chat(owned.gateway, marker, streamed, "scope-yaml-chat") + _assert_raw_call(rig, marker, response, streamed, int(streamed), name) + + +def test_c2_key_attached_rail_respects_stream_scope(rig: MatrixRig) -> None: + marker0: Final = f"audit-c2-0-{uuid.uuid4().hex}" + response0: Final = _raw_chat(rig.candidate, marker0, False, rig.models["chat"], key=rig.key) + _assert_raw_call(rig, marker0, response0, False, 0, rig.rails["c2_key"]) + marker1: Final = f"audit-c2-1-{uuid.uuid4().hex}" + response1: Final = _raw_chat(rig.candidate, marker1, True, rig.models["chat"], key=rig.key) + _assert_raw_call(rig, marker1, response1, True, 1, rig.rails["c2_key"]) + + +def test_c3_team_attached_rail_respects_stream_scope(rig: MatrixRig) -> None: + team: Final = rig.scenario.team(guardrails=[rig.rails["c3_team"]]) + key: Final = rig.scenario.key(team_id=team) + marker0: Final = f"audit-c3-0-{uuid.uuid4().hex}" + response0: Final = _raw_chat(rig.candidate, marker0, False, rig.models["chat"], key=key) + _assert_raw_call(rig, marker0, response0, False, 0, rig.rails["c3_team"]) + marker1: Final = f"audit-c3-1-{uuid.uuid4().hex}" + response1: Final = _raw_chat(rig.candidate, marker1, True, rig.models["chat"], key=key) + _assert_raw_call(rig, marker1, response1, True, 1, rig.rails["c3_team"]) + + +@pytest.mark.parametrize("streamed", (False, True), ids=("S0", "S1")) +def test_c4_per_mode_scope_map_selects_each_mode(rig: MatrixRig, streamed: bool) -> None: + marker: Final = f"audit-c4-{int(streamed)}-{uuid.uuid4().hex}" + response: Final = _raw_chat( + rig.candidate, + marker, + streamed, + rig.models["chat"], + rails=(rig.rails["c4_modes"],), + ) + assert response.status_code == 200, response.text + assert PROVIDER_TEXT in response.text and marker in response.text, response.text + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows) + sink_rows: Final = _rail_scans(rig.sink.drain(), rig.rails["c4_modes"], marker) + request_bodies: Final = tuple(JSON_OBJECT.validate_json(row.body) for row in sink_rows) + request_texts: Final = tuple(chain.from_iterable(_strings(body.get("texts", [])) for body in request_bodies)) + assert len(sink_rows) == 1, (marker, streamed, sink_rows) + assert any(marker in text for text in request_texts), (marker, request_texts) + assert (any(PROVIDER_TEXT in text for text in request_texts)) == streamed, ( + marker, + streamed, + request_texts, + ) + + +@pytest.mark.parametrize("streamed", (False, True), ids=("S0", "S1")) +def test_c5_mode_omitted_from_scope_map_means_both(rig: MatrixRig, streamed: bool) -> None: + marker: Final = f"audit-c5-{int(streamed)}-{uuid.uuid4().hex}" + response: Final = _raw_chat( + rig.candidate, + marker, + streamed, + rig.models["chat"], + rails=(rig.rails["c5_omitted"],), + ) + sink_rows: Final = _assert_raw_call( + rig, + marker, + response, + streamed, + 1 + int(streamed), + rig.rails["c5_omitted"], + ) + request_bodies: Final = tuple(JSON_OBJECT.validate_json(row.body) for row in sink_rows) + request_texts: Final = tuple(chain.from_iterable(_strings(body.get("texts", [])) for body in request_bodies)) + assert sum(marker in text and PROVIDER_TEXT not in text for text in request_texts) == 1, (marker, request_texts) + assert sum(PROVIDER_TEXT in text for text in request_texts) == int(streamed), (marker, request_texts) + + +@pytest.mark.parametrize("streamed", (False, True), ids=("S0", "S1")) +def test_c6_explicit_both_matches_unset_scope(rig: MatrixRig, streamed: bool) -> None: + marker: Final = f"audit-c6-{int(streamed)}-{uuid.uuid4().hex}" + response: Final = _raw_chat( + rig.candidate, + marker, + streamed, + rig.models["chat"], + rails=(rig.rails["c6_both"], rig.rails["c6_unset"]), + ) + assert response.status_code == 200, response.text + assert PROVIDER_TEXT in response.text and marker in response.text, response.text + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows) + sink_rows: Final = rig.sink.drain() + both_rows: Final = _rail_scans(sink_rows, rig.rails["c6_both"], marker) + unset_rows: Final = _rail_scans(sink_rows, rig.rails["c6_unset"], marker) + assert (len(both_rows), len(unset_rows)) == (1, 1), (marker, streamed, both_rows, unset_rows) + + +@pytest.mark.parametrize( + ("case", "scope"), + ( + ("missing", None), + ("null", None), + ("empty", ""), + ("invalid-scalar", "sometimes"), + ("invalid-map", {"unknown_mode": "streaming"}), + ), + ids=("missing", "null", "empty", "invalid-scalar", "invalid-map"), +) +def test_c7_yaml_unset_and_invalid_scope_values_are_tolerated( + rig: MatrixRig, + tmp_path: Path, + case: str, + scope: JsonValue | None, +) -> None: + name: Final = f"c7-yaml-{case}-{uuid.uuid4().hex}" + parameters: Final = { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": True, + "api_base": f"{rig.sink.url}/{name}", + "api_key": "synthetic-guardrail-key", + **({} if case == "missing" else {"stream_scope": scope}), + } + config_path: Final = tmp_path / f"{case}-stream-scope.yaml" + config_path.write_text(yaml.safe_dump(_yaml_proxy_config(rig, name, parameters))) + with owned_proxy_process(rig.candidate, tmp_path, {}, config=config_path, workers=1) as owned: + for streamed in (False, True): + marker: Final = f"audit-c7-{case}-{int(streamed)}-{uuid.uuid4().hex}" + response: Final = _raw_chat(owned.gateway, marker, streamed, "scope-yaml-chat") + _assert_raw_call(rig, marker, response, streamed, 1, name) + if case == "empty": + assert "Ignoring invalid stored stream_scope value of type str" in owned.log.read_text() + + +def test_c8_identical_requests_have_one_scan_and_spend_each(rig: MatrixRig) -> None: + marker: Final = f"audit-c8-identical-{uuid.uuid4().hex}" + call_ids: Final = tuple(f"matrix-c8-{uuid.uuid4().hex}" for _ in range(3)) + responses: Final = tuple( + _raw_chat( + rig.candidate, + marker, + True, + rig.models["chat"], + rails=(rig.rails["a_post_call_streaming"],), + call_id=call_id, + ) + for call_id in call_ids + ) + assert tuple(response.status_code for response in responses) == (200, 200, 200), tuple( + response.text for response in responses + ) + for call_id, response in zip(call_ids, responses): + _spend_row_for_call(call_id, response.content) + sink_rows: Final = _rail_scans(rig.sink.drain(), rig.rails["a_post_call_streaming"], marker) + assert len(sink_rows) == 3, (marker, sink_rows) + guardrail_call_ids: Final = tuple( + cast(str, JSON_OBJECT.validate_json(row.body).get("litellm_call_id")) for row in sink_rows + ) + assert set(guardrail_call_ids) == set(call_ids), (call_ids, guardrail_call_ids) + + +def _create_management_rail(rig: MatrixRig, name: str, params: Mapping[str, JsonValue]) -> str: + created: Final = rig.candidate.request( + "POST", + "/guardrails", + {"guardrail": {"guardrail_name": name, "litellm_params": dict(params)}}, + ) + assert created.status_code == 200, created.text + payload: Final = JSON_OBJECT.validate_json(created.content) + identity: Final = payload.get("guardrail_id") + assert isinstance(identity, str), payload + return identity + + +def _management_observation( + rig: MatrixRig, + name: str, + streamed: bool, +) -> tuple[int, int, int]: + marker: Final = f"audit-d-management-{int(streamed)}-{uuid.uuid4().hex}" + call_id: Final = f"matrix-d1-{uuid.uuid4().hex}" + response: Final = _raw_chat( + rig.candidate, + marker, + streamed, + rig.models["chat"], + rails=(name,), + call_id=call_id, + ) + provider_rows: Final = _matching_requests(rig.provider, marker) + sink_rows: Final = _rail_scans(rig.sink.drain(), name, marker) + if response.status_code == 200 and len(provider_rows) == 1: + assert marker in response.text, response.text + _spend_row_for_call(call_id, response.content) + return response.status_code, len(provider_rows), len(sink_rows) + + +def _eventually_management_scope( + rig: MatrixRig, + name: str, + expected: tuple[int, int], +) -> tuple[tuple[int, int, int], tuple[int, int, int]]: + return eventually( + lambda: ( + _management_observation(rig, name, False), + _management_observation(rig, name, True), + ), + lambda values: ( + (values[0][0], values[0][2]) == (200, expected[0]) + and (values[1][0], values[1][2]) == (200, expected[1]) + and values[0][1] == values[1][1] == 1 + ), + seconds=30, + ) + + +def test_d1_management_create_read_and_runtime_scope(rig: MatrixRig) -> None: + name: Final = f"d1-management-{uuid.uuid4().hex}" + params: Final = _management_params(name, rig) + identity: Final = _create_management_rail(rig, name, params) + try: + info: Final = rig.candidate.request("GET", f"/guardrails/{identity}/info") + listing: Final = rig.candidate.request("GET", "/v2/guardrails/list") + info_payload: Final = JSON_OBJECT.validate_json(info.content) + list_payload: Final = JSON_OBJECT.validate_json(listing.content) + list_rows: Final = list_payload.get("guardrails") + listed: Final = ( + tuple(row for row in list_rows if isinstance(row, dict) and row.get("guardrail_id") == identity) + if isinstance(list_rows, list) + else () + ) + info_params: Final = info_payload.get("litellm_params") + list_params: Final = listed[0].get("litellm_params") if listed else None + assert ( + info.status_code == 200 + and listing.status_code == 200 + and isinstance(info_params, dict) + and info_params.get("stream_scope") == "streaming" + and len(listed) == 1 + and isinstance(list_params, dict) + and list_params.get("stream_scope") == "streaming" + ), (info.status_code, info.text, listing.status_code, listing.text) + observed: Final = _eventually_management_scope(rig, name, (0, 1)) + assert observed[0][2] == 0 and observed[1][2] == 1, observed + finally: + deleted: Final = rig.candidate.request("DELETE", f"/guardrails/{identity}") + assert deleted.status_code == 200, deleted.text + + +def test_d2_patch_then_put_updates_runtime_scope(rig: MatrixRig) -> None: + name: Final = f"d2-management-{uuid.uuid4().hex}" + identity: Final = _create_management_rail(rig, name, _management_params(name, rig)) + try: + patched: Final = rig.candidate.request( + "PATCH", + f"/guardrails/{identity}", + {"litellm_params": {"stream_scope": "non_streaming"}}, + ) + assert patched.status_code == 200, patched.text + after_patch: Final = _eventually_management_scope(rig, name, (1, 0)) + assert after_patch[0][2] == 1 and after_patch[1][2] == 0, after_patch + params: Final = _management_params(name, rig) + updated: Final = rig.candidate.request( + "PUT", + f"/guardrails/{identity}", + {"guardrail": {"guardrail_name": name, "litellm_params": params}}, + ) + assert updated.status_code == 200, updated.text + after_put: Final = _eventually_management_scope(rig, name, (0, 1)) + assert after_put[0][2] == 0 and after_put[1][2] == 1, after_put + finally: + deleted: Final = rig.candidate.request("DELETE", f"/guardrails/{identity}") + assert deleted.status_code == 200, deleted.text + + +HOSTILE_SCOPE_VALUES: Final[tuple[JsonValue, ...]] = ( + 1, + [], + "", + "x" * 5000, + {"unknown_mode": "streaming"}, + {"pre_call": 1}, +) + + +@pytest.mark.parametrize("operation", ("POST", "PUT", "PATCH")) +@pytest.mark.parametrize( + "scope", HOSTILE_SCOPE_VALUES, ids=("integer", "list", "empty", "long", "unknown-key", "wrong-value") +) +@pytest.mark.parametrize("repeat", (1, 2), ids=("first", "second")) +def test_d3_management_rejects_invalid_scope_values( + rig: MatrixRig, + operation: str, + scope: JsonValue, + repeat: int, +) -> None: + name: Final = f"d3-management-{operation.lower()}-{repeat}-{uuid.uuid4().hex}" + identity: Final = _create_management_rail(rig, name, _management_params(name, rig)) if operation != "POST" else None + payload: Final = ( + {"litellm_params": {"stream_scope": scope}} + if operation == "PATCH" + else { + "guardrail": { + "guardrail_name": name, + "litellm_params": _management_params(name, rig, scope=scope), + } + } + ) + path: Final = "/guardrails" if operation == "POST" else f"/guardrails/{identity}" + try: + response: Final = rig.candidate.request(operation, path, payload) + if operation == "POST" and response.status_code == 200: + created_payload: Final = JSON_OBJECT.validate_json(response.content) + created_identity: Final = created_payload.get("guardrail_id") + if isinstance(created_identity, str): + rig.candidate.request("DELETE", f"/guardrails/{created_identity}") + assert response.status_code == 422 and "stream_scope" in response.text, ( + operation, + scope, + repeat, + response.status_code, + response.text, + ) + if operation == "POST": + stored: Final = read_rows( + 'SELECT guardrail_id FROM "LiteLLM_GuardrailsTable" WHERE guardrail_name=%s', + (name,), + ) + assert stored == [], (operation, scope, stored) + else: + existing: Final = rig.candidate.request("GET", f"/guardrails/{identity}/info") + existing_payload: Final = JSON_OBJECT.validate_json(existing.content) + existing_params: Final = existing_payload.get("litellm_params") + assert ( + existing.status_code == 200 + and isinstance(existing_params, dict) + and existing_params.get("stream_scope") == "streaming" + ), (operation, scope, existing.status_code, existing.text) + finally: + if isinstance(identity, str): + deleted: Final = rig.candidate.request("DELETE", f"/guardrails/{identity}") + assert deleted.status_code == 200, deleted.text + + +def test_d3_management_normalizes_uppercase_scope(rig: MatrixRig) -> None: + name: Final = f"d3-uppercase-{uuid.uuid4().hex}" + identity: Final = _create_management_rail(rig, name, _management_params(name, rig, scope="STREAMING")) + try: + info: Final = rig.candidate.request("GET", f"/guardrails/{identity}/info") + payload: Final = JSON_OBJECT.validate_json(info.content) + parameters: Final = payload.get("litellm_params") + assert ( + info.status_code == 200 and isinstance(parameters, dict) and parameters.get("stream_scope") == "streaming" + ), (info.status_code, info.text) + finally: + deleted: Final = rig.candidate.request("DELETE", f"/guardrails/{identity}") + assert deleted.status_code == 200, deleted.text + + +def test_d4_unauthenticated_management_create_is_rejected(rig: MatrixRig) -> None: + name: Final = f"d4-unauthenticated-{uuid.uuid4().hex}" + response: Final = rig.candidate.client.request( + "POST", + "/guardrails", + json={"guardrail": {"guardrail_name": name, "litellm_params": _management_params(name, rig)}}, + ) + assert response.status_code == 401, response.text + + +def _scoped_rows( + rows: Sequence[Request], + marker: str, + streaming_name: str, + non_streaming_name: str, +) -> tuple[tuple[Request, ...], tuple[Request, ...]]: + return ( + _rail_scans(rows, streaming_name, marker), + _rail_scans(rows, non_streaming_name, marker), + ) + + +HOSTILE_CLASSIFICATION_VALUES: Final[tuple[JsonValue, ...]] = ( + True, + 1, + [], + "", + "litellm-server-streaming", + "x" * 5000, +) + + +@pytest.mark.parametrize( + "value", + HOSTILE_CLASSIFICATION_VALUES, + ids=("true", "integer", "list", "empty", "marker", "long"), +) +def test_ea_is_streaming_request_body_does_not_change_chat_classification( + rig: MatrixRig, + value: JsonValue, +) -> None: + marker: Final = f"audit-ea-{uuid.uuid4().hex}" + response: Final = rig.candidate.request( + "POST", + "/v1/chat/completions", + { + "model": rig.models["chat"], + "messages": [{"role": "user", "content": marker}], + "guardrails": [rig.rails["e_stream"], rig.rails["e_non_stream"]], + "is_streaming_request": value, + }, + ) + assert response.status_code == 200, response.text + assert PROVIDER_TEXT in response.text and marker in response.text, response.text + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows) + provider_body: Final = JSON_OBJECT.validate_json(provider_rows[0].body) + assert provider_body.get("is_streaming_request") == value, provider_rows[0].body.decode() + sink_rows: Final = rig.sink.drain() + streaming_rows, non_streaming_rows = _scoped_rows( + sink_rows, + marker, + rig.rails["e_stream"], + rig.rails["e_non_stream"], + ) + assert (len(streaming_rows), len(non_streaming_rows)) == (0, 1), (marker, streaming_rows, non_streaming_rows) + + +@pytest.mark.parametrize("endpoint", ("chat", "passthrough"), ids=("chat", "configured-pass-through")) +@pytest.mark.parametrize("value", (True, "litellm-server-streaming"), ids=("true", "marker")) +def test_eb_namespaced_caller_body_field_cannot_flip_classification( + rig: MatrixRig, + endpoint: str, + value: JsonValue, +) -> None: + marker: Final = f"audit-eb-{endpoint}-{uuid.uuid4().hex}" + body: Final = ( + { + "model": rig.models["chat"], + "messages": [{"role": "user", "content": marker}], + "guardrails": [rig.rails["e_stream"], rig.rails["e_non_stream"]], + "litellm_server_streaming_classification": value, + } + if endpoint == "chat" + else { + "marker": marker, + "litellm_server_streaming_classification": value, + } + ) + response: Final = rig.candidate.request( + "POST", + "/v1/chat/completions" if endpoint == "chat" else "/pt", + body, + ) + assert response.status_code == 200 and marker in response.text, response.text + actual_streamed: Final = response.headers.get("content-type", "").lower().startswith("text/event-stream") + assert not actual_streamed, (endpoint, response.headers, response.text) + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows) + provider_body: Final = JSON_OBJECT.validate_json(provider_rows[0].body) + if endpoint == "passthrough" and value == "litellm-server-streaming": + assert "litellm_server_streaming_classification" not in provider_body, provider_rows[0].body.decode() + else: + assert provider_body.get("litellm_server_streaming_classification") == value, provider_rows[0].body.decode() + sink_rows: Final = rig.sink.drain() + streaming_rows, non_streaming_rows = _scoped_rows( + sink_rows, + marker, + rig.rails["e_stream"], + rig.rails["e_non_stream"], + ) + assert (len(streaming_rows), len(non_streaming_rows)) == (0, 1), (marker, streaming_rows, non_streaming_rows) + + +HOSTILE_STREAM_VALUES: Final[tuple[JsonValue, ...]] = ("true", 1, [], "") + + +@pytest.mark.parametrize("value", HOSTILE_STREAM_VALUES, ids=("string", "integer", "list", "empty")) +def test_ec_hostile_stream_values_follow_observed_response_shape(rig: MatrixRig, value: JsonValue) -> None: + marker: Final = f"audit-ec-{uuid.uuid4().hex}" + response: Final = rig.candidate.request( + "POST", + "/v1/chat/completions", + { + "model": rig.models["chat"], + "messages": [{"role": "user", "content": marker}], + "guardrails": [rig.rails["e_stream"], rig.rails["e_non_stream"]], + "stream": value, + }, + ) + assert response.status_code == 200, response.text + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows) + observed_stream: Final = response.headers.get("content-type", "").startswith("text/event-stream") + assert marker in response.text, response.text + sink_rows: Final = rig.sink.drain() + streaming_rows, non_streaming_rows = _scoped_rows( + sink_rows, + marker, + rig.rails["e_stream"], + rig.rails["e_non_stream"], + ) + assert (len(streaming_rows), len(non_streaming_rows)) == ( + int(observed_stream), + int(not observed_stream), + ), (marker, value, response.headers, streaming_rows, non_streaming_rows) + + +@pytest.mark.parametrize("streamed", (False, True), ids=("stream-absent", "stream-true")) +def test_ed_configured_passthrough_forwards_caller_flag_and_uses_route_body_stream( + rig: MatrixRig, + streamed: bool, +) -> None: + marker: Final = f"audit-ed-{int(streamed)}-{uuid.uuid4().hex}" + body: Final = { + "marker": marker, + "is_streaming_request": True, + **({"stream": True} if streamed else {}), + } + response: Final = rig.candidate.request("POST", "/pt", body) + assert response.status_code == 200 and marker in response.text, response.text + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows) + provider_body: Final = JSON_OBJECT.validate_json(provider_rows[0].body) + assert provider_body == body, (body, provider_body) + + +@pytest.mark.parametrize( + ("route", "body", "streamed"), + ( + ( + "/anthropic/v1/messages", + { + "model": "synthetic-anthropic-model", + "max_tokens": 32, + "messages": [{"role": "user", "content": "MARKER"}], + }, + False, + ), + ( + "/anthropic/v1/messages", + { + "model": "synthetic-anthropic-model", + "max_tokens": 32, + "messages": [{"role": "user", "content": "MARKER"}], + "stream": True, + }, + True, + ), + ( + "/gemini/v1beta/models/audit-model:generateContent", + {"contents": [{"parts": [{"text": "MARKER"}]}]}, + False, + ), + ( + "/gemini/v1beta/models/audit-model:streamGenerateContent?alt=sse", + {"contents": [{"parts": [{"text": "MARKER"}]}]}, + True, + ), + ), + ids=("anthropic-S0", "anthropic-S1", "gemini-generate", "gemini-stream"), +) +def test_eh_provider_passthrough_routes_classify_effective_streaming( + rig: MatrixRig, + route: str, + body: dict[str, JsonValue], + streamed: bool, +) -> None: + marker: Final = f"audit-eh-{uuid.uuid4().hex}" + request_body: Final = JSON_OBJECT.validate_python(json.loads(json.dumps(body).replace("MARKER", marker))) + team: Final = ( + rig.scenario.team( + metadata={ + "allowed_passthrough_routes": [ + "/gemini/v1beta/models/audit-model:generateContent", + "/gemini/v1beta/models/audit-model:streamGenerateContent", + ] + } + ) + if route.startswith("/gemini/") + else None + ) + key: Final = ( + rig.scenario.key(team_id=team, guardrails=[rig.rails["e_stream"], rig.rails["e_non_stream"]]) + if team is not None + else rig.scenario.key(guardrails=[rig.rails["e_stream"], rig.rails["e_non_stream"]]) + ) + headers: Final = {"x-goog-api-key": key} if route.startswith("/gemini/") else None + response: Final = rig.candidate.request( + "POST", + route, + request_body, + headers=headers, + ) + assert response.status_code == 200 and marker in response.text, response.text + actual_streamed: Final = response.headers.get("content-type", "").lower().startswith("text/event-stream") + assert actual_streamed is streamed, (route, streamed, response.headers, response.text) + provider_rows: Final = _matching_requests(rig.provider, marker) + assert len(provider_rows) == 1, (marker, provider_rows) + sink_rows: Final = rig.sink.drain() + streaming_rows, non_streaming_rows = _scoped_rows( + sink_rows, + marker, + rig.rails["e_stream"], + rig.rails["e_non_stream"], + ) + assert (len(streaming_rows), len(non_streaming_rows)) == ((1, 0) if streamed else (0, 1)), ( + marker, + route, + streaming_rows, + non_streaming_rows, + ) + + +def test_ei_unauthenticated_scoped_request_has_no_upstream_or_guardrail_call(rig: MatrixRig) -> None: + marker: Final = f"audit-ei-{uuid.uuid4().hex}" + response: Final = rig.candidate.client.request( + "POST", + "/v1/chat/completions", + json={ + "model": rig.models["chat"], + "messages": [{"role": "user", "content": marker}], + "stream": True, + "guardrails": [rig.rails["e_stream"]], + }, + ) + assert response.status_code == 401, response.text + assert _matching_requests(rig.provider, marker) == () + assert _rail_scans(rig.sink.drain(), rig.rails["e_stream"], marker) == () + + +@pytest.mark.parametrize( + ("rail_key", "blocked"), + (("f1_stream", False), ("f1_non_stream", True), ("f1_unset", True)), + ids=("streaming", "non_streaming", "unset"), +) +def test_f1_mcp_content_filter_respects_stream_scope( + rig: MatrixRig, + rail_key: Literal["f1_stream", "f1_non_stream", "f1_unset"], + blocked: bool, +) -> None: + with scripted_peer(echo_tool("echo")) as peer: + alias: Final = f"streamscope{uuid.uuid4().hex}" + identity: Final = register_mcp(rig.scenario, peer, alias) + key: Final = rig.scenario.key( + object_permission={"mcp_servers": [identity]}, + guardrails=[rig.rails[rail_key]], + ) + marker: Final = f"audit-f1-{uuid.uuid4().hex}" + caller: Final = McpCaller(rig.candidate, key, "rest", alias) + outcome: Final = caller.call( + f"{alias}-echo", + {"text": f"{F1_BLOCKED_WORD} {marker}"}, + identity, + ) + peer_requests: Final = peer.drain() + marker_requests: Final = tuple(request for request in peer_requests if marker in json.dumps(request)) + if blocked: + assert outcome.error is not None, outcome.raw + assert "block" in outcome.raw.casefold() or F1_BLOCKED_WORD in outcome.raw.casefold(), outcome.raw + assert marker_requests == (), (marker, marker_requests) + return + assert outcome.ok and marker in (outcome.text or ""), outcome.raw + assert len(marker_requests) == 1, (marker, marker_requests) + + +@pytest.mark.parametrize("streamed", (False, True), ids=("S0", "S1")) +def test_f2_mismatched_policy_step_skips_matching_sibling_still_enforces( + rig: MatrixRig, + streamed: bool, +) -> None: + marker: Final = f"audit-f2-{int(streamed)}-{uuid.uuid4().hex}" + response: Final = _raw_chat( + rig.candidate, + marker, + streamed, + "matrix-pipeline-model", + ) + provider_rows: Final = _matching_requests(rig.provider, marker) + sink_rows: Final = rig.sink.drain() + streaming_rows, non_streaming_rows = _scoped_rows( + sink_rows, + marker, + rig.rails["f2_stream"], + rig.rails["f2_non_stream"], + ) + assert len(provider_rows) == 1, (marker, provider_rows) + if streamed: + assert response.status_code == 200 and marker in response.text, response.text + assert (len(streaming_rows), len(non_streaming_rows)) == (1, 0), (marker, streaming_rows, non_streaming_rows) + return + assert response.status_code == 400 and "synthetic policy block" in response.text, response.text + assert (len(streaming_rows), len(non_streaming_rows)) == (0, 1), (marker, streaming_rows, non_streaming_rows) + + +def test_f3_realtime_transcription_uses_streaming_scope(rig: MatrixRig) -> None: + streaming_scenario_id: Final = f"matrix-f3-stream-{uuid.uuid4().hex}" + streaming_upstream: Final = register_scenario( + streaming_scenario_id, + _realtime_transcription_response(F3_STREAM_BLOCKED_WORD), + ) + rig.scenario.cleanups.callback(delete_scenario, streaming_upstream) + streaming_model: Final = rig.scenario.model( + model="openai/gpt-realtime-2", + api_key=streaming_scenario_id, + api_base=rig.candidate.upstream_url, + ) + non_streaming_scenario_id: Final = f"matrix-f3-non-stream-{uuid.uuid4().hex}" + non_streaming_upstream: Final = register_scenario( + non_streaming_scenario_id, + _realtime_transcription_response(F3_NON_STREAM_BLOCKED_WORD), + ) + rig.scenario.cleanups.callback(delete_scenario, non_streaming_upstream) + non_streaming_model: Final = rig.scenario.model( + model="openai/gpt-realtime-2", + api_key=non_streaming_scenario_id, + api_base=rig.candidate.upstream_url, + ) + key: Final = rig.scenario.key() + proxy_url: Final = str(rig.candidate.client.base_url).rstrip("/") + + streaming_events: Final = run(_realtime_transcription_events(proxy_url, key, streaming_model)) + streaming_transcriptions: Final = tuple( + event + for event in streaming_events + if event.get("type") == "conversation.item.input_audio_transcription.completed" + ) + streaming_errors: Final = tuple(event for event in streaming_events if event.get("type") == "error") + assert len(streaming_transcriptions) == 1, streaming_events + assert streaming_transcriptions[0].get("transcript") == F3_STREAM_BLOCKED_WORD, streaming_events + assert len(streaming_errors) == 1, streaming_events + streaming_error: Final = streaming_errors[0].get("error") + assert isinstance(streaming_error, dict) and streaming_error.get("type") == "guardrail_violation", streaming_events + + non_streaming_events: Final = run(_realtime_transcription_events(proxy_url, key, non_streaming_model)) + non_streaming_transcriptions: Final = tuple( + event + for event in non_streaming_events + if event.get("type") == "conversation.item.input_audio_transcription.completed" + ) + non_streaming_errors: Final = tuple(event for event in non_streaming_events if event.get("type") == "error") + completions: Final = tuple(event for event in non_streaming_events if event.get("type") == "response.done") + assert len(non_streaming_transcriptions) == 1, non_streaming_events + assert non_streaming_transcriptions[0].get("transcript") == F3_NON_STREAM_BLOCKED_WORD, non_streaming_events + assert non_streaming_errors == (), non_streaming_events + assert len(completions) == 1, non_streaming_events + + +def test_g1_sink_failure_fails_closed_only_when_rail_is_in_scope(rig: MatrixRig) -> None: + marker_s0: Final = f"audit-g1-S0-{uuid.uuid4().hex}" + response_s0: Final = _raw_chat( + rig.candidate, + marker_s0, + False, + rig.models["chat"], + rails=(rig.rails["g1_failure"],), + ) + assert response_s0.status_code == 200 and marker_s0 in response_s0.text, response_s0.text + provider_rows_s0: Final = _matching_requests(rig.provider, marker_s0) + sink_rows_s0: Final = _rail_scans(rig.sink.drain(), rig.rails["g1_failure"], marker_s0) + assert len(provider_rows_s0) == 1 and sink_rows_s0 == (), (marker_s0, provider_rows_s0, sink_rows_s0) + marker_s1: Final = f"audit-g1-S1-{uuid.uuid4().hex}" + response_s1: Final = _raw_chat( + rig.candidate, + marker_s1, + True, + rig.models["chat"], + rails=(rig.rails["g1_failure"],), + ) + assert response_s1.status_code == 500 and "Generic Guardrail API failed" in response_s1.text, response_s1.text + assert _matching_requests(rig.provider, marker_s1) == () + sink_rows: Final = _rail_scans(rig.sink.drain(), rig.rails["g1_failure"], marker_s1) + assert len(sink_rows) == 1, (marker_s1, sink_rows) + + +def test_g2_blocked_verdict_blocks_only_when_rail_is_in_scope(rig: MatrixRig) -> None: + marker_s0: Final = f"audit-g2-S0-{uuid.uuid4().hex}" + response_s0: Final = _raw_chat( + rig.candidate, + marker_s0, + False, + rig.models["chat"], + rails=(rig.rails["g2_block"],), + ) + assert response_s0.status_code == 200 and marker_s0 in response_s0.text, response_s0.text + assert len(_matching_requests(rig.provider, marker_s0)) == 1 + assert _rail_scans(rig.sink.drain(), rig.rails["g2_block"], marker_s0) == () + marker_s1: Final = f"audit-g2-S1-{uuid.uuid4().hex}" + response_s1: Final = _raw_chat( + rig.candidate, + marker_s1, + True, + rig.models["chat"], + rails=(rig.rails["g2_block"],), + ) + assert response_s1.status_code == 400 and "synthetic policy block" in response_s1.text, response_s1.text + assert _matching_requests(rig.provider, marker_s1) == () + sink_rows: Final = _rail_scans(rig.sink.drain(), rig.rails["g2_block"], marker_s1) + assert len(sink_rows) == 1, (marker_s1, sink_rows) + + +def test_g3_provider_errors_reach_caller_and_proxy_remains_usable(rig: MatrixRig) -> None: + marker: Final = f"audit-g3-401-{uuid.uuid4().hex}" + unauthorized: Final = _raw_chat( + rig.candidate, + marker, + False, + rig.models["chat"], + rails=(rig.rails["g3_provider"],), + ) + assert unauthorized.status_code == 401 and "synthetic provider unauthorized" in unauthorized.text, unauthorized.text + assert len(_matching_requests(rig.provider, marker)) == 1 + assert len(_rail_scans(rig.sink.drain(), rig.rails["g3_provider"], marker)) == 1 + unknown_marker: Final = f"audit-g3-unknown-{uuid.uuid4().hex}" + unknown: Final = _raw_chat( + rig.candidate, + unknown_marker, + False, + "audit-unknown-model", + rails=(rig.rails["g3_provider"],), + ) + assert unknown.status_code in (400, 404) and "model" in unknown.text.lower(), unknown.text + assert _matching_requests(rig.provider, unknown_marker) == () + healthy_marker: Final = f"audit-g3-healthy-{uuid.uuid4().hex}" + healthy: Final = _raw_chat(rig.candidate, healthy_marker, False, rig.models["chat"], key=rig.key) + assert healthy.status_code == 200 and healthy_marker in healthy.text, healthy.text + assert len(_matching_requests(rig.provider, healthy_marker)) == 1 + + +VERTEX_LIVE_HOST: Final = "us-central1-aiplatform.googleapis.com" +VERTEX_LIVE_PATH: Final = "/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent" +VERTEX_LIVE_AUTHORITY: Final = f"{VERTEX_LIVE_HOST}:443" + + +@dataclass(frozen=True, slots=True) +class _VertexLivePeer: + port: int + paths: SimpleQueue[str] + authorizations: SimpleQueue[str] + frames: SimpleQueue[dict[str, JsonValue]] + + +@dataclass(frozen=True, slots=True) +class _VertexConnectTunnel: + url: str + authorities: SimpleQueue[str] + + +async def _answer_vertex_live_peer( + connection: ServerConnection, + paths: SimpleQueue[str], + authorizations: SimpleQueue[str], + frames: SimpleQueue[dict[str, JsonValue]], +) -> None: + request: Final = connection.request + assert request is not None + paths.put(request.path) + authorizations.put(request.headers.get("Authorization", "")) + async for raw_frame in connection: + frame: Final = JSON_OBJECT.validate_python(json.loads(raw_frame)) + frames.put(frame) + if frames.qsize() == 1: + await connection.send(json.dumps({"setupComplete": {}})) + + +async def _serve_vertex_live_peer( + tls: ssl.SSLContext, + paths: SimpleQueue[str], + authorizations: SimpleQueue[str], + frames: SimpleQueue[dict[str, JsonValue]], + ports: SimpleQueue[int], + stop: asyncio.Event, +) -> None: + async with serve( + lambda connection: _answer_vertex_live_peer(connection, paths, authorizations, frames), + "127.0.0.1", + 0, + ssl=tls, + ) as server: + ports.put(next(iter(server.sockets)).getsockname()[1]) + await stop.wait() + + +@contextmanager +def _vertex_live_peer(cert: tuple[Path, Path]) -> Iterator[_VertexLivePeer]: + loop: Final = asyncio.new_event_loop() + stop: Final = asyncio.Event() + paths: Final = SimpleQueue[str]() + authorizations: Final = SimpleQueue[str]() + frames: Final = SimpleQueue[dict[str, JsonValue]]() + ports: Final = SimpleQueue[int]() + thread: Final = threading.Thread( + target=loop.run_until_complete, + args=(_serve_vertex_live_peer(server_context(*cert), paths, authorizations, frames, ports, stop),), + daemon=True, + ) + thread.start() + try: + yield _VertexLivePeer(ports.get(timeout=10), paths, authorizations, frames) + finally: + loop.call_soon_threadsafe(stop.set) + thread.join(timeout=10) + loop.close() + + +def _pipe_vertex_socket(source: socket.socket, sink: socket.socket) -> None: + with contextlib.suppress(OSError): + for chunk in iter(lambda: source.recv(65536), b""): + sink.sendall(chunk) + with contextlib.suppress(OSError): + sink.shutdown(socket.SHUT_WR) + + +@contextmanager +def _vertex_connect_tunnel(peer: _VertexLivePeer) -> Generator[_VertexConnectTunnel, None, None]: + authorities: Final = SimpleQueue[str]() + + class ConnectHandler(socketserver.StreamRequestHandler): + rbufsize = 0 + request: socket.socket + + def handle(self) -> None: + request_line: Final = self.rfile.readline().decode().split() + authority: Final = request_line[1] if len(request_line) > 1 else "" + while self.rfile.readline() not in (b"\r\n", b""): + pass + authorities.put(authority) + if authority != VERTEX_LIVE_AUTHORITY: + self.wfile.write(b"HTTP/1.1 403 Forbidden\r\ncontent-length: 0\r\n\r\n") + return + self.wfile.write(b"HTTP/1.1 200 Connection established\r\n\r\n") + self.request.settimeout(10) + with socket.create_connection(("127.0.0.1", peer.port), timeout=10) as upstream: + outbound: Final = threading.Thread(target=_pipe_vertex_socket, args=(self.request, upstream)) + outbound.start() + _pipe_vertex_socket(upstream, self.request) + outbound.join(timeout=12) + + with socketserver.ThreadingTCPServer(("127.0.0.1", 0), ConnectHandler) as server: + thread: Final = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.05}) + thread.start() + try: + yield _VertexConnectTunnel(f"http://127.0.0.1:{server.server_address[1]}", authorities) + finally: + server.shutdown() + thread.join(timeout=6) + + +async def _exchange_vertex_live_frame( + url: str, key: str, vertex_project: str, vertex_location: str, frame: str +) -> str | bytes: + websocket_url: Final = ( + f"{url.replace('http://', 'ws://', 1).rstrip('/')}/vertex_ai/live" + f"?vertex_project={vertex_project}&vertex_location={vertex_location}" + ) + async with websockets.connect( + websocket_url, + additional_headers={"Authorization": f"Bearer {key}"}, + open_timeout=10, + ) as websocket: + await websocket.send(frame) + reply: Final = await wait_for(websocket.recv(), timeout=10) + await websocket.close() + return reply + + +def test_ej_websocket_passthrough_runs_only_streaming_scoped_rails(rig: MatrixRig, tmp_path: Path) -> None: + vertex_project: Final = f"matrix-ej-{uuid.uuid4().hex}" + vertex_location: Final = "us-central1" + streaming_name: Final = f"ej-streaming-{uuid.uuid4().hex}" + non_streaming_name: Final = f"ej-non-streaming-{uuid.uuid4().hex}" + + def _token_response(_request: Request) -> Reply: + return Reply( + body=_json( + { + "access_token": "synthetic-vertex-access-token", + "expires_in": 3600, + "token_type": "Bearer", + } + ) + ) + + private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + private_key_pem: Final = private_key.private_bytes( + serialization.Encoding.PEM, + serialization.PrivateFormat.PKCS8, + serialization.NoEncryption(), + ).decode("utf-8") + credentials_path: Final = tmp_path / "vertex-service-account.json" + + with ( + wire_server(_token_response) as token_double, + wire_server(lambda _request: Reply(body=_json({"action": "NONE"}))) as sink, + ): + credentials_path.write_text( + json.dumps( + { + "type": "service_account", + "project_id": vertex_project, + "private_key_id": uuid.uuid4().hex, + "private_key": private_key_pem, + "client_email": f"integration-test@{vertex_project}.iam.gserviceaccount.com", + "client_id": "123456789012345678901", + "token_uri": f"{token_double.url}/token", + "auth_uri": "https://accounts.google.com/o/oauth2/auth", + "auth_provider_x509_cert_url": "https://www.googleapis.com/oauth2/v1/certs", + "client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/integration-test", + } + ) + ) + rails: Final = ( + _rail(streaming_name, sink, mode="pre_call", scope="streaming", default_on=True), + _rail(non_streaming_name, sink, mode="pre_call", scope="non_streaming", default_on=True), + ) + config_data: Final = cast(object, yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())) + config_base: Final = JSON_OBJECT.validate_python(config_data) + config: Final = cast(dict[str, JsonValue], {**config_base, "guardrails": list(rails)}) + config_path: Final = tmp_path / "ej-vertex-live-stream-scope.yaml" + config_path.write_text(yaml.safe_dump(config)) + certificate: Final = write_self_signed_cert(tmp_path, (VERTEX_LIVE_HOST,)) + with _vertex_live_peer(certificate) as peer, _vertex_connect_tunnel(peer) as tunnel: + with owned_proxy_process( + rig.candidate, + tmp_path, + { + "DEFAULT_VERTEXAI_PROJECT": vertex_project, + "DEFAULT_VERTEXAI_LOCATION": vertex_location, + "DEFAULT_GOOGLE_APPLICATION_CREDENTIALS": str(credentials_path), + "HTTPS_PROXY": tunnel.url, + "https_proxy": tunnel.url, + "NO_PROXY": "127.0.0.1,localhost", + "no_proxy": "127.0.0.1,localhost", + "SSL_CERT_FILE": str(certificate[0]), + }, + config=config_path, + workers=1, + ) as owned: + with owned.gateway.scenario() as scenario: + key: Final = scenario.key() + marker: Final = f"vertex-live-{uuid.uuid4().hex}" + frame: Final = json.dumps( + { + "clientContent": { + "turns": [{"role": "user", "parts": [{"text": marker}]}], + "turnComplete": True, + } + } + ) + client_reply_raw: Final = run( + _exchange_vertex_live_frame( + str(owned.gateway.client.base_url), + key, + vertex_project, + vertex_location, + frame, + ) + ) + proxy_log: Final = owned.log + + tunnel_authorities: Final = tuple( + tunnel.authorities.get_nowait() for _ in range(tunnel.authorities.qsize()) + ) + peer_path: Final = peer.paths.get_nowait() + peer_authorization: Final = peer.authorizations.get_nowait() + peer_frames: Final = tuple(peer.frames.get_nowait() for _ in range(peer.frames.qsize())) + client_reply: Final = JSON_OBJECT.validate_python(json.loads(client_reply_raw)) + expected_frame: Final = JSON_OBJECT.validate_python(json.loads(frame)) + token_requests: Final = token_double.drain() + sink_rows: Final = sink.drain() + streaming_rows: Final = tuple(row for row in sink_rows if row.target.startswith(f"/{streaming_name}/")) + non_streaming_rows: Final = tuple( + row for row in sink_rows if row.target.startswith(f"/{non_streaming_name}/") + ) + token_request_count: Final = len(token_requests) + streaming_count: Final = len(streaming_rows) + non_streaming_count: Final = len(non_streaming_rows) + assert tunnel_authorities == (VERTEX_LIVE_AUTHORITY,), ( + tunnel_authorities, + proxy_log, + ) + assert peer_path == VERTEX_LIVE_PATH, ( + peer_path, + proxy_log, + ) + assert peer_authorization == "Bearer synthetic-vertex-access-token", (peer_authorization, proxy_log) + assert peer_frames == (expected_frame,), (marker, peer_frames, proxy_log) + assert client_reply == {"setupComplete": {}}, (client_reply, proxy_log) + assert token_request_count == 1, (token_request_count, proxy_log) + assert streaming_count == 1, (streaming_count, non_streaming_count, proxy_log) + assert non_streaming_count == 0, (streaming_count, non_streaming_count, proxy_log) diff --git a/tests/unit/experimental_mcp_client/test_mcp_client.py b/tests/unit/experimental_mcp_client/test_mcp_client.py index b06ed9468b0..9a30e1990b1 100644 --- a/tests/unit/experimental_mcp_client/test_mcp_client.py +++ b/tests/unit/experimental_mcp_client/test_mcp_client.py @@ -2936,8 +2936,11 @@ async def test_cancellation_delivers_termination_over_tcp( listener: Final = await asyncio.start_server(handle_connection, "127.0.0.1", 0) port: Final = listener.sockets[0].getsockname()[1] + client_timeout: Final = 2 if cancel_mode == "read_timeout" else 30 client: Final = MCPClient( - server_url=f"http://127.0.0.1:{port}/mcp", protocol_version=protocol_version, timeout=2 if cancel_mode == "read_timeout" else 30 + server_url=f"http://127.0.0.1:{port}/mcp", + protocol_version=protocol_version, + timeout=client_timeout, ) async def calls(): diff --git a/tests/unit/integrations/test_custom_guardrail.py b/tests/unit/integrations/test_custom_guardrail.py index 72c362425bb..e04f8af5a66 100644 --- a/tests/unit/integrations/test_custom_guardrail.py +++ b/tests/unit/integrations/test_custom_guardrail.py @@ -1,5 +1,9 @@ import asyncio +import copy import datetime as dt +import json +import pickle +import threading from typing import TYPE_CHECKING, ClassVar, Final, Literal, Optional from unittest.mock import AsyncMock @@ -8,7 +12,10 @@ import pytest from litellm.integrations.custom_guardrail import ( DEFAULT_ADVISORY_MESSAGE, CustomGuardrail, + _request_is_streaming, + guardrail_request_data_with_streaming, log_guardrail_information, + without_server_streaming_classification, ) from litellm.litellm_core_utils.litellm_logging import Logging from litellm.proxy._types import CallTypes, UserAPIKeyAuth @@ -531,6 +538,233 @@ class TestCustomGuardrailShouldRunGuardrail: assert always_on.should_run_guardrail(data=forged, event_type=GuardrailEventHooks.pre_call) is True +_STREAM_SCOPE_HOOKS: Final = ( + GuardrailEventHooks.pre_call, + GuardrailEventHooks.during_call, + GuardrailEventHooks.post_call, +) + + +class TestCustomGuardrailStreamScope: + @pytest.mark.parametrize("event_type", _STREAM_SCOPE_HOOKS) + @pytest.mark.parametrize("stream", [True, False]) + @pytest.mark.parametrize("stream_scope", [None, "both"]) + def test_both_and_omitted_run_on_streaming_and_non_streaming( + self, + event_type: GuardrailEventHooks, + stream: bool, + stream_scope: str | None, + ): + guardrail = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=True, + event_hook=event_type, + stream_scope=stream_scope, + ) + assert guardrail.should_run_guardrail({"stream": stream}, event_type) is True + + @pytest.mark.parametrize("event_type", _STREAM_SCOPE_HOOKS) + def test_scalar_streaming_skips_non_streaming(self, event_type: GuardrailEventHooks): + guardrail = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=True, + event_hook=event_type, + stream_scope="streaming", + ) + assert guardrail.should_run_guardrail({"stream": True}, event_type) is True + assert guardrail.should_run_guardrail({"stream": False}, event_type) is False + assert guardrail.should_run_guardrail({}, event_type) is False + + @pytest.mark.parametrize("event_type", _STREAM_SCOPE_HOOKS) + def test_scalar_non_streaming_skips_streaming(self, event_type: GuardrailEventHooks): + guardrail = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=True, + event_hook=event_type, + stream_scope="non_streaming", + ) + assert guardrail.should_run_guardrail({"stream": False}, event_type) is True + assert guardrail.should_run_guardrail({}, event_type) is True + assert guardrail.should_run_guardrail({"stream": True}, event_type) is False + + def test_per_mode_map_applies_to_named_hooks_only(self): + guardrail = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=True, + event_hook=[GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call], + stream_scope={"pre_call": "both", "post_call": "streaming"}, + ) + assert guardrail.should_run_guardrail({"stream": True}, GuardrailEventHooks.pre_call) is True + assert guardrail.should_run_guardrail({"stream": False}, GuardrailEventHooks.pre_call) is True + assert guardrail.should_run_guardrail({"stream": True}, GuardrailEventHooks.post_call) is True + assert guardrail.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) is False + + def test_default_on_early_return_still_honors_stream_scope(self): + guardrail = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=True, + event_hook=GuardrailEventHooks.post_call, + stream_scope="non_streaming", + ) + assert guardrail.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) is True + assert guardrail.should_run_guardrail({"stream": True}, GuardrailEventHooks.post_call) is False + + def test_apply_stream_scope_overwrites_constructor_default(self): + guardrail = CustomGuardrail( + guardrail_name="scoped", + default_on=True, + event_hook=GuardrailEventHooks.post_call, + ) + assert guardrail.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) is True + guardrail.apply_stream_scope("streaming") + assert guardrail.should_run_guardrail({"stream": True}, GuardrailEventHooks.post_call) is True + assert guardrail.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) is False + + def test_direct_constructor_normalizes_mixed_case_map_keys(self): + guardrail = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=True, + event_hook=GuardrailEventHooks.pre_call, + stream_scope={"Pre_Call": "streaming"}, + ) + assert guardrail.should_run_guardrail({"stream": True}, GuardrailEventHooks.pre_call) is True + assert guardrail.should_run_guardrail({"stream": False}, GuardrailEventHooks.pre_call) is False + + def test_realtime_transcription_counts_as_streaming(self): + streaming_only = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=True, + event_hook=GuardrailEventHooks.realtime_input_transcription, + stream_scope="streaming", + ) + assert ( + streaming_only.should_run_guardrail( + {"litellm_metadata": {}}, GuardrailEventHooks.realtime_input_transcription + ) + is True + ) + non_streaming_only = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=True, + event_hook=GuardrailEventHooks.realtime_input_transcription, + stream_scope="non_streaming", + ) + assert ( + non_streaming_only.should_run_guardrail( + {"litellm_metadata": {}}, GuardrailEventHooks.realtime_input_transcription + ) + is False + ) + + def test_path_defined_streaming_classification_cannot_be_spoofed(self): + generate_content_body: Final = {"contents": [{"parts": [{"text": "hi"}]}]} + streaming_only = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=True, + event_hook=GuardrailEventHooks.pre_call, + stream_scope="streaming", + ) + assert streaming_only.should_run_guardrail(generate_content_body, GuardrailEventHooks.pre_call) is False + assert ( + streaming_only.should_run_guardrail( + {**generate_content_body, "is_streaming_request": True}, + GuardrailEventHooks.pre_call, + ) + is False + ) + assert ( + streaming_only.should_run_guardrail( + {**generate_content_body, "is_streaming_request": "litellm-server-streaming"}, + GuardrailEventHooks.pre_call, + ) + is False + ) + assert ( + streaming_only.should_run_guardrail( + {**generate_content_body, "litellm_server_streaming_classification": True}, + GuardrailEventHooks.pre_call, + ) + is False + ) + server_streaming_data: Final = guardrail_request_data_with_streaming( + generate_content_body, + is_streaming=True, + ) + assert ( + streaming_only.should_run_guardrail( + server_streaming_data, + GuardrailEventHooks.pre_call, + ) + is True + ) + non_streaming_only = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=True, + event_hook=GuardrailEventHooks.pre_call, + stream_scope="non_streaming", + ) + assert ( + non_streaming_only.should_run_guardrail( + server_streaming_data, + GuardrailEventHooks.pre_call, + ) + is False + ) + + def test_streaming_classification_is_json_serializable_without_spoofing(self): + d: Final = guardrail_request_data_with_streaming({}, is_streaming=True) + serialized: Final = json.dumps(d) + assert _request_is_streaming(d) is True + + round_tripped: Final = json.loads(serialized) + assert _request_is_streaming(round_tripped) is False + assert round_tripped["litellm_server_streaming_classification"] == "litellm-server-streaming" + assert isinstance(round_tripped["litellm_server_streaming_classification"], str) + assert "litellm_server_streaming_classification" not in without_server_streaming_classification(round_tripped) + + def test_streaming_classification_preserves_caller_fields_and_removes_only_server_marker(self): + caller_data: Final = { + "contents": [{"parts": [{"text": "hi"}]}], + "is_streaming_request": "caller-value", + "litellm_server_streaming_classification": True, + } + non_streaming_data: Final = guardrail_request_data_with_streaming(caller_data, is_streaming=False) + server_streaming_data: Final = guardrail_request_data_with_streaming(caller_data, is_streaming=True) + + assert non_streaming_data is not caller_data + assert server_streaming_data is not caller_data + assert non_streaming_data == caller_data + assert server_streaming_data["is_streaming_request"] == "caller-value" + assert server_streaming_data["litellm_server_streaming_classification"] is not True + assert without_server_streaming_classification(caller_data) == caller_data + assert "litellm_server_streaming_classification" not in without_server_streaming_classification( + server_streaming_data + ) + + def test_server_streaming_classification_survives_scan_raw_request_snapshot(self): + from litellm.litellm_core_utils.core_helpers import independent_snapshot + + generate_content_body: Final = {"contents": [{"parts": [{"text": "hi"}]}]} + snapshot: Final = independent_snapshot( + guardrail_request_data_with_streaming(generate_content_body, is_streaming=True) + ) + streaming_only = CustomGuardrail( + guardrail_name="test_guardrail", + default_on=True, + event_hook=GuardrailEventHooks.pre_call, + stream_scope="streaming", + scan_raw_request=True, + ) + assert streaming_only.should_run_guardrail(snapshot, GuardrailEventHooks.pre_call) is True + assert ( + streaming_only.should_run_guardrail( + independent_snapshot({**generate_content_body, "is_streaming_request": True}), + GuardrailEventHooks.pre_call, + ) + is False + ) + + class TestApplyGuardrailCheck: def test_apply_guardrail_check_only_on_direct_implementation(self): """ @@ -3307,3 +3541,113 @@ class TestCustomGuardrailTimeout: ) assert guardrail.timeout == 7.0 + + +@pytest.mark.parametrize("stream_scope", [None, "both", {"post_call": "streaming"}]) +def test_guardrail_survives_deepcopy_and_pickle_with_its_stream_scope(stream_scope): + guardrail: Final = CustomGuardrail( + guardrail_name="copyable", + default_on=True, + event_hook=GuardrailEventHooks.post_call, + stream_scope=stream_scope, + ) + expected: Final = guardrail.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) + for clone in (copy.deepcopy(guardrail), pickle.loads(pickle.dumps(guardrail))): + assert dict(clone.stream_scope_by_hook) == dict(guardrail.stream_scope_by_hook) + assert clone.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) is expected + + +class _LockHoldingGuardrail(CustomGuardrail): + def __init__(self, **kwargs): + self.lock = threading.Lock() + super().__init__(**kwargs) + + def __getstate__(self): + state: Final = dict(self.__dict__) + state.pop("lock") + return state + + def __setstate__(self, state): + self.__dict__.update(state) + self.lock = threading.Lock() + + +class _SetstateOnlyGuardrail(CustomGuardrail): + def __setstate__(self, state): + self.__dict__.update(state) + self.restored = True + + +class _SlotsGuardrail(CustomGuardrail): + __slots__ = ("vendor_client",) + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.vendor_client = "vendor-client-object" + + +def _subclass_guardrails(): + kwargs: Final = { + "guardrail_name": "vendor", + "default_on": True, + "event_hook": GuardrailEventHooks.post_call, + "stream_scope": {"post_call": "streaming"}, + } + return ( + pytest.param(_LockHoldingGuardrail(**kwargs), id="dict-getstate"), + pytest.param(_SetstateOnlyGuardrail(**kwargs), id="setstate-only"), + pytest.param(_SlotsGuardrail(**kwargs), id="slots"), + ) + + +@pytest.mark.parametrize("guardrail", _subclass_guardrails()) +@pytest.mark.parametrize( + "cloner", + [copy.copy, copy.deepcopy, lambda g: pickle.loads(pickle.dumps(g))], + ids=["copy", "deepcopy", "pickle"], +) +def test_out_of_tree_guardrail_subclasses_survive_copy_and_pickle(guardrail, cloner): + clone = cloner(guardrail) + + if isinstance(clone, _LockHoldingGuardrail): + assert isinstance(clone.lock, type(threading.Lock())) + elif isinstance(clone, _SetstateOnlyGuardrail): + assert clone.restored is True + else: + assert clone.vendor_client == "vendor-client-object" + assert clone.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) is False + assert clone.should_run_guardrail({"stream": True}, GuardrailEventHooks.post_call) is True + + +def test_router_constructs_with_a_dict_getstate_guardrail_in_deployment_callbacks(): + import litellm + + litellm.Router( + model_list=[ + { + "model_name": "m", + "litellm_params": { + "model": "openai/gpt-5.4-mini", + "api_key": "sk-test", + "callbacks": [ + _LockHoldingGuardrail( + guardrail_name="vendor", + default_on=True, + event_hook=GuardrailEventHooks.post_call, + stream_scope={"post_call": "streaming"}, + ) + ], + }, + } + ] + ) + + +def test_subclass_that_skips_super_init_still_runs_with_default_scope(): + class NoSuperInit(CustomGuardrail): + def __init__(self) -> None: + self.guardrail_name = "no-super" + self.event_hook = None + self.default_on = True + + assert NoSuperInit().should_run_guardrail({"stream": True}, GuardrailEventHooks.pre_call) is True diff --git a/tests/unit/llms/bedrock/passthrough/test_bedrock_passthrough_transformation.py b/tests/unit/llms/bedrock/passthrough/test_bedrock_passthrough_transformation.py index fbcbda183bc..875a2fc1b1a 100644 --- a/tests/unit/llms/bedrock/passthrough/test_bedrock_passthrough_transformation.py +++ b/tests/unit/llms/bedrock/passthrough/test_bedrock_passthrough_transformation.py @@ -10,9 +10,14 @@ from datetime import datetime from typing import Final from unittest.mock import patch +import pytest + from litellm.litellm_core_utils.litellm_logging import Logging from litellm.llms.base_llm.passthrough.transformation import PassthroughStreamCollector -from litellm.llms.bedrock.passthrough.transformation import BedrockPassthroughConfig +from litellm.llms.bedrock.passthrough.transformation import ( + BedrockPassthroughConfig, + is_bedrock_streaming_endpoint, +) from litellm.types.utils import ModelResponse CONVERSE_MODEL = "anthropic.claude-sonnet-4-5-20250929-v1:0" @@ -20,6 +25,22 @@ CONVERSE_STREAM_ENDPOINT = f"/model/{CONVERSE_MODEL}/converse-stream" INVOKE_STREAM_ENDPOINT = f"/model/{CONVERSE_MODEL}/invoke-with-response-stream" +@pytest.mark.parametrize( + ("endpoint", "expected"), + [ + ("converse-stream", True), + ("invoke-with-response-stream", True), + ("converse", False), + ("invoke", False), + ("model/my-converse-stream-model/converse", False), + ("model/x/converse-stream?foo=1", True), + ("model/x/converse-stream/", True), + ], +) +def test_is_bedrock_streaming_endpoint_matches_final_action_segment(endpoint: str, expected: bool) -> None: + assert is_bedrock_streaming_endpoint(endpoint) is expected + + def test_bedrock_passthrough_get_complete_url_default_endpoint(): """Test get_complete_url with default AWS endpoint (no override)""" config = BedrockPassthroughConfig() diff --git a/tests/unit/llms/pass_through/guardrail_translation/test_handler.py b/tests/unit/llms/pass_through/guardrail_translation/test_handler.py new file mode 100644 index 00000000000..82ad7702052 --- /dev/null +++ b/tests/unit/llms/pass_through/guardrail_translation/test_handler.py @@ -0,0 +1,52 @@ +import json +from typing import Final + +import pytest + +from litellm.constants import SERVER_STREAMING_CLASSIFICATION_KEY, SERVER_STREAMING_CLASSIFICATION_MARKER +from litellm.llms.pass_through.guardrail_translation.handler import PassThroughEndpointHandler + + +def test_full_payload_guardrail_text_excludes_the_server_streaming_marker(): + body: Final = {"model": "m", "stream": True, "messages": [{"role": "user", "content": "hi"}]} + + text: Final = PassThroughEndpointHandler()._extract_text_for_guardrail( + {**body, SERVER_STREAMING_CLASSIFICATION_KEY: SERVER_STREAMING_CLASSIFICATION_MARKER}, + None, + ) + + assert json.loads(text) == body, text + + +@pytest.mark.parametrize( + "marker", + [ + SERVER_STREAMING_CLASSIFICATION_MARKER, + json.loads(json.dumps(SERVER_STREAMING_CLASSIFICATION_MARKER)), + ], + ids=["enum", "json-string"], +) +def test_full_payload_guardrail_text_scans_caller_value_but_not_the_marker(marker: str): + text: Final = PassThroughEndpointHandler()._extract_text_for_guardrail( + { + "model": "m", + "stream": True, + "messages": [{"role": "user", "content": "hi"}], + SERVER_STREAMING_CLASSIFICATION_KEY: "BLOCKME caller content", + }, + None, + ) + + assert "BLOCKME caller content" in text, text + + marker_text: Final = PassThroughEndpointHandler()._extract_text_for_guardrail( + { + "model": "m", + "stream": True, + "messages": [{"role": "user", "content": "hi"}], + SERVER_STREAMING_CLASSIFICATION_KEY: marker, + }, + None, + ) + + assert SERVER_STREAMING_CLASSIFICATION_KEY not in json.loads(marker_text), marker_text diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py index 23defa6f513..5ae3c595c53 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server.py @@ -464,7 +464,7 @@ async def test_sse_mcp_handler_mock(): mock_sse = MagicMock() mock_sse.connect_sse.side_effect = connect_sse - run = AsyncMock() + serve = AsyncMock() # Mock scope, receive, send with proper ASGI scope format mock_scope = { @@ -489,7 +489,7 @@ async def test_sse_mcp_handler_mock(): ) with ( - patch("litellm.proxy._experimental.mcp_server.server.serve_loop", run), + patch("litellm.proxy._experimental.mcp_server.server.serve_loop", serve), patch( "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", True, @@ -505,13 +505,21 @@ async def test_sse_mcp_handler_mock(): patch( "litellm.proxy._experimental.mcp_server.server.set_auth_context", ), + patch( + "litellm.proxy._experimental.mcp_server.server._raise_preemptive_401_for_unauthenticated_servers", + new=AsyncMock(), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._check_passthrough_upstream_auth", + new=AsyncMock(), + ), ): from litellm.proxy._experimental.mcp_server.server import handle_sse_mcp # Call the handler await handle_sse_mcp(mock_scope, mock_receive, mock_send) - assert run.await_args.args[1:3] == (read_stream, write_stream) + assert serve.await_args.args[1:3] == (read_stream, write_stream) assert mock_sse.connect_sse.call_args.args[0]["path"] == "/mcp/sse" diff --git a/tests/unit/proxy/guardrails/test_guardrail_endpoints.py b/tests/unit/proxy/guardrails/test_guardrail_endpoints.py index fa33c2d462b..b7bc5f3f8ec 100644 --- a/tests/unit/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/unit/proxy/guardrails/test_guardrail_endpoints.py @@ -575,7 +575,7 @@ async def test_get_guardrail_info_from_db(mocker, mock_prisma_client): """Test getting guardrail info from DB""" mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - response = await get_guardrail_info("test-db-guardrail") + response: Final = await get_guardrail_info("test-db-guardrail") assert response.guardrail_id == "test-db-guardrail" assert response.guardrail_name == "Test DB Guardrail" @@ -584,6 +584,21 @@ async def test_get_guardrail_info_from_db(mocker, mock_prisma_client): @pytest.mark.asyncio +async def test_get_guardrail_info_tolerates_invalid_stored_stream_scope(mocker, mock_prisma_client): + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mock_prisma_client.db.litellm_guardrailstable.find_unique = AsyncMock( + return_value={ + **MOCK_DB_GUARDRAIL, + "litellm_params": { + **MOCK_DB_GUARDRAIL["litellm_params"], + "stream_scope": "sometimes", + }, + } + ) + + response = await get_guardrail_info("test-db-guardrail") + + assert response.litellm_params.stream_scope is None async def test_get_guardrail_info_normalizes_invalid_scope_from_db( mocker, mock_guardrail_registry, mock_in_memory_handler ): @@ -745,6 +760,40 @@ def test_get_guardrails_list_response_includes_guardrail_id(): assert response.guardrails[0].guardrail_id == "stable-config-id" +def test_get_guardrails_list_response_tolerates_invalid_config_stream_scope(): + from litellm.proxy.guardrails.guardrail_endpoints import ( + _get_guardrails_list_response, + ) + + response = _get_guardrails_list_response( + [ + { + "guardrail_id": "invalid-scope", + "guardrail_name": "invalid-scope", + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "stream_scope": "sometimes", + }, + }, + { + "guardrail_id": "valid-scope", + "guardrail_name": "valid-scope", + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "stream_scope": "STREAMING", + }, + }, + ] + ) + + assert response.guardrails[0].litellm_params is not None + assert response.guardrails[0].litellm_params.stream_scope is None + assert response.guardrails[1].litellm_params is not None + assert response.guardrails[1].litellm_params.stream_scope == "streaming" + + def test_get_provider_specific_params(): """Test getting provider-specific parameters""" from litellm.proxy.guardrails.guardrail_endpoints import _get_fields_from_model diff --git a/tests/unit/proxy/guardrails/test_guardrail_registry.py b/tests/unit/proxy/guardrails/test_guardrail_registry.py index c7b13f0bd00..3bc6d163e5d 100644 --- a/tests/unit/proxy/guardrails/test_guardrail_registry.py +++ b/tests/unit/proxy/guardrails/test_guardrail_registry.py @@ -1,7 +1,7 @@ import json from collections.abc import Iterable, Iterator from typing import ClassVar, Final -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, patch import pytest from pydantic import ValidationError @@ -160,6 +160,43 @@ def test_duplicate_config_guardrail_names_get_distinct_stable_ids(): registry_module.guardrail_initializer_registry.pop("dup_name_test", None) +def test_initialize_guardrail_treats_invalid_stored_scope_as_both(): + from litellm.proxy.guardrails import guardrail_registry as registry_module + + guardrail_type: Final = "invalid_stored_scope_test" + + def _initializer(litellm_params: LitellmParams, guardrail: Guardrail) -> CustomGuardrail: + return CustomGuardrail( + guardrail_name=guardrail["guardrail_name"], + event_hook=GuardrailEventHooks(litellm_params.mode), + default_on=True, + ) + + registry_module.guardrail_initializer_registry[guardrail_type] = _initializer + try: + handler: Final = InMemoryGuardrailHandler() + guardrail: Final = Guardrail( + guardrail_id="invalid-stored-scope", + guardrail_name="invalid-stored-scope", + litellm_params={ + "guardrail": guardrail_type, + "mode": "pre_call", + "default_on": True, + "stream_scope": "sometimes", + }, + ) + + parsed_guardrail: Final = handler.initialize_guardrail(guardrail=guardrail, source="db") + callback: Final = handler.guardrail_id_to_custom_guardrail["invalid-stored-scope"] + + assert parsed_guardrail["litellm_params"].stream_scope is None + assert callback is not None + assert callback.should_run_guardrail(data={}, event_type=GuardrailEventHooks.pre_call) is True + assert callback.should_run_guardrail(data={"stream": True}, event_type=GuardrailEventHooks.pre_call) is True + finally: + registry_module.guardrail_initializer_registry.pop(guardrail_type, None) + + def _register_mode_following_initializer(guardrail_type: str): """Registers like the shipped initializers do: construct, then add the instance to litellm's callbacks.""" import litellm @@ -1605,6 +1642,35 @@ def test_sync_guardrail_from_db_applies_db_dict_params_to_live_instance(): cb_list[:] = snapshot +def test_configure_callback_scoping_copies_stream_scope_when_constructor_omits_it(): + from litellm.proxy.guardrails.guardrail_registry import _configure_callback_scoping + + class _CtorWithoutStreamScope(CustomGuardrail): + def __init__(self) -> None: + super().__init__( + guardrail_name="scoped", + event_hook=GuardrailEventHooks.post_call, + default_on=True, + ) + + instance = _CtorWithoutStreamScope() + params = LitellmParams(guardrail="bedrock", mode="post_call", stream_scope="streaming") + _configure_callback_scoping(instance, "scoped", params) + + assert instance.stream_scope_default == "streaming" + assert instance.should_run_guardrail({"stream": True}, GuardrailEventHooks.post_call) is True + assert instance.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) is False + + +def test_configure_callback_scoping_tolerates_a_custom_logger_callback(): + from litellm.integrations.custom_logger import CustomLogger + from litellm.proxy.guardrails.guardrail_registry import _configure_callback_scoping + + callback: Final = CustomLogger() + _configure_callback_scoping(callback, "logger-backed", LitellmParams(guardrail="custom", mode="pre_call")) # pyright: ignore[reportArgumentType] # module-path guardrails may be plain CustomLogger + assert "stream_scope_by_hook" not in vars(callback) + + _ENCRYPTED_PREFIX = "litellm_enc::" diff --git a/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 1887021a53a..d329db75fb4 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -23,6 +23,7 @@ from starlette.datastructures import FormData import litellm +from litellm.integrations.custom_guardrail import CustomGuardrail from tests._master_key import MASTER_KEY as SHARED_MASTER_KEY from litellm.caching.caching import DualCache from litellm.types.utils import CallTypesLiteral @@ -47,6 +48,7 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( get_vertex_base_url, is_azure_ai_search_service_level_index_create, gigachat_proxy_route, + handle_bedrock_passthrough_router_model, llm_passthrough_factory_proxy_route, milvus_proxy_route, mistral_proxy_route, @@ -61,9 +63,32 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( from litellm.proxy._types import LitellmUserRoles, SpecialHeaders, UserAPIKeyAuth from litellm.proxy.auth.handle_jwt import JWTHandler from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.types.guardrails import GuardrailEventHooks from litellm.types.passthrough_endpoints.vertex_ai import VertexPassThroughCredentials +def _assert_bedrock_processing_data_classification(data: dict[str, object], is_streaming: bool) -> None: + streaming_guardrail: Final = CustomGuardrail( + guardrail_name="streaming-only", + default_on=True, + event_hook=GuardrailEventHooks.pre_call, + stream_scope="streaming", + ) + non_streaming_guardrail: Final = CustomGuardrail( + guardrail_name="non-streaming-only", + default_on=True, + event_hook=GuardrailEventHooks.pre_call, + stream_scope="non_streaming", + ) + provider_body: Final = data["data"] + + assert streaming_guardrail.should_run_guardrail(data, GuardrailEventHooks.pre_call) is is_streaming + assert non_streaming_guardrail.should_run_guardrail(data, GuardrailEventHooks.pre_call) is not is_streaming + assert isinstance(provider_body, dict) + assert "is_streaming_request" not in provider_body + assert "litellm_server_streaming_classification" not in provider_body + + class TestVertexPassthroughGetVertexBaseUrl: """Module-local get_vertex_base_url (trailing slash); rules match common_utils.""" @@ -413,9 +438,7 @@ class TestVertexAIPassThroughHandler: # Mock the vertex handler for global location mock_handler = Mock() - mock_handler.get_default_base_target_url.return_value = ( - "https://aiplatform.googleapis.com/" - ) + mock_handler.get_default_base_target_url.return_value = "https://aiplatform.googleapis.com/" mock_get_handler.return_value = mock_handler # Mock create_pass_through_route to return a function that returns a mock response @@ -1237,9 +1260,7 @@ class TestVertexAIDiscoveryPassThroughHandler: # Mock the discovery handler mock_handler = Mock() - mock_handler.get_default_base_target_url.return_value = ( - "https://discoveryengine.googleapis.com" - ) + mock_handler.get_default_base_target_url.return_value = "https://discoveryengine.googleapis.com" mock_get_handler.return_value = mock_handler # Mock create_pass_through_route to return a function that returns a mock response @@ -1459,6 +1480,170 @@ class TestBedrockLLMProxyRoute: assert call_kwargs["model"] == "anthropic.claude-3-sonnet-20240229-v1:0" assert result == "success" + @pytest.mark.asyncio + @pytest.mark.parametrize( + "action, is_streaming", + [ + ("converse-stream", True), + ("invoke-with-response-stream", True), + ("converse", False), + ("invoke", False), + ], + ) + async def test_bedrock_direct_actions_classify_guardrail_stream_scope( + self, action: str, is_streaming: bool + ) -> None: + mock_request: Final = Mock() + mock_request.method = "POST" + mock_processor: Final = Mock() + mock_processor.base_passthrough_process_llm_request = AsyncMock(return_value="success") + request_body: Final = {"messages": [{"role": "user", "content": "test"}]} + + with ( + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._read_request_body", + return_value=request_body, + ), + patch( + "litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing", + return_value=mock_processor, + ) as processor_constructor, + ): + result: Final = await bedrock_llm_proxy_route( + endpoint=f"model/test-model/{action}", + request=mock_request, + fastapi_response=Mock(), + user_api_key_dict=Mock(), + ) + + assert result == "success" + processing_data: Final = processor_constructor.call_args.kwargs["data"] + _assert_bedrock_processing_data_classification(processing_data, is_streaming) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "action, is_streaming", + [ + ("converse-stream", True), + ("invoke-with-response-stream", True), + ("converse", False), + ("invoke", False), + ], + ) + async def test_bedrock_router_actions_classify_guardrail_stream_scope( + self, action: str, is_streaming: bool + ) -> None: + mock_request: Final = Mock() + mock_request.method = "POST" + mock_processor: Final = Mock() + mock_processor.base_passthrough_process_llm_request = AsyncMock(return_value="success") + request_body: Final = {"messages": [{"role": "user", "content": "test"}]} + + with patch( + "litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing", + return_value=mock_processor, + ) as processor_constructor: + result: Final = await handle_bedrock_passthrough_router_model( + model="test-model", + endpoint=f"model/test-model/{action}", + request=mock_request, + request_body=request_body, + llm_router=Mock(), + user_api_key_dict=Mock(), + proxy_logging_obj=Mock(), + general_settings={}, + proxy_config=None, + select_data_generator=None, + user_model=None, + user_temperature=None, + user_request_timeout=None, + user_max_tokens=None, + user_api_base=None, + version=None, + ) + + assert result == "success" + processing_data: Final = processor_constructor.call_args.kwargs["data"] + _assert_bedrock_processing_data_classification(processing_data, is_streaming) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("model_id", "action"), + [ + ("my-converse-stream-model", "converse"), + ("my-invoke-with-response-stream-model", "invoke"), + ], + ) + async def test_bedrock_direct_model_id_does_not_imply_streaming(self, model_id: str, action: str) -> None: + mock_request: Final = Mock() + mock_request.method = "POST" + mock_processor: Final = Mock() + mock_processor.base_passthrough_process_llm_request = AsyncMock(return_value="success") + request_body: Final = {"messages": [{"role": "user", "content": "test"}]} + + with ( + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._read_request_body", + return_value=request_body, + ), + patch( + "litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing", + return_value=mock_processor, + ) as processor_constructor, + ): + result: Final = await bedrock_llm_proxy_route( + endpoint=f"/model/{model_id}/{action}", + request=mock_request, + fastapi_response=Mock(), + user_api_key_dict=Mock(), + ) + + assert result == "success" + processing_data: Final = processor_constructor.call_args.kwargs["data"] + _assert_bedrock_processing_data_classification(processing_data, is_streaming=False) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + ("model_id", "action"), + [ + ("my-converse-stream-model", "converse"), + ("my-invoke-with-response-stream-model", "invoke"), + ], + ) + async def test_bedrock_router_model_id_does_not_imply_streaming(self, model_id: str, action: str) -> None: + mock_request: Final = Mock() + mock_request.method = "POST" + mock_processor: Final = Mock() + mock_processor.base_passthrough_process_llm_request = AsyncMock(return_value="success") + request_body: Final = {"messages": [{"role": "user", "content": "test"}]} + + with patch( + "litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing", + return_value=mock_processor, + ) as processor_constructor: + result: Final = await handle_bedrock_passthrough_router_model( + model=model_id, + endpoint=f"/model/{model_id}/{action}", + request=mock_request, + request_body=request_body, + llm_router=Mock(), + user_api_key_dict=Mock(), + proxy_logging_obj=Mock(), + general_settings={}, + proxy_config=None, + select_data_generator=None, + user_model=None, + user_temperature=None, + user_request_timeout=None, + user_max_tokens=None, + user_api_base=None, + version=None, + ) + + assert result == "success" + processing_data: Final = processor_constructor.call_args.kwargs["data"] + _assert_bedrock_processing_data_classification(processing_data, is_streaming=False) + @pytest.mark.asyncio async def test_bedrock_error_handling_returns_actual_error(self): """ @@ -1879,7 +2064,6 @@ class TestBedrockAgentRuntimePassthroughToggle: class TestBedrockAgentRuntimePassthroughVirtualKeyLeak: - VKEY: Final = "sk-litellm-victim-key" MASTER_KEY: Final = SHARED_MASTER_KEY ENDPOINT: Final = "knowledgebases/KB1234567/retrieve" @@ -2075,10 +2259,10 @@ class TestVLLMProxyRoute: "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_passthrough_request_using_router_model", return_value=True, ) - @patch("litellm.proxy.proxy_server.llm_router") # test-quality-ok: patching litellm internal for unit test isolation - async def test_vllm_proxy_route_with_router_model( - self, mock_llm_router, mock_is_router, mock_get_body - ): + @patch( + "litellm.proxy.proxy_server.llm_router" + ) # test-quality-ok: patching litellm internal for unit test isolation + async def test_vllm_proxy_route_with_router_model(self, mock_llm_router, mock_is_router, mock_get_body): mock_request = MagicMock(spec=Request) mock_request.method = "POST" mock_request.headers = {"content-type": "application/json"} @@ -2111,9 +2295,7 @@ class TestVLLMProxyRoute: @patch( # test-quality-ok: patching litellm internal for unit test isolation "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.llm_passthrough_factory_proxy_route" ) - async def test_vllm_proxy_route_fallback_to_factory( - self, mock_factory_route, mock_is_router, mock_get_body - ): + async def test_vllm_proxy_route_fallback_to_factory(self, mock_factory_route, mock_is_router, mock_get_body): mock_request = MagicMock(spec=Request) mock_fastapi_response = MagicMock(spec=Response) mock_user_api_key_dict = MagicMock() @@ -2140,10 +2322,10 @@ class TestGigachatProxyRoute: "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_passthrough_request_using_router_model", return_value=True, ) - @patch("litellm.proxy.proxy_server.llm_router") # test-quality-ok: patching litellm internal for unit test isolation - async def test_gigachat_proxy_route_with_router_model( - self, mock_llm_router, mock_is_router, mock_get_body - ): + @patch( + "litellm.proxy.proxy_server.llm_router" + ) # test-quality-ok: patching litellm internal for unit test isolation + async def test_gigachat_proxy_route_with_router_model(self, mock_llm_router, mock_is_router, mock_get_body): mock_request = MagicMock(spec=Request) mock_request.method = "POST" mock_request.headers = {"content-type": "application/json"} @@ -2396,21 +2578,25 @@ class TestGigachatProxyRoute: return _inner() - with patch.object( - processor, - "common_processing_pre_call_logic", - new=AsyncMock( - return_value=( - processor.data, - processor.data["litellm_logging_obj"], - ) + with ( + patch.object( + processor, + "common_processing_pre_call_logic", + new=AsyncMock( + return_value=( + processor.data, + processor.data["litellm_logging_obj"], + ) + ), + ), + patch( # test-quality-ok: patching litellm internal for unit test isolation + "litellm.proxy.common_request_processing.route_request", + new=_fake_route_request, + ), + patch( # test-quality-ok: patching litellm internal for unit test isolation + "litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing.get_custom_headers", + return_value={"x-litellm-call-id": "call-123"}, ), - ), patch( # test-quality-ok: patching litellm internal for unit test isolation - "litellm.proxy.common_request_processing.route_request", - new=_fake_route_request, - ), patch( # test-quality-ok: patching litellm internal for unit test isolation - "litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing.get_custom_headers", - return_value={"x-litellm-call-id": "call-123"}, ): result = await processor.base_passthrough_process_llm_request( request=mock_request, @@ -4223,7 +4409,9 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak: @pytest.mark.parametrize( ("credential", "authenticated"), [ - pytest.param("modified_key", UserAPIKeyAuth(api_key="modified_key"), id="custom-auth-echoing-opaque-credential"), + pytest.param( + "modified_key", UserAPIKeyAuth(api_key="modified_key"), id="custom-auth-echoing-opaque-credential" + ), pytest.param( LITELLM_JWT, UserAPIKeyAuth(api_key=LITELLM_JWT, user_id="jwt-subject"), @@ -4298,7 +4486,9 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak: (b"x-goog-api-key", b"AIza-real-google-api-key"), (b"content-type", b"application/json"), ], - authenticated=UserAPIKeyAuth(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS, user_role=LitellmUserRoles.PROXY_ADMIN), + authenticated=UserAPIKeyAuth( + api_key=LITELLM_PROXY_MASTER_KEY_ALIAS, user_role=LitellmUserRoles.PROXY_ADMIN + ), ) assert raised is None assert forwarded is not None @@ -4440,10 +4630,14 @@ class TestAnthropicPassthroughVirtualKeyLeak: raised, forwarded = await self._run( monkeypatch, [(header, value), (b"anthropic-version", b"2023-06-01"), (b"content-type", b"application/json")], - authenticated=UserAPIKeyAuth(api_key="sk-ant-api03-callers-own-key", user_role=LitellmUserRoles.INTERNAL_USER), + authenticated=UserAPIKeyAuth( + api_key="sk-ant-api03-callers-own-key", user_role=LitellmUserRoles.INTERNAL_USER + ), master_key=None, ) - assert raised is None, "with no master key the proxy authenticated nothing, so nothing of the caller's is a LiteLLM secret" + assert raised is None, ( + "with no master key the proxy authenticated nothing, so nothing of the caller's is a LiteLLM secret" + ) assert forwarded is not None assert forwarded.get(header.decode()) == value.decode() @@ -5231,7 +5425,9 @@ class TestTranscribeProxyRoute: ) -> None: with respx.mock(assert_all_called=False) as upstream: route = upstream.post(TRANSCRIBE_UPSTREAM) - response = transcribe_client.post("/transcribe/StartTranscriptionJob", json={**dict(self.START_JOB_BODY), **body}) + response = transcribe_client.post( + "/transcribe/StartTranscriptionJob", json={**dict(self.START_JOB_BODY), **body} + ) assert response.status_code == 403 assert member in response.json()["detail"] @@ -6873,9 +7069,7 @@ class TestAzureBodyModelGroupRelay: AZURE_SPEECH_SHORT_AUDIO_ENDPOINT: Final = "/speech/recognition/conversation/cognitiveservices/v1" AZURE_SPEECH_BATCH_ENDPOINT: Final = "/speechtotext/v3.2/transcriptions" AZURE_SPEECH_FAST_ENDPOINT: Final = "/speechtotext/transcriptions:transcribe" -AZURE_SPEECH_PCM16_HEADER: Final = ( - b"RIFF\x24\x0c\x00\x00WAVEfmt \x10\x00\x00\x00\x01\x00\x01\x00\x80\x3e\x00\x00\x00\x7d\x00\x00\x02\x00\x10\x00data\x00\x0c\x00\x00" -) +AZURE_SPEECH_PCM16_HEADER: Final = b"RIFF\x24\x0c\x00\x00WAVEfmt \x10\x00\x00\x00\x01\x00\x01\x00\x80\x3e\x00\x00\x00\x7d\x00\x00\x02\x00\x10\x00data\x00\x0c\x00\x00" AZURE_SPEECH_WAV_BYTES: Final = AZURE_SPEECH_PCM16_HEADER + b"\x00" * 3072 AZURE_SPEECH_WAV_SECONDS: Final = 3072 / (16000 * 2) AZURE_SPEECH_NON_UTF8_WAV_BYTES: Final = AZURE_SPEECH_PCM16_HEADER + bytes(range(256)) * 12 @@ -6912,9 +7106,9 @@ class TestAzureSpeechProxyRoute: def test_short_audio_forwards_raw_wav_bytes_with_server_key(self, azure_speech_client: TestClient) -> None: with respx.mock(assert_all_called=True) as upstream: - route = upstream.post( - f"https://eastus.stt.speech.microsoft.com{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}" - ).mock(return_value=httpx.Response(200, json=AZURE_SPEECH_TRANSCRIPT)) + route = upstream.post(f"https://eastus.stt.speech.microsoft.com{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}").mock( + return_value=httpx.Response(200, json=AZURE_SPEECH_TRANSCRIPT) + ) response = azure_speech_client.post( f"/azure_speech{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}", @@ -7390,9 +7584,9 @@ class TestAzureSpeechRawBodyThroughRealAuth: self, monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture, body: bytes ) -> None: with respx.mock(assert_all_called=True) as upstream, caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"): - route = upstream.post( - f"https://eastus.stt.speech.microsoft.com{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}" - ).mock(return_value=httpx.Response(200, json=AZURE_SPEECH_TRANSCRIPT)) + route = upstream.post(f"https://eastus.stt.speech.microsoft.com{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}").mock( + return_value=httpx.Response(200, json=AZURE_SPEECH_TRANSCRIPT) + ) response = self._post_wav( monkeypatch, f"/azure_speech{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}", "sk-master-key", body=body @@ -7416,9 +7610,9 @@ class TestAzureSpeechRawBodyThroughRealAuth: ) -> None: boundary: Final = "lit7939boundary" multipart_body: Final = ( - f"--{boundary}\r\nContent-Disposition: form-data; name=\"definition\"\r\n\r\n".encode() + f'--{boundary}\r\nContent-Disposition: form-data; name="definition"\r\n\r\n'.encode() + json.dumps({"locales": ["en-US"]}).encode() - + f"\r\n--{boundary}\r\nContent-Disposition: form-data; name=\"audio\"; filename=\"eagle.wav\"\r\n" + + f'\r\n--{boundary}\r\nContent-Disposition: form-data; name="audio"; filename="eagle.wav"\r\n' "Content-Type: audio/wav\r\n\r\n".encode() + AZURE_SPEECH_NON_UTF8_WAV_BYTES + f"\r\n--{boundary}--\r\n".encode() @@ -8170,9 +8364,7 @@ class TestTinyFishProxyRoute: assert response.status_code == 200 assert json.loads(route.calls.last.request.content)["use_vault"] is True - def test_returns_401_on_missing_api_key( - self, tinyfish_client: TestClient, monkeypatch: pytest.MonkeyPatch - ) -> None: + def test_returns_401_on_missing_api_key(self, tinyfish_client: TestClient, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.delenv("TINYFISH_API_KEY") with respx.mock: diff --git a/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py index cd279a6ed4b..a6ce8ee6fbc 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -5,7 +5,7 @@ import logging import os import sys import zlib -from collections.abc import Callable, Mapping +from collections.abc import AsyncIterator, Callable, Mapping from contextlib import ExitStack, contextmanager from dataclasses import dataclass from io import BytesIO @@ -27,6 +27,7 @@ from starlette.datastructures import UploadFile as StarletteUploadFile import litellm from litellm._logging import verbose_proxy_logger from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS +from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy._lazy_features import LazyFeature, attach_lazy_features @@ -55,6 +56,7 @@ from litellm.proxy.pass_through_endpoints.success_handler import ( ) from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.types import utils as types_utils +from litellm.types.guardrails import GuardrailEventHooks from litellm.types.passthrough_endpoints.pass_through_endpoints import ( LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, @@ -1638,6 +1640,223 @@ async def test_pass_through_request_streaming_marks_logging_obj_as_stream(): assert logging_obj.model_call_details["stream"] is True +@pytest.mark.asyncio +@pytest.mark.parametrize("body_stream", [None, True], ids=["stream-absent", "stream-true"]) +async def test_pass_through_request_preserves_caller_streaming_request_field(body_stream): + captured_hook_data: dict[str, object] = {} + + async def capture_pre_call(user_api_key_dict, data, call_type, endpoint_type: EndpointType): + captured_hook_data.update(data) + return data + + with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging: + with patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client" + ) as mock_get_client: + with patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.PassThroughStreamingHandler.chunk_processor" + ) as mock_chunk_processor: + mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=capture_pre_call) + mock_proxy_logging.post_call_failure_hook = AsyncMock() + mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={}) + + upstream_response = MagicMock() + upstream_response.status_code = 200 + upstream_response.headers = {} + upstream_response.raise_for_status = MagicMock() + + async_client = MagicMock() + async_client.build_request = MagicMock(return_value=MagicMock()) + async_client.send = AsyncMock(return_value=upstream_response) + mock_get_client.return_value = MagicMock(client=async_client) + + async def _empty_chunks(*args, **kwargs): + return + yield # pragma: no cover + + mock_chunk_processor.return_value = _empty_chunks() + + mock_request = MagicMock(spec=Request) + mock_request.method = "POST" + mock_request.url = httpx.URL( + "http://test-proxy.com/gemini/v1beta/models/gemini-pro:streamGenerateContent" + ) + mock_request.scope = {"path": "/gemini/v1beta/models/gemini-pro:streamGenerateContent"} + request_body: Final = { + "contents": [{"parts": [{"text": "hi"}]}], + "is_streaming_request": "caller-value", + **({"stream": True} if body_stream is True else {}), + } + mock_request.body = AsyncMock(return_value=json.dumps(request_body).encode()) + mock_request.headers = Headers({"content-type": "application/json"}) + mock_request.query_params = QueryParams({}) + + await pass_through_request( + request=mock_request, + target="https://generativelanguage.googleapis.com/v1beta/models/gemini-pro:streamGenerateContent", + custom_headers={}, + user_api_key_dict=MagicMock(), + stream=True, + ) + + assert captured_hook_data.get("is_streaming_request") == "caller-value" + assert captured_hook_data.get("stream") is body_stream + + upstream_json = async_client.build_request.call_args.kwargs["json"] + assert upstream_json["is_streaming_request"] == "caller-value" + assert "litellm_server_streaming_classification" not in upstream_json + assert upstream_json["contents"] == request_body["contents"] + + +@pytest.mark.asyncio +async def test_streaming_pass_through_drops_marker_after_hook_rebuilds_body_from_json(): + async def json_rebuilding_pre_call(user_api_key_dict, data, call_type, endpoint_type: EndpointType): + rebuilt = json.loads(json.dumps({k: v for k, v in data.items() if k != "litellm_logging_obj"})) + return {**rebuilt, "litellm_logging_obj": data["litellm_logging_obj"]} + + with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging: + with patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client" + ) as mock_get_client: + with patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.PassThroughStreamingHandler.chunk_processor" + ) as mock_chunk_processor: + mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=json_rebuilding_pre_call) + mock_proxy_logging.post_call_failure_hook = AsyncMock() + mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={}) + + upstream_response = MagicMock() + upstream_response.status_code = 200 + upstream_response.headers = {} + upstream_response.raise_for_status = MagicMock() + + async_client = MagicMock() + async_client.build_request = MagicMock(return_value=MagicMock()) + async_client.send = AsyncMock(return_value=upstream_response) + mock_get_client.return_value = MagicMock(client=async_client) + + async def _empty_chunks(*args, **kwargs): + return + yield # pragma: no cover + + mock_chunk_processor.return_value = _empty_chunks() + + mock_request = MagicMock(spec=Request) + mock_request.method = "POST" + mock_request.url = httpx.URL("http://test-proxy.com/openai/v1/chat/completions") + mock_request.scope = {"path": "/openai/v1/chat/completions"} + request_body: Final = { + "model": "gpt-5-mini", + "stream": True, + "messages": [{"role": "user", "content": "hi"}], + } + mock_request.body = AsyncMock(return_value=json.dumps(request_body).encode()) + mock_request.headers = Headers({"content-type": "application/json"}) + mock_request.query_params = QueryParams({}) + + await pass_through_request( + request=mock_request, + target="https://api.openai.com/v1/chat/completions", + custom_headers={}, + user_api_key_dict=MagicMock(), + ) + + upstream_json = async_client.build_request.call_args.kwargs["json"] + assert upstream_json == request_body, upstream_json + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("route_stream", "body_stream", "expected_streaming"), + [ + (True, False, False), + (False, True, True), + (True, None, True), + (None, None, False), + ], + ids=["body-disables-route-stream", "body-enables-streaming", "route-enables-absent-body", "both-absent"], +) +async def test_passthrough_guardrails_follow_effective_relay_stream_decision( + route_stream: bool | None, + body_stream: bool | None, + expected_streaming: bool, +): + async def return_pre_call_data( + user_api_key_dict: UserAPIKeyAuth, + data: dict[str, object], + call_type: str, + endpoint_type: EndpointType, + ) -> dict[str, object]: + return dict(data) + + with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging: + with patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client" + ) as mock_get_client: + with patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.PassThroughStreamingHandler.chunk_processor" + ) as mock_chunk_processor: + mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=return_pre_call_data) + mock_proxy_logging.post_call_failure_hook = AsyncMock() + mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={}) + + upstream_response = MagicMock() + upstream_response.status_code = 200 + upstream_response.headers = {} + upstream_response.aread = AsyncMock(return_value=b"{}") + upstream_response.text = "{}" + upstream_response.raise_for_status = MagicMock() + + async_client = MagicMock() + async_client.build_request = MagicMock(return_value=MagicMock()) + async_client.send = AsyncMock(return_value=upstream_response) + async_client.request = AsyncMock(return_value=upstream_response) + mock_get_client.return_value = MagicMock(client=async_client) + + async def empty_chunks() -> AsyncIterator[bytes]: + yield b"" + + mock_chunk_processor.return_value = empty_chunks() + + request_body: Final = { + "message": "hello", + **({"stream": body_stream} if body_stream is not None else {}), + } + mock_request = MagicMock(spec=Request) + mock_request.method = "POST" + mock_request.url = httpx.URL("http://test-proxy.com/guardrail-stream-scope") + mock_request.scope = {"path": "/guardrail-stream-scope"} + mock_request.body = AsyncMock(return_value=json.dumps(request_body).encode()) + mock_request.headers = Headers({"content-type": "application/json"}) + mock_request.query_params = QueryParams({}) + + await pass_through_request( + request=mock_request, + target="http://upstream.test/guardrail-stream-scope", + custom_headers={}, + user_api_key_dict=MagicMock(), + stream=route_stream, + ) + hook_data: Final[dict[str, object]] = mock_proxy_logging.pre_call_hook.call_args.kwargs["data"] + + streaming_guardrail: Final = CustomGuardrail( + guardrail_name="streaming-only", + default_on=True, + event_hook=GuardrailEventHooks.pre_call, + stream_scope="streaming", + ) + non_streaming_guardrail: Final = CustomGuardrail( + guardrail_name="non-streaming-only", + default_on=True, + event_hook=GuardrailEventHooks.pre_call, + stream_scope="non_streaming", + ) + assert streaming_guardrail.should_run_guardrail(hook_data, GuardrailEventHooks.pre_call) is expected_streaming + assert ( + non_streaming_guardrail.should_run_guardrail(hook_data, GuardrailEventHooks.pre_call) is not expected_streaming + ) + + @pytest.mark.asyncio async def test_pass_through_request_sse_response_marks_logging_obj_as_stream(): """ diff --git a/tests/unit/proxy/policy_engine/test_pipeline_executor.py b/tests/unit/proxy/policy_engine/test_pipeline_executor.py index 6aa7eca0f15..d13abbd379b 100644 --- a/tests/unit/proxy/policy_engine/test_pipeline_executor.py +++ b/tests/unit/proxy/policy_engine/test_pipeline_executor.py @@ -96,6 +96,20 @@ class AlwaysPassGuardrail(CustomGuardrail): return None +class StreamScopedPassGuardrail(CustomGuardrail): + def __init__(self, guardrail_name: str, stream_scope: object): + super().__init__( + guardrail_name=guardrail_name, + event_hook="pre_call", + default_on=True, + stream_scope=stream_scope, + ) + self.calls = 0 + + async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type): + self.calls += 1 + + class PassthroughBlockGuardrail(CustomGuardrail): """Mock guardrail that blocks using the legacy passthrough contract.""" @@ -378,6 +392,46 @@ async def test_block_carries_original_guardrail_exception(monkeypatch): assert result.original_exception.detail == "Content policy violation" +@pytest.mark.asyncio +async def test_pipeline_step_honors_stream_scope(monkeypatch): + stream_only = StreamScopedPassGuardrail(guardrail_name="stream-only", stream_scope="streaming") + later = AlwaysFailGuardrail(guardrail_name="later-block") + monkeypatch.setattr(litellm, "callbacks", [stream_only, later]) + steps = [ + PipelineStep(guardrail="stream-only", on_fail="block", on_pass="allow"), + PipelineStep(guardrail="later-block", on_fail="block", on_pass="allow"), + ] + + skipped = await PipelineExecutor.execute_steps( + steps=steps, + mode="pre_call", + data={"messages": [{"role": "user", "content": "hi"}]}, + user_api_key_dict=MagicMock(), + call_type="completion", + policy_name="stream-scope", + ) + assert stream_only.calls == 0 + assert later.calls == 1 + assert skipped.step_results[0].outcome == "skip" + assert skipped.step_results[0].action_taken == "next" + assert skipped.terminal_action == "block" + + stream_only.calls = 0 + later.calls = 0 + ran = await PipelineExecutor.execute_steps( + steps=steps, + mode="pre_call", + data={"messages": [{"role": "user", "content": "hi"}], "stream": True}, + user_api_key_dict=MagicMock(), + call_type="completion", + policy_name="stream-scope", + ) + assert stream_only.calls == 1 + assert later.calls == 0 + assert ran.step_results[0].outcome == "pass" + assert ran.terminal_action == "allow" + + @pytest.mark.asyncio async def test_unsupported_mode_yields_error_outcome_without_exception(monkeypatch): """An unexpected hook mode must surface as an error outcome (carrying no diff --git a/tests/unit/proxy/test_litellm_pre_call_utils.py b/tests/unit/proxy/test_litellm_pre_call_utils.py index c6c473ae71e..818ca50fed9 100644 --- a/tests/unit/proxy/test_litellm_pre_call_utils.py +++ b/tests/unit/proxy/test_litellm_pre_call_utils.py @@ -16,6 +16,7 @@ from pydantic import ValidationError as PydanticValidationError from starlette.datastructures import Headers import litellm +from litellm.constants import SERVER_STREAMING_CLASSIFICATION_KEY, SERVER_STREAMING_CLASSIFICATION_MARKER from litellm.proxy._types import AddTeamCallback, ProxyException, TeamCallbackMetadata, UserAPIKeyAuth from litellm.proxy.litellm_pre_call_utils import ( KeyAndTeamLoggingSettings, @@ -850,8 +851,13 @@ def test_initial_snapshot_refresh_clears_a_previous_guardrail_checkpoint() -> No from litellm.proxy.litellm_pre_call_utils import refresh_proxy_server_request_body_snapshot logging_obj: Final = Logging( - model="test-model", messages=[], stream=False, call_type="acompletion", - start_time=datetime.now(), litellm_call_id="new-request", function_id="new-request", + model="test-model", + messages=[], + stream=False, + call_type="acompletion", + start_time=datetime.now(), + litellm_call_id="new-request", + function_id="new-request", ) logging_obj.shadow_eval_request_snapshot = GuardrailRequestSnapshot.capture( {"messages": [{"role": "user", "content": "previous request"}]}, @@ -913,8 +919,14 @@ async def test_post_guardrail_snapshot_preserves_logging_only_masking_in_spend_l } data: Final = {"messages": messages, "metadata": metadata, "proxy_server_request": {}} logging_obj: Final = Logging( - model="test-model", messages=messages, stream=False, call_type="acompletion", - start_time=datetime.now(), litellm_call_id="mask-spend", function_id="mask-spend", kwargs=data, + model="test-model", + messages=messages, + stream=False, + call_type="acompletion", + start_time=datetime.now(), + litellm_call_id="mask-spend", + function_id="mask-spend", + kwargs=data, ) data["litellm_logging_obj"] = logging_obj refresh_proxy_server_request_body_snapshot(data, guardrails_applied=True) @@ -928,9 +940,13 @@ async def test_post_guardrail_snapshot_preserves_logging_only_masking_in_spend_l kwargs, _ = await guardrail.async_logging_hook( kwargs=logging_obj.model_call_details, result=None, call_type="acompletion" ) - stored: Final = json.loads(_get_proxy_server_request_for_spend_logs_payload( - metadata={}, litellm_params=kwargs["litellm_params"], kwargs=kwargs, - )) + stored: Final = json.loads( + _get_proxy_server_request_for_spend_logs_payload( + metadata={}, + litellm_params=kwargs["litellm_params"], + kwargs=kwargs, + ) + ) assert kwargs["messages"] == [{"role": "user", "content": "email [EMAIL]"}] assert stored["messages"] == kwargs["messages"] @@ -1089,6 +1105,7 @@ async def test_add_litellm_data_to_request_strips_user_control_fields(): "messages": [{"role": "user", "content": "hello"}], "mock_response": "free response", "mock_tool_calls": [{"id": "call_1"}], + "is_streaming_request": "caller-value", "disable_global_guardrails": True, "enable_prompt_caching": True, "routing_decision": {"cause": "forged", "routed_model": "spoofed"}, @@ -1114,6 +1131,7 @@ async def test_add_litellm_data_to_request_strips_user_control_fields(): assert "enable_prompt_caching" not in updated assert "routing_decision" not in updated assert "litellm_gateway_injected_cache" not in updated + assert updated["is_streaming_request"] == "caller-value" assert "weights" not in updated assert "_router_weights" not in updated assert "weights" not in updated["proxy_server_request"]["body"] @@ -8803,20 +8821,27 @@ async def test_mcp_credentials_only_removed_from_logging_copies(path: str, custo request.headers = Headers(request.headers) settings: Final = {"mcp_client_side_auth_header_name": custom_auth, "user_header_name": "x-user-id"} server: Final = MCPServer( - server_id="header-test", name="header-test", transport="http", url="https://example.com/mcp", + server_id="header-test", + name="header-test", + transport="http", + url="https://example.com/mcp", extra_headers=["x-service-token", "x-user-id"], ) with ( patch("litellm.proxy.proxy_server.general_settings", settings), patch.dict( "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.config_mcp_servers", - {"header-test": server}, clear=True, + {"header-test": server}, + clear=True, ), ): updated: Final = await add_litellm_data_to_request( data={"model": "test-model", "messages": [{"role": "user", "content": "hello"}]}, - request=request, user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), - proxy_config=MagicMock(), general_settings=settings, version="test", + request=request, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + proxy_config=MagicMock(), + general_settings=settings, + version="test", ) for header_dict in _all_header_dicts(updated, metadata_name): assert not any(value in json.dumps(header_dict) for value in secrets.values()) @@ -8836,7 +8861,10 @@ def test_signoz_callback_vars_are_scoped_to_the_signoz_callback(): data=AddTeamCallback( callback_name="signoz", callback_type="success", - callback_vars={"signoz_ingestion_key": "team-key", "signoz_ingestion_endpoint": "https://ingest.eu.signoz.cloud:443"}, + callback_vars={ + "signoz_ingestion_key": "team-key", + "signoz_ingestion_endpoint": "https://ingest.eu.signoz.cloud:443", + }, ), team_callback_settings_obj=None, ) @@ -8856,6 +8884,54 @@ def test_signoz_callback_vars_are_scoped_to_the_signoz_callback(): assert under_other.callback_vars == {"langfuse_host": "https://cloud.langfuse.com"} +def test_body_snapshot_excludes_the_server_streaming_marker() -> None: + from litellm.constants import SERVER_STREAMING_CLASSIFICATION_KEY, SERVER_STREAMING_CLASSIFICATION_MARKER + from litellm.proxy.litellm_pre_call_utils import refresh_proxy_server_request_body_snapshot + + proxy_request: Final = {"body": {}} + data: Final = { + "messages": [{"role": "user", "content": "hi"}], + SERVER_STREAMING_CLASSIFICATION_KEY: SERVER_STREAMING_CLASSIFICATION_MARKER, + "proxy_server_request": proxy_request, + } + + refresh_proxy_server_request_body_snapshot(data) + + assert proxy_request == {"body": {"messages": [{"role": "user", "content": "hi"}]}} + + +@pytest.mark.parametrize( + "marker", + [ + SERVER_STREAMING_CLASSIFICATION_MARKER, + json.loads(json.dumps(SERVER_STREAMING_CLASSIFICATION_MARKER)), + ], + ids=["enum", "json-string"], +) +def test_body_snapshot_drops_only_the_marker_and_keeps_caller_value(marker: str) -> None: + from litellm.proxy.litellm_pre_call_utils import refresh_proxy_server_request_body_snapshot + + marker_request: Final = {"body": {}} + marker_data: Final = { + "messages": [{"role": "user", "content": "hi"}], + SERVER_STREAMING_CLASSIFICATION_KEY: marker, + "proxy_server_request": marker_request, + } + + refresh_proxy_server_request_body_snapshot(marker_data) + + assert SERVER_STREAMING_CLASSIFICATION_KEY not in marker_request["body"], marker_request + + caller_request: Final = {"body": {}} + caller_data: Final = { + "messages": [{"role": "user", "content": "hi"}], + SERVER_STREAMING_CLASSIFICATION_KEY: "caller-value", + "proxy_server_request": caller_request, + } + + refresh_proxy_server_request_body_snapshot(caller_data) + + assert caller_request["body"][SERVER_STREAMING_CLASSIFICATION_KEY] == "caller-value", caller_request def test_arize_otlp_protocol_on_a_key_logging_entry_reaches_the_destination(monkeypatch): from litellm.integrations.otel.model.config import is_otel_v2_enabled from litellm.proxy.litellm_pre_call_utils import resolve_tenant_otel_destinations diff --git a/tests/unit/proxy/test_pricing_field_strip.py b/tests/unit/proxy/test_pricing_field_strip.py index a0e25e91f37..395fdc3a16e 100644 --- a/tests/unit/proxy/test_pricing_field_strip.py +++ b/tests/unit/proxy/test_pricing_field_strip.py @@ -26,7 +26,6 @@ from litellm.proxy.litellm_pre_call_utils import ( from litellm.types.utils import CustomPricingLiteLLMParams - def _make_request_mock() -> Request: request_mock = MagicMock(spec=Request) request_mock.url.path = "/v1/chat/completions" @@ -58,9 +57,7 @@ class TestStripClientPricingOverrides: # The strip set is built from the model so additions are picked up # automatically — this test guards against the model and the strip # set drifting apart if someone replaces the auto-derivation later. - assert _CLIENT_PRICING_CONTROL_FIELDS == frozenset( - CustomPricingLiteLLMParams.model_fields.keys() - ) + assert _CLIENT_PRICING_CONTROL_FIELDS == frozenset(CustomPricingLiteLLMParams.model_fields.keys()) # Sanity: the obvious top-level pricing fields are in the set. for field in ( "input_cost_per_token", @@ -184,9 +181,7 @@ class TestStripClientPricingOverrides: verbose_proxy_logger.setLevel(logging.DEBUG) with caplog.at_level(logging.DEBUG, logger=verbose_proxy_logger.name): _strip_client_pricing_overrides({"model": "gpt-4", "temperature": 0.7}) - assert not any( - "pricing" in record.getMessage().lower() for record in caplog.records - ) + assert not any("pricing" in record.getMessage().lower() for record in caplog.records) @pytest.mark.asyncio @@ -211,6 +206,26 @@ async def test_add_litellm_data_to_request_strips_root_pricing_fields(): assert "output_cost_per_token" not in updated +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_preserves_caller_streaming_request(): + data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "hi"}], + "is_streaming_request": True, + } + + updated = await add_litellm_data_to_request( + data=data, + request=_make_request_mock(), + user_api_key_dict=_user_api_key_auth(), + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated["is_streaming_request"] is True + + @pytest.mark.asyncio async def test_add_litellm_data_to_request_strips_client_disconnect_metadata(): data = { @@ -318,9 +333,7 @@ async def test_add_litellm_data_to_request_skips_strip_with_team_opt_in(): "input_cost_per_token": 0.0001, } - user_auth = _user_api_key_auth( - team_metadata={"allow_client_pricing_override": True} - ) + user_auth = _user_api_key_auth(team_metadata={"allow_client_pricing_override": True}) updated = await add_litellm_data_to_request( data=data, request=_make_request_mock(), diff --git a/tests/unit/types/test_guardrails_case_normalization.py b/tests/unit/types/test_guardrails_case_normalization.py index 26c1d395320..8c8192be3be 100644 --- a/tests/unit/types/test_guardrails_case_normalization.py +++ b/tests/unit/types/test_guardrails_case_normalization.py @@ -2,12 +2,19 @@ Test case normalization in LitellmParams for all guardrail types """ -from typing import Literal +import logging +from typing import Final, Literal import pytest from pydantic import ValidationError -from litellm.types.guardrails import BaseLitellmParams, LitellmParams +from litellm.types.guardrails import ( + BaseLitellmParams, + LitellmParams, + runtime_stream_scope, + stored_stream_scope, + with_tolerated_stream_scope, +) class TestLitellmParamsCaseNormalization: @@ -184,3 +191,80 @@ class TestSensitiveDataRoutingValidation: on_sensitive_data="BLOCK", ) assert params.on_sensitive_data == "block" + + +class TestStreamScopeValidation: + def test_scalar_is_case_normalized(self): + params = LitellmParams(guardrail="bedrock", mode="post_call", stream_scope="Streaming") + assert params.stream_scope == "streaming" + + def test_map_keys_and_values_are_normalized(self): + params = LitellmParams( + guardrail="bedrock", + mode=["pre_call", "post_call"], + stream_scope={"Pre_Call": "Both", "POST_CALL": "Non_Streaming"}, + ) + assert params.stream_scope == {"pre_call": "both", "post_call": "non_streaming"} + + def test_invalid_scalar_is_rejected(self): + with pytest.raises(ValidationError, match="stream_scope must be one of"): + LitellmParams(guardrail="bedrock", mode="post_call", stream_scope="chunks") + + def test_invalid_map_key_is_rejected(self): + with pytest.raises(ValidationError, match="stream_scope keys must be guardrail modes"): + LitellmParams(guardrail="bedrock", mode="post_call", stream_scope={"not_a_mode": "both"}) + + def test_invalid_map_value_is_rejected(self): + with pytest.raises(ValidationError, match="stream_scope must be one of"): + LitellmParams(guardrail="bedrock", mode="post_call", stream_scope={"post_call": "sometimes"}) + + def test_runtime_stream_scope_normalizes_direct_constructor_maps(self): + default, by_hook = runtime_stream_scope({"Pre_Call": "streaming"}) + assert default == "both" + assert dict(by_hook) == {"pre_call": "streaming"} + + def test_runtime_stream_scope_rejects_invalid_direct_input(self): + with pytest.raises(ValueError, match="stream_scope must be one of"): + runtime_stream_scope("chunks") + + @pytest.mark.parametrize( + "value, expected", + [ + ("streaming", "streaming"), + ("non_streaming", "non_streaming"), + ("both", "both"), + ({"pre_call": "streaming"}, {"pre_call": "streaming"}), + ("sometimes", None), + ({"pre_call": "sometimes"}, None), + ], + ) + def test_stored_stream_scope_tolerates_invalid_values( + self, + value: object, + expected: object, + caplog: pytest.LogCaptureFixture, + ) -> None: + caplog.set_level(logging.WARNING) + + result: Final = stored_stream_scope(value) + + assert result == expected + if expected is None: + assert f"Ignoring invalid stored stream_scope value of type {type(value).__name__}" in caplog.text + assert "sometimes" not in caplog.text + + def test_tolerated_stream_scope_rewrites_only_the_scope_field(self) -> None: + params: Final = { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "stream_scope": "sometimes", + } + + tolerated: Final = with_tolerated_stream_scope(params) + + assert tolerated == { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "stream_scope": None, + } + assert params["stream_scope"] == "sometimes" diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailReadOnlyDetails.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailReadOnlyDetails.tsx new file mode 100644 index 00000000000..5fe88b4870c --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailReadOnlyDetails.tsx @@ -0,0 +1,68 @@ +import { Badge } from "@/components/ui/badge"; +import { GuardrailModeRows } from "./GuardrailModeDisplay"; +import { GuardrailStreamScopeDetail } from "./StreamScopeFields"; +import ToolPermissionRulesEditor, { type ToolPermissionConfig } from "./tool_permission/ToolPermissionRulesEditor"; + +export const GuardrailReadOnlyDetails = ({ + guardrailId, + guardrailName, + displayName, + litellmParams, + streamScope, + defaultOn, + piiEntityCount, + createdAt, + updatedAt, + showToolPermission, + toolPermissionConfig, +}: { + guardrailId: string; + guardrailName: string; + displayName: string; + litellmParams: { mode?: unknown; logging_only_scope?: string | null }; + streamScope: unknown; + defaultOn: boolean | undefined; + piiEntityCount: number; + createdAt: string; + updatedAt: string; + showToolPermission: boolean; + toolPermissionConfig: ToolPermissionConfig; +}) => ( +
+
+

Guardrail ID

+
{guardrailId}
+
+
+

Guardrail Name

+
{guardrailName || "Unnamed Guardrail"}
+
+
+

Provider

+
{displayName}
+
+ + +
+

Default On

+ {defaultOn ? "Yes" : "No"} +
+ {piiEntityCount > 0 && ( +
+

PII Protection

+
+ {piiEntityCount} PII entities configured +
+
+ )} +
+

Created At

+
{createdAt}
+
+
+

Last Updated

+
{updatedAt}
+
+ {showToolPermission && } +
+); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/StreamScopeFields.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/StreamScopeFields.tsx new file mode 100644 index 00000000000..1a8f55ce477 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/StreamScopeFields.tsx @@ -0,0 +1,95 @@ +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { GuardrailField, labelWithHint, type GuardrailFormControl } from "./GuardrailFormField"; +import { + STREAM_SCOPE_OPTIONS, + formatGuardrailStreamScope, + type GuardrailStreamScope, + isGuardrailStreamScope, +} from "./guardrail_info_helpers"; + +const STREAM_SCOPE_ITEMS = STREAM_SCOPE_OPTIONS.map((option) => ({ + label: option.label, + value: option.value, +})); + +const REQUEST_SHAPE_HINT = + "Run this guardrail on streaming requests, non-streaming requests, or both, for each selected mode."; + +export const StreamScopeFields = ({ + modes, + value, + onChange, +}: { + modes: string[]; + value: Record; + onChange: (next: Record) => void; +}) => { + if (modes.length === 0) return null; + + return ( +
+ {modes.map((mode) => { + const selected = value[mode] ?? "both"; + return ( +
+ + +
+ ); + })} +
+ ); +}; + +export const StreamScopeFormField = ({ control, modes }: { control: GuardrailFormControl; modes: string[] }) => ( + + {({ value, onChange }) => ( + | undefined) ?? {}} + onChange={onChange} + /> + )} + +); + +export const GuardrailStreamScopeCaption = ({ raw }: { raw: unknown }) => { + const label = formatGuardrailStreamScope(raw); + if (!label) return null; + return

{label}

; +}; + +export const GuardrailStreamScopeDetail = ({ raw }: { raw: unknown }) => { + const label = formatGuardrailStreamScope(raw); + if (!label) return null; + return ( +
+

Request shape

+
{label}
+
+ ); +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.characterization.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.characterization.test.tsx index 230be13ef33..9628c284bd2 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.characterization.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.characterization.test.tsx @@ -89,6 +89,23 @@ describe("AddGuardrailForm create payload characterization", () => { }); }); + it("sends stream_scope when a mode is restricted to streaming requests", async () => { + const user = userEvent.setup({ delay: null }); + renderForm(); + + await user.type(await screen.findByLabelText("Guardrail Name"), "my-bedrock"); + await pickProvider(user, "Bedrock Guardrail"); + await chooseSelectOption(user, screen.getByLabelText("pre_call applies to"), "Streaming only"); + await user.type(await screen.findByPlaceholderText("The guardrail id on Bedrock"), "gr-123"); + await user.click(screen.getByRole("button", { name: "Next" })); + await user.click(await screen.findByRole("button", { name: "Create Guardrail" })); + + await waitFor(() => expect(networking.createGuardrailCall).toHaveBeenCalledTimes(1)); + expect(payload()).toMatchObject({ + litellm_params: { stream_scope: "streaming" }, + }); + }); + it("switches mode from the seeded string to an array once the user touches the multi select", async () => { const user = userEvent.setup({ delay: null }); renderForm(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx index 43a542289f6..098245fa34b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx @@ -23,11 +23,14 @@ import { shouldRenderContentFilterConfigSettings, shouldRenderLLMJudgeFields, shouldRenderPIIConfigSettings, - supportsDirectionalLoggingOnlyScope, + streamScopePayload, toModeArray, + type GuardrailStreamScope, + supportsDirectionalLoggingOnlyScope, type LoggingOnlyScope, type LoggingOnlyScopeChoice, } from "./guardrail_info_helpers"; +import { StreamScopeFormField } from "./StreamScopeFields"; import { Logo } from "@/components/molecules/logo/Logo"; import { MultiSelect } from "@/components/shared/MultiSelect"; import { FieldGroup } from "@/components/ui/field"; @@ -172,6 +175,7 @@ const INITIAL_VALUES: GuardrailFormValues = { logging_only_scope_choice: "default", skip_system_message_choice: "inherit", skip_tool_message_choice: "inherit", + stream_scope_by_mode: {}, }; const ALWAYS_ON_ITEMS = [ @@ -466,6 +470,14 @@ const AddGuardrailForm: React.FC = ({ visible, onClose, a guardrail_info: {}, }; + const streamScope = streamScopePayload( + toModeArray(values.mode), + (values.stream_scope_by_mode as Record | undefined) ?? {}, + ); + if (streamScope !== undefined) { + guardrailData.litellm_params.stream_scope = streamScope; + } + const skipForCreate = choiceToSkipSystemForCreate(asSkipChoice(values.skip_system_message_choice)); if (skipForCreate !== undefined) { guardrailData.litellm_params.skip_system_message_in_guardrail = skipForCreate; @@ -773,6 +785,8 @@ const AddGuardrailForm: React.FC = ({ visible, onClose, a )} + + { expect(screen.getByRole("button", { name: /update guardrail/i })).toBeInTheDocument(); }); + it("should open without crashing for a tag-scoped mode and show it read-only", async () => { + renderModal({ + editData: { + guardrail_id: "g-tag", + guardrail_name: "tag-mode-guardrail", + litellm_params: { + mode: { tags: { "team-a": "pre_call" }, default: "post_call" }, + default_on: true, + custom_code: "def apply_guardrail(): pass", + }, + }, + }); + + expect(await screen.findByText("Edit Custom Guardrail")).toBeInTheDocument(); + const modeInput = screen.getByLabelText("Mode (tag-scoped)"); + expect(modeInput).toBeDisabled(); + expect(modeInput).toHaveValue("post_call, pre_call (tag-based)"); + expect(screen.getByText("Mode (tag-scoped, read-only)")).toBeInTheDocument(); + expect(screen.queryByText(/applies to/)).not.toBeInTheDocument(); + }); + + it("should omit mode and stream_scope from the update payload for a tag-scoped guardrail", async () => { + const user = userEvent.setup(); + renderModal({ + editData: { + guardrail_id: "g-tag", + guardrail_name: "tag-mode-guardrail", + litellm_params: { + mode: { tags: { "team-a": "pre_call" }, default: "post_call" }, + default_on: true, + custom_code: "def apply_guardrail(): pass", + }, + }, + }); + + const nameInput = await screen.findByDisplayValue("tag-mode-guardrail"); + await user.clear(nameInput); + await user.type(nameInput, "renamed-guardrail"); + await user.click(screen.getByRole("button", { name: /update guardrail/i })); + + await waitFor(() => { + expect(mockUpdate).toHaveBeenCalledTimes(1); + }); + const [token, guardrailId, payload] = mockUpdate.mock.calls[0] as [ + string, + string, + Record>, + ]; + expect(token).toBe("test-token"); + expect(guardrailId).toBe("g-tag"); + expect(payload.guardrail_name).toBe("renamed-guardrail"); + expect(payload.litellm_params.custom_code).toBe("def apply_guardrail(): pass"); + expect(payload.litellm_params).not.toHaveProperty("mode"); + expect(payload.litellm_params).not.toHaveProperty("stream_scope"); + }); + it("should keep save disabled until a guardrail name is entered", async () => { const user = userEvent.setup(); renderModal(); @@ -111,8 +167,7 @@ describe("CustomCodeModal", () => { expect(await screen.findByDisplayValue(/async def apply_guardrail/)).toBeInTheDocument(); - const comboboxes = screen.getAllByRole("combobox"); - await user.click(comboboxes[comboboxes.length - 1]); + await user.click(screen.getByRole("combobox", { name: "Template" })); const options = await screen.findAllByText("Block SSN"); await user.click(options[options.length - 1]); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/CustomCodeModal.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/CustomCodeModal.tsx index 9deef4aeee3..c89540596bd 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/CustomCodeModal.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/CustomCodeModal.tsx @@ -37,152 +37,30 @@ import { import { Switch } from "@/components/ui/switch"; import { Textarea } from "@/components/ui/textarea"; import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; - -// Code templates -const CODE_TEMPLATES = { - empty: { - name: "Empty Template", - code: `async def apply_guardrail(inputs, request_data, input_type): - # inputs: {texts, images, tools, tool_calls, structured_messages, model} - # request_data: {model, user_id, team_id, end_user_id, metadata} - # input_type: "request" or "response" - return allow()`, - }, - blockSSN: { - name: "Block SSN", - code: `def apply_guardrail(inputs, request_data, input_type): - for text in inputs["texts"]: - if regex_match(text, r"\\d{3}-\\d{2}-\\d{4}"): - return block("SSN detected") - return allow()`, - }, - redactEmail: { - name: "Redact Emails", - code: `def apply_guardrail(inputs, request_data, input_type): - pattern = r"[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\\.[a-zA-Z]{2,}" - modified = [] - for text in inputs["texts"]: - modified.append(regex_replace(text, pattern, "[EMAIL REDACTED]")) - return modify(texts=modified)`, - }, - blockSQL: { - name: "Block SQL Injection", - code: `def apply_guardrail(inputs, request_data, input_type): - if input_type != "request": - return allow() - for text in inputs["texts"]: - if contains_code_language(text, ["sql"]): - return block("SQL code not allowed") - return allow()`, - }, - validateJSON: { - name: "Validate JSON", - code: `def apply_guardrail(inputs, request_data, input_type): - if input_type != "response": - return allow() - - schema = {"type": "object", "required": ["name", "value"]} - - for text in inputs["texts"]: - obj = json_parse(text) - if obj is None: - return block("Invalid JSON response") - if not json_schema_valid(obj, schema): - return block("Response missing required fields") - return allow()`, - }, - externalAPI: { - name: "External API Check (async)", - code: `async def apply_guardrail(inputs, request_data, input_type): - # Call an external moderation API (async for non-blocking) - for text in inputs["texts"]: - response = await http_post( - "https://api.example.com/moderate", - body={"text": text, "user_id": request_data["user_id"]}, - headers={"Authorization": "Bearer YOUR_API_KEY"}, - timeout=10 - ) - - if not response["success"]: - # API call failed, allow by default or block - return allow() - - if response["body"].get("flagged"): - return block(response["body"].get("reason", "Content flagged")) - - return allow()`, - }, -}; - -// Available primitives organized by category -const PRIMITIVES = { - "Return Values": [ - { name: "allow()", desc: "Let request/response through" }, - { name: "block(reason)", desc: "Reject with message" }, - { name: "flag(reason, metadata={})", desc: "Let through, record a non-blocking violation" }, - { name: "modify(texts=[], images=[], tool_calls=[])", desc: "Transform content" }, - ], - "HTTP Requests (async)": [ - { name: "await http_request(url, method, headers, body)", desc: "Make async HTTP request" }, - { name: "await http_get(url, headers)", desc: "Async GET request" }, - { name: "await http_post(url, body, headers)", desc: "Async POST request" }, - ], - "Regex Functions": [ - { name: "regex_match(text, pattern)", desc: "Returns True if pattern found" }, - { name: "regex_replace(text, pattern, replacement)", desc: "Replace all matches" }, - { name: "regex_find_all(text, pattern)", desc: "Return list of matches" }, - ], - "JSON Functions": [ - { name: "json_parse(text)", desc: "Parse JSON string, returns None on error" }, - { name: "json_stringify(obj)", desc: "Convert to JSON string" }, - { name: "json_schema_valid(obj, schema)", desc: "Validate against JSON schema" }, - ], - "URL Functions": [ - { name: "extract_urls(text)", desc: "Extract all URLs from text" }, - { name: "is_valid_url(url)", desc: "Check if URL is valid" }, - { name: "all_urls_valid(text)", desc: "Check all URLs in text are valid" }, - ], - "Code Detection": [ - { name: "detect_code(text)", desc: "Returns True if code detected" }, - { name: "detect_code_languages(text)", desc: "Returns list of detected languages" }, - { name: 'contains_code_language(text, ["sql"])', desc: "Check for specific languages" }, - ], - "Text Utilities": [ - { name: "contains(text, substring)", desc: "Check if substring exists" }, - { name: "contains_any(text, [substr1, substr2])", desc: "Check if any substring exists" }, - { name: "word_count(text)", desc: "Count words" }, - { name: "char_count(text)", desc: "Count characters" }, - { name: "lower(text) / upper(text) / trim(text)", desc: "String transforms" }, - ], -}; - -const MODE_OPTIONS = [ - { value: "pre_call", label: "pre_call (Request)" }, - { value: "post_call", label: "post_call (Response)" }, - { value: "during_call", label: "during_call (Parallel)" }, - { value: "logging_only", label: "logging_only" }, - { value: "pre_mcp_call", label: "pre_mcp_call (Before MCP Tool Call)" }, - { value: "post_mcp_call", label: "post_mcp_call (After MCP Tool Call)" }, - { value: "during_mcp_call", label: "during_mcp_call (During MCP Tool Call)" }, -]; - -const TEMPLATE_ITEMS = Object.entries(CODE_TEMPLATES).map(([key, template]) => ({ - value: key, - label: template.name, -})); - -type ModeOption = (typeof MODE_OPTIONS)[number]; - -const MODE_OPTION_BY_VALUE: Record = Object.fromEntries( - MODE_OPTIONS.map((option) => [option.value, option]), -); +import { StreamScopeFields } from "../StreamScopeFields"; +import { + formatGuardrailMode, + streamScopeByModeFromConfig, + streamScopeForUpdate, + streamScopePayload, + type GuardrailStreamScope, +} from "../guardrail_info_helpers"; +import { + CODE_TEMPLATES, + MODE_OPTION_BY_VALUE, + MODE_OPTIONS, + PRIMITIVES, + TEMPLATE_ITEMS, + type ModeOption, +} from "./custom_code_catalog"; // Data for editing an existing guardrail + export interface EditGuardrailData { guardrail_id: string; guardrail_name: string; litellm_params: { - mode?: string | string[]; + mode?: string | string[] | Record; default_on?: boolean; custom_code?: string; logging_only_scope?: LoggingOnlyScope | null; @@ -204,6 +82,7 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS const isEditMode = !!editData; const [guardrailName, setGuardrailName] = useState(""); const [mode, setMode] = useState(["pre_call"]); + const [streamScopeByMode, setStreamScopeByMode] = useState>({}); const [loggingOnlyScopeChoice, setLoggingOnlyScopeChoice] = useState("default"); const [defaultOn, setDefaultOn] = useState(false); const [selectedTemplate, setSelectedTemplate] = useState("empty"); @@ -315,11 +194,14 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS setCode(CODE_TEMPLATES[templateKey as keyof typeof CODE_TEMPLATES].code); }; - // Normalize mode from API (string or string[]) to string[] - const normalizeMode = (m: string | string[] | undefined): string[] => { + // Normalize mode from API (string or string[]) to string[]. + // A tag-scoped mode dict ({ tags, default }) is managed outside this editor, so it + // contributes no editable modes and is displayed read-only instead. + const normalizeMode = (m: string | string[] | Record | undefined): string[] => { if (m === undefined || m === null) return ["pre_call"]; if (Array.isArray(m)) return m.length ? m : ["pre_call"]; - return [m]; + if (typeof m === "string") return [m]; + return []; }; // Reset form when modal opens or editData changes @@ -329,6 +211,12 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS // Edit mode: populate with existing data setGuardrailName(editData.guardrail_name || ""); setMode(normalizeMode(editData.litellm_params?.mode)); + setStreamScopeByMode( + streamScopeByModeFromConfig( + editData.litellm_params?.stream_scope, + normalizeMode(editData.litellm_params?.mode), + ), + ); setLoggingOnlyScopeChoice(loggingOnlyScopeToChoice(editData.litellm_params?.logging_only_scope)); setDefaultOn(editData.litellm_params?.default_on || false); setCode(editData.litellm_params?.custom_code || CODE_TEMPLATES.empty.code); @@ -337,6 +225,7 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS // Create mode: reset to defaults setGuardrailName(""); setMode(["pre_call"]); + setStreamScopeByMode({}); setLoggingOnlyScopeChoice("default"); setDefaultOn(false); setSelectedTemplate("empty"); @@ -411,11 +300,21 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS if (defaultOn !== editData.litellm_params?.default_on) { updateData.litellm_params.default_on = defaultOn; } + const nextStreamScope = streamScopeForUpdate( + mode, + streamScopeByMode, + editData.litellm_params?.stream_scope, + existingMode, + ); + if (nextStreamScope !== undefined) { + updateData.litellm_params.stream_scope = nextStreamScope; + } await updateGuardrailCall(accessToken, editData.guardrail_id, updateData); toast.success("Custom code guardrail updated successfully"); } else { // Create new guardrail + const streamScope = streamScopePayload(mode, streamScopeByMode); const guardrailData = { guardrail_name: guardrailName, litellm_params: { @@ -423,6 +322,7 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS mode: mode, default_on: defaultOn, custom_code: code, + ...(streamScope !== undefined ? { stream_scope: streamScope } : {}), ...getCustomCodeLoggingOnlyScopeCreate(mode, loggingOnlyScopeChoice), }, guardrail_info: {}, @@ -511,6 +411,11 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS const lineCount = code.split("\n").length; const selectedModeOptions = mode.map((value) => MODE_OPTION_BY_VALUE[value]).filter(Boolean); + const rawEditMode = editData?.litellm_params?.mode; + const tagScopedModeLabel = + rawEditMode !== null && typeof rawEditMode === "object" && !Array.isArray(rawEditMode) + ? formatGuardrailMode(rawEditMode) || "-" + : null; return ( !open && onClose()}> @@ -533,32 +438,38 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS />
- - setMode(options.map((option) => option.value))} - multiple - > - } className="w-full"> - {selectedModeOptions.map((option) => ( - - {option.label} - - ))} - - - - No matching modes - - {(option: ModeOption) => ( - + + {tagScopedModeLabel ? ( + + ) : ( + setMode(options.map((option) => option.value))} + multiple + > + } className="w-full"> + {selectedModeOptions.map((option) => ( + {option.label} - - )} - - - + + ))} + + + + No matching modes + + {(option: ModeOption) => ( + + {option.label} + + )} + + + + )}
{mode.includes("logging_only") && ( @@ -600,6 +511,11 @@ const CustomCodeModal: React.FC = ({ visible, onClose, onS + {mode.length > 0 && ( +
+ +
+ )} {/* Main Content */}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/custom_code_catalog.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/custom_code_catalog.ts new file mode 100644 index 00000000000..aad3df181a7 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/custom_code_catalog.ts @@ -0,0 +1,136 @@ +export const CODE_TEMPLATES = { + empty: { + name: "Empty Template", + code: `async def apply_guardrail(inputs, request_data, input_type): + # inputs: {texts, images, tools, tool_calls, structured_messages, model} + # request_data: {model, user_id, team_id, end_user_id, metadata} + # input_type: "request" or "response" + return allow()`, + }, + blockSSN: { + name: "Block SSN", + code: `def apply_guardrail(inputs, request_data, input_type): + for text in inputs["texts"]: + if regex_match(text, r"\\d{3}-\\d{2}-\\d{4}"): + return block("SSN detected") + return allow()`, + }, + redactEmail: { + name: "Redact Emails", + code: `def apply_guardrail(inputs, request_data, input_type): + pattern = r"[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\\.[a-zA-Z]{2,}" + modified = [] + for text in inputs["texts"]: + modified.append(regex_replace(text, pattern, "[EMAIL REDACTED]")) + return modify(texts=modified)`, + }, + blockSQL: { + name: "Block SQL Injection", + code: `def apply_guardrail(inputs, request_data, input_type): + if input_type != "request": + return allow() + for text in inputs["texts"]: + if contains_code_language(text, ["sql"]): + return block("SQL code not allowed") + return allow()`, + }, + validateJSON: { + name: "Validate JSON", + code: `def apply_guardrail(inputs, request_data, input_type): + if input_type != "response": + return allow() + + schema = {"type": "object", "required": ["name", "value"]} + + for text in inputs["texts"]: + obj = json_parse(text) + if obj is None: + return block("Invalid JSON response") + if not json_schema_valid(obj, schema): + return block("Response missing required fields") + return allow()`, + }, + externalAPI: { + name: "External API Check (async)", + code: `async def apply_guardrail(inputs, request_data, input_type): + # Call an external moderation API (async for non-blocking) + for text in inputs["texts"]: + response = await http_post( + "https://api.example.com/moderate", + body={"text": text, "user_id": request_data["user_id"]}, + headers={"Authorization": "Bearer YOUR_API_KEY"}, + timeout=10 + ) + + if not response["success"]: + # API call failed, allow by default or block + return allow() + + if response["body"].get("flagged"): + return block(response["body"].get("reason", "Content flagged")) + + return allow()`, + }, +}; + +export const PRIMITIVES = { + "Return Values": [ + { name: "allow()", desc: "Let request/response through" }, + { name: "block(reason)", desc: "Reject with message" }, + { name: "flag(reason, metadata={})", desc: "Let through, record a non-blocking violation" }, + { name: "modify(texts=[], images=[], tool_calls=[])", desc: "Transform content" }, + ], + "HTTP Requests (async)": [ + { name: "await http_request(url, method, headers, body)", desc: "Make async HTTP request" }, + { name: "await http_get(url, headers)", desc: "Async GET request" }, + { name: "await http_post(url, body, headers)", desc: "Async POST request" }, + ], + "Regex Functions": [ + { name: "regex_match(text, pattern)", desc: "Returns True if pattern found" }, + { name: "regex_replace(text, pattern, replacement)", desc: "Replace all matches" }, + { name: "regex_find_all(text, pattern)", desc: "Return list of matches" }, + ], + "JSON Functions": [ + { name: "json_parse(text)", desc: "Parse JSON string, returns None on error" }, + { name: "json_stringify(obj)", desc: "Convert to JSON string" }, + { name: "json_schema_valid(obj, schema)", desc: "Validate against JSON schema" }, + ], + "URL Functions": [ + { name: "extract_urls(text)", desc: "Extract all URLs from text" }, + { name: "is_valid_url(url)", desc: "Check if URL is valid" }, + { name: "all_urls_valid(text)", desc: "Check all URLs in text are valid" }, + ], + "Code Detection": [ + { name: "detect_code(text)", desc: "Returns True if code detected" }, + { name: "detect_code_languages(text)", desc: "Returns list of detected languages" }, + { name: 'contains_code_language(text, ["sql"])', desc: "Check for specific languages" }, + ], + "Text Utilities": [ + { name: "contains(text, substring)", desc: "Check if substring exists" }, + { name: "contains_any(text, [substr1, substr2])", desc: "Check if any substring exists" }, + { name: "word_count(text)", desc: "Count words" }, + { name: "char_count(text)", desc: "Count characters" }, + { name: "lower(text) / upper(text) / trim(text)", desc: "String transforms" }, + ], +}; + +export const MODE_OPTIONS = [ + { value: "pre_call", label: "pre_call (Request)" }, + { value: "post_call", label: "post_call (Response)" }, + { value: "during_call", label: "during_call (Parallel)" }, + { value: "logging_only", label: "logging_only" }, + { value: "pre_mcp_call", label: "pre_mcp_call (Before MCP Tool Call)" }, + { value: "post_mcp_call", label: "post_mcp_call (After MCP Tool Call)" }, + { value: "during_mcp_call", label: "during_mcp_call (During MCP Tool Call)" }, +]; + +export const TEMPLATE_ITEMS = Object.entries(CODE_TEMPLATES).map(([key, template]) => ({ + value: key, + label: template.name, +})); + +export type ModeOption = (typeof MODE_OPTIONS)[number]; + +export const MODE_OPTION_BY_VALUE: Record = Object.fromEntries( + MODE_OPTIONS.map((option) => [option.value, option]), +); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx index 73d9fb8123e..c67758c0d98 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx @@ -33,7 +33,8 @@ import { } from "./GuardrailFormField"; import ContentFilterManager, { formatContentFilterDataForAPI } from "./content_filter/ContentFilterManager"; import CustomCodeModal, { EditGuardrailData } from "./custom_code/CustomCodeModal"; -import { GuardrailModeCard, GuardrailModeRows } from "./GuardrailModeDisplay"; +import { GuardrailModeCard } from "./GuardrailModeDisplay"; +import { GuardrailReadOnlyDetails } from "./GuardrailReadOnlyDetails"; import { getLoggingOnlyScopeUpdate, getGuardrailLogoAndName, @@ -41,10 +42,15 @@ import { loggingOnlyScopeToChoice, skipSystemMessageToChoice, skipToolMessageToChoice, + streamScopeByModeFromConfig, + streamScopeForUpdate, supportsDirectionalLoggingOnlyScope, + toModeArray, type SkipSystemMessageChoice, type SkipToolMessageChoice, + type GuardrailStreamScope, } from "./guardrail_info_helpers"; +import { GuardrailStreamScopeCaption, StreamScopeFormField } from "./StreamScopeFields"; import GuardrailOptionalParams from "./guardrail_optional_params"; import GuardrailProviderFields from "./guardrail_provider_fields"; import PiiConfiguration from "./pii_configuration"; @@ -236,6 +242,13 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose, "skip_tool_message_choice", skipToolMessageToChoice(guardrailData.litellm_params?.skip_tool_message_in_guardrail), ); + form.setValue( + "stream_scope_by_mode", + streamScopeByModeFromConfig( + guardrailData.litellm_params?.stream_scope, + toModeArray(guardrailData.litellm_params?.mode), + ), + ); form.setValue( "guardrail_info", guardrailData.guardrail_info ? JSON.stringify(guardrailData.guardrail_info, null, 2) : "", @@ -304,6 +317,16 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose, updateData.litellm_params.default_on = values.default_on; } + const modes = toModeArray(guardrailData.litellm_params?.mode); + const nextStreamScope = streamScopeForUpdate( + modes, + (values.stream_scope_by_mode as Record | undefined) ?? {}, + guardrailData.litellm_params?.stream_scope, + ); + if (nextStreamScope !== undefined) { + updateData.litellm_params.stream_scope = nextStreamScope; + } + const prevSkipChoice = skipSystemMessageToChoice(guardrailData.litellm_params?.skip_system_message_in_guardrail); const nextSkipChoice = values.skip_system_message_choice as SkipSystemMessageChoice | undefined; if (nextSkipChoice !== undefined && nextSkipChoice !== prevSkipChoice) { @@ -566,6 +589,7 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose, +

Created At

@@ -723,6 +747,11 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose, )} + + = ({ guardrailId, onClose, ) : ( -
-
-

Guardrail ID

-
{guardrailData.guardrail_id}
-
-
-

Guardrail Name

-
{guardrailData.guardrail_name || "Unnamed Guardrail"}
-
-
-

Provider

-
{displayName}
-
- -
-

Default On

- - {guardrailData.litellm_params?.default_on ? "Yes" : "No"} - -
- - {guardrailData.litellm_params?.pii_entities_config && - Object.keys(guardrailData.litellm_params.pii_entities_config).length > 0 && ( -
-

PII Protection

-
- - {Object.keys(guardrailData.litellm_params.pii_entities_config).length} PII entities - configured - -
-
- )} - -
-

Created At

-
{formatDate(guardrailData.created_at)}
-
-
-

Last Updated

-
{formatDate(guardrailData.updated_at)}
-
- - {guardrailData.litellm_params?.guardrail === "tool_permission" && ( - - )} -
+ )}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.test.tsx index 96ade4b8c1b..d149ef98328 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.test.tsx @@ -15,6 +15,11 @@ import { skipToolMessageToChoice, choiceToSkipToolForCreate, formatGuardrailMode, + formatGuardrailStreamScope, + streamScopeByModeFromConfig, + streamScopeForMode, + streamScopeForUpdate, + streamScopePayload, loggingOnlyScopeToChoice, choiceToLoggingOnlyScope, getLoggingOnlyScopeUpdate, @@ -369,4 +374,56 @@ describe("guardrail_info_helpers", () => { expect(choiceToSkipToolForCreate("no")).toBe(false); }); }); + + describe("stream_scope helpers", () => { + it("treats omitted config as both for every mode", () => { + expect(streamScopeForMode(undefined, "pre_call")).toBe("both"); + expect(streamScopeByModeFromConfig(undefined, ["pre_call", "post_call"])).toEqual({ + pre_call: "both", + post_call: "both", + }); + }); + + it("applies a scalar to every selected mode and omits both-only payloads", () => { + expect(streamScopeForMode("streaming", "post_call")).toBe("streaming"); + expect(streamScopePayload(["pre_call", "post_call"], { pre_call: "both", post_call: "both" })).toBeUndefined(); + expect(streamScopePayload(["pre_call", "post_call"], { pre_call: "streaming", post_call: "streaming" })).toBe( + "streaming", + ); + }); + + it("keeps a mixed map instead of collapsing it to a scalar", () => { + expect(streamScopePayload(["pre_call", "post_call"], { pre_call: "both", post_call: "streaming" })).toEqual({ + post_call: "streaming", + }); + }); + + it("formats scalar and per-mode stream scopes for display", () => { + expect(formatGuardrailStreamScope(undefined)).toBe(""); + expect(formatGuardrailStreamScope("both")).toBe("Streaming and non-streaming"); + expect(formatGuardrailStreamScope("non_streaming")).toBe("Non-streaming only"); + expect(formatGuardrailStreamScope({ post_call: "streaming", pre_call: "both" })).toBe( + "post_call: Streaming only, pre_call: Streaming and non-streaming", + ); + }); + + it("emits both on update only when a prior restriction is cleared", () => { + expect(streamScopeForUpdate(["post_call"], { post_call: "streaming" }, undefined)).toBe("streaming"); + expect(streamScopeForUpdate(["post_call"], { post_call: "both" }, "streaming")).toBe("both"); + expect(streamScopeForUpdate(["post_call"], { post_call: "both" }, undefined)).toBeUndefined(); + }); + + it("keeps stored restrictions for modes outside the current selection", () => { + expect( + streamScopeForUpdate( + ["pre_call"], + { pre_call: "streaming" }, + { pre_call: "streaming", post_call: "non_streaming" }, + ), + ).toBeUndefined(); + expect( + streamScopeForUpdate(["pre_call"], { pre_call: "both" }, { pre_call: "streaming", post_call: "non_streaming" }), + ).toEqual({ post_call: "non_streaming" }); + }); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx index f55148b7315..a898f3b9abd 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx @@ -343,3 +343,86 @@ export function choiceToSkipToolForCreate(choice: SkipToolMessageChoice | undefi if (choice === "no") return false; return undefined; } + +export const GUARDRAIL_STREAM_SCOPES = ["both", "streaming", "non_streaming"] as const; +export type GuardrailStreamScope = (typeof GUARDRAIL_STREAM_SCOPES)[number]; + +export const STREAM_SCOPE_OPTIONS: { value: GuardrailStreamScope; label: string }[] = [ + { value: "both", label: "Streaming and non-streaming" }, + { value: "streaming", label: "Streaming only" }, + { value: "non_streaming", label: "Non-streaming only" }, +]; + +export const isGuardrailStreamScope = (value: unknown): value is GuardrailStreamScope => + value === "both" || value === "streaming" || value === "non_streaming"; + +export const streamScopeForMode = (raw: unknown, mode: string): GuardrailStreamScope => { + if (isGuardrailStreamScope(raw)) return raw; + if (raw !== null && typeof raw === "object" && !Array.isArray(raw)) { + const value = (raw as Record)[mode]; + if (isGuardrailStreamScope(value)) return value; + } + return "both"; +}; + +export const streamScopeByModeFromConfig = (raw: unknown, modes: string[]): Record => + Object.fromEntries(modes.map((mode) => [mode, streamScopeForMode(raw, mode)])); + +export const streamScopePayload = ( + modes: string[], + scopes: Record, +): GuardrailStreamScope | Record | undefined => { + const perMode: Record = Object.fromEntries( + modes.map((mode) => [mode, scopes[mode] ?? "both"]), + ); + const values = Object.values(perMode); + if (values.length === 0 || values.every((scope) => scope === "both")) return undefined; + const unique = new Set(values); + if (unique.size === 1) return values[0]; + return Object.fromEntries(Object.entries(perMode).filter((entry) => entry[1] !== "both")); +}; + +export const formatGuardrailStreamScope = (raw: unknown): string => { + if (isGuardrailStreamScope(raw)) { + return STREAM_SCOPE_OPTIONS.find((option) => option.value === raw)?.label ?? raw; + } + if (raw !== null && typeof raw === "object" && !Array.isArray(raw)) { + const entries = Object.entries(raw as Record).filter( + (entry): entry is [string, GuardrailStreamScope] => isGuardrailStreamScope(entry[1]), + ); + if (entries.length === 0) return ""; + return entries.map(([mode, scope]) => `${mode}: ${formatGuardrailStreamScope(scope)}`).join(", "); + } + return ""; +}; + +export const streamScopeForUpdate = ( + modes: string[], + nextByMode: Record, + previousRaw: unknown, + previousModes: string[] = modes, +): GuardrailStreamScope | Record | undefined => { + const previousMap: Record = + previousRaw !== null && typeof previousRaw === "object" && !Array.isArray(previousRaw) + ? (previousRaw as Record) + : {}; + const preserved: Record = Object.fromEntries( + Object.entries(previousMap).filter( + (entry): entry is [string, GuardrailStreamScope] => !modes.includes(entry[0]) && isGuardrailStreamScope(entry[1]), + ), + ); + const nextModes = [...modes, ...Object.keys(preserved)]; + const nextStreamScope = streamScopePayload(nextModes, { + ...preserved, + ...Object.fromEntries(modes.map((mode) => [mode, nextByMode[mode] ?? "both"])), + }); + const previousCompareModes = Object.keys(previousMap).length > 0 ? Object.keys(previousMap) : previousModes; + const previousStreamScope = streamScopePayload( + previousCompareModes, + streamScopeByModeFromConfig(previousRaw, previousCompareModes), + ); + if (JSON.stringify(nextStreamScope ?? "both") === JSON.stringify(previousStreamScope ?? "both")) { + return undefined; + } + return nextStreamScope ?? "both"; +}; diff --git a/ui/litellm-dashboard/src/components/policies/types.ts b/ui/litellm-dashboard/src/components/policies/types.ts index 430864f93df..2830a22950c 100644 --- a/ui/litellm-dashboard/src/components/policies/types.ts +++ b/ui/litellm-dashboard/src/components/policies/types.ts @@ -102,7 +102,7 @@ export interface PolicyAttachmentListResponse { export interface PipelineStepResult { guardrail_name: string; - outcome: "pass" | "fail" | "error"; + outcome: "pass" | "fail" | "error" | "skip"; action_taken: string; modified_data: Record | null; error_detail: string | null; diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 2a2b5acae78..cd85908979e 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -27393,6 +27393,13 @@ export interface components { * @default true */ sticky_session_routing: boolean | null; + /** + * Stream Scope + * @description Whether this guardrail runs on streaming requests, non-streaming requests, or both. A string applies to every configured mode. A map overrides named modes (pre_call, during_call, post_call, ...); omitted keys default to both. Unset means both, matching historical behavior. + */ + stream_scope?: ("streaming" | "non_streaming" | "both") | { + [key: string]: "streaming" | "non_streaming" | "both"; + } | null; /** * Template Id * @description The ID of your Model Armor template @@ -37786,6 +37793,13 @@ export interface components { * @default true */ sticky_session_routing: boolean | null; + /** + * Stream Scope + * @description Whether this guardrail runs on streaming requests, non-streaming requests, or both. A string applies to every configured mode. A map overrides named modes (pre_call, during_call, post_call, ...); omitted keys default to both. Unset means both, matching historical behavior. + */ + stream_scope?: ("streaming" | "non_streaming" | "both") | { + [key: string]: "streaming" | "non_streaming" | "both"; + } | null; /** * Template Id * @description The ID of your Model Armor template