From 78ff5ac9cd144770f9e53931d190ca2d480d9d98 Mon Sep 17 00:00:00 2001 From: tin-berri Date: Wed, 2 Sep 2026 18:31:25 -0700 Subject: [PATCH 1/4] feat(router): arm safeguard-refusal fallback on generic chains when no content-policy list exists (#39274) --- litellm/router.py | 30 ++++- .../router_utils/fallback_event_handlers.py | 32 ++++++ ...test_router_anthropic_messages_fallback.py | 106 ++++++++++++++++++ 3 files changed, 166 insertions(+), 2 deletions(-) diff --git a/litellm/router.py b/litellm/router.py index 0af514fe8a2..303b22c9484 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -150,8 +150,10 @@ from litellm.router_utils.fallback_event_handlers import ( _check_non_standard_fallback_format, clear_pre_routing_selection, fallback_lookup_groups, + fallbacks_disabled_for_request, get_fallback_model_group_for_lookup_groups, get_pre_routing_selection, + record_disable_fallbacks, record_pre_routing_selection, run_async_fallback, ) @@ -5193,7 +5195,7 @@ class Router: if not has_generated_content and error_event is None else None ) - if refusal_stop_details is not None and self._has_content_policy_fallback(model, initial_kwargs): + if refusal_stop_details is not None and self._refusal_fallback_available(model, initial_kwargs): refusal_error = safeguard_refusal_error(model=model, stop_details=refusal_stop_details) raise MidStreamFallbackError( message=refusal_error.message, @@ -7266,6 +7268,7 @@ class Router: _fallback_metadata["original_model_group"] = model_group include_fallback_errors: Final = kwargs.get("include_fallback_errors", False) is True disable_fallbacks: Final[bool | None] = kwargs.pop("disable_fallbacks", False) + record_disable_fallbacks(kwargs, disable_fallbacks is True) fallbacks: Final[list | None] = kwargs.get("fallbacks", self.fallbacks) context_window_fallbacks: list | None = kwargs.get("context_window_fallbacks", self.context_window_fallbacks) content_policy_fallbacks: list | None = kwargs.get("content_policy_fallbacks", self.content_policy_fallbacks) @@ -8131,6 +8134,29 @@ class Router: ) return False + def _refusal_fallback_available(self, model_group: str, kwargs: Mapping[str, Any]) -> bool: + """ + Whether a safeguard refusal can actually be recovered by the dispatcher. A configured + content-policy list is authoritative; with none configured at all, the dispatcher falls + through to the generic fallbacks lookup, so the gate mirrors that reachability and arms + on a resolving generic chain (tier first, then the requested group, then "*"). + """ + if fallbacks_disabled_for_request(kwargs): + return False + content_policy_fallbacks: Final = kwargs.get("content_policy_fallbacks", self.content_policy_fallbacks) + if content_policy_fallbacks is not None: + return self._has_content_policy_fallback(model_group, kwargs) + if self._has_default_fallbacks(): + return True + fallbacks: Final = kwargs.get("fallbacks", self.fallbacks) + if fallbacks is None: + return False + resolved, _ = get_fallback_model_group_for_lookup_groups( + fallbacks=fallbacks, + lookup_groups=fallback_lookup_groups(kwargs, model_group), + ) + return resolved is not None + def _should_raise_content_policy_error(self, model: str, response: ModelResponse, kwargs: dict) -> bool: """ Determines if a content policy error should be raised. @@ -8162,7 +8188,7 @@ class Router: return False if get_safeguard_refusal_stop_details(response) is None: return False - return self._has_content_policy_fallback(model, kwargs) + return self._refusal_fallback_available(model, kwargs) def _get_healthy_deployments(self, model: str, parent_otel_span: Span | None): _all_deployments: list = [] diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py index f7855cb38ff..601b32c4386 100644 --- a/litellm/router_utils/fallback_event_handlers.py +++ b/litellm/router_utils/fallback_event_handlers.py @@ -263,6 +263,38 @@ def get_pre_routing_selection(kwargs: Mapping[str, Any]) -> str | None: return next((selected for selected in selections if isinstance(selected, str) and selected), None) +DISABLE_FALLBACKS_METADATA_KEY: Final = "_disable_fallbacks" + + +def record_disable_fallbacks(request_kwargs: Mapping[str, Any] | None, disabled: bool) -> None: + """ + Write-or-clear the request's disable_fallbacks verdict into the router-internal metadata + bucket. The wrapper pops the raw kwarg before any downstream frame runs, so the refusal + gate (which decides whether to convert a refusal into a recoverable error) needs this + carrier to know recovery is impossible. + """ + from litellm.litellm_core_utils.core_helpers import get_metadata_variable_name_from_kwargs + + if request_kwargs is None: + return + bucket: Final = request_kwargs.get(get_metadata_variable_name_from_kwargs(request_kwargs)) + if not isinstance(bucket, dict): + return + if disabled: + bucket[DISABLE_FALLBACKS_METADATA_KEY] = True + else: + bucket.pop(DISABLE_FALLBACKS_METADATA_KEY, None) + + +def fallbacks_disabled_for_request(kwargs: Mapping[str, Any]) -> bool: + """True when this request opted out of fallbacks, read from the raw kwarg (pre-pop + snapshots keep it) or the router-internal bucket the wrapper stamps after popping it.""" + if kwargs.get("disable_fallbacks") is True: + return True + buckets: Final = (kwargs.get(name) for name in _ROUTER_METADATA_BUCKETS) + return any(isinstance(bucket, dict) and bucket.get(DISABLE_FALLBACKS_METADATA_KEY) is True for bucket in buckets) + + def fallback_lookup_groups(kwargs: Mapping[str, Any], model_group: str | None) -> tuple[str, ...]: """ Ordered keys for resolving a fallback chain: the tier a pre-routing hook selected wins, diff --git a/tests/router_unit_tests/test_router_anthropic_messages_fallback.py b/tests/router_unit_tests/test_router_anthropic_messages_fallback.py index 0c4d1dfc21e..4812d199c06 100644 --- a/tests/router_unit_tests/test_router_anthropic_messages_fallback.py +++ b/tests/router_unit_tests/test_router_anthropic_messages_fallback.py @@ -338,6 +338,112 @@ def test_record_pre_routing_selection_writes_only_the_internal_bucket(): assert kwargs["metadata"] == {"user_id": "u1"} +@pytest.mark.asyncio +@pytest.mark.parametrize("stream", [False, True], ids=["non-streaming", "streaming"]) +async def test_generic_only_row_recovers_safeguard_refusal(stream): + """With no content-policy list configured, a generic fallback row covers safeguard refusals, + so the dashboard's generic fallbacks work without config-only content_policy rows.""" + fake = FakeAnthropicUpstream() + router = Router(model_list=[FABLE_TIER, OPUS_TARGET], fallbacks=[{"fable-tier": ["opus-target"]}]) + + with fake.install(): + response = await router.aanthropic_messages( + model="fable-tier", max_tokens=16, stream=stream, messages=[{"role": "user", "content": "hi"}] + ) + body = await _collect(response) if stream else response + + if stream: + assert b'"refusal"' not in body + assert b"text_delta" in body + else: + assert body["stop_reason"] == "end_turn" + assert len(fake.calls) == 2 + assert "claude-opus-5" in fake.calls[1] + + +@pytest.mark.asyncio +async def test_configured_content_policy_list_stays_authoritative_over_generic_rows(): + fake = FakeAnthropicUpstream() + router = Router( + model_list=[FABLE_TIER, OPUS_TARGET], + fallbacks=[{"fable-tier": ["opus-target"]}], + content_policy_fallbacks=[{"unrelated-group": ["opus-target"]}], + ) + + with fake.install(): + response = await router.aanthropic_messages( + model="fable-tier", max_tokens=16, messages=[{"role": "user", "content": "hi"}] + ) + + assert response["stop_reason"] == "refusal" + assert len(fake.calls) == 1 + + +def test_refusal_fallback_available_arms_on_generic_rows_only_without_content_policy(): + router = Router(model_list=[FABLE_TIER, OPUS_TARGET], fallbacks=[{"tier-group": ["opus-target"]}]) + stamped = {"litellm_metadata": {PRE_ROUTING_SELECTED_MODEL_KEY: "tier-group"}} + + assert router._refusal_fallback_available("router-group", stamped) is True + assert router._refusal_fallback_available("router-group", {}) is False + assert router._refusal_fallback_available("router-group", {"content_policy_fallbacks": [{"other": ["x"]}]}) is False + + +def test_chat_content_filter_gate_unchanged_by_generic_rows(): + """The generic-row arming is scoped to /v1/messages safeguard refusals; the chat surface's + content_filter gate keeps its long-standing content-policy-only semantics.""" + from litellm.types.utils import Choices, ModelResponse + + router = Router(model_list=[FABLE_TIER, OPUS_TARGET], fallbacks=[{"fable-tier": ["opus-target"]}]) + response = ModelResponse(choices=[Choices(finish_reason="content_filter")]) + + assert router._should_raise_content_policy_error(model="fable-tier", response=response, kwargs={}) is False + + +@pytest.mark.asyncio +@pytest.mark.parametrize("stream", [False, True], ids=["non-streaming", "streaming"]) +async def test_disable_fallbacks_returns_the_refusal_instead_of_raising(stream): + """A request that opted out of fallbacks must receive the provider's refusal response, + never a ContentPolicyViolationError the dispatcher refuses to recover.""" + fake = FakeAnthropicUpstream() + router = Router(model_list=[FABLE_TIER, OPUS_TARGET], fallbacks=[{"fable-tier": ["opus-target"]}]) + + with fake.install(): + response = await router.aanthropic_messages( + model="fable-tier", + max_tokens=16, + stream=stream, + disable_fallbacks=True, + messages=[{"role": "user", "content": "hi"}], + ) + body = await _collect(response) if stream else response + + if stream: + assert b'"stop_reason": "refusal"' in body + else: + assert body["stop_reason"] == "refusal" + assert len(fake.calls) == 1 + + +@pytest.mark.asyncio +async def test_disable_fallbacks_beats_a_content_policy_row_too(): + fake = FakeAnthropicUpstream() + router = Router( + model_list=[FABLE_TIER, OPUS_TARGET], + content_policy_fallbacks=[{"fable-tier": ["opus-target"]}], + ) + + with fake.install(): + response = await router.aanthropic_messages( + model="fable-tier", + max_tokens=16, + disable_fallbacks=True, + messages=[{"role": "user", "content": "hi"}], + ) + + assert response["stop_reason"] == "refusal" + assert len(fake.calls) == 1 + + def test_refusal_gate_keys_on_pre_routing_tier_stamp(): router = _router(content_policy_fallbacks=[{"tier-group": ["opus-target"]}]) From e0e249225b1fd2d6620eaad90a13a5a8e62f713d Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Wed, 2 Sep 2026 18:55:22 -0700 Subject: [PATCH 2/4] feat(azure): support credential chain for storage (#39229) * feat(azure): support credential chain for storage * test(azure): clarify credential seam suppressions * fix(azure): read chain tokens in a worker thread The credential chain walk (IMDS probe, CLI subprocess) is blocking I/O, so reading the provider inline in async set_valid_azure_ad_token stalls every request on the worker's event loop --- .../azure_storage/azure_storage.py | 49 +++- .../azure_storage/test_azure_storage.py | 241 +++++++++++++++--- .../files/test_azure_blob_storage_backend.py | 47 +++- 3 files changed, 282 insertions(+), 55 deletions(-) diff --git a/litellm/integrations/azure_storage/azure_storage.py b/litellm/integrations/azure_storage/azure_storage.py index cb7175691df..16ef6920114 100644 --- a/litellm/integrations/azure_storage/azure_storage.py +++ b/litellm/integrations/azure_storage/azure_storage.py @@ -1,7 +1,9 @@ import asyncio import os import time +from collections.abc import Callable from datetime import datetime, timedelta +from functools import cache from typing import Final from litellm._logging import verbose_logger @@ -19,21 +21,40 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) +from litellm.secret_managers.get_azure_ad_token_provider import ( + get_azure_ad_token_provider, +) +from litellm.types.secret_managers.get_azure_ad_token_provider import ( + AzureCredentialType, +) from litellm.types.utils import StandardLoggingPayload +AZURE_STORAGE_TOKEN_SCOPE: Final = "https://storage.azure.com/.default" + + +@cache +def _cached_credential_chain_token_provider() -> Callable[[], str]: + return get_azure_ad_token_provider( + azure_scope=AZURE_STORAGE_TOKEN_SCOPE, + azure_credential=AzureCredentialType.DefaultAzureCredential, + ) + class AzureBlobStorageLogger(CustomBatchLogger): def __init__( self, + build_credential_chain_token_provider: Callable[ + [], Callable[[], str] + ] = _cached_credential_chain_token_provider, **kwargs, ): try: verbose_logger.debug("AzureBlobStorageLogger: in init azure blob storage logger") # Env Variables used for Azure Storage Authentication - self.tenant_id = os.getenv("AZURE_STORAGE_TENANT_ID") - self.client_id = os.getenv("AZURE_STORAGE_CLIENT_ID") - self.client_secret = os.getenv("AZURE_STORAGE_CLIENT_SECRET") + self.tenant_id = os.getenv("AZURE_STORAGE_TENANT_ID") or None + self.client_id = os.getenv("AZURE_STORAGE_CLIENT_ID") or None + self.client_secret = os.getenv("AZURE_STORAGE_CLIENT_SECRET") or None self.azure_storage_account_key: str | None = os.getenv("AZURE_STORAGE_ACCOUNT_KEY") # Required Env Variables for Azure Storage @@ -55,6 +76,9 @@ class AzureBlobStorageLogger(CustomBatchLogger): # Internal variables used for Token based authentication self.azure_auth_token: str | None = None # the Azure AD token to use for Azure Storage API requests self.token_expiry: datetime | None = None # the expiry time of the currentAzure AD token + self._build_credential_chain_token_provider: Callable[[], Callable[[], str]] = ( + build_credential_chain_token_provider + ) asyncio.create_task(self.periodic_flush()) self.flush_lock = asyncio.Lock() @@ -231,10 +255,15 @@ class AzureBlobStorageLogger(CustomBatchLogger): """ Wrapper to set self.azure_auth_token to a valid Azure AD token, refreshing if necessary - Refreshes the token when: - - Token is expired - - Token is not set + Without a service principal configured, the credential chain provider is read every + time; it caches internally and refreshes against the token's real expiry. The read runs + in a worker thread because the chain walk (IMDS probe, CLI subprocess) is blocking """ + if self.tenant_id is None and self.client_id is None and self.client_secret is None: + token_provider: Final = self._build_credential_chain_token_provider() + self.azure_auth_token = await asyncio.to_thread(token_provider) + return + # Check if token needs refresh if self._azure_ad_token_is_expired() or self.azure_auth_token is None: verbose_logger.debug("Azure AD token needs refresh") @@ -273,13 +302,9 @@ class AzureBlobStorageLogger(CustomBatchLogger): tenant_id=tenant_id, client_id=client_id, client_secret=client_secret, - scope="https://storage.azure.com/.default", + scope=AZURE_STORAGE_TOKEN_SCOPE, ) - token: Final = token_provider() - - verbose_logger.debug("azure auth token %s", token) - - return token + return token_provider() def _azure_ad_token_is_expired(self): """ diff --git a/tests/test_litellm/integrations/azure_storage/test_azure_storage.py b/tests/test_litellm/integrations/azure_storage/test_azure_storage.py index 16c518ff412..a96eae0f9c3 100644 --- a/tests/test_litellm/integrations/azure_storage/test_azure_storage.py +++ b/tests/test_litellm/integrations/azure_storage/test_azure_storage.py @@ -1,10 +1,15 @@ +import asyncio import sys +import threading from unittest.mock import AsyncMock, MagicMock, patch import pytest - -from litellm.integrations.azure_storage.azure_storage import AzureBlobStorageLogger +from litellm.integrations.azure_storage.azure_storage import ( + AzureBlobStorageLogger, + _cached_credential_chain_token_provider, +) +from litellm.types.secret_managers.get_azure_ad_token_provider import AzureCredentialType from litellm.types.utils import StandardLoggingPayload @@ -25,6 +30,26 @@ def mock_gov_env_vars(mock_env_vars, monkeypatch): monkeypatch.setenv("AZURE_STORAGE_ENDPOINT_SUFFIX", "core.usgovcloudapi.net") +@pytest.fixture +def workload_identity_env_vars(monkeypatch): + monkeypatch.setenv("AZURE_STORAGE_ACCOUNT_NAME", "test-account") + monkeypatch.setenv("AZURE_STORAGE_FILE_SYSTEM", "test-container") + for unset in ( + "AZURE_STORAGE_TENANT_ID", + "AZURE_STORAGE_CLIENT_ID", + "AZURE_STORAGE_CLIENT_SECRET", + "AZURE_STORAGE_ACCOUNT_KEY", + "AZURE_STORAGE_ENDPOINT_SUFFIX", + "AZURE_CLIENT_SECRET", + "AZURE_CREDENTIAL", + "AZURE_SCOPE", + ): + monkeypatch.delenv(unset, raising=False) + monkeypatch.setenv("AZURE_CLIENT_ID", "workload-identity-client-id") + monkeypatch.setenv("AZURE_TENANT_ID", "workload-identity-tenant-id") + monkeypatch.setenv("AZURE_FEDERATED_TOKEN_FILE", "/var/run/secrets/azure/tokens/azure-identity-token") + + @pytest.mark.asyncio async def test_async_upload_payload_to_azure_blob_storage(mock_env_vars): """ @@ -32,17 +57,12 @@ async def test_async_upload_payload_to_azure_blob_storage(mock_env_vars): a payload to Azure Blob Storage using the 3-step process (create, append, flush). """ with ( - patch( - "litellm.integrations.azure_storage.azure_storage.get_async_httpx_client" - ) as mock_get_client, - patch( - "litellm.llms.azure.common_utils.get_azure_ad_token_from_entra_id" - ) as mock_get_token, + patch("litellm.integrations.azure_storage.azure_storage.get_async_httpx_client") as mock_get_client, + patch("litellm.integrations.azure_storage.azure_storage.get_azure_ad_token_from_entra_id") as mock_get_token, ): # Create mock HTTP client mock_http_client = AsyncMock() - mock_response = AsyncMock() - mock_response.raise_for_status = AsyncMock() + mock_response = MagicMock() mock_http_client.put.return_value = mock_response mock_http_client.patch.return_value = mock_response mock_get_client.return_value = mock_http_client @@ -79,9 +99,7 @@ async def test_async_upload_payload_to_azure_blob_storage(mock_env_vars): put_call_args = mock_http_client.put.call_args assert put_call_args[0][0] == f"{expected_base_url}?resource=file" assert put_call_args[1]["headers"]["x-ms-version"] is not None - assert ( - put_call_args[1]["headers"]["Authorization"] == "Bearer mock-azure-ad-token" - ) + assert put_call_args[1]["headers"]["Authorization"] == "Bearer mock-azure-ad-token" # Step 2: Append data assert mock_http_client.patch.call_count == 2 # Called for append and flush @@ -89,9 +107,7 @@ async def test_async_upload_payload_to_azure_blob_storage(mock_env_vars): assert append_call[0][0] == f"{expected_base_url}?action=append&position=0" assert append_call[1]["headers"]["x-ms-version"] is not None assert append_call[1]["headers"]["Content-Type"] == "application/json" - assert ( - append_call[1]["headers"]["Authorization"] == "Bearer mock-azure-ad-token" - ) + assert append_call[1]["headers"]["Authorization"] == "Bearer mock-azure-ad-token" assert "test-log-id-123" in append_call[1]["data"] # Step 3: Flush data @@ -110,9 +126,7 @@ async def test_async_upload_payload_uses_configured_endpoint_suffix(mock_gov_env AZURE_STORAGE_ENDPOINT_SUFFIX must reach the Entra-ID REST upload path so a sovereign-cloud account is addressed instead of the commercial dfs host. """ - with patch( - "litellm.integrations.azure_storage.azure_storage.get_async_httpx_client" - ) as mock_get_client: + with patch("litellm.integrations.azure_storage.azure_storage.get_async_httpx_client") as mock_get_client: mock_http_client = AsyncMock() mock_response = MagicMock() mock_http_client.put.return_value = mock_response @@ -127,17 +141,10 @@ async def test_async_upload_payload_uses_configured_endpoint_suffix(mock_gov_env await logger.async_upload_payload_to_azure_blob_storage(test_payload) - expected_base_url = ( - "https://test-account.dfs.core.usgovcloudapi.net/test-container/gov-log-id.json" - ) + expected_base_url = "https://test-account.dfs.core.usgovcloudapi.net/test-container/gov-log-id.json" assert mock_http_client.put.call_args[0][0] == f"{expected_base_url}?resource=file" - assert ( - mock_http_client.patch.call_args_list[0][0][0] - == f"{expected_base_url}?action=append&position=0" - ) - assert mock_http_client.patch.call_args_list[1][0][0].startswith( - f"{expected_base_url}?action=flush" - ) + assert mock_http_client.patch.call_args_list[0][0][0] == f"{expected_base_url}?action=append&position=0" + assert mock_http_client.patch.call_args_list[1][0][0].startswith(f"{expected_base_url}?action=flush") @pytest.mark.asyncio @@ -148,9 +155,7 @@ async def test_service_client_uses_configured_endpoint_suffix(mock_gov_env_vars) """ fake_aio_module = MagicMock() - with patch.dict( - sys.modules, {"azure.storage.filedatalake.aio": fake_aio_module} - ): + with patch.dict(sys.modules, {"azure.storage.filedatalake.aio": fake_aio_module}): logger = AzureBlobStorageLogger() await logger.get_service_client() @@ -160,14 +165,180 @@ async def test_service_client_uses_configured_endpoint_suffix(mock_gov_env_vars) ) +@pytest.mark.asyncio +async def test_upload_authenticates_through_the_credential_chain_under_workload_identity( + workload_identity_env_vars, +): + build_provider = MagicMock(return_value=lambda: "workload-identity-token") + with patch( # test-quality-ok: REST client is created inside the method; assert emitted request headers + "litellm.integrations.azure_storage.azure_storage.get_async_httpx_client" + ) as mock_get_client: + mock_http_client = AsyncMock() + mock_http_client.put.return_value = MagicMock() + mock_http_client.patch.return_value = MagicMock() + mock_get_client.return_value = mock_http_client + + logger = AzureBlobStorageLogger(build_credential_chain_token_provider=build_provider) + await logger.async_upload_payload_to_azure_blob_storage({"id": "wif-log-id"}) + + build_provider.assert_called_once_with() + assert logger.azure_auth_token == "workload-identity-token" + sent_headers = [mock_http_client.put.call_args[1]["headers"]] + [ + call[1]["headers"] for call in mock_http_client.patch.call_args_list + ] + assert len(sent_headers) == 3 + assert all(headers["Authorization"] == "Bearer workload-identity-token" for headers in sent_headers) + + +def test_default_chain_provider_is_storage_scoped_and_built_once_per_process(): + _cached_credential_chain_token_provider.cache_clear() + with ( + patch( # test-quality-ok: assert the default factory's fixed scope and credential type without constructing Azure SDK credentials + "litellm.integrations.azure_storage.azure_storage.get_azure_ad_token_provider", + return_value=lambda: "chain-token", + ) as mock_builder + ): + first = _cached_credential_chain_token_provider() + second = _cached_credential_chain_token_provider() + _cached_credential_chain_token_provider.cache_clear() + + assert first is second + assert first() == "chain-token" + mock_builder.assert_called_once_with( + azure_scope="https://storage.azure.com/.default", + azure_credential=AzureCredentialType.DefaultAzureCredential, + ) + + +@pytest.mark.asyncio +async def test_chain_tokens_are_read_from_the_provider_on_every_refresh( + workload_identity_env_vars, +): + provider = MagicMock(side_effect=["chain-token-1", "chain-token-2"]) + logger = AzureBlobStorageLogger(build_credential_chain_token_provider=MagicMock(return_value=provider)) + await logger.set_valid_azure_ad_token() + first_token = logger.azure_auth_token + await logger.set_valid_azure_ad_token() + + assert first_token == "chain-token-1" + assert logger.azure_auth_token == "chain-token-2" + assert provider.call_count == 2 + + +@pytest.mark.asyncio +async def test_chain_token_read_yields_to_the_event_loop(workload_identity_env_vars): + """ + The chain walk is blocking I/O (IMDS probe, CLI subprocess), so reading the provider + inline would stall every request on the worker. Prove other coroutines run during the read. + """ + loop_was_free = threading.Event() + + def provider() -> str: + if not loop_was_free.wait(timeout=5): + raise TimeoutError("the event loop never ran the observer while the token was being read") + return "chain-token" + + async def observer(): + loop_was_free.set() + + logger = AzureBlobStorageLogger(build_credential_chain_token_provider=MagicMock(return_value=provider)) + observer_task = asyncio.create_task(observer()) + await logger.set_valid_azure_ad_token() + await observer_task + + assert logger.azure_auth_token == "chain-token" + + +@pytest.mark.asyncio +async def test_empty_string_service_principal_vars_still_use_the_credential_chain( + workload_identity_env_vars, monkeypatch +): + for name in ("AZURE_STORAGE_TENANT_ID", "AZURE_STORAGE_CLIENT_ID", "AZURE_STORAGE_CLIENT_SECRET"): + monkeypatch.setenv(name, "") + + logger = AzureBlobStorageLogger( + build_credential_chain_token_provider=MagicMock(return_value=lambda: "workload-identity-token") + ) + await logger.set_valid_azure_ad_token() + + assert logger.azure_auth_token == "workload-identity-token" + + +@pytest.mark.asyncio +async def test_client_secret_auth_still_uses_the_storage_scoped_service_principal(mock_env_vars): + build_provider = MagicMock() + with ( + patch( # test-quality-ok: assert the storage scope passed to the shared token factory without making an external auth call + "litellm.integrations.azure_storage.azure_storage.get_azure_ad_token_from_entra_id", + return_value=lambda: "client-secret-token", + ) as mock_entra_id + ): + logger = AzureBlobStorageLogger(build_credential_chain_token_provider=build_provider) + await logger.set_valid_azure_ad_token() + + assert logger.azure_auth_token == "client-secret-token" + build_provider.assert_not_called() + assert mock_entra_id.call_args.kwargs == { + "tenant_id": "test-tenant-id", + "client_id": "test-client-id", + "client_secret": "test-client-secret", + "scope": "https://storage.azure.com/.default", + } + + +@pytest.mark.parametrize( + "missing_var", + ["AZURE_STORAGE_TENANT_ID", "AZURE_STORAGE_CLIENT_ID", "AZURE_STORAGE_CLIENT_SECRET"], +) +@pytest.mark.asyncio +async def test_partially_configured_service_principal_still_names_the_missing_variable( + mock_env_vars, monkeypatch, missing_var +): + monkeypatch.delenv(missing_var) + + build_provider = MagicMock() + logger = AzureBlobStorageLogger(build_credential_chain_token_provider=build_provider) + with pytest.raises(ValueError, match=f"Missing required environment variable: {missing_var}"): + await logger.set_valid_azure_ad_token() + + build_provider.assert_not_called() + + +@pytest.mark.asyncio +async def test_account_key_auth_never_requests_a_token(workload_identity_env_vars, monkeypatch): + monkeypatch.setenv("AZURE_STORAGE_ACCOUNT_KEY", "dGVzdC1rZXk=") + + file_client = MagicMock() + file_client.create_file = AsyncMock() + file_client.append_data = AsyncMock() + file_client.flush_data = AsyncMock() + directory_client = MagicMock() + directory_client.exists = AsyncMock(return_value=True) + directory_client.get_file_client = MagicMock(return_value=file_client) + file_system_client = MagicMock() + file_system_client.get_directory_client = MagicMock(return_value=directory_client) + service_client = MagicMock() + service_client.get_file_system_client = MagicMock(return_value=file_system_client) + fake_aio_module = MagicMock() + fake_aio_module.DataLakeServiceClient = MagicMock(return_value=service_client) + + build_provider = MagicMock() + with patch.dict(sys.modules, {"azure.storage.filedatalake.aio": fake_aio_module}): + logger = AzureBlobStorageLogger(build_credential_chain_token_provider=build_provider) + await logger.async_upload_payload_to_azure_blob_storage({"id": "account-key-log-id"}) + + build_provider.assert_not_called() + assert logger.azure_auth_token is None + file_client.flush_data.assert_awaited_once() + assert fake_aio_module.DataLakeServiceClient.call_args.kwargs["credential"] == "dGVzdC1rZXk=" + + @pytest.mark.asyncio async def test_service_client_defaults_to_commercial_endpoint(mock_env_vars): """Unset AZURE_STORAGE_ENDPOINT_SUFFIX keeps the pre-existing commercial host""" fake_aio_module = MagicMock() - with patch.dict( - sys.modules, {"azure.storage.filedatalake.aio": fake_aio_module} - ): + with patch.dict(sys.modules, {"azure.storage.filedatalake.aio": fake_aio_module}): logger = AzureBlobStorageLogger() await logger.get_service_client() diff --git a/tests/test_litellm/llms/base_llm/files/test_azure_blob_storage_backend.py b/tests/test_litellm/llms/base_llm/files/test_azure_blob_storage_backend.py index 7e1139ca79a..222ead92b54 100644 --- a/tests/test_litellm/llms/base_llm/files/test_azure_blob_storage_backend.py +++ b/tests/test_litellm/llms/base_llm/files/test_azure_blob_storage_backend.py @@ -27,6 +27,20 @@ def mock_gov_env_vars(mock_env_vars, monkeypatch): monkeypatch.setenv("AZURE_STORAGE_ENDPOINT_SUFFIX", GOV_SUFFIX) +@pytest.fixture +def credential_chain_env_vars(monkeypatch): + monkeypatch.setenv("AZURE_STORAGE_ACCOUNT_NAME", "test-account") + monkeypatch.setenv("AZURE_STORAGE_FILE_SYSTEM", "test-container") + for name in ( + "AZURE_STORAGE_TENANT_ID", + "AZURE_STORAGE_CLIENT_ID", + "AZURE_STORAGE_CLIENT_SECRET", + "AZURE_STORAGE_ACCOUNT_KEY", + "AZURE_STORAGE_ENDPOINT_SUFFIX", + ): + monkeypatch.delenv(name, raising=False) + + def _make_backend() -> AzureBlobStorageBackend: backend = AzureBlobStorageBackend() backend.azure_auth_token = "mock-azure-ad-token" @@ -42,6 +56,29 @@ def _mock_upload_client() -> AsyncMock: return client +@pytest.mark.asyncio +async def test_upload_file_with_credential_chain(credential_chain_env_vars): + client = _mock_upload_client() + build_provider = MagicMock(return_value=lambda: "workload-identity-token") + + with patch( # test-quality-ok: the backend creates its REST client internally; assert the emitted authorization header + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client", return_value=client + ): + backend = AzureBlobStorageBackend(build_credential_chain_token_provider=build_provider) + storage_url = await backend.upload_file( + file_content=b"hello", + filename="report.json", + content_type="application/json", + path_prefix="logs", + file_naming_strategy="original_filename", + ) + + build_provider.assert_called_once_with() + assert storage_url == "https://test-account.blob.core.windows.net/test-container/logs/report.json" + assert client.put.call_args[1]["headers"]["Authorization"] == "Bearer workload-identity-token" + assert client.patch.call_count == 2 + + @pytest.mark.parametrize( "env_fixture, expected_suffix", [("mock_env_vars", "core.windows.net"), ("mock_gov_env_vars", GOV_SUFFIX)], @@ -125,10 +162,7 @@ async def test_download_file_accepts_url_persisted_before_the_suffix_was_set(moc ) assert content == b"file-bytes" - assert ( - client.get.call_args[0][0] - == f"https://test-account.blob.{GOV_SUFFIX}/test-container/logs/report.json" - ) + assert client.get.call_args[0][0] == f"https://test-account.blob.{GOV_SUFFIX}/test-container/logs/report.json" @pytest.mark.parametrize( @@ -178,10 +212,7 @@ async def test_download_file_drops_query_string_from_the_stored_url(mock_env_var "https://test-account.blob.core.windows.net/test-container/logs/report.json?sig=redacted&se=2026" ) - assert ( - client.get.call_args[0][0] - == "https://test-account.blob.core.windows.net/test-container/logs/report.json" - ) + assert client.get.call_args[0][0] == "https://test-account.blob.core.windows.net/test-container/logs/report.json" @pytest.mark.parametrize( From 8065ede40bb380d8947939120f95a740f970614f Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 3 Sep 2026 02:11:08 +0000 Subject: [PATCH 3/4] test(guardrails): expect the deduped end-of-stream scan in crowdstrike cadence test (#39467) Co-authored-by: yucheng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../guardrails/guardrail_hooks/test_crowdstrike_aidr.py | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py index 1b50ea53db2..a07157396df 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py @@ -1703,8 +1703,8 @@ async def _guard_calls_for_stream(handler: CrowdStrikeAIDRHandler, chunk_texts: @pytest.mark.parametrize( ("configured", "expected_calls"), [ - ({}, 3), - ({"streaming_sampling_rate": 2}, 6), + ({}, 2), + ({"streaming_sampling_rate": 2}, 5), ({"streaming_end_of_stream_only": True}, 1), ({"streaming_end_of_stream_only": True, "streaming_sampling_rate": 2}, 1), ], @@ -1712,7 +1712,10 @@ async def _guard_calls_for_stream(handler: CrowdStrikeAIDRHandler, chunk_texts: async def test_streaming_params_from_config_control_output_scan_cadence( configured: dict[str, object], expected_calls: int ) -> None: - """10 chunks: default samples at 5 and 10 plus the final pass, rate 2 samples 5 times plus final, end-of-stream scans once.""" + """10 chunks: default samples at 5 and 10, rate 2 samples 5 times, end-of-stream scans once. + + The final pass is skipped because chunk 10 already scanned the complete output. + """ handler = _initialize_from_config(mode="post_call", **configured) assert await _guard_calls_for_stream(handler, list("ABCDEFGHIJ")) == expected_calls From 7a81ae98e68746aab8aadad58ee6706a2b04de00 Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Wed, 2 Sep 2026 19:15:14 -0700 Subject: [PATCH 4/4] fix(model_armor): handle Anthropic Messages and Responses streams in post_call (#39181) * fix(model_armor): handle Anthropic Messages and Responses streams in post_call The post_call streaming hook buffered every chunk and fed it to stream_chunk_builder, which only understands chat-completion deltas. /v1/messages streams raw Anthropic SSE bytes and /v1/responses streams typed Responses events, so both raised litellm.APIError and surfaced to the client as a 500 on every streamed request. Assemble each surface with its own reader, frame guardrail failures as terminal items in that surface's wire format, and pass the stream through unscanned when it cannot be assembled instead of raising. * fix(model_armor): classify the stream surface and fail closed when it cannot be assembled Decide the wire format explicitly instead of inferring it from a boolean pair, so an opaque raw SSE stream (the Google :streamGenerateContent route) is never refused in Anthropic framing, and a stream that cannot be assembled is blocked rather than released unscanned unless fail_on_error is disabled. Also scan Responses tool-call arguments, read the body only off a terminal Responses event, and record the applied guardrail on the fail-closed path. * test(model_armor): pin the error-only stream predicate against content-carrying streams is_sse_error_stream decides whether a buffered stream is forwarded to the client untouched, so a stream that still carries content must not qualify: the frames-only join drops typed chunks, an empty stream is not a refusal, and a content event may carry an empty error field. * fix(model_armor): let a streamed de-identify match mask instead of blocking A de-identify template reports MATCH_FOUND for every redaction it makes. The streaming block check omitted allow_sanitization, so with mask_response_content enabled that match read as a refusal and the client got a 400 where the non-streaming sibling returned the redacted text. Pass the flag through, as the non-streaming hook already does, and stamp the logged status from the same decision so the spend row agrees with what the client received. Also drop Any from the chat-completion assembler's parameter; stream_chunk_builder takes a bare list, so list[object] carries the mutability requirement without erasing the element type. * fix(model_armor): fail closed when a streamed de-identify match cannot be applied Allowing sanitization past the streaming block check is a promise to apply the redaction Model Armor asked for. Two paths broke that promise and released the buffered original instead: a match that comes back with no sanitized text, and a surface with no assembled body to rewrite. The outcome is now resolved once, before it is recorded, so the status stamped on request metadata agrees with what the client receives rather than reporting the success the block check alone would have implied. * fix: scan the deltas when a Responses stream ends without a body response.failed and response.incomplete are terminal events like response.completed, but a turn that broke mid-generation reports an empty output while the deltas ahead of it already spelled the answer out to the client. Reading only the terminal body found nothing to scan there, and the empty-content shortcut then forwarded every buffered delta past the guardrail. Fall back to the text the delta events carry whenever a Responses stream assembles to nothing. * fix: read the Responses delta event types off the event enum The hand-listed set left out response.mcp_call_arguments.delta, so a turn that streamed only MCP tool arguments and then reported an empty body still took the no-content shortcut and forwarded those chunks unscanned. Deriving the set from ResponsesAPIStreamEvents keeps it complete as the enum grows, and the str guard in the reader already covers any event whose delta is not text. * fix(model_armor): scan responses deltas alongside the terminal body A /v1/responses stream spells out reasoning summaries and tool-call arguments in delta events that its terminal body never repeats, so scanning the body alone handed every summary delta to the client unscanned whenever the body carried text. * fix(model_armor): scan responses delta fields apart from each other A Responses turn spells out its reasoning summary, its visible answer and its tool-call arguments in separate delta events. Joining every delta into one string let a finding form across the boundary between two fields that each carry nothing to find, so a safe stream could be blocked. Group the deltas by the field they belong to, join a field's own deltas as they streamed, and keep the fields apart. * fix(model_armor): scan each responses field once, not twice Separating delta fields stopped the terminal body from matching the delta text, so a turn with two visible fields sent Model Armor both copies. Only the delta fields the body does not already carry are appended now. --------- Co-authored-by: yassin --- litellm/proxy/guardrails/anthropic_sse.py | 64 +- .../model_armor/model_armor.py | 479 +++++-- .../guardrail_hooks/test_model_armor.py | 1151 +++++++++++++++++ 3 files changed, 1589 insertions(+), 105 deletions(-) diff --git a/litellm/proxy/guardrails/anthropic_sse.py b/litellm/proxy/guardrails/anthropic_sse.py index 50c05daee11..28220f09f00 100644 --- a/litellm/proxy/guardrails/anthropic_sse.py +++ b/litellm/proxy/guardrails/anthropic_sse.py @@ -13,6 +13,19 @@ from typing import Final from litellm.types.utils import Choices, ModelResponse +_ANTHROPIC_EVENT_TYPES: Final = frozenset( + { + "message_start", + "message_delta", + "message_stop", + "content_block_start", + "content_block_delta", + "content_block_stop", + "ping", + "error", + } +) + def is_raw_sse_stream(all_chunks: Sequence[object]) -> bool: return any(isinstance(chunk, (str, bytes)) for chunk in all_chunks) @@ -30,23 +43,43 @@ def _joined_sse_stream(all_chunks: Sequence[object]) -> str | None: return None -def _anthropic_message_start(sse_stream: str) -> Mapping[str, object] | None: +def _parsed_sse_events(sse_stream: str) -> tuple[Mapping[str, object], ...]: from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import ( AnthropicPassthroughLoggingHandler, ) + return tuple( + event_data + for event in AnthropicPassthroughLoggingHandler._split_sse_chunk_into_events(sse_stream) # pyright: ignore[reportPrivateUsage] # same parser the assembler uses + if (event_data := AnthropicPassthroughLoggingHandler._extract_sse_data(event)) is not None # pyright: ignore[reportPrivateUsage] # same parser the assembler uses; a private import beats forking SSE parsing + ) + + +def _anthropic_message_start(sse_stream: str) -> Mapping[str, object] | None: return next( ( message - for event in AnthropicPassthroughLoggingHandler._split_sse_chunk_into_events(sse_stream) # pyright: ignore[reportPrivateUsage] # same parser the assembler uses - if (event_data := AnthropicPassthroughLoggingHandler._extract_sse_data(event)) is not None # pyright: ignore[reportPrivateUsage] # same parser the assembler uses; a private import beats forking SSE parsing - and event_data.get("type") == "message_start" - and isinstance(message := event_data.get("message"), dict) + for event_data in _parsed_sse_events(sse_stream) + if event_data.get("type") == "message_start" and isinstance(message := event_data.get("message"), dict) ), None, ) +def is_anthropic_sse_stream(all_chunks: Sequence[object]) -> bool: + """Whether raw SSE frames are Anthropic Messages events. + + ``is_raw_sse_stream`` only says the chunks are unparsed bytes, and ``/v1/messages`` is not the + only endpoint that streams those: the Google ``:streamGenerateContent`` route marks its own + stream raw too. Reading its frames as Anthropic ones would refuse the response in a wire format + its client cannot parse, so the surface is decided on the event types actually present. + """ + sse_stream: Final = _joined_sse_stream(all_chunks) + if sse_stream is None: + return False + return any(event.get("type") in _ANTHROPIC_EVENT_TYPES for event in _parsed_sse_events(sse_stream)) + + def assemble_anthropic_sse_stream( all_chunks: Sequence[object], *, restore_identity: bool = False ) -> ModelResponse | None: @@ -111,6 +144,27 @@ def anthropic_sse_error_frames(message: str) -> tuple[bytes, ...]: ) +def is_sse_error_stream(all_chunks: Sequence[object]) -> bool: + """Whether the buffered stream carries nothing but error frames. + + post_call guardrails run in a chain, so a hook can be handed the terminal error frames an + earlier guardrail emitted when it blocked. Those carry no message to assemble, and replacing + them would hide the refusal the client is owed. Covers both wire forms a guardrail emits: the + Anthropic ``error`` event and the chat-completions ``{"error": ...}`` payload. + """ + if not all(isinstance(chunk, (str, bytes)) for chunk in all_chunks): + # A stream mixing typed chunks with an error frame still carries content to scan, and the + # frames-only join below would drop exactly the part that has to be scanned + return False + sse_stream: Final = _joined_sse_stream(all_chunks) + if sse_stream is None: + return False + events: Final = _parsed_sse_events(sse_stream) + return len(events) > 0 and all( + event.get("type") == "error" or isinstance(event.get("error"), Mapping) for event in events + ) + + def anthropic_sse_chunks_from_response(assembled: ModelResponse) -> tuple[bytes, ...]: from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import ( LiteLLMAnthropicMessagesAdapter, diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py index d187b5b12e9..7d88a037f4f 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py @@ -1,4 +1,5 @@ from collections.abc import AsyncGenerator, Mapping, Sequence +from enum import Enum, auto from typing import TYPE_CHECKING, Any, Final, Literal import httpx @@ -27,12 +28,25 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.llms.vertex_ai.vertex_llm_base import VertexBase from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails.anthropic_sse import ( + anthropic_sse_chunks_from_response, + anthropic_sse_error_frames, + assemble_anthropic_sse_stream, + is_anthropic_sse_stream, + is_raw_sse_stream, + is_sse_error_stream, +) from litellm.proxy.guardrails.guardrail_hooks.model_armor.file_scanning import ( MODEL_ARMOR_MAX_FILE_SIZE_BYTES, plan_file_scans, ) from litellm.types.guardrails import GuardrailEventHooks, LitellmParams -from litellm.types.llms.openai import AllMessageValues +from litellm.types.llms.openai import ( + AllMessageValues, + ChatCompletionToolCallChunk, + ResponsesAPIResponse, + ResponsesAPIStreamEvents, +) from litellm.types.utils import ( CallTypes, CallTypesLiteral, @@ -41,10 +55,33 @@ from litellm.types.utils import ( ModelResponse, ModelResponseStream, StandardLoggingGuardrailInformation, + TextCompletionResponse, ) GUARDRAIL_NAME: Final = "model_armor" +# Only these carry the finished output; response.created carries an empty body +_RESPONSES_TERMINAL_EVENT_TYPES: Final = frozenset({"response.completed", "response.incomplete", "response.failed"}) + +# Every event whose ``delta`` is model output already on its way to the client. Read off the event +# enum rather than listed, so an event added there cannot quietly fall out of the scan +_RESPONSES_DELTA_EVENT_TYPES: Final = frozenset( + event.value for event in ResponsesAPIStreamEvents if event.value.endswith(".delta") +) + +# What makes two delta events part of the same field of the turn, rather than two fields that merely +# streamed next to each other +_RESPONSES_DELTA_FIELD_ATTRS: Final = ("type", "item_id", "output_index", "content_index", "summary_index") + + +class _StreamSurface(Enum): + """Wire format of a buffered streaming response, which decides how it is read and how it is refused.""" + + CHAT_COMPLETIONS = auto() + ANTHROPIC_MESSAGES = auto() + RESPONSES = auto() + OPAQUE_SSE = auto() + class ModelArmorAPIError(Exception): """Model Armor API failure (non-2xx), distinct from a content-block decision so @@ -322,19 +359,9 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): else: return {"modelResponseData": {"byteItem": {"byteDataType": file_type, "byteData": base64_data}}} - def _should_block_content(self, armor_response: dict, allow_sanitization: bool = False) -> bool: + def _should_block_content(self, armor_response: Mapping[str, Any], allow_sanitization: bool = False) -> bool: """Check if Model Armor response indicates content should be blocked, including both inspectResult and deidentifyResult.""" - sanitization_result: Final = armor_response.get("sanitizationResult", {}) - filter_results: Final = sanitization_result.get("filterResults", {}) - - # filterResults can be a dict (named keys) or a list (array of filter result dicts) - filter_result_items = [] - if isinstance(filter_results, dict): - filter_result_items = list(filter_results.values()) - elif isinstance(filter_results, list): - filter_result_items = filter_results - - for filt in filter_result_items: + for filt in self._filter_result_items(armor_response): # Check RAI, PI/Jailbreak, Malicious URI, CSAM, Virus scan as before if filt.get("raiFilterResult", {}).get("matchState") == "MATCH_FOUND": return True @@ -358,22 +385,12 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): # Fallback dict code removed; all cases handled above return False - def _get_sanitized_content(self, armor_response: dict) -> str | None: + def _get_sanitized_content(self, armor_response: Mapping[str, Any]) -> str | None: """ Get the sanitized content from a Model Armor response, if available. Looks for sanitized text in deidentifyResult, and falls back to root-level fields if not found. """ - result: Final = armor_response.get("sanitizationResult", {}) - filter_results: Final = result.get("filterResults", {}) - - # filterResults can be a dict (single filter) or a list (multiple filters) - filters: Final = ( - list(filter_results.values()) - if isinstance(filter_results, dict) - else filter_results - if isinstance(filter_results, list) - else [] - ) + filters: Final = self._filter_result_items(armor_response) # Prefer sanitized text from deidentifyResult if present for filter_entry in filters: @@ -397,6 +414,61 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): # Fallback: if Model Armor put sanitized text at the root, use it return armor_response.get("sanitizedText") or armor_response.get("text") + @staticmethod + def _filter_result_items(armor_response: Mapping[str, Any]) -> Sequence[Any]: + """Every filter result in a scan response. + + filterResults is a dict of named filters on most templates and a list on some, so both + shapes are flattened to the same list of filter entries. + """ + filter_results: Final = armor_response.get("sanitizationResult", {}).get("filterResults", {}) + if isinstance(filter_results, dict): + return list(filter_results.values()) + if isinstance(filter_results, list): + return filter_results + return [] + + def _has_deidentify_match(self, armor_response: Mapping[str, Any]) -> bool: + """Whether an SDP de-identify filter matched, i.e. Model Armor owes this response a redaction.""" + for filter_entry in self._filter_result_items(armor_response): + sdp = filter_entry.get("sdpFilterResult") + if sdp and sdp.get("deidentifyResult", {}).get("matchState") == "MATCH_FOUND": + return True + return False + + def _resolve_streaming_outcome( + self, + armor_response: Mapping[str, Any], + assembled_response: object, + content: str, + ) -> tuple[bool, str | None]: + """Whether to block the buffered stream, and the rewrite to emit when it is not blocked. + + A de-identify match only reaches here unblocked because masking is on, so the redaction it + stands for has to be both resolvable and emittable. Where it is neither, the buffered + original still carries what Model Armor matched on, so this fails closed instead of + releasing it. + """ + if self._should_block_content(armor_response, allow_sanitization=self.mask_response_content): + return True, None + if not self.mask_response_content: + return False, None + + sanitized_content: Final = self._get_sanitized_content(armor_response) + if not sanitized_content: + # No rewrite to apply. Harmless unless a match is outstanding, in which case applying + # nothing would hand back the very content that matched + return self._has_deidentify_match(armor_response), None + if sanitized_content == content: + return False, None + if not isinstance(assembled_response, ModelResponse): + verbose_proxy_logger.warning( + "Model Armor: sanitized content cannot be re-emitted on this streaming endpoint, " + "blocking the response instead" + ) + return True, None + return False, sanitized_content + @staticmethod def _append_armor_response(existing: object, armor_response: Mapping[str, object]) -> object: """Accumulate scan responses so a later text scan does not drop an earlier file scan. @@ -831,6 +903,185 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): return response + @staticmethod + def _is_terminal_error_stream(all_chunks: Sequence[object]) -> bool: + """Whether the buffered stream is only the refusal an earlier guardrail in the chain emitted. + + post_call guardrails are composed, so this hook can be handed the terminal error items a + preceding one produced. They carry no message to scan, and replacing them would hide the + refusal the client is owed. + """ + if all(getattr(chunk, "type", None) == "error" for chunk in all_chunks): + return True + return is_sse_error_stream(all_chunks) + + @staticmethod + def _classify_stream(all_chunks: Sequence[object]) -> _StreamSurface: + """Wire format the buffered chunks belong to.""" + if is_raw_sse_stream(all_chunks): + return ( + _StreamSurface.ANTHROPIC_MESSAGES if is_anthropic_sse_stream(all_chunks) else _StreamSurface.OPAQUE_SSE + ) + if any( + isinstance(event_type := getattr(chunk, "type", None), str) and event_type.startswith("response.") + for chunk in all_chunks + ): + return _StreamSurface.RESPONSES + return _StreamSurface.CHAT_COMPLETIONS + + @staticmethod + def _final_responses_api_response(all_chunks: Sequence[object]) -> ResponsesAPIResponse | None: + """Response body carried by a terminal ``/v1/responses`` event. + + A stream cut short before it completes has to read as unassembled rather than as a clean + empty response: ``response.created`` also carries a body, but an empty one, and scanning + that would release every buffered delta unscanned. + """ + return next( + ( + body + for chunk in reversed(all_chunks) + if getattr(chunk, "type", None) in _RESPONSES_TERMINAL_EVENT_TYPES + and isinstance(body := getattr(chunk, "response", None), ResponsesAPIResponse) + ), + None, + ) + + @staticmethod + def _responses_api_response_text(response: ResponsesAPIResponse) -> str: + """Text to scan in a Responses API response, tool-call arguments included. + + Tool calls are folded in because ``get_content_from_model_response`` folds them into what + the chat surface scans, and a Responses turn can carry its whole payload in them. + """ + from litellm.llms.openai.responses.guardrail_translation.handler import ( + OpenAIResponsesHandler, + ) + + texts: Final[list[str]] = [] # mutable-ok: the shared extractor below appends into caller-owned lists + tool_calls: Final[list[ChatCompletionToolCallChunk]] = [] # mutable-ok: the same extractor's tool-call sink + handler: Final = OpenAIResponsesHandler() + for output_idx, output_item in enumerate(response.output or ()): + handler._extract_output_text_and_images( # pyright: ignore[reportPrivateUsage] # the shared Responses output extractor; forking it would duplicate per-item parsing + output_item=output_item, + output_idx=output_idx, + texts_to_check=texts, + images_to_check=[], # mutable-ok: the extractor's images sink, unused here + task_mappings=[], # mutable-ok: the extractor's task-mapping sink, unused here + tool_calls_to_check=tool_calls, + ) + return "".join((*texts, *(json.dumps(tool_call) for tool_call in tool_calls))) + + def _extract_streaming_content(self, assembled_response: object) -> str: + """Text to scan from an assembled stream, for every endpoint shape this hook serves.""" + if isinstance(assembled_response, ResponsesAPIResponse): + return self._responses_api_response_text(assembled_response) + return self._extract_content_from_response(assembled_response) + + @staticmethod + def _responses_delta_field(chunk: object) -> tuple[str, ...]: + """Which field of the turn a delta event belongs to.""" + return tuple(str(getattr(chunk, attr, None)) for attr in _RESPONSES_DELTA_FIELD_ATTRS) + + @staticmethod + def _responses_delta_field_texts(all_chunks: Sequence[object]) -> tuple[str, ...]: + """Text each field of a ``/v1/responses`` turn has already spelled out in its delta events. + + One field's deltas are joined as they streamed, since a finding can be split across them, + and separate fields stay apart, so a reasoning summary running into the visible answer + cannot spell out a finding that neither of them carries. + """ + deltas: Final = tuple( + (ModelArmorGuardrail._responses_delta_field(chunk), delta) + for chunk in all_chunks + if getattr(chunk, "type", None) in _RESPONSES_DELTA_EVENT_TYPES + and isinstance(delta := getattr(chunk, "delta", None), str) + ) + return tuple( + "".join(delta for field, delta in deltas if field == streamed_field) + for streamed_field in dict.fromkeys(field for field, _ in deltas) + ) + + def _streaming_content_to_scan( + self, + assembled_response: object, + all_chunks: Sequence[object], + surface: _StreamSurface, + ) -> str: + """Text to scan for a buffered stream, which is everything the client is about to receive. + + A ``/v1/responses`` stream also spells out reasoning summaries and tool-call arguments in + delta events that its terminal body never repeats, so every delta field the body does not + already carry is scanned after it. + """ + content: Final = self._extract_streaming_content(assembled_response) + if surface is not _StreamSurface.RESPONSES: + return content + unscanned: Final = tuple(text for text in self._responses_delta_field_texts(all_chunks) if text not in content) + return "\n".join(part for part in (content, *unscanned) if part) + + @staticmethod + def _apply_sanitized_content(assembled_response: ModelResponse, sanitized_content: str) -> None: + """Replace every non-empty choice message with the Model Armor sanitized text.""" + for choice in assembled_response.choices: + if isinstance(choice, Choices) and choice.message.content: + choice.message.content = sanitized_content + + @staticmethod + def _assemble_chat_completion_stream( + all_chunks: list[object], # mutable-ok: stream_chunk_builder only accepts a mutable list + ) -> ModelResponse | TextCompletionResponse | None: + """Assemble chat-completion chunks, returning ``None`` when they cannot be assembled.""" + from litellm.main import stream_chunk_builder + + try: + return stream_chunk_builder(chunks=all_chunks) + except Exception as exc: + verbose_proxy_logger.warning("Model Armor: chat-completion stream assembly failed (%s)", exc) + return None + + def _assemble_stream( + self, all_chunks: Sequence[object], surface: _StreamSurface + ) -> ModelResponse | TextCompletionResponse | ResponsesAPIResponse | None: + """Assemble the buffered stream into the scannable response its surface produces.""" + if surface is _StreamSurface.ANTHROPIC_MESSAGES: + return assemble_anthropic_sse_stream(all_chunks, restore_identity=True) + if surface is _StreamSurface.RESPONSES: + return self._final_responses_api_response(all_chunks) + if surface is _StreamSurface.OPAQUE_SSE: + return None + return self._assemble_chat_completion_stream(list(all_chunks)) + + @staticmethod + def _error_payload(exc: HTTPException) -> Mapping[str, object]: + """Error object for a terminal stream item, carrying the status the frame would otherwise lose.""" + detail: Final = exc.detail if isinstance(exc.detail, Mapping) else {"message": str(exc.detail)} + error_value: Final = detail.get("error", detail) + return { + **(dict(error_value) if isinstance(error_value, Mapping) else {"message": str(error_value)}), + "code": str(exc.status_code), + } + + @staticmethod + def _build_responses_error_items(exc: HTTPException) -> Sequence[object] | None: + """Responses API error events for a failure discovered after the stream started.""" + from litellm.llms.openai.responses.guardrail_translation.handler import ( + OpenAIResponsesHandler, + ) + + return OpenAIResponsesHandler().build_stream_error_items(exc, responses_so_far=None) + + def _stream_error_items(self, exc: HTTPException, *, surface: _StreamSurface) -> Sequence[object]: + """Frame a guardrail failure as terminal stream items in this endpoint's wire format.""" + payload: Final = self._error_payload(exc) + if surface is _StreamSurface.ANTHROPIC_MESSAGES: + return anthropic_sse_error_frames(str(payload.get("message", ""))) + if surface is _StreamSurface.RESPONSES and (responses_items := self._build_responses_error_items(exc)): + return responses_items + # Also the fallback when a surface cannot frame its own error: create_response() reads the + # status back out of this form, so the refusal keeps its code instead of arriving as a 200 + return (f"data: {json.dumps({'error': payload})}\n\n",) + async def async_post_call_streaming_iterator_hook( self, user_api_key_dict: UserAPIKeyAuth, @@ -840,97 +1091,125 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): """Process streaming response chunks.""" from litellm.llms.base_llm.base_model_iterator import MockResponseIterator - from litellm.main import stream_chunk_builder + from litellm.proxy.common_utils.callback_utils import ( + add_guardrail_to_applied_guardrails_header, + ) # Collect all chunks - all_chunks: Final[list[ModelResponseStream]] = [] + all_chunks: Final[list[Any]] = [] async for chunk in response: all_chunks.append(chunk) + if not all_chunks or self._is_terminal_error_stream(all_chunks): + for chunk in all_chunks: + yield chunk + return + + surface: Final = self._classify_stream(all_chunks) + # Build complete response - assembled_response: Final = stream_chunk_builder(chunks=all_chunks) + assembled_response: Final = self._assemble_stream(all_chunks, surface) - if isinstance(assembled_response, ModelResponse): - # Extract content - content: Final = self._extract_content_from_response(assembled_response) + if assembled_response is None: + if not self.optional_params.get("fail_on_error", True): + verbose_proxy_logger.warning( + "Model Armor: streamed response could not be assembled for scanning, " + "forwarding it unscanned because fail_on_error is disabled" + ) + for chunk in all_chunks: + yield chunk + return - if content: - try: - # Check with Model Armor - armor_response: Final = await self.make_model_armor_request( - content=content, - source="model_response", - request_data=request_data, - ) + # Forwarding an unscannable stream would silently disable the guardrail, so fail closed + add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name) + for error_item in self._stream_error_items( + HTTPException( + status_code=500, + detail=f"{self.guardrail_name}: streamed response could not be assembled for scanning, blocking it", + ), + surface=surface, + ): + yield error_item + return - # Attach Model Armor response & status to this request's metadata to avoid race conditions - if isinstance(request_data, dict): - _, metadata = get_or_create_metadata_bucket(request_data) - metadata["_model_armor_response"] = self._build_logging_response(armor_response) - metadata["_model_armor_status"] = ( - "blocked" if self._should_block_content(armor_response) else "success" - ) + # Extract content + content: Final = self._streaming_content_to_scan( + assembled_response=assembled_response, all_chunks=all_chunks, surface=surface + ) - # Add guardrail to applied_guardrails BEFORE potential blocking - # This ensures guardrail is recorded even when it blocks the request - from litellm.proxy.common_utils.callback_utils import ( - add_guardrail_to_applied_guardrails_header, - ) + if not content: + verbose_proxy_logger.debug("Model Armor: No text content in streaming response, skipping guardrail") + for chunk in all_chunks: + yield chunk + return - add_guardrail_to_applied_guardrails_header( - request_data=request_data, guardrail_name=self.guardrail_name - ) + try: + # Check with Model Armor + armor_response: Final = await self.make_model_armor_request( + content=content, + source="model_response", + request_data=request_data, + ) - # Check if blocked - if self._should_block_content(armor_response): - raise HTTPException( - status_code=400, - detail=self._build_block_error_detail( - "Streaming response blocked by Model Armor", - armor_response, - ), - ) + # Decide the outcome before recording it. Mirrors the non-streaming sibling: with + # masking on, a de-identify match is a redaction to apply rather than a refusal, but + # that only holds while the redaction can actually be delivered + blocked, sanitized_content = self._resolve_streaming_outcome( + armor_response=armor_response, + assembled_response=assembled_response, + content=content, + ) - # Apply sanitization if enabled - if self.mask_response_content: - sanitized_content: Final = self._get_sanitized_content(armor_response) - if sanitized_content and sanitized_content != content: - # Update assembled response - for choice in assembled_response.choices: - if isinstance(choice, Choices): - if choice.message.content: - choice.message.content = sanitized_content + # Attach Model Armor response & status to this request's metadata to avoid race conditions + if isinstance(request_data, dict): + _, metadata = get_or_create_metadata_bucket(request_data) + metadata["_model_armor_response"] = self._build_logging_response(armor_response) + metadata["_model_armor_status"] = "blocked" if blocked else "success" - # Return sanitized stream - mock_response: Final = MockResponseIterator(model_response=assembled_response) - async for chunk in mock_response: - yield chunk - return + # Add guardrail to applied_guardrails BEFORE potential blocking + # This ensures guardrail is recorded even when it blocks the request + add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name) - except ModelArmorAPIError as e: - if self.optional_params.get("fail_on_error", True): - error_obj = {"message": e.detail, "code": "500"} - yield f"data: {json.dumps({'error': error_obj})}\n\n" - return - except HTTPException as e: - # Yield error as SSE event so create_response() detects it and - # returns a proper JSON error response with the correct status code. - # (Raising from a generator hits create_response's generic except → 500.) - detail: Final = e.detail if isinstance(e.detail, dict) else {"message": str(e.detail)} - error_value: Final = detail.get("error", detail) - if isinstance(error_value, dict): - error_obj = dict(error_value) - else: - error_obj = {"message": str(error_value)} - error_obj["code"] = str(e.status_code) - yield f"data: {json.dumps({'error': error_obj})}\n\n" + if blocked: + raise HTTPException( + status_code=400, + detail=self._build_block_error_detail( + "Streaming response blocked by Model Armor", + armor_response, + ), + ) + + if sanitized_content is not None and isinstance(assembled_response, ModelResponse): + self._apply_sanitized_content(assembled_response, sanitized_content) + + # Return sanitized stream + if surface is _StreamSurface.ANTHROPIC_MESSAGES: + for sse_chunk in anthropic_sse_chunks_from_response(assembled_response): + yield sse_chunk return - except Exception as e: - verbose_proxy_logger.error("Model Armor streaming error: %s", str(e), exc_info=True) - if self.optional_params.get("fail_on_error", True): - raise - else: - verbose_proxy_logger.debug("Model Armor: No text content in streaming response, skipping guardrail") + mock_response: Final = MockResponseIterator(model_response=assembled_response) + async for chunk in mock_response: + yield chunk + return + + except ModelArmorAPIError as e: + if self.optional_params.get("fail_on_error", True): + for error_item in self._stream_error_items( + HTTPException(status_code=500, detail=e.detail), surface=surface + ): + yield error_item + return + except HTTPException as e: + # Yield the error as a terminal stream item so create_response() detects it and returns + # a proper JSON error response with the correct status code. Raising from a generator + # instead hits create_response's generic except and becomes a 500. + for error_item in self._stream_error_items(e, surface=surface): + yield error_item + return + except Exception as e: + verbose_proxy_logger.error("Model Armor streaming error: %s", str(e), exc_info=True) + if self.optional_params.get("fail_on_error", True): + raise # Return original chunks if no sanitization needed for chunk in all_chunks: diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py index da66c36328e..47089b7b1b1 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py @@ -15,6 +15,7 @@ import litellm.types.utils from litellm._logging import verbose_proxy_logger from litellm.caching import DualCache from litellm.llms.custom_httpx.http_handler import MaskedHTTPStatusError +from litellm.proxy.guardrails.anthropic_sse import anthropic_sse_error_frames from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.guardrails.guardrail_hooks.model_armor import ModelArmorGuardrail from litellm.proxy.guardrails.guardrail_hooks.model_armor.model_armor import ( @@ -3778,3 +3779,1153 @@ async def test_moderation_hook_skips_chat_traffic_when_configured_for_during_mcp assert result == data mock_post.assert_not_called() + + +_ANTHROPIC_SSE_CHUNKS = ( + b'event: message_start\ndata: {"type":"message_start","message":{"id":"msg_1","type":"message",' + b'"role":"assistant","model":"claude","content":[],"usage":{"input_tokens":5,"output_tokens":0}}}\n\n', + b'event: content_block_start\ndata: {"type":"content_block_start","index":0,' + b'"content_block":{"type":"text","text":""}}\n\n', + b'event: content_block_delta\ndata: {"type":"content_block_delta","index":0,' + b'"delta":{"type":"text_delta","text":"my card is 4111-1111-1111-1111"}}\n\n', + b'event: content_block_stop\ndata: {"type":"content_block_stop","index":0}\n\n', + b'event: message_delta\ndata: {"type":"message_delta","delta":{"stop_reason":"end_turn"},' + b'"usage":{"output_tokens":9}}\n\n', + b'event: message_stop\ndata: {"type":"message_stop"}\n\n', +) + +_MODEL_ARMOR_CLEAN = {"sanitizationResult": {"filterMatchState": "NO_MATCH_FOUND"}} + +_MODEL_ARMOR_BLOCK = { + "sanitizationResult": { + "filterMatchState": "MATCH_FOUND", + "filterResults": { + "sdp": { + "sdpFilterResult": { + "inspectResult": { + "matchState": "MATCH_FOUND", + "findings": [ + {"infoType": "CREDIT_CARD_NUMBER", "likelihood": "VERY_LIKELY"} + ], + } + } + } + }, + } +} + +# The root-level sanitizedText fallback in _get_sanitized_content, i.e. a rewrite that trips no +# named filter +_MODEL_ARMOR_SANITIZED = { + "sanitizedText": "my card is [REDACTED]", + "sanitizationResult": {"filterMatchState": "NO_MATCH_FOUND"}, +} + +# The shape a real de-identify template returns: the SDP filter both matches and hands back the +# rewritten text, so whether it blocks or masks is decided by allow_sanitization alone +_MODEL_ARMOR_DEIDENTIFIED = { + "sanitizationResult": { + "filterMatchState": "MATCH_FOUND", + "filterResults": { + "sdp": { + "sdpFilterResult": { + "deidentifyResult": { + "matchState": "MATCH_FOUND", + "data": {"text": "my card is [REDACTED]"}, + } + } + } + }, + } +} + + +def _chat_completion_chunks(): + """The chat-completions surface: typed ModelResponseStream chunks.""" + return ( + litellm.types.utils.ModelResponseStream( + choices=[ + litellm.types.utils.StreamingChoices( + index=0, + delta=litellm.types.utils.Delta(content="my card is 4111-1111-1111-1111"), + ) + ] + ), + litellm.types.utils.ModelResponseStream( + choices=[ + litellm.types.utils.StreamingChoices( + index=0, + delta=litellm.types.utils.Delta(content=""), + finish_reason="stop", + ) + ] + ), + ) + + +def _surface_guardrail(**kwargs): + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-test", + **kwargs, + ) + guardrail._ensure_access_token_async = AsyncMock( + return_value=("test-token", "test-project") + ) + return guardrail + + +def _armor_post_mock(payload): + mock_response = AsyncMock() + mock_response.status_code = 200 + mock_response.json = AsyncMock(return_value=payload) + return AsyncMock(return_value=mock_response) + + +async def _anthropic_sse_stream(): + for chunk in _ANTHROPIC_SSE_CHUNKS: + yield chunk + + +def _responses_api_events(): + from litellm.types.llms.openai import ( + OutputTextDeltaEvent, + ResponseCompletedEvent, + ResponsesAPIResponse, + ResponsesAPIStreamEvents, + ) + + completed = ResponsesAPIResponse( + id="resp_1", + created_at=0, + model="gpt-4o-mini", + object="response", + output=[ + { + "type": "message", + "id": "msg_1", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "my card is 4111-1111-1111-1111"}], + } + ], + parallel_tool_calls=False, + tool_choice="auto", + tools=[], + ) + return ( + OutputTextDeltaEvent( + type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, + item_id="msg_1", + output_index=0, + content_index=0, + delta="my card is 4111-1111-1111-1111", + ), + ResponseCompletedEvent( + type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + response=completed, + ), + ) + + +async def _drain_surface_hook(guardrail, chunks, request_data=None): + async def _stream(): + for chunk in chunks: + yield chunk + + return [ + item + async for item in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(), + response=_stream(), + request_data=request_data + if request_data is not None + else { + "model": "claude-haiku", + "messages": [{"role": "user", "content": "show me a card"}], + "metadata": {"guardrails": ["model-armor-test"]}, + }, + ) + ] + + +@pytest.mark.asyncio +async def test_streaming_hook_scans_raw_anthropic_sse_instead_of_crashing(): + """A /v1/messages stream arrives as raw SSE bytes and must be assembled, then scanned. + + Regression for the 500 `Error building chunks for logging/streaming usage calculation`: + stream_chunk_builder calls .get() on each chunk, which raises on bytes. + """ + guardrail = _surface_guardrail() + post = _armor_post_mock(_MODEL_ARMOR_CLEAN) + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook(guardrail, _ANTHROPIC_SSE_CHUNKS) + + post.assert_called_once() + scanned = post.call_args.kwargs["json"]["modelResponseData"]["text"] + assert "my card is 4111-1111-1111-1111" in scanned + assert tuple(delivered) == _ANTHROPIC_SSE_CHUNKS + + +@pytest.mark.asyncio +async def test_streaming_hook_scans_responses_api_events_instead_of_crashing(): + """A /v1/responses stream arrives as typed Responses events, which stream_chunk_builder + cannot subscript. The final response.completed event carries the text to scan.""" + guardrail = _surface_guardrail() + post = _armor_post_mock(_MODEL_ARMOR_CLEAN) + events = _responses_api_events() + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook(guardrail, events) + + post.assert_called_once() + scanned = post.call_args.kwargs["json"]["modelResponseData"]["text"] + assert scanned == "my card is 4111-1111-1111-1111" + assert tuple(delivered) == events + + +@pytest.mark.asyncio +async def test_streaming_block_emits_anthropic_error_frame(): + """A block on /v1/messages must terminate the stream in Anthropic's error format. + + The OpenAI-shaped `data: {"error": ...}` frame the chat surface uses is rejected by + Anthropic clients. + """ + guardrail = _surface_guardrail() + + with patch.object( + guardrail.async_handler, "post", _armor_post_mock(_MODEL_ARMOR_BLOCK) + ): + delivered = await _drain_surface_hook(guardrail, _ANTHROPIC_SSE_CHUNKS) + + body = b"".join(delivered) + assert b"event: error" in body + assert b'"type": "error"' in body + assert b"guardrail_error" in body + assert b"Streaming response blocked by Model Armor" in body + assert b"4111-1111-1111-1111" not in body + + +@pytest.mark.asyncio +async def test_streaming_block_emits_responses_api_error_event(): + """A block on /v1/responses must terminate the stream with a Responses ErrorEvent.""" + from litellm.types.llms.openai import ErrorEvent + + guardrail = _surface_guardrail() + + with patch.object( + guardrail.async_handler, "post", _armor_post_mock(_MODEL_ARMOR_BLOCK) + ): + delivered = await _drain_surface_hook(guardrail, _responses_api_events()) + + assert len(delivered) == 1 + error_event = delivered[0] + assert isinstance(error_event, ErrorEvent) + assert error_event.error.type == "guardrail_error" + assert error_event.error.code == "400" + assert error_event.error.message == "Streaming response blocked by Model Armor" + + +@pytest.mark.asyncio +async def test_streaming_masking_re_emits_anthropic_sse_with_sanitized_text(): + """mask_response_content on /v1/messages must ship the sanitized text, not the original.""" + guardrail = _surface_guardrail(mask_response_content=True) + + with patch.object( + guardrail.async_handler, "post", _armor_post_mock(_MODEL_ARMOR_SANITIZED) + ): + delivered = await _drain_surface_hook(guardrail, _ANTHROPIC_SSE_CHUNKS) + + body = b"".join(delivered) + assert b"[REDACTED]" in body + assert b"4111-1111-1111-1111" not in body + + +@pytest.mark.asyncio +async def test_streaming_masking_blocks_responses_api_stream(): + """A Responses event stream cannot be rebuilt from sanitized text, so releasing it would + ship the content the guardrail just rewrote. It is blocked instead.""" + from litellm.types.llms.openai import ErrorEvent + + guardrail = _surface_guardrail(mask_response_content=True) + + with patch.object( + guardrail.async_handler, "post", _armor_post_mock(_MODEL_ARMOR_SANITIZED) + ): + delivered = await _drain_surface_hook(guardrail, _responses_api_events()) + + assert len(delivered) == 1 + assert isinstance(delivered[0], ErrorEvent) + assert delivered[0].error.code == "400" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("surface", ["anthropic_sse", "responses"]) +async def test_streaming_api_failure_frames_error_per_surface(surface): + """A Model Armor outage with fail_on_error must terminate the stream in the endpoint's + own error format rather than leaking an OpenAI SSE frame onto it.""" + from litellm.types.llms.openai import ErrorEvent + + guardrail = _surface_guardrail(fail_on_error=True) + chunks = _ANTHROPIC_SSE_CHUNKS if surface == "anthropic_sse" else _responses_api_events() + + mock_response = AsyncMock() + mock_response.status_code = 500 + mock_response.text = "Internal Server Error" + + with patch.object( + guardrail.async_handler, "post", AsyncMock(return_value=mock_response) + ): + delivered = await _drain_surface_hook(guardrail, chunks) + + assert len(delivered) >= 1 + if surface == "anthropic_sse": + assert b"event: error" in b"".join(delivered) + else: + assert isinstance(delivered[0], ErrorEvent) + assert delivered[0].error.code == "500" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "chunks", + [ + pytest.param( + (b'event: content_block_delta\ndata: {"type":"content_block_delta","index":0,' + b'"delta":{"type":"text_delta","text":"hi"}}\n\n',), + id="anthropic-sse-without-message-start", + ), + pytest.param(None, id="responses-stream-without-completed-event"), + pytest.param("created", id="responses-stream-cut-off-after-response-created"), + ], +) +async def test_streaming_hook_fails_closed_when_a_surface_stream_cannot_be_assembled(chunks): + """Forwarding an unscannable /v1/messages or /v1/responses stream would silently disable the + guardrail, so the stream is refused in its own wire format instead of released unscanned.""" + from litellm.types.llms.openai import ( + ErrorEvent, + OutputTextDeltaEvent, + ResponsesAPIStreamEvents, + ) + + if chunks is None or chunks == "created": + delta = OutputTextDeltaEvent( + type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, + item_id="msg_1", + output_index=0, + content_index=0, + delta="my card is 4111-1111-1111-1111", + ) + # response.created carries a ResponsesAPIResponse too, but an empty one: reading the body + # off it would scan "" and release every buffered delta unscanned + chunks = (delta,) if chunks is None else (_responses_created_event(), delta) + guardrail = _surface_guardrail() + post = _armor_post_mock(_MODEL_ARMOR_CLEAN) + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook(guardrail, chunks) + + post.assert_not_called() + assert tuple(delivered) != tuple(chunks) + if isinstance(chunks[0], bytes): + joined = b"".join(item.encode() if isinstance(item, str) else item for item in delivered).decode() + assert "event: error" in joined + assert "could not be assembled for scanning" in joined + return + assert len(delivered) == 1 + assert isinstance(delivered[0], ErrorEvent) + assert "could not be assembled for scanning" in delivered[0].error.message + + +@pytest.mark.asyncio +async def test_streaming_hook_forwards_a_preceding_guardrails_error_item(): + """A guardrail earlier in the post_call chain replaces the stream with its own terminal + error item. That item is not a chat delta, and feeding it to stream_chunk_builder is what + surfaced the ticket's 500, so it has to be forwarded untouched instead.""" + from litellm.types.llms.openai import ( + ErrorEvent, + ErrorEventError, + ResponsesAPIStreamEvents, + ) + + chunks = ( + ErrorEvent( + type=ResponsesAPIStreamEvents.ERROR, + sequence_number=1, + error=ErrorEventError( + type="guardrail_error", + code="400", + message="Streaming response blocked by Model Armor", + param=None, + ), + ), + ) + guardrail = _surface_guardrail() + post = _armor_post_mock(_MODEL_ARMOR_CLEAN) + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook(guardrail, chunks) + + post.assert_not_called() + assert tuple(delivered) == chunks + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "chunks", + [ + pytest.param(None, id="anthropic-error-event"), + pytest.param( + ('data: {"error": {"message": "Streaming response blocked by the first guardrail", "code": "400"}}\n\n',), + id="chat-completions-error-payload", + ), + ], +) +async def test_streaming_hook_forwards_a_preceding_guardrails_error_frame(chunks): + """Chained post_call guardrails hand each other their output. An earlier guardrail's error + frame carries no message to assemble, and replacing it would hide the real refusal.""" + if chunks is None: + chunks = anthropic_sse_error_frames("Streaming response blocked by the first guardrail") + guardrail = _surface_guardrail() + post = _armor_post_mock(_MODEL_ARMOR_CLEAN) + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook(guardrail, chunks) + + post.assert_not_called() + assert tuple(delivered) == chunks + + +@pytest.mark.asyncio +async def test_streaming_responses_error_falls_back_to_sse_when_the_handler_declines(): + """build_stream_error_items may return None, which must not swallow the block into a clean + 200: the refusal falls back to the chat-completions SSE form that still carries the status.""" + from litellm.proxy.guardrails.guardrail_hooks.model_armor.model_armor import ( + _StreamSurface, + ) + + class _DecliningGuardrail(ModelArmorGuardrail): + @staticmethod + def _build_responses_error_items(exc): + return None + + guardrail = _DecliningGuardrail( + template_id="test-template", + project_id="test-project", + location="us-central1", + guardrail_name="model-armor-test", + ) + exc = HTTPException(status_code=400, detail={"message": "blocked"}) + + items = guardrail._stream_error_items(exc, surface=_StreamSurface.RESPONSES) + + assert len(items) == 1 + assert '"code": "400"' in items[0] + assert "blocked" in items[0] + + +def _responses_created_event(): + from litellm.types.llms.openai import ( + ResponseCreatedEvent, + ResponsesAPIResponse, + ResponsesAPIStreamEvents, + ) + + return ResponseCreatedEvent( + type=ResponsesAPIStreamEvents.RESPONSE_CREATED, + response=ResponsesAPIResponse( + id="resp_1", + created_at=0, + model="gpt-4o-mini", + object="response", + output=[], + parallel_tool_calls=False, + tool_choice="auto", + tools=[], + ), + ) + + +@pytest.mark.asyncio +async def test_streaming_hook_refuses_an_opaque_sse_stream_without_anthropic_framing(): + """/v1/messages is not the only endpoint that streams raw SSE: the Google generateContent + route marks its own stream raw too. Its frames carry no Anthropic event types, so refusing + them in Anthropic's format would hand a Google client a body it cannot parse.""" + guardrail = _surface_guardrail() + post = _armor_post_mock(_MODEL_ARMOR_CLEAN) + chunks = (b'data: {"candidates":[{"content":{"parts":[{"text":"my card is 4111"}]}}]}\n\n',) + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook(guardrail, chunks) + + post.assert_not_called() + assert tuple(delivered) != chunks + body = "".join(item.decode() if isinstance(item, bytes) else item for item in delivered) + assert "could not be assembled for scanning" in body + assert "event: error" not in body + assert '"code": "500"' in body + + +@pytest.mark.asyncio +async def test_streaming_unassemblable_stream_is_forwarded_when_fail_on_error_is_disabled(): + """fail_on_error: false is a deliberate choice to degrade open, and it governs every other + path in this hook. The fail-closed refusal has to honour it too.""" + guardrail = _surface_guardrail(fail_on_error=False) + post = _armor_post_mock(_MODEL_ARMOR_CLEAN) + chunks = ( + b'event: content_block_delta\ndata: {"type":"content_block_delta","index":0,' + b'"delta":{"type":"text_delta","text":"hi"}}\n\n', + ) + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook(guardrail, chunks) + + post.assert_not_called() + assert tuple(delivered) == chunks + + +@pytest.mark.asyncio +async def test_streaming_fail_closed_records_the_applied_guardrail(): + """A refusal that no header or log attributes to the guardrail leaves on-call unable to tell + a guardrail block apart from a provider failure.""" + guardrail = _surface_guardrail() + request_data = { + "model": "claude-haiku", + "messages": [{"role": "user", "content": "show me a card"}], + "metadata": {"guardrails": ["model-armor-test"]}, + } + chunks = ( + b'event: content_block_delta\ndata: {"type":"content_block_delta","index":0,' + b'"delta":{"type":"text_delta","text":"hi"}}\n\n', + ) + + with patch.object(guardrail.async_handler, "post", _armor_post_mock(_MODEL_ARMOR_CLEAN)): + await _drain_surface_hook(guardrail, chunks, request_data=request_data) + + assert request_data["metadata"]["applied_guardrails"] == ["model-armor-test"] + + +@pytest.mark.asyncio +async def test_streaming_responses_tool_call_output_is_scanned(): + """An agentic /v1/responses turn can carry its whole payload in tool-call arguments, which + is what the chat surface already folds into the scanned text.""" + from litellm.types.llms.openai import ( + ResponseCompletedEvent, + ResponsesAPIResponse, + ResponsesAPIStreamEvents, + ) + + guardrail = _surface_guardrail() + post = _armor_post_mock(_MODEL_ARMOR_CLEAN) + completed = ResponseCompletedEvent( + type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + response=ResponsesAPIResponse( + id="resp_1", + created_at=0, + model="gpt-4o-mini", + object="response", + output=[ + { + "type": "function_call", + "id": "fc_1", + "call_id": "call_1", + "name": "send_email", + "arguments": '{"body": "my card is 4111-1111-1111-1111"}', + } + ], + parallel_tool_calls=False, + tool_choice="auto", + tools=[], + ), + ) + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook(guardrail, (completed,)) + + post.assert_called_once() + scanned = post.call_args.kwargs["json"]["modelResponseData"]["text"] + assert "4111-1111-1111-1111" in scanned + assert tuple(delivered) == (completed,) + + +@pytest.mark.asyncio +async def test_streaming_hook_refuses_a_content_stream_that_ends_with_an_error_frame(): + """The chain-aware passthrough must stay narrow. A stream carrying real content plus a + trailing error frame is not a bare refusal to forward: the assembler cannot read it, and + releasing it would ship the buffered content unscanned.""" + guardrail = _surface_guardrail() + post = _armor_post_mock(_MODEL_ARMOR_CLEAN) + chunks = (*_ANTHROPIC_SSE_CHUNKS, *anthropic_sse_error_frames("upstream gave up")) + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook(guardrail, chunks) + + post.assert_not_called() + body = b"".join(delivered) + assert b"4111-1111-1111-1111" not in body + assert b"could not be assembled for scanning" in body + + +@pytest.mark.parametrize( + "chunks, expected, case", + [ + (anthropic_sse_error_frames("blocked upstream"), True, "anthropic-error-frames-only"), + ((f"data: {json.dumps({'error': {'message': 'blocked'}})}\n\n",), True, "chat-error-payload-only"), + ((), False, "empty-stream"), + ( + (b'event: message_delta\ndata: {"type":"message_delta","error":null}\n\n',), + False, + "content-event-carrying-a-null-error-field", + ), + ( + ( + litellm.types.utils.ModelResponseStream( + choices=[ + litellm.types.utils.StreamingChoices( + index=0, + delta=litellm.types.utils.Delta(content="my card is 4111-1111-1111-1111"), + ) + ] + ), + *anthropic_sse_error_frames("upstream gave up"), + ), + False, + "typed-content-chunks-plus-a-trailing-error-frame", + ), + ], +) +def test_is_sse_error_stream_only_matches_a_stream_that_is_nothing_but_refusals(chunks, expected, case): + """The chain-aware passthrough turns on this predicate, so anything it calls error-only is + forwarded to the client untouched. A stream that still carries content must not qualify: the + frames-only join drops typed chunks, and a content event may carry an empty ``error`` field.""" + from litellm.proxy.guardrails.anthropic_sse import is_sse_error_stream + + assert is_sse_error_stream(chunks) is expected, case + + +@pytest.mark.asyncio +async def test_streaming_hook_does_not_forward_typed_chunks_that_end_with_an_error_frame(): + """A stream mixing buffered content with a trailing refusal is not the bare refusal the chain + passthrough exists for. Forwarding it would release the content no scanner ever saw.""" + guardrail = _surface_guardrail() + post = _armor_post_mock(_MODEL_ARMOR_CLEAN) + chunks = ( + litellm.types.utils.ModelResponseStream( + choices=[ + litellm.types.utils.StreamingChoices( + index=0, + delta=litellm.types.utils.Delta(content="my card is 4111-1111-1111-1111"), + ) + ] + ), + *anthropic_sse_error_frames("upstream gave up"), + ) + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook(guardrail, chunks) + + post.assert_not_called() + body = b"".join(item if isinstance(item, bytes) else str(item).encode() for item in delivered) + assert b"4111-1111-1111-1111" not in body + assert b"could not be assembled for scanning" in body + + +def _delivered_bytes(delivered): + return b"".join( + item + if isinstance(item, bytes) + else item.encode() + if isinstance(item, str) + else str(item.model_dump() if hasattr(item, "model_dump") else item).encode() + for item in delivered + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("chunks, case", [(None, "chat_completions"), (_ANTHROPIC_SSE_CHUNKS, "anthropic_sse")]) +async def test_streaming_deidentify_match_masks_when_masking_is_enabled(chunks, case): + """A de-identify template reports MATCH_FOUND for every redaction it makes, so reading that + match as a refusal makes mask_response_content unusable on a stream: the client gets an error + where its non-streaming sibling gets redacted text. The block check has to allow sanitization + exactly as the non-streaming hook does.""" + guardrail = _surface_guardrail(mask_response_content=True) + post = _armor_post_mock(_MODEL_ARMOR_DEIDENTIFIED) + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook( + guardrail, _chat_completion_chunks() if chunks is None else chunks + ) + + body = _delivered_bytes(delivered) + assert b"[REDACTED]" in body, case + assert b"4111-1111-1111-1111" not in body, case + assert b"blocked by Model Armor" not in body, case + + +@pytest.mark.asyncio +@pytest.mark.parametrize("chunks, case", [(None, "chat_completions"), (_ANTHROPIC_SSE_CHUNKS, "anthropic_sse")]) +async def test_streaming_deidentify_match_still_blocks_when_masking_is_disabled(chunks, case): + """Without mask_response_content there is nowhere to put the rewritten text, so the same + de-identify match must still end the stream rather than release the original.""" + guardrail = _surface_guardrail() + post = _armor_post_mock(_MODEL_ARMOR_DEIDENTIFIED) + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook( + guardrail, _chat_completion_chunks() if chunks is None else chunks + ) + + body = _delivered_bytes(delivered) + assert b"Streaming response blocked by Model Armor" in body, case + assert b"4111-1111-1111-1111" not in body, case + + +@pytest.mark.asyncio +async def test_streaming_deidentify_match_logs_masked_run_as_success_not_blocked(): + """The status stamped on request metadata feeds the spend log, so it has to agree with what + the client actually received: a masked stream is a success, not a block.""" + guardrail = _surface_guardrail(mask_response_content=True) + request_data = { + "model": "claude-haiku", + "messages": [{"role": "user", "content": "show me a card"}], + "metadata": {"guardrails": ["model-armor-test"]}, + } + + with patch.object(guardrail.async_handler, "post", _armor_post_mock(_MODEL_ARMOR_DEIDENTIFIED)): + await _drain_surface_hook(guardrail, _chat_completion_chunks(), request_data=request_data) + + assert request_data["metadata"]["_model_armor_status"] == "success" + + +# A de-identify template that matched but handed back no rewrite, e.g. because the transformation +# itself failed. The match still says the buffered original carries what it matched on +_MODEL_ARMOR_DEIDENTIFIED_NO_TEXT = { + "sanitizationResult": { + "filterMatchState": "MATCH_FOUND", + "filterResults": { + "sdp": {"sdpFilterResult": {"deidentifyResult": {"matchState": "MATCH_FOUND"}}} + }, + } +} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("chunks, case", [(None, "chat_completions"), (_ANTHROPIC_SSE_CHUNKS, "anthropic_sse")]) +async def test_streaming_deidentify_match_without_a_rewrite_fails_closed(chunks, case): + """Allowing sanitization past the block check is a promise to apply the redaction. When Model + Armor matches but returns no sanitized text there is nothing to apply, and yielding the + buffered chunks would hand back exactly what it matched on.""" + guardrail = _surface_guardrail(mask_response_content=True) + post = _armor_post_mock(_MODEL_ARMOR_DEIDENTIFIED_NO_TEXT) + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook( + guardrail, _chat_completion_chunks() if chunks is None else chunks + ) + + body = _delivered_bytes(delivered) + assert b"4111-1111-1111-1111" not in body, case + assert b"Streaming response blocked by Model Armor" in body, case + + +@pytest.mark.asyncio +async def test_streaming_status_records_a_surface_that_cannot_carry_the_rewrite_as_blocked(): + """The Responses surface has no assembled body to rewrite, so a de-identify match ends as a + refusal. The status stamped on metadata feeds the spend log and has to say so rather than + reporting the success the block check alone would have implied.""" + guardrail = _surface_guardrail(mask_response_content=True) + request_data = { + "model": "gpt-4o-mini", + "input": "show me a card", + "metadata": {"guardrails": ["model-armor-test"]}, + } + + with patch.object(guardrail.async_handler, "post", _armor_post_mock(_MODEL_ARMOR_DEIDENTIFIED)): + delivered = await _drain_surface_hook( + guardrail, _responses_api_events(), request_data=request_data + ) + + body = _delivered_bytes(delivered) + assert b"4111-1111-1111-1111" not in body + assert b"Streaming response blocked by Model Armor" in body + assert request_data["metadata"]["_model_armor_status"] == "blocked" + + +def _responses_api_events_truncated(terminal: str): + """A /v1/responses stream whose text went out as deltas and whose terminal event reports no body. + + ``response.failed`` and ``response.incomplete`` are terminal like ``response.completed``, but a + turn that broke mid-generation reports an empty ``output`` while the deltas ahead of it already + spelled the answer out to the client. + """ + from litellm.types.llms.openai import ( + OutputTextDeltaEvent, + ResponseFailedEvent, + ResponseIncompleteEvent, + ResponsesAPIResponse, + ResponsesAPIStreamEvents, + ) + + empty_body = ResponsesAPIResponse( + id="resp_1", + created_at=0, + model="gpt-4o-mini", + object="response", + output=[], + parallel_tool_calls=False, + tool_choice="auto", + tools=[], + ) + terminal_event = ( + ResponseFailedEvent(type=ResponsesAPIStreamEvents.RESPONSE_FAILED, response=empty_body) + if terminal == "failed" + else ResponseIncompleteEvent(type=ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE, response=empty_body) + ) + return ( + OutputTextDeltaEvent( + type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, + item_id="msg_1", + output_index=0, + content_index=0, + delta="my card is 4111-1111-1111-1111", + ), + terminal_event, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("terminal", ["failed", "incomplete"]) +async def test_streaming_responses_terminal_event_without_a_body_still_scans_the_deltas(terminal): + """A /v1/responses turn that broke mid-generation has still delivered its deltas. + + Reading only the terminal body would find nothing to scan and hand every buffered delta to the + client untouched, so the deltas themselves are what gets scanned. + """ + guardrail = _surface_guardrail() + post = _armor_post_mock(_MODEL_ARMOR_BLOCK) + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook(guardrail, _responses_api_events_truncated(terminal)) + + post.assert_called_once() + assert "4111-1111-1111-1111" in post.call_args.kwargs["json"]["modelResponseData"]["text"] + rendered = "".join(str(item) for item in delivered) + assert "4111-1111-1111-1111" not in rendered + assert "Streaming response blocked by Model Armor" in rendered + + +@pytest.mark.asyncio +async def test_streaming_responses_mcp_argument_deltas_are_scanned_when_the_body_is_empty(): + """A turn that only streamed MCP tool arguments still handed the client a payload. + + The delta fallback is read off the event enum rather than listed by hand, so an argument event + that carries no `output_text` cannot fall out of the scan. + """ + from litellm.types.llms.openai import ( + MCPCallArgumentsDeltaEvent, + ResponseIncompleteEvent, + ResponsesAPIResponse, + ResponsesAPIStreamEvents, + ) + + chunks = ( + MCPCallArgumentsDeltaEvent( + type=ResponsesAPIStreamEvents.MCP_CALL_ARGUMENTS_DELTA, + output_index=0, + item_id="mcp_1", + delta='{"note": "my card is 4111-1111-1111-1111"}', + sequence_number=0, + ), + ResponseIncompleteEvent( + type=ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE, + response=ResponsesAPIResponse( + id="resp_1", + created_at=0, + model="gpt-4o-mini", + object="response", + output=[], + parallel_tool_calls=False, + tool_choice="auto", + tools=[], + ), + ), + ) + guardrail = _surface_guardrail() + post = _armor_post_mock(_MODEL_ARMOR_BLOCK) + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook(guardrail, chunks) + + post.assert_called_once() + assert "4111-1111-1111-1111" in post.call_args.kwargs["json"]["modelResponseData"]["text"] + rendered = "".join(str(item) for item in delivered) + assert "4111-1111-1111-1111" not in rendered + assert "Streaming response blocked by Model Armor" in rendered + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "with_output_text_delta", + [True, False], + ids=["summary-and-text-deltas", "summary-delta-only"], +) +async def test_streaming_responses_reasoning_summary_deltas_are_scanned_alongside_the_body(with_output_text_delta): + """A reasoning turn streams its summary in deltas the terminal body never repeats. + + Reading only the body scans the visible answer and hands the client every summary delta + unscanned, so the body and the deltas are scanned together. + """ + from litellm.types.llms.openai import ( + OutputTextDeltaEvent, + ReasoningSummaryTextDeltaEvent, + ResponseCompletedEvent, + ResponsesAPIResponse, + ResponsesAPIStreamEvents, + ) + + answer = "the weather is fine" + summary_delta = ReasoningSummaryTextDeltaEvent( + type=ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DELTA, + item_id="rs_1", + output_index=0, + delta="the user said my card is 4111-1111-1111-1111", + ) + text_deltas = ( + ( + OutputTextDeltaEvent( + type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, + item_id="msg_1", + output_index=1, + content_index=0, + delta=answer, + ), + ) + if with_output_text_delta + else () + ) + completed = ResponseCompletedEvent( + type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + response=ResponsesAPIResponse( + id="resp_1", + created_at=0, + model="gpt-5-mini", + object="response", + output=[ + { + "type": "message", + "id": "msg_1", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": answer, "annotations": []}], + } + ], + parallel_tool_calls=False, + tool_choice="auto", + tools=[], + ), + ) + guardrail = _surface_guardrail() + post = _armor_post_mock(_MODEL_ARMOR_BLOCK) + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook(guardrail, (summary_delta, *text_deltas, completed)) + + post.assert_called_once() + scanned = post.call_args.kwargs["json"]["modelResponseData"]["text"] + assert "4111-1111-1111-1111" in scanned + assert answer in scanned + assert scanned.count(answer) == 1 + rendered = "".join(str(item) for item in delivered) + assert "4111-1111-1111-1111" not in rendered + assert "Streaming response blocked by Model Armor" in rendered + + +@pytest.mark.asyncio +async def test_streaming_responses_deltas_of_separate_fields_do_not_form_a_finding_across_their_boundary(): + """Two fields of a turn are separate text, so what runs across their boundary is not model output. + + A reasoning summary ending in half a card number and an answer opening with the other half + each carry nothing to find, and joining them without a break would invent one. + """ + from litellm.types.llms.openai import ( + OutputTextDeltaEvent, + ReasoningSummaryTextDeltaEvent, + ResponseCompletedEvent, + ResponsesAPIResponse, + ResponsesAPIStreamEvents, + ) + + answer = "1111-1111 is not a full card" + summary_delta = ReasoningSummaryTextDeltaEvent( + type=ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DELTA, + item_id="rs_1", + output_index=0, + delta="the prefix they gave me is 4111-1111-", + ) + text_delta = OutputTextDeltaEvent( + type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, + item_id="msg_1", + output_index=1, + content_index=0, + delta=answer, + ) + completed = ResponseCompletedEvent( + type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + response=ResponsesAPIResponse( + id="resp_1", + created_at=0, + model="gpt-5-mini", + object="response", + output=[ + { + "type": "message", + "id": "msg_1", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": answer, "annotations": []}], + } + ], + parallel_tool_calls=False, + tool_choice="auto", + tools=[], + ), + ) + guardrail = _surface_guardrail() + post = _armor_post_mock(_MODEL_ARMOR_CLEAN) + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook(guardrail, (summary_delta, text_delta, completed)) + + post.assert_called_once() + scanned = post.call_args.kwargs["json"]["modelResponseData"]["text"] + assert "4111-1111-" in scanned + assert answer in scanned + assert "4111-1111-1111-1111" not in scanned + rendered = "".join(str(item) for item in delivered) + assert "Streaming response blocked by Model Armor" not in rendered + + +@pytest.mark.asyncio +async def test_streaming_responses_one_fields_deltas_still_join_into_a_single_finding(): + """A card number split across two deltas of one field is still one card number to scan.""" + from litellm.types.llms.openai import ( + OutputTextDeltaEvent, + ResponseCompletedEvent, + ResponsesAPIResponse, + ResponsesAPIStreamEvents, + ) + + halves = ("my card is 4111-1111-", "1111-1111") + text_deltas = tuple( + OutputTextDeltaEvent( + type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, + item_id="msg_1", + output_index=0, + content_index=0, + delta=half, + ) + for half in halves + ) + completed = ResponseCompletedEvent( + type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + response=ResponsesAPIResponse( + id="resp_1", + created_at=0, + model="gpt-5-mini", + object="response", + output=[], + parallel_tool_calls=False, + tool_choice="auto", + tools=[], + ), + ) + guardrail = _surface_guardrail() + post = _armor_post_mock(_MODEL_ARMOR_BLOCK) + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook(guardrail, (*text_deltas, completed)) + + post.assert_called_once() + assert "4111-1111-1111-1111" in post.call_args.kwargs["json"]["modelResponseData"]["text"] + rendered = "".join(str(item) for item in delivered) + assert "4111-1111-1111-1111" not in rendered + assert "Streaming response blocked by Model Armor" in rendered + + +@pytest.mark.asyncio +async def test_streaming_responses_fields_the_body_repeats_are_not_scanned_a_second_time(): + """A turn whose visible fields all reach the terminal body is scanned once, not twice. + + Two output_text fields stream as deltas and come back in the completed body, so scanning the + deltas on top of the body would send Model Armor two copies of everything the client sees. + """ + from litellm.types.llms.openai import ( + OutputTextDeltaEvent, + ResponseCompletedEvent, + ResponsesAPIResponse, + ResponsesAPIStreamEvents, + ) + + paragraphs = ("the first thing to know", "a second and separate point") + text_deltas = tuple( + OutputTextDeltaEvent( + type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, + item_id=f"msg_{index}", + output_index=index, + content_index=0, + delta=paragraph, + ) + for index, paragraph in enumerate(paragraphs) + ) + completed = ResponseCompletedEvent( + type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + response=ResponsesAPIResponse( + id="resp_1", + created_at=0, + model="gpt-5-mini", + object="response", + output=[ + { + "type": "message", + "id": f"msg_{index}", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": paragraph, "annotations": []}], + } + for index, paragraph in enumerate(paragraphs) + ], + parallel_tool_calls=False, + tool_choice="auto", + tools=[], + ), + ) + guardrail = _surface_guardrail() + post = _armor_post_mock(_MODEL_ARMOR_CLEAN) + + with patch.object(guardrail.async_handler, "post", post): + delivered = await _drain_surface_hook(guardrail, (*text_deltas, completed)) + + post.assert_called_once() + scanned = post.call_args.kwargs["json"]["modelResponseData"]["text"] + assert [scanned.count(paragraph) for paragraph in paragraphs] == [1, 1] + rendered = "".join(str(item) for item in delivered) + assert all(paragraph in rendered for paragraph in paragraphs) + + +def test_every_responses_delta_event_is_in_the_scanned_set(): + """Every ``.delta`` the Responses event enum defines is model output on its way to the client.""" + from litellm.proxy.guardrails.guardrail_hooks.model_armor.model_armor import ( + _RESPONSES_DELTA_EVENT_TYPES, + ) + from litellm.types.llms.openai import ResponsesAPIStreamEvents + + missing = { + event.value + for event in ResponsesAPIStreamEvents + if event.value.endswith(".delta") and event.value not in _RESPONSES_DELTA_EVENT_TYPES + } + assert not missing + assert "response.mcp_call_arguments.delta" in _RESPONSES_DELTA_EVENT_TYPES