From 232007e7f92a037bfb48f16edde32e02b403f6a2 Mon Sep 17 00:00:00 2001 From: Dominic White Date: Mon, 7 Sep 2026 09:03:45 +0200 Subject: [PATCH 1/7] feat(proxy): enforce trusted safety identifiers --- litellm/proxy/common_request_processing.py | 32 +++++ litellm/utils.py | 8 ++ proxy_server_config.yaml | 3 +- .../test_safety_identifier.py | 131 ++++++++++++++++++ 4 files changed, 173 insertions(+), 1 deletion(-) create mode 100644 tests/proxy_unit_tests/test_safety_identifier.py diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 5e6c9b34332..074782c3ec7 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1,8 +1,10 @@ import asyncio import contextlib +import hashlib import json import logging import math +import os from collections.abc import AsyncGenerator, Awaitable, Callable, Coroutine, Mapping, Sequence from datetime import datetime from functools import lru_cache @@ -67,6 +69,7 @@ from litellm.proxy.utils import ProxyLogging, _check_and_merge_model_level_guard from litellm.router import Router from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict from litellm.router_utils.common_utils import resolve_model_group_alias +from litellm.secret_managers.main import str_to_bool from litellm.types.guardrails import GuardrailEventHooks from litellm.types.router import RouterRateLimitError @@ -1532,6 +1535,23 @@ class ProxyBaseLLMRequestProcessing: def __init__(self, data: dict): self.data = data + @staticmethod + def _enforce_safety_identifier( + *, + data: dict[str, object], + route_type: ProxyRouteType, + user_api_key_dict: UserAPIKeyAuth, + ) -> dict[str, object]: + if route_type not in {"acompletion", "aresponses"}: + return data + if str_to_bool(os.getenv("LITELLM_ENFORCE_SAFETY_IDENTIFIER")) is not True: + return data + user_id: Final = user_api_key_dict.user_id + if not user_id: + return data + safety_identifier: Final = hashlib.sha256(user_id.encode("utf-8")).hexdigest() + return {**data, "safety_identifier": safety_identifier} + @staticmethod def _merge_passthrough_streaming_headers( response_headers: httpx.Headers | dict | None, @@ -2005,6 +2025,12 @@ class ProxyBaseLLMRequestProcessing: trust_client_model_info=False, ) + self.data = self._enforce_safety_identifier( + data=self.data, + route_type=route_type, + user_api_key_dict=user_api_key_dict, + ) + # An auto router with its own compression policy is authoritative for this # request: suppress every other compression guardrail and arm whichever one # the policy names for the model call, before those guardrails get a chance @@ -2017,6 +2043,12 @@ class ProxyBaseLLMRequestProcessing: call_type=route_type, ) + self.data = self._enforce_safety_identifier( + data=self.data, + route_type=route_type, + user_api_key_dict=user_api_key_dict, + ) + # Refresh AFTER pre_call_hook: guardrails (e.g. Presidio PII masking) may # have mutated `self.data` in place, and the audit-trail snapshot taken in # add_litellm_data_to_request predates that mutation. diff --git a/litellm/utils.py b/litellm/utils.py index bc2f4a86f12..941854e3075 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -4254,6 +4254,14 @@ def get_optional_params( allowed_openai_params = allowed_openai_params or [] supported_params.extend(allowed_openai_params) + # safety_identifier is injected by the proxy for trusted attribution. It is + # optional and provider-specific, so do not make providers that do not + # advertise it reject the entire request. Providers that support it still + # receive it through their normal parameter mapping, and callers can opt + # into an unlisted provider parameter via allowed_openai_params. + if "safety_identifier" in non_default_params and "safety_identifier" not in supported_params: + non_default_params.pop("safety_identifier") + _check_valid_arg( supported_params=supported_params or [], ) diff --git a/proxy_server_config.yaml b/proxy_server_config.yaml index 73990153227..28346279b24 100644 --- a/proxy_server_config.yaml +++ b/proxy_server_config.yaml @@ -246,7 +246,8 @@ general_settings: forward_headers: True # environment_variables: + # LITELLM_ENFORCE_SAFETY_IDENTIFIER: "true" # Hash the authenticated user_id and overwrite client safety_identifier values on chat/responses requests # settings for using redis caching # REDIS_HOST: redis-16337.c322.us-east-1-2.ec2.cloud.redislabs.com # REDIS_PORT: "16337" - # REDIS_PASSWORD: \ No newline at end of file + # REDIS_PASSWORD: diff --git a/tests/proxy_unit_tests/test_safety_identifier.py b/tests/proxy_unit_tests/test_safety_identifier.py new file mode 100644 index 00000000000..83ede9be79e --- /dev/null +++ b/tests/proxy_unit_tests/test_safety_identifier.py @@ -0,0 +1,131 @@ +import hashlib +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi import Request + +import litellm +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing + + +def test_enforce_safety_identifier_hashes_authenticated_user(monkeypatch): + monkeypatch.setenv("LITELLM_ENFORCE_SAFETY_IDENTIFIER", "true") + + result = ProxyBaseLLMRequestProcessing._enforce_safety_identifier( + data={"safety_identifier": "caller-value"}, + route_type="acompletion", + user_api_key_dict=UserAPIKeyAuth(user_id="user-123"), + ) + + assert result["safety_identifier"] == hashlib.sha256(b"user-123").hexdigest() + + +@pytest.mark.parametrize("setting", [None, "false"]) +def test_enforce_safety_identifier_is_opt_in(monkeypatch, setting): + if setting is None: + monkeypatch.delenv("LITELLM_ENFORCE_SAFETY_IDENTIFIER", raising=False) + else: + monkeypatch.setenv("LITELLM_ENFORCE_SAFETY_IDENTIFIER", setting) + data = {"safety_identifier": "caller-value"} + + result = ProxyBaseLLMRequestProcessing._enforce_safety_identifier( + data=data, + route_type="acompletion", + user_api_key_dict=UserAPIKeyAuth(user_id="user-123"), + ) + + assert result == data + + +def test_enforce_safety_identifier_skips_missing_user_id(monkeypatch): + monkeypatch.setenv("LITELLM_ENFORCE_SAFETY_IDENTIFIER", "true") + data = {"safety_identifier": "caller-value"} + + result = ProxyBaseLLMRequestProcessing._enforce_safety_identifier( + data=data, + route_type="aresponses", + user_api_key_dict=UserAPIKeyAuth(user_id=None), + ) + + assert result == data + + +def test_enforce_safety_identifier_only_applies_to_openai_generation_routes(monkeypatch): + monkeypatch.setenv("LITELLM_ENFORCE_SAFETY_IDENTIFIER", "true") + data = {"safety_identifier": "caller-value"} + + result = ProxyBaseLLMRequestProcessing._enforce_safety_identifier( + data=data, + route_type="aembedding", + user_api_key_dict=UserAPIKeyAuth(user_id="user-123"), + ) + + assert result == data + + +@pytest.mark.parametrize( + ("provider", "model"), + [("anthropic", "claude-3-5-sonnet-20241022"), ("gemini", "gemini-2.0-flash")], +) +def test_unsupported_safety_identifier_is_dropped_by_provider_translation(provider, model): + result = litellm.get_optional_params( + model=model, + custom_llm_provider=provider, + safety_identifier="trusted-value", + ) + + assert "safety_identifier" not in result + + +def test_supported_safety_identifier_is_preserved_by_provider_translation(): + result = litellm.get_optional_params( + model="gpt-4o", + custom_llm_provider="openai", + safety_identifier="trusted-value", + ) + + assert result["safety_identifier"] == "trusted-value" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("route_type", ["acompletion", "aresponses"]) +async def test_pre_call_hook_cannot_override_enforced_safety_identifier(monkeypatch, route_type): + monkeypatch.setenv("LITELLM_ENFORCE_SAFETY_IDENTIFIER", "true") + request = MagicMock(spec=Request) + request.headers.get.return_value = "call-id" + logging_obj = MagicMock() + proxy_logging_obj = MagicMock() + proxy_logging_obj.pre_call_hook = AsyncMock(return_value={"model": "gpt-5", "safety_identifier": "hook-value"}) + user_api_key_dict = UserAPIKeyAuth(user_id="user-123") + processor = ProxyBaseLLMRequestProcessing(data={"model": "gpt-5", "safety_identifier": "caller-value"}) + + with ( + patch( # test-quality-ok: isolate shared pre-call ordering without making an upstream request + "litellm.proxy.common_request_processing.add_litellm_data_to_request", + new=AsyncMock(side_effect=lambda **kwargs: kwargs["data"]), + ), + patch( # test-quality-ok: isolate shared pre-call ordering without initializing logging callbacks + "litellm.proxy.common_request_processing.litellm.utils.function_setup", + return_value=(logging_obj, processor.data), + ), + patch( # test-quality-ok: isolate shared pre-call ordering from router configuration + "litellm.proxy.common_request_processing._check_and_merge_model_level_guardrails", + side_effect=lambda **kwargs: kwargs["data"], + ), + patch( # test-quality-ok: isolate shared pre-call ordering from optional compression hooks + "litellm.proxy.common_request_processing._arm_auto_router_compression", + new=AsyncMock(), + ), + ): + result, _ = await processor.common_processing_pre_call_logic( + request=request, + general_settings={}, + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=proxy_logging_obj, + proxy_config=MagicMock(), + route_type=route_type, + version="test", + ) + + assert result["safety_identifier"] == hashlib.sha256(b"user-123").hexdigest() From 74174c87d8699b26dc2d72088d05077cc8dbaef4 Mon Sep 17 00:00:00 2001 From: Dominic White Date: Mon, 7 Sep 2026 09:28:36 +0200 Subject: [PATCH 2/7] test(proxy): register safety identifier coverage --- .github/workflows/test-unit-proxy-db.yml | 1 + litellm/proxy/common_request_processing.py | 7 +++++-- 2 files changed, 6 insertions(+), 2 deletions(-) diff --git a/.github/workflows/test-unit-proxy-db.yml b/.github/workflows/test-unit-proxy-db.yml index 3725e0f5805..527bc2e5c0b 100644 --- a/.github/workflows/test-unit-proxy-db.yml +++ b/.github/workflows/test-unit-proxy-db.yml @@ -195,6 +195,7 @@ jobs: tests/proxy_unit_tests/test_check_responses_cost.py tests/proxy_unit_tests/test_response_polling_handler.py tests/proxy_unit_tests/test_response_polling_pre_call_checks.py + tests/proxy_unit_tests/test_safety_identifier.py tests/proxy_unit_tests/test_realtime_cache.py tests/proxy_unit_tests/test_proxy_exception_mapping.py tests/proxy_unit_tests/test_custom_tokenizer_bug.py diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 074782c3ec7..720e8dc86b4 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1542,7 +1542,7 @@ class ProxyBaseLLMRequestProcessing: route_type: ProxyRouteType, user_api_key_dict: UserAPIKeyAuth, ) -> dict[str, object]: - if route_type not in {"acompletion", "aresponses"}: + if route_type not in ("acompletion", "aresponses"): return data if str_to_bool(os.getenv("LITELLM_ENFORCE_SAFETY_IDENTIFIER")) is not True: return data @@ -1550,7 +1550,10 @@ class ProxyBaseLLMRequestProcessing: if not user_id: return data safety_identifier: Final = hashlib.sha256(user_id.encode("utf-8")).hexdigest() - return {**data, "safety_identifier": safety_identifier} + return { # mutable-ok: downstream request processing mutates payloads + **data, + "safety_identifier": safety_identifier, + } @staticmethod def _merge_passthrough_streaming_headers( From 0d8c2e3bc4631ab9ea7c662149076bbcc7fce8a1 Mon Sep 17 00:00:00 2001 From: Dominic White Date: Mon, 7 Sep 2026 10:16:27 +0200 Subject: [PATCH 3/7] fix(proxy): preserve request payload typing --- litellm/proxy/common_request_processing.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 720e8dc86b4..f21e105ecdc 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1538,10 +1538,10 @@ class ProxyBaseLLMRequestProcessing: @staticmethod def _enforce_safety_identifier( *, - data: dict[str, object], + data: dict[str, Any], route_type: ProxyRouteType, user_api_key_dict: UserAPIKeyAuth, - ) -> dict[str, object]: + ) -> dict[str, Any]: if route_type not in ("acompletion", "aresponses"): return data if str_to_bool(os.getenv("LITELLM_ENFORCE_SAFETY_IDENTIFIER")) is not True: From 541bb66120080fbd2bbfe4be2a9c6d89973d69db Mon Sep 17 00:00:00 2001 From: Dominic White Date: Mon, 7 Sep 2026 10:49:35 +0200 Subject: [PATCH 4/7] fix(proxy): close safety identifier enforcement gaps --- .../litellm_core_utils/safety_identifier.py | 23 +++++++ litellm/proxy/common_request_processing.py | 29 ++++----- litellm/responses/streaming_iterator.py | 28 +++++++-- litellm/responses/utils.py | 6 ++ litellm/utils.py | 5 -- .../test_safety_identifier.py | 46 +++++++++----- .../test_responses_websocket_all_providers.py | 63 ++++++++++++++++++- 7 files changed, 159 insertions(+), 41 deletions(-) create mode 100644 litellm/litellm_core_utils/safety_identifier.py diff --git a/litellm/litellm_core_utils/safety_identifier.py b/litellm/litellm_core_utils/safety_identifier.py new file mode 100644 index 00000000000..b231e096fc5 --- /dev/null +++ b/litellm/litellm_core_utils/safety_identifier.py @@ -0,0 +1,23 @@ +import hashlib +from collections.abc import MutableMapping +from typing import Final + + +def enforce_safety_identifier( + *, + data: MutableMapping[str, object], # mutable-ok: trusted enforcement rewrites the request payload in place + user_id: str | None, + enabled: bool, +) -> bool: + if not enabled: + return False + if user_id: + safety_identifier: Final = hashlib.sha256(user_id.encode("utf-8")).hexdigest() + if data.get("safety_identifier") == safety_identifier: + return False + data["safety_identifier"] = safety_identifier # rebind-ok: enforce the trusted request identity in place + return True + if "safety_identifier" not in data: + return False + data.pop("safety_identifier", None) # rebind-ok: remove the untrusted client value when no identity exists + return True diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index f21e105ecdc..19d02c90d84 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1,11 +1,10 @@ import asyncio import contextlib -import hashlib import json import logging import math import os -from collections.abc import AsyncGenerator, Awaitable, Callable, Coroutine, Mapping, Sequence +from collections.abc import AsyncGenerator, Awaitable, Callable, Coroutine, Mapping, MutableMapping, Sequence from datetime import datetime from functools import lru_cache from types import MappingProxyType @@ -46,6 +45,7 @@ from litellm.litellm_core_utils.llm_response_utils.get_headers import ( get_response_headers, ) from litellm.litellm_core_utils.safe_json_dumps import safe_dumps +from litellm.litellm_core_utils.safety_identifier import enforce_safety_identifier from litellm.litellm_core_utils.streaming_handler import ( backfill_missing_cache_usage_fields, ) @@ -1538,22 +1538,19 @@ class ProxyBaseLLMRequestProcessing: @staticmethod def _enforce_safety_identifier( *, - data: dict[str, Any], + data: MutableMapping[str, object], route_type: ProxyRouteType, user_api_key_dict: UserAPIKeyAuth, - ) -> dict[str, Any]: + ) -> None: if route_type not in ("acompletion", "aresponses"): - return data + return if str_to_bool(os.getenv("LITELLM_ENFORCE_SAFETY_IDENTIFIER")) is not True: - return data - user_id: Final = user_api_key_dict.user_id - if not user_id: - return data - safety_identifier: Final = hashlib.sha256(user_id.encode("utf-8")).hexdigest() - return { # mutable-ok: downstream request processing mutates payloads - **data, - "safety_identifier": safety_identifier, - } + return + enforce_safety_identifier( + data=data, + user_id=user_api_key_dict.user_id, + enabled=True, + ) @staticmethod def _merge_passthrough_streaming_headers( @@ -2028,7 +2025,7 @@ class ProxyBaseLLMRequestProcessing: trust_client_model_info=False, ) - self.data = self._enforce_safety_identifier( + self._enforce_safety_identifier( data=self.data, route_type=route_type, user_api_key_dict=user_api_key_dict, @@ -2046,7 +2043,7 @@ class ProxyBaseLLMRequestProcessing: call_type=route_type, ) - self.data = self._enforce_safety_identifier( + self._enforce_safety_identifier( data=self.data, route_type=route_type, user_api_key_dict=user_api_key_dict, diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 9f9016c5a7f..49968451e28 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -2,14 +2,15 @@ from __future__ import annotations import asyncio import json +import os import time import traceback import uuid -from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence +from collections.abc import Awaitable, Callable, Iterable, Mapping, MutableMapping, Sequence from datetime import datetime from functools import lru_cache from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, overload, runtime_checkable +from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast, overload, runtime_checkable import httpx from openai._streaming import SSEDecoder @@ -29,9 +30,11 @@ from litellm.litellm_core_utils.llm_response_utils.get_api_base import get_api_b from litellm.litellm_core_utils.llm_response_utils.response_metadata import ( update_response_metadata, ) +from litellm.litellm_core_utils.safety_identifier import enforce_safety_identifier from litellm.litellm_core_utils.thread_pool_executor import executor from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig from litellm.responses.utils import ResponseAPILoggingUtils, ResponsesAPIRequestUtils +from litellm.secret_managers.main import str_to_bool from litellm.types.llms.openai import ( PART_UNION_TYPES, ResponseAPIUsage, @@ -92,6 +95,19 @@ def _is_str_mapping(value: object) -> TypeIs[dict[str, str]]: # guard-ok: verif return _is_json_object(value) and all(isinstance(item, str) for item in value.values()) +def _enforce_responses_ws_safety_identifier( + msg_obj: _MutableJsonObject, + user_api_key_dict: UserAPIKeyAuth | None, +) -> bool: + return enforce_safety_identifier( + data=cast( # cast-ok: JSON protocol is backed by a mutable response.create dictionary + MutableMapping[str, object], 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, + ) + + class _MutableJsonObject(Protocol): @overload def get(self, key: str, /) -> object | None: ... @@ -1744,16 +1760,18 @@ class ResponsesWebSocketStreaming: if msg_obj.get("type") != "response.create": return message + safety_identifier_modified: Final = _enforce_responses_ws_safety_identifier(msg_obj, self.user_api_key_dict) + # Always enforce the authorized model, even when PII masking is off. model_modified: Final = self._enforce_authorized_model(msg_obj) if not self.guardrail_callbacks: - return json.dumps(msg_obj) if model_modified else message + return json.dumps(msg_obj) if model_modified or safety_identifier_modified else message if "metadata" not in self.request_data: self.request_data["metadata"] = {} - modified = model_modified + modified = model_modified or safety_identifier_modified guardrail_cbs: Final[tuple[PresidioGuardrailCallback, ...]] = tuple(self.guardrail_callbacks) for cb in guardrail_cbs: presidio_config = cb.get_presidio_settings_from_request_data(self.request_data) @@ -2521,6 +2539,8 @@ class ManagedResponsesWebSocketHandler: if msg_obj is None: return + _enforce_responses_ws_safety_identifier(msg_obj, self.user_api_key_dict) + # generate=false is a prompt-cache warmup hint (sent by codex prewarm). # Native provider sockets handle it server-side, but there is no HTTP # equivalent and the frame carries empty input. Managed providers must diff --git a/litellm/responses/utils.py b/litellm/responses/utils.py index 3ca7b0503bf..b03194af7c3 100644 --- a/litellm/responses/utils.py +++ b/litellm/responses/utils.py @@ -210,6 +210,12 @@ class ResponsesAPIRequestUtils: should_drop_params: Final = litellm.drop_params or drop_params is True non_default_params: Final = cast(dict, response_api_optional_params) + if ( + "safety_identifier" in non_default_params + and "safety_identifier" not in supported_params + and (allowed_openai_params is None or "safety_identifier" not in allowed_openai_params) + ): + non_default_params.pop("safety_identifier") # Check for unsupported parameters ResponsesAPIRequestUtils._check_valid_arg( supported_params=supported_params + (allowed_openai_params or []), diff --git a/litellm/utils.py b/litellm/utils.py index 941854e3075..c3db7862057 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -4254,11 +4254,6 @@ def get_optional_params( allowed_openai_params = allowed_openai_params or [] supported_params.extend(allowed_openai_params) - # safety_identifier is injected by the proxy for trusted attribution. It is - # optional and provider-specific, so do not make providers that do not - # advertise it reject the entire request. Providers that support it still - # receive it through their normal parameter mapping, and callers can opt - # into an unlisted provider parameter via allowed_openai_params. if "safety_identifier" in non_default_params and "safety_identifier" not in supported_params: non_default_params.pop("safety_identifier") diff --git a/tests/proxy_unit_tests/test_safety_identifier.py b/tests/proxy_unit_tests/test_safety_identifier.py index 83ede9be79e..6072a88241b 100644 --- a/tests/proxy_unit_tests/test_safety_identifier.py +++ b/tests/proxy_unit_tests/test_safety_identifier.py @@ -1,74 +1,78 @@ import hashlib +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 from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing +from litellm.responses.utils import ResponsesAPIRequestUtils -def test_enforce_safety_identifier_hashes_authenticated_user(monkeypatch): +def test_enforce_safety_identifier_hashes_authenticated_user(monkeypatch: pytest.MonkeyPatch): monkeypatch.setenv("LITELLM_ENFORCE_SAFETY_IDENTIFIER", "true") - result = ProxyBaseLLMRequestProcessing._enforce_safety_identifier( - data={"safety_identifier": "caller-value"}, + data = {"safety_identifier": "caller-value"} + ProxyBaseLLMRequestProcessing._enforce_safety_identifier( + data=data, route_type="acompletion", user_api_key_dict=UserAPIKeyAuth(user_id="user-123"), ) - assert result["safety_identifier"] == hashlib.sha256(b"user-123").hexdigest() + assert data["safety_identifier"] == hashlib.sha256(b"user-123").hexdigest() @pytest.mark.parametrize("setting", [None, "false"]) -def test_enforce_safety_identifier_is_opt_in(monkeypatch, setting): +def test_enforce_safety_identifier_is_opt_in(monkeypatch: pytest.MonkeyPatch, setting: str | None): if setting is None: monkeypatch.delenv("LITELLM_ENFORCE_SAFETY_IDENTIFIER", raising=False) else: monkeypatch.setenv("LITELLM_ENFORCE_SAFETY_IDENTIFIER", setting) data = {"safety_identifier": "caller-value"} - result = ProxyBaseLLMRequestProcessing._enforce_safety_identifier( + ProxyBaseLLMRequestProcessing._enforce_safety_identifier( data=data, route_type="acompletion", user_api_key_dict=UserAPIKeyAuth(user_id="user-123"), ) - assert result == data + assert data == {"safety_identifier": "caller-value"} -def test_enforce_safety_identifier_skips_missing_user_id(monkeypatch): +def test_enforce_safety_identifier_removes_untrusted_identifier(monkeypatch: pytest.MonkeyPatch): monkeypatch.setenv("LITELLM_ENFORCE_SAFETY_IDENTIFIER", "true") data = {"safety_identifier": "caller-value"} - result = ProxyBaseLLMRequestProcessing._enforce_safety_identifier( + ProxyBaseLLMRequestProcessing._enforce_safety_identifier( data=data, route_type="aresponses", user_api_key_dict=UserAPIKeyAuth(user_id=None), ) - assert result == data + assert data == {} -def test_enforce_safety_identifier_only_applies_to_openai_generation_routes(monkeypatch): +def test_enforce_safety_identifier_only_applies_to_openai_generation_routes(monkeypatch: pytest.MonkeyPatch): monkeypatch.setenv("LITELLM_ENFORCE_SAFETY_IDENTIFIER", "true") data = {"safety_identifier": "caller-value"} - result = ProxyBaseLLMRequestProcessing._enforce_safety_identifier( + ProxyBaseLLMRequestProcessing._enforce_safety_identifier( data=data, route_type="aembedding", user_api_key_dict=UserAPIKeyAuth(user_id="user-123"), ) - assert result == data + assert data == {"safety_identifier": "caller-value"} @pytest.mark.parametrize( ("provider", "model"), [("anthropic", "claude-3-5-sonnet-20241022"), ("gemini", "gemini-2.0-flash")], ) -def test_unsupported_safety_identifier_is_dropped_by_provider_translation(provider, model): +def test_unsupported_safety_identifier_is_dropped_by_provider_translation(provider: str, model: str): result = litellm.get_optional_params( model=model, custom_llm_provider=provider, @@ -88,9 +92,21 @@ def test_supported_safety_identifier_is_preserved_by_provider_translation(): assert result["safety_identifier"] == "trusted-value" +def test_unsupported_safety_identifier_is_dropped_by_responses_translation(): + result = ResponsesAPIRequestUtils.get_optional_params_responses_api( + model="sonar", + responses_api_provider_config=PerplexityResponsesConfig(), + response_api_optional_params={"safety_identifier": "trusted-value"}, + ) + + assert "safety_identifier" not in result + + @pytest.mark.asyncio @pytest.mark.parametrize("route_type", ["acompletion", "aresponses"]) -async def test_pre_call_hook_cannot_override_enforced_safety_identifier(monkeypatch, route_type): +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.headers.get.return_value = "call-id" 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 e8333214ea8..fec56fe7ab1 100644 --- a/tests/test_litellm/responses/test_responses_websocket_all_providers.py +++ b/tests/test_litellm/responses/test_responses_websocket_all_providers.py @@ -7,8 +7,9 @@ Tests that: 3. Providers without native websocket support use ManagedResponsesWebSocketHandler """ +import hashlib import json -from unittest.mock import MagicMock +from unittest.mock import AsyncMock, MagicMock import pytest @@ -1196,6 +1197,66 @@ class TestWebSocketProjectQuotaEnforcement: class TestNativeWebSocketGuardrails: + @pytest.mark.asyncio + async def test_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", "safety_identifier": "caller-value"}) + ) + + assert json.loads(masked)["safety_identifier"] == hashlib.sha256(b"user-123").hexdigest() + + @pytest.mark.asyncio + async def test_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", "safety_identifier": "caller-value"}) + ) + + assert "safety_identifier" not in json.loads(masked) + + @pytest.mark.asyncio + async def test_managed_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", "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 From 94193b34c8a39d648565a57c5fe801c7d197a5d1 Mon Sep 17 00:00:00 2001 From: Dominic White Date: Mon, 7 Sep 2026 11:14:27 +0200 Subject: [PATCH 5/7] fix: remove unchecked safety identifier cast --- litellm/litellm_core_utils/safety_identifier.py | 15 ++++++++++++--- litellm/responses/streaming_iterator.py | 8 +++----- 2 files changed, 15 insertions(+), 8 deletions(-) diff --git a/litellm/litellm_core_utils/safety_identifier.py b/litellm/litellm_core_utils/safety_identifier.py index b231e096fc5..f7e7954a8fb 100644 --- a/litellm/litellm_core_utils/safety_identifier.py +++ b/litellm/litellm_core_utils/safety_identifier.py @@ -1,11 +1,20 @@ import hashlib -from collections.abc import MutableMapping -from typing import Final +from typing import Final, Protocol + + +class _SafetyIdentifierPayload(Protocol): + def get(self, key: str, /) -> object | None: ... + + def __setitem__(self, key: str, value: object, /) -> None: ... + + def __contains__(self, key: object, /) -> bool: ... + + def pop(self, key: str, default: object | None = None, /) -> object | None: ... def enforce_safety_identifier( *, - data: MutableMapping[str, object], # mutable-ok: trusted enforcement rewrites the request payload in place + data: _SafetyIdentifierPayload, user_id: str | None, enabled: bool, ) -> bool: diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 49968451e28..14f1af19d23 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -6,11 +6,11 @@ import os import time import traceback import uuid -from collections.abc import Awaitable, Callable, Iterable, Mapping, MutableMapping, Sequence +from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence from datetime import datetime from functools import lru_cache from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast, overload, runtime_checkable +from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, overload, runtime_checkable import httpx from openai._streaming import SSEDecoder @@ -100,9 +100,7 @@ def _enforce_responses_ws_safety_identifier( user_api_key_dict: UserAPIKeyAuth | None, ) -> bool: return enforce_safety_identifier( - data=cast( # cast-ok: JSON protocol is backed by a mutable response.create dictionary - MutableMapping[str, object], msg_obj - ), + 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, ) From 0d5a3bf2c3b108f96ecc12e3746d999542f9dec7 Mon Sep 17 00:00:00 2001 From: Dominic White Date: Mon, 7 Sep 2026 11:42:20 +0200 Subject: [PATCH 6/7] fix: enforce nested response safety identifiers --- litellm/responses/streaming_iterator.py | 12 ++-- .../test_safety_identifier.py | 3 +- .../test_responses_websocket_all_providers.py | 64 +++++++++++++++++++ 3 files changed, 72 insertions(+), 7 deletions(-) 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 From fcdff6ed11f5bfa27cc2a67c494ea5fc2288c0ce Mon Sep 17 00:00:00 2001 From: Dominic White Date: Mon, 7 Sep 2026 12:09:43 +0200 Subject: [PATCH 7/7] test: cover idempotent safety identifier enforcement --- tests/proxy_unit_tests/test_safety_identifier.py | 11 +++++++++++ 1 file changed, 11 insertions(+) diff --git a/tests/proxy_unit_tests/test_safety_identifier.py b/tests/proxy_unit_tests/test_safety_identifier.py index e4b79ef77ef..b57ffc09ad4 100644 --- a/tests/proxy_unit_tests/test_safety_identifier.py +++ b/tests/proxy_unit_tests/test_safety_identifier.py @@ -5,6 +5,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest import litellm +from litellm.litellm_core_utils.safety_identifier import enforce_safety_identifier from litellm.llms.perplexity.responses.transformation import PerplexityResponsesConfig from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing @@ -24,6 +25,16 @@ def test_enforce_safety_identifier_hashes_authenticated_user(monkeypatch: pytest assert data["safety_identifier"] == hashlib.sha256(b"user-123").hexdigest() +def test_enforce_safety_identifier_is_idempotent(): + safety_identifier = hashlib.sha256(b"user-123").hexdigest() + data = {"safety_identifier": safety_identifier} + + modified = enforce_safety_identifier(data=data, user_id="user-123", enabled=True) + + assert modified is False + assert data == {"safety_identifier": safety_identifier} + + @pytest.mark.parametrize("setting", [None, "false"]) def test_enforce_safety_identifier_is_opt_in(monkeypatch: pytest.MonkeyPatch, setting: str | None): if setting is None: