fix(responses): encrypt the response id on every streamed event

Background streaming emits event types with no typed model, which arrive
as GenericEvent holding a plain dict. Only typed events had their nested
response id rewritten, so those frames advertised the raw internal id
while their siblings advertised the encrypted one. The raw shape skips
the ownership check, so any other key could retrieve or cancel that
response.

Rewrite the advertised id wherever an event carries one, whichever shape
it arrives in, so a future event type cannot reopen this.
This commit is contained in:
mateo-berri 2026-09-03 02:41:04 -07:00
parent 658f50663d
commit d3b6ce98d6
4 changed files with 179 additions and 43 deletions

View file

@ -5,10 +5,11 @@ This hook uses the DBSpendUpdateWriter to batch-write response IDs to the databa
instead of writing immediately on each request.
"""
from collections.abc import AsyncGenerator
from collections.abc import AsyncGenerator, Callable, Mapping
from typing import TYPE_CHECKING, Any, Final, cast
from fastapi import HTTPException
from pydantic import TypeAdapter, ValidationError
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_logger import CustomLogger
@ -32,6 +33,44 @@ _RESPONSES_API_PROVIDER_PREFIX: Final = "/openai"
_RESPONSES_API_CREATE_ROUTES: Final = frozenset({"/v1/responses", "/responses"})
_RESPONSE_PAYLOAD_ADAPTER: Final = TypeAdapter(Mapping[str, object])
def _response_payload(response_obj: object) -> Mapping[str, object] | None:
try:
return _RESPONSE_PAYLOAD_ADAPTER.validate_python(response_obj)
except ValidationError:
return None
def _rewrite_advertised_id(
event: BaseLiteLLMOpenAIResponseObject,
rewrite: Callable[[str], str],
) -> BaseLiteLLMOpenAIResponseObject:
event_id: Final = getattr(event, "id", None)
if isinstance(event_id, str) and event_id.startswith("resp_"):
setattr(event, "id", rewrite(event_id))
return event
nested: Final = getattr(event, "response", None)
if isinstance(nested, ResponsesAPIResponse):
setattr(nested, "id", rewrite(nested.id))
setattr(event, "response", nested)
return event
payload: Final = _response_payload(nested)
if payload is None:
return event
payload_id: Final = payload.get("id")
if not isinstance(payload_id, str):
return event
rewritten: Final = {**payload, "id": rewrite(payload_id)} # mutable-ok: pydantic cannot serialize a frozen map
setattr(event, "response", rewritten)
return event
def _is_responses_api_create_route(request_route: str | None) -> bool:
if request_route is None:
return False
@ -196,10 +235,6 @@ class ResponsesIDSecurity(CustomLogger):
user_api_key_dict: "UserAPIKeyAuth",
request_cache: dict[str, str] | None = None,
) -> BaseLiteLLMOpenAIResponseObject:
# encrypt the response id using the symmetric key
# encrypt the response id, and encode the user id and response id in base64
# Check if signing key is available
signing_key: Final = self._get_signing_key()
if signing_key is None:
verbose_proxy_logger.debug(
@ -210,43 +245,22 @@ class ResponsesIDSecurity(CustomLogger):
)
return response
response_id: Final = getattr(response, "id", None)
response_obj: Final = getattr(response, "response", None)
def encrypt(original_id: str) -> str:
cached: Final = request_cache.get(original_id) if request_cache is not None else None
if cached is not None:
return cached
if response_id and isinstance(response_id, str) and response_id.startswith("resp_"):
# Check request-scoped cache first (for streaming consistency)
if request_cache is not None and response_id in request_cache:
setattr(response, "id", request_cache[response_id])
else:
encrypted_response_id = SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format(
response_id,
user_api_key_dict.user_id or "",
user_api_key_dict.team_id or "",
)
managed_id: Final = SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format(
original_id,
user_api_key_dict.user_id or "",
user_api_key_dict.team_id or "",
)
encrypted_id: Final = f"resp_{encrypt_value_helper(value=managed_id)}"
if request_cache is not None:
request_cache[original_id] = encrypted_id
return encrypted_id
encoded_user_id_and_response_id = encrypt_value_helper(value=encrypted_response_id)
encrypted_id = f"resp_{encoded_user_id_and_response_id}"
if request_cache is not None:
request_cache[response_id] = encrypted_id
setattr(response, "id", encrypted_id)
elif response_obj and isinstance(response_obj, ResponsesAPIResponse):
# Check request-scoped cache first (for streaming consistency)
if request_cache is not None and response_obj.id in request_cache:
setattr(response_obj, "id", request_cache[response_obj.id])
else:
encrypted_response_id = SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format(
response_obj.id,
user_api_key_dict.user_id or "",
user_api_key_dict.team_id or "",
)
encoded_user_id_and_response_id = encrypt_value_helper(value=encrypted_response_id)
encrypted_id = f"resp_{encoded_user_id_and_response_id}"
if request_cache is not None:
request_cache[response_obj.id] = encrypted_id
setattr(response_obj, "id", encrypted_id)
setattr(response, "response", response_obj)
return response
return _rewrite_advertised_id(response, encrypt)
async def async_post_call_success_hook(
self,

View file

@ -42,7 +42,7 @@
"limit": 52
},
"B010": {
"limit": 190
"limit": 189
},
"B018": {
"limit": 2

View file

@ -14,7 +14,9 @@ from litellm.proxy.hooks.responses_id_security import (
_is_responses_api_create_route,
)
from litellm.types.llms.openai import (
GenericEvent,
ResponseCompletedEvent,
ResponseCreatedEvent,
ResponsesAPIResponse,
ResponsesAPIStreamEvents,
)
@ -691,6 +693,126 @@ class TestAsyncPostCallStreamingIteratorHook:
assert not responses_id_security._is_encrypted_response_id(streamed_id)
class TestStreamedGenericEventIdEncryption:
"""A background stream carries event types with no typed model, which arrive as
GenericEvent holding a plain dict. Those used to skip encryption while their typed
siblings were encrypted, so one stream advertised two ids and the unencrypted one
skipped the ownership check. Asserts the property rather than one event type: every
id a client can see is the same encrypted id, and the raw one appears in no frame."""
RAW_ID = "resp_rawprovider123"
@staticmethod
async def _agen(chunks):
for chunk in chunks:
yield chunk
@classmethod
def _typed_event(cls, event_type):
return {
ResponsesAPIStreamEvents.RESPONSE_CREATED: ResponseCreatedEvent,
ResponsesAPIStreamEvents.RESPONSE_COMPLETED: ResponseCompletedEvent,
}[event_type](
type=event_type,
response=ResponsesAPIResponse(
id=cls.RAW_ID,
created_at=0,
model="gpt-5.1",
object="response",
output=[],
parallel_tool_calls=False,
tool_choice="auto",
tools=[],
),
)
@classmethod
def _background_stream(cls):
return [
cls._typed_event(ResponsesAPIStreamEvents.RESPONSE_CREATED),
GenericEvent(
type="response.queued",
response={"id": cls.RAW_ID, "status": "queued"},
),
GenericEvent(type="keepalive"),
GenericEvent(
type="response.some_event_openai_adds_later",
response={"id": cls.RAW_ID, "status": "in_progress"},
),
cls._typed_event(ResponsesAPIStreamEvents.RESPONSE_COMPLETED),
]
@staticmethod
def _advertised_ids(events):
nested = (getattr(event, "response", None) for event in events)
return [
payload["id"] if isinstance(payload, dict) else payload.id
for payload in nested
if payload is not None
] + [
event.id for event in events if isinstance(getattr(event, "id", None), str)
]
async def _drain(self, responses_id_security, monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-key-abcdefghij")
mock_auth = MagicMock()
mock_auth.user_id = "user-a"
mock_auth.team_id = "team-a"
mock_auth.request_route = "/v1/responses"
return [
out
async for out in responses_id_security.async_post_call_streaming_iterator_hook(
user_api_key_dict=mock_auth,
response=self._agen(self._background_stream()),
request_data={},
)
]
@pytest.mark.asyncio
async def test_every_event_advertises_the_same_encrypted_id(
self, responses_id_security, monkeypatch
):
events = await self._drain(responses_id_security, monkeypatch)
advertised = self._advertised_ids(events)
assert len(advertised) == 4
assert len(set(advertised)) == 1
streamed_id = advertised[0]
assert streamed_id != self.RAW_ID
assert responses_id_security._is_encrypted_response_id(streamed_id)
assert responses_id_security._decrypt_response_id(streamed_id) == (
self.RAW_ID,
"user-a",
"team-a",
)
@pytest.mark.asyncio
async def test_raw_provider_id_never_reaches_the_client(
self, responses_id_security, monkeypatch
):
events = await self._drain(responses_id_security, monkeypatch)
assert [self.RAW_ID in event.model_dump_json() for event in events] == [
False
] * len(events)
@pytest.mark.asyncio
async def test_sibling_fields_survive_the_rewrite(
self, responses_id_security, monkeypatch
):
_, queued, keepalive, later, _ = await self._drain(
responses_id_security, monkeypatch
)
assert queued.response["status"] == "queued"
assert later.response["status"] == "in_progress"
assert keepalive.type == "keepalive"
assert getattr(keepalive, "response", None) is None
class TestAsyncPostCallSuccessHook:
"""Test async_post_call_success_hook function"""

View file

@ -27,10 +27,10 @@
"limit": 0
},
"LIT010": {
"limit": 16477
"limit": 16471
},
"LIT011": {
"limit": 5519
"limit": 5517
},
"LIT012": {
"limit": 4489