fix: enforce nested response safety identifiers

This commit is contained in:
Dominic White 2026-09-07 11:42:20 +02:00
parent 94193b34c8
commit 0d5a3bf2c3
3 changed files with 72 additions and 7 deletions

View file

@ -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):

View file

@ -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()

View file

@ -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