diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 14f1af19d23..0788a6e317e 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -99,11 +99,13 @@ def _enforce_responses_ws_safety_identifier( msg_obj: _MutableJsonObject, user_api_key_dict: UserAPIKeyAuth | None, ) -> bool: - return enforce_safety_identifier( - data=msg_obj, - user_id=user_api_key_dict.user_id if user_api_key_dict is not None else None, - enabled=str_to_bool(os.getenv("LITELLM_ENFORCE_SAFETY_IDENTIFIER")) is True, - ) + user_id: Final[str | None] = user_api_key_dict.user_id if user_api_key_dict is not None else None + enabled: Final[bool] = str_to_bool(os.getenv("LITELLM_ENFORCE_SAFETY_IDENTIFIER")) is True + modified = enforce_safety_identifier(data=msg_obj, user_id=user_id, enabled=enabled) + nested_candidate: Final = msg_obj.get("response") + if _is_json_object(nested_candidate): + modified = enforce_safety_identifier(data=nested_candidate, user_id=user_id, enabled=enabled) or modified + return modified class _MutableJsonObject(Protocol): diff --git a/tests/proxy_unit_tests/test_safety_identifier.py b/tests/proxy_unit_tests/test_safety_identifier.py index 6072a88241b..e4b79ef77ef 100644 --- a/tests/proxy_unit_tests/test_safety_identifier.py +++ b/tests/proxy_unit_tests/test_safety_identifier.py @@ -3,7 +3,6 @@ from typing import Literal from unittest.mock import AsyncMock, MagicMock, patch import pytest -from fastapi import Request import litellm from litellm.llms.perplexity.responses.transformation import PerplexityResponsesConfig @@ -108,7 +107,7 @@ async def test_pre_call_hook_cannot_override_enforced_safety_identifier( monkeypatch: pytest.MonkeyPatch, route_type: Literal["acompletion", "aresponses"] ): monkeypatch.setenv("LITELLM_ENFORCE_SAFETY_IDENTIFIER", "true") - request = MagicMock(spec=Request) + request = MagicMock() request.headers.get.return_value = "call-id" logging_obj = MagicMock() proxy_logging_obj = MagicMock() diff --git a/tests/test_litellm/responses/test_responses_websocket_all_providers.py b/tests/test_litellm/responses/test_responses_websocket_all_providers.py index fec56fe7ab1..f911e0cc23c 100644 --- a/tests/test_litellm/responses/test_responses_websocket_all_providers.py +++ b/tests/test_litellm/responses/test_responses_websocket_all_providers.py @@ -1235,6 +1235,46 @@ class TestNativeWebSocketGuardrails: assert "safety_identifier" not in json.loads(masked) + @pytest.mark.asyncio + async def test_nested_response_create_overwrites_safety_identifier(self, monkeypatch: pytest.MonkeyPatch): + from litellm.proxy._types import UserAPIKeyAuth + from litellm.responses.streaming_iterator import ResponsesWebSocketStreaming + + monkeypatch.setenv("LITELLM_ENFORCE_SAFETY_IDENTIFIER", "true") + handler = ResponsesWebSocketStreaming( + websocket=MagicMock(), + backend_ws=MagicMock(), + logging_obj=MagicMock(), + user_api_key_dict=UserAPIKeyAuth(user_id="user-123"), + ) + + masked = await handler._mask_response_create( + json.dumps({"type": "response.create", "response": {"safety_identifier": "caller-value"}}) + ) + + assert json.loads(masked)["response"]["safety_identifier"] == hashlib.sha256(b"user-123").hexdigest() + + @pytest.mark.asyncio + async def test_nested_response_create_removes_safety_identifier_without_user_id( + self, monkeypatch: pytest.MonkeyPatch + ): + from litellm.proxy._types import UserAPIKeyAuth + from litellm.responses.streaming_iterator import ResponsesWebSocketStreaming + + monkeypatch.setenv("LITELLM_ENFORCE_SAFETY_IDENTIFIER", "true") + handler = ResponsesWebSocketStreaming( + websocket=MagicMock(), + backend_ws=MagicMock(), + logging_obj=MagicMock(), + user_api_key_dict=UserAPIKeyAuth(user_id=None), + ) + + masked = await handler._mask_response_create( + json.dumps({"type": "response.create", "response": {"safety_identifier": "caller-value"}}) + ) + + assert "safety_identifier" not in json.loads(masked)["response"] + @pytest.mark.asyncio async def test_managed_response_create_forwards_trusted_safety_identifier(self, monkeypatch: pytest.MonkeyPatch): from litellm.proxy._types import UserAPIKeyAuth @@ -1257,6 +1297,30 @@ class TestNativeWebSocketGuardrails: call_kwargs = stream_and_forward.call_args.args[1] assert call_kwargs["safety_identifier"] == hashlib.sha256(b"user-123").hexdigest() + @pytest.mark.asyncio + async def test_managed_nested_response_create_forwards_trusted_safety_identifier( + self, monkeypatch: pytest.MonkeyPatch + ): + from litellm.proxy._types import UserAPIKeyAuth + from litellm.responses.streaming_iterator import ManagedResponsesWebSocketHandler + + monkeypatch.setenv("LITELLM_ENFORCE_SAFETY_IDENTIFIER", "true") + handler = ManagedResponsesWebSocketHandler( + websocket=MagicMock(), + model="gpt-4o", + logging_obj=MagicMock(), + user_api_key_dict=UserAPIKeyAuth(user_id="user-123"), + ) + stream_and_forward = AsyncMock(return_value=None) + monkeypatch.setattr(handler, "_stream_and_forward", stream_and_forward) + + await handler._process_response_create( + json.dumps({"type": "response.create", "response": {"input": "hi", "safety_identifier": "caller-value"}}) + ) + + call_kwargs = stream_and_forward.call_args.args[1] + assert call_kwargs["safety_identifier"] == hashlib.sha256(b"user-123").hexdigest() + @pytest.mark.asyncio async def test_response_create_injects_authorized_model(self): import json