mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix: enforce nested response safety identifiers
This commit is contained in:
parent
94193b34c8
commit
0d5a3bf2c3
3 changed files with 72 additions and 7 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue