From 601c75a475feb5120f6b046bee37ef2e2acd3de0 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 25 Sep 2026 13:45:43 -0700 Subject: [PATCH 1/5] fix(proxy): record streamed /v1/responses container ownership before the response.completed frame (#43140) * test(e2e): cover Azure code_interpreter container files by native id with a service-account key * test(e2e): require the code_interpreter tool, skip at collection, and scope the container call timeout * fix(e2e): fail the containers suite when the Azure credentials are missing instead of skipping * fix(proxy): record streamed /v1/responses container ownership before the response.completed frame --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/proxy/common_request_processing.py | 87 +++++++++---------- .../coverage_registry/llm_conversational.yaml | 1 + .../LLM_TRANSLATION_COVERAGE_MATRIX.md | 3 +- .../llm_translation/test_containers_e2e.py | 43 +++++++-- .../proxy/test_common_request_processing.py | 64 ++++++++++++++ 5 files changed, 144 insertions(+), 54 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 7e1989b0aab..4f9b6b3a96f 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -2781,15 +2781,6 @@ class ProxyBaseLLMRequestProcessing: request=request, ) if route_type == "aresponses": - # Streaming /v1/responses returns here without - # reaching the non-streaming ownership tail below. - # Wrap the SSE generator so container ownership is - # written once the upstream iterator finishes - # assembling ``completed_response`` — otherwise - # code-interpreter containers created during the - # stream stay unregistered and follow-up file API - # calls 403. Covers the background-polling path - # too, which loops ``body_iterator`` end-to-end. selected_data_generator = ( ProxyBaseLLMRequestProcessing._wrap_responses_stream_for_container_ownership( original_stream_response=response, @@ -3011,50 +3002,50 @@ class ProxyBaseLLMRequestProcessing: wrapped_generator: Any, user_api_key_dict: UserAPIKeyAuth, ): - """Forward SSE chunks, then record container ownership at stream end. + """Forward SSE chunks and record container ownership before the terminal chunk goes out. Streaming ``/v1/responses`` short-circuits out of ``base_process_llm_request`` before the non-streaming ownership - tail runs, so without this wrap the - ``LiteLLM_ManagedObjectTable`` row for any container created - during the stream is never written and follow-up file API calls - return 403. + tail runs. The OpenAI SDK closes the connection at ``data: [DONE]`` + and starlette cancels the body task on disconnect, so a write that + waits for the generator to finish never lands. The iterator sets + ``completed_response`` before it hands over its terminal chunk, so + the ``LiteLLM_ManagedObjectTable`` row is written the moment it + appears, ahead of the chunk carrying ``response.completed``. """ - try: - async for chunk in wrapped_generator: + async for chunk in wrapped_generator: + completed_obj = ProxyBaseLLMRequestProcessing._extract_completed_responses_response( + original_stream_response + ) + if completed_obj is None: yield chunk - finally: - try: - completed_obj: Final = ProxyBaseLLMRequestProcessing._extract_completed_responses_response( - original_stream_response - ) - if completed_obj is not None: - await ProxyBaseLLMRequestProcessing._record_container_owners_from_responses_if_needed( - response=completed_obj, - user_api_key_dict=user_api_key_dict, - ) - else: - # Silent skip caused #30210: the proxy's Router wrapper - # of the responses streaming iterator wasn't propagating - # ``completed_response``, so this hook recorded nothing - # and follow-up /v1/containers//files calls 403'd - # for non-admin keys with no proxy-side hint. Log a - # warning so future regressions of the same shape - # surface in operator logs. - verbose_proxy_logger.warning( - "Container ownership recording skipped on streaming " - "/v1/responses: no completed_response on stream " - "iterator %s. If this stream created any tool " - "container (e.g. code_interpreter), follow-up " - "/v1/containers//files calls will 403 for " - "non-admin keys.", - type(original_stream_response).__name__, - ) - except Exception as e: - verbose_proxy_logger.exception( - "Container ownership recording failed after streaming responses call: %s", - e, - ) + continue + await ProxyBaseLLMRequestProcessing._record_container_owners_from_responses_if_needed( + response=completed_obj, + user_api_key_dict=user_api_key_dict, + ) + yield chunk + async for remaining_chunk in wrapped_generator: + yield remaining_chunk + return + late_completed_obj: Final = ProxyBaseLLMRequestProcessing._extract_completed_responses_response( + original_stream_response + ) + if late_completed_obj is not None: + await ProxyBaseLLMRequestProcessing._record_container_owners_from_responses_if_needed( + response=late_completed_obj, + user_api_key_dict=user_api_key_dict, + ) + return + verbose_proxy_logger.warning( + "Container ownership recording skipped on streaming " + "/v1/responses: no completed_response on stream " + "iterator %s. If this stream created any tool " + "container (e.g. code_interpreter), follow-up " + "/v1/containers//files calls will 403 for " + "non-admin keys.", + type(original_stream_response).__name__, + ) async def base_passthrough_process_llm_request( self, diff --git a/tests/e2e/coverage_registry/llm_conversational.yaml b/tests/e2e/coverage_registry/llm_conversational.yaml index bc19668e2c9..20c87dbbd74 100644 --- a/tests/e2e/coverage_registry/llm_conversational.yaml +++ b/tests/e2e/coverage_registry/llm_conversational.yaml @@ -85,6 +85,7 @@ - {id: llm.responses.azure_openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "response_api_endpoints/endpoints.py:26", rationale: "Responses w/ Azure OpenAI (smoke)"} - {id: llm.responses.azure_openai.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: tool_use, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Responses tool calls w/ Azure OpenAI"} - {id: llm.responses.azure_openai.code_interpreter.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: code_interpreter, streaming: nonstream, assertions: [works], source: "llm_translation/test_containers_e2e.py", rationale: "An implicit code_interpreter container on an Azure deployment that carries its own api_base must serve GET /v1/containers/{id}/files/{fid}/content by its native cntr_ id to a team service-account key, the customer's shape (#27921, #28990)", fail_before_fix: proven} +- {id: llm.responses.azure_openai.code_interpreter.stream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: code_interpreter, streaming: stream, assertions: [works], source: "llm_translation/test_containers_e2e.py", rationale: "A container created by a streamed /v1/responses code_interpreter call must serve /v1/containers/{id}/files to the same service-account key right after the OpenAI SDK closes at [DONE]; the ownership row used to be written after the stream and the disconnect cancelled it (LIT-8612)", fail_before_fix: proven} - {id: llm.chat_completions.together_ai.thinking.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: thinking, streaming: nonstream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together reasoning surfaces as reasoning_content (LIT-5960)"} - {id: llm.chat_completions.together_ai.thinking.stream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: thinking, streaming: stream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together reasoning deltas stream as reasoning_content"} - {id: llm.chat_completions.together_ai.thinking.nonstream.template_kwargs_forwarded, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: thinking, streaming: nonstream, assertions: [template_kwargs_forwarded], source: "llm_translation/test_together_ai_e2e.py", rationale: "chat_template_kwargs reaches Together and turns thinking off"} diff --git a/tests/e2e/llm_translation/LLM_TRANSLATION_COVERAGE_MATRIX.md b/tests/e2e/llm_translation/LLM_TRANSLATION_COVERAGE_MATRIX.md index a18c81fa01d..92330db0530 100644 --- a/tests/e2e/llm_translation/LLM_TRANSLATION_COVERAGE_MATRIX.md +++ b/tests/e2e/llm_translation/LLM_TRANSLATION_COVERAGE_MATRIX.md @@ -48,7 +48,7 @@ most likely to silently break and the one a mock can't prove works. |----------|---------------|-----------|------------|-------------|--------| | Chat | live (spend suite) | live (spend suite) | gap | live | partial | | Embeddings | live (spend suite) | n/a | n/a | live | covered | -| Responses (Azure code_interpreter container files) | live | gap | live | gap | partial | +| Responses (Azure code_interpreter container files) | live | live | live | gap | partial | | Image / audio / rerank / realtime | - | - | - | - | gap | ## This suite's files @@ -63,6 +63,7 @@ most likely to silently break and the one a mock can't prove works. | `test_anthropic_passthrough_tool_call_logs_cost` | anthropic native, tool call, cost | | `test_vertex_passthrough_via_managed_model_logs_cost` | vertex_ai native, non-stream, cost | | `test_service_account_key_reads_container_file_by_native_id` | azure responses code_interpreter, non-stream, native container id, service-account key | +| `test_service_account_key_reads_container_file_created_by_a_streamed_response` | azure responses code_interpreter, stream, native container id, service-account key, upload right after `[DONE]` | Vertex keeps the credential on the proxy like gemini/anthropic, but the deployment is added at runtime instead of declared in the gateway config: the test POSTs `/model/new` diff --git a/tests/e2e/llm_translation/test_containers_e2e.py b/tests/e2e/llm_translation/test_containers_e2e.py index 887aecb8df1..3048a830810 100644 --- a/tests/e2e/llm_translation/test_containers_e2e.py +++ b/tests/e2e/llm_translation/test_containers_e2e.py @@ -31,10 +31,11 @@ A proxy whose env carries ``AZURE_API_BASE`` for the same resource masks the second regression, since the global-credential fallback then reaches the container anyway. -The streaming variant is not here: a streamed ``/v1/responses`` writes the -container ownership row only after the ``[DONE]`` frame, and the OpenAI SDK -closes the connection at ``[DONE]``, so the write is cancelled and every -follow-up container call 403s (LIT-8612). That cell comes with its fix. +The streaming cell repeats the flow with ``stream=True`` and uploads right +after the last event. The OpenAI SDK closes the connection at ``[DONE]``, so an +ownership row written after the stream is cancelled with the body task and every +follow-up container call 403s (LIT-8612); the row has to land before the +``response.completed`` frame goes out. """ from __future__ import annotations @@ -52,7 +53,7 @@ from lifecycle import ResourceManager from management.management_client import ManagementClient, build_client from models import KeyGenerateBody, KeyGenerateResponse, LiteLLMParamsBody, TeamNewBody, UserNewBody from openai import OpenAI -from openai.types.responses import Response, ResponseCodeInterpreterToolCall +from openai.types.responses import Response, ResponseCodeInterpreterToolCall, ResponseCompletedEvent from openai.types.responses.tool_param import CodeInterpreter from proxy_client import ProxyClient from sdk_clients import NO_PROXY_CACHE, SdkClients @@ -120,6 +121,24 @@ def _response_with_code_interpreter(client: OpenAI, model: str) -> Response: ) +def _streamed_response_with_code_interpreter(client: OpenAI, model: str) -> Response: + events: Final = tuple( + client.with_options(timeout=CODE_INTERPRETER_TIMEOUT).responses.create( + model=model, + input=PROMPT, + tools=[CODE_INTERPRETER], + tool_choice="required", + stream=True, + extra_body=NO_PROXY_CACHE, + ) + ) + assert events, "responses stream returned no events" + assert isinstance(events[-1], ResponseCompletedEvent), ( + f"responses stream did not terminate with response.completed: {events[-1].type}" + ) + return events[-1].response + + def _container_id(response: Response) -> str: calls: Final = tuple(item for item in response.output if isinstance(item, ResponseCodeInterpreterToolCall)) assert calls, f"no code_interpreter_call in the responses output: {response.output!r}" @@ -165,3 +184,17 @@ class TestAzureContainerFiles: f"container id is not the provider's own id: {native_id}" ) _assert_file_round_trip(client, native_id, marker) + + @pytest.mark.covers("llm.responses.azure_openai.code_interpreter.stream.works") + def test_service_account_key_reads_container_file_created_by_a_streamed_response( + self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients + ) -> None: + marker: Final = unique_marker() + model: Final = _register_two_azure_deployments(proxy, resources, marker) + key: Final = _service_account_key(proxy, resources, build_client(proxy), marker, model) + client: Final = sdk.openai(key) + native_id: Final = _native_container_id( + _container_id(_streamed_response_with_code_interpreter(client, model)) + ) + resources.defer(lambda: client.containers.delete(native_id, extra_query=AZURE_PROVIDER_QUERY)) + _assert_file_round_trip(client, native_id, marker) diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 7b74e69685c..bdf003085ef 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -9899,3 +9899,67 @@ class TestErrorLogCarriesCallId: record: Final = caplog.records[-1] assert record.litellm_call_id == call_id assert call_id in record.getMessage() + + +class TestStreamingContainerOwnershipRecordedBeforeDone: + """Regression for LIT-8612: the OpenAI SDK closes the connection at + ``data: [DONE]`` and starlette cancels the body task, so an ownership row + written after the SSE generator is exhausted never lands. The row must be + written before the chunk carrying ``response.completed`` is handed to the + client.""" + + CHUNKS: Final = ( + 'data: {"type":"response.created"}\n\n', + 'data: {"type":"response.output_text.delta"}\n\n', + 'data: {"type":"response.completed"}\n\n', + "data: [DONE]\n\n", + ) + TERMINAL_INDEX: Final = 2 + + @staticmethod + def _completed_event() -> SimpleNamespace: + return SimpleNamespace( + type="response.completed", + response=SimpleNamespace( + id="resp_lit8612", + output=[SimpleNamespace(type="code_interpreter_call", container_id="cntr_lit8612")], + ), + ) + + async def _sse(self, stream: SimpleNamespace, populate_at: int) -> AsyncGenerator[str, None]: + for index, chunk in enumerate(self.CHUNKS): + if index == populate_at: + stream.completed_response = self._completed_event() + yield chunk + if populate_at == len(self.CHUNKS): + stream.completed_response = self._completed_event() + + async def _await_counts_per_chunk(self, populate_at: int) -> tuple[tuple[tuple[str, int], ...], AsyncMock]: + stream: Final = SimpleNamespace(completed_response=None, _hidden_params={"custom_llm_provider": "azure"}) + recorder: Final = AsyncMock(return_value=None) + with patch( + "litellm.proxy.container_endpoints.ownership.record_container_owners_from_responses_response", recorder + ): + wrapped: Final = ProxyBaseLLMRequestProcessing._wrap_responses_stream_for_container_ownership( + original_stream_response=stream, + wrapped_generator=self._sse(stream, populate_at), + user_api_key_dict=ProxyUserAPIKeyAuth(api_key="sk-test", team_id="team-1"), + ) + observed: Final = tuple([(chunk, recorder.await_count) async for chunk in wrapped]) + return observed, recorder + + async def test_row_is_written_before_the_terminal_chunk_reaches_the_client(self) -> None: + observed, recorder = await self._await_counts_per_chunk(populate_at=self.TERMINAL_INDEX) + + assert tuple(chunk for chunk, _ in observed) == self.CHUNKS + assert tuple(count for _, count in observed) == (0, 0, 1, 1) + recorder.assert_awaited_once() + assert recorder.await_args.kwargs["response"].output[0].container_id == "cntr_lit8612" + assert recorder.await_args.kwargs["user_api_key_dict"].team_id == "team-1" + + async def test_row_is_still_written_when_the_iterator_completes_only_at_exhaustion(self) -> None: + observed, recorder = await self._await_counts_per_chunk(populate_at=len(self.CHUNKS)) + + assert tuple(chunk for chunk, _ in observed) == self.CHUNKS + assert tuple(count for _, count in observed) == (0, 0, 0, 0) + recorder.assert_awaited_once() From 25de1ab2b02ce724439a1da538dad32ebf2a1bc5 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 25 Sep 2026 14:01:27 -0700 Subject: [PATCH 2/5] fix(tests): drop the repeated UNIT_FLAG key in test_unit_shard_missing_paths (#43212) Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- tests/unit/test_unit_shard_missing_paths.py | 1 - 1 file changed, 1 deletion(-) diff --git a/tests/unit/test_unit_shard_missing_paths.py b/tests/unit/test_unit_shard_missing_paths.py index b528a75d0df..0360a227142 100644 --- a/tests/unit/test_unit_shard_missing_paths.py +++ b/tests/unit/test_unit_shard_missing_paths.py @@ -40,7 +40,6 @@ def _run_shard(tmp_path: Path, test_path: str, workers: str) -> subprocess.Compl "TEST_PATH": test_path, "UNIT_FLAG": "", "WORKERS": workers, - "UNIT_FLAG": "", }, capture_output=True, text=True, From 8327cd6d4724b6218f85562ba164f0cb2cb71ef6 Mon Sep 17 00:00:00 2001 From: daqiangganjun <93830914+daqiangganjun@users.noreply.github.com> Date: Sat, 26 Sep 2026 05:06:49 +0800 Subject: [PATCH 3/5] fix(router): count provider budget spend on every API surface (#38172) * fix(router): count provider budget spend on every API surface RouterBudgetLimiting read custom_llm_provider from litellm_params, which only chat completions populates. Responses, anthropic_messages, embedding and rerank calls raised inside the success callback before any spend was recorded, so those budgets never moved and a ceiling made up mostly of that traffic was never hit. Read the provider from the standard logging payload, which every surface fills in. Dropping the raise also stops one missing field from taking the deployment and tag budgets down with it. * chore(router): drop the inline comment and type the budget limiter test helper --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/router_strategy/budget_limiter.py | 10 +- .../router_strategy/test_budget_limiter.py | 137 ++++++++++++++++++ 2 files changed, 142 insertions(+), 5 deletions(-) create mode 100644 tests/test_litellm/router_strategy/test_budget_limiter.py diff --git a/litellm/router_strategy/budget_limiter.py b/litellm/router_strategy/budget_limiter.py index 4e84bded9de..64252cbbfb3 100644 --- a/litellm/router_strategy/budget_limiter.py +++ b/litellm/router_strategy/budget_limiter.py @@ -502,12 +502,12 @@ class RouterBudgetLimiting(CustomLogger): response_cost: Final[float] = standard_logging_payload.get("response_cost", 0) model_id: Final[str] = str(standard_logging_payload.get("model_id", "")) - custom_llm_provider: Final[str] = kwargs.get("litellm_params", {}).get("custom_llm_provider", None) - if custom_llm_provider is None: - raise ValueError("custom_llm_provider is required") + custom_llm_provider: Final[str | None] = standard_logging_payload.get("custom_llm_provider") - budget_config: Final = self._get_budget_config_for_provider(custom_llm_provider) - if budget_config: + budget_config: Final = ( + self._get_budget_config_for_provider(custom_llm_provider) if custom_llm_provider is not None else None + ) + if custom_llm_provider is not None and budget_config is not None: # increment spend for provider spend_key: Final = f"provider_spend:{custom_llm_provider}:{budget_config.budget_duration}" start_time_key: Final = f"provider_budget_start_time:{custom_llm_provider}" diff --git a/tests/test_litellm/router_strategy/test_budget_limiter.py b/tests/test_litellm/router_strategy/test_budget_limiter.py new file mode 100644 index 00000000000..62de1586fdd --- /dev/null +++ b/tests/test_litellm/router_strategy/test_budget_limiter.py @@ -0,0 +1,137 @@ +""" +Spend tracking in RouterBudgetLimiting.async_log_success_event. + +Only chat completions puts custom_llm_provider into litellm_params. The responses, +anthropic_messages, embedding and rerank surfaces leave it unset, which used to make +the callback raise before any spend was recorded, so those budgets never moved. +""" + +from typing import Final + +import pytest + +from litellm.caching.caching import DualCache +from litellm.router_strategy.budget_limiter import RouterBudgetLimiting + + +@pytest.fixture +def disable_budget_sync(monkeypatch): + async def noop(*args, **kwargs): + return None + + monkeypatch.setattr( + "litellm.router_strategy.budget_limiter.RouterBudgetLimiting.periodic_sync_in_memory_spend_with_redis", + noop, + ) + + +def _success_kwargs( + *, + provider_in_litellm_params: str | None, + provider_in_payload: str | None, + call_type: str = "aresponses", + response_cost: float = 0.25, + model_id: str = "deployment-1", +) -> dict[str, object]: + provider_params: Final[dict[str, str]] = ( + {} if provider_in_litellm_params is None else {"custom_llm_provider": provider_in_litellm_params} + ) + litellm_params: Final[dict[str, str]] = {"model": "openai/gpt-4o", **provider_params} + + return { + "call_type": call_type, + "litellm_params": litellm_params, + "standard_logging_object": { + "response_cost": response_cost, + "model_id": model_id, + "custom_llm_provider": provider_in_payload, + }, + } + + +async def _log_success(limiter: RouterBudgetLimiting, kwargs: dict[str, object]) -> None: + await limiter.async_log_success_event(kwargs=kwargs, response_obj=None, start_time=None, end_time=None) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("call_type", ["aresponses", "anthropic_messages", "aembedding", "arerank"]) +async def test_provider_spend_tracked_when_litellm_params_omits_provider(disable_budget_sync, call_type): + """Non-chat surfaces carry the provider only on the standard logging payload.""" + limiter = RouterBudgetLimiting( + dual_cache=DualCache(), + provider_budget_config={"openai": {"budget_limit": 10.0, "time_period": "1d"}}, + ) + + await _log_success( + limiter, + _success_kwargs( + provider_in_litellm_params=None, + provider_in_payload="openai", + call_type=call_type, + ), + ) + + assert await limiter.dual_cache.async_get_cache("provider_spend:openai:1d") == 0.25 + + +@pytest.mark.asyncio +async def test_chat_completions_spend_still_tracked(disable_budget_sync): + """Chat completions fills in both sources and must keep accumulating.""" + limiter = RouterBudgetLimiting( + dual_cache=DualCache(), + provider_budget_config={"openai": {"budget_limit": 10.0, "time_period": "1d"}}, + ) + + await _log_success( + limiter, + _success_kwargs( + provider_in_litellm_params="openai", + provider_in_payload="openai", + call_type="acompletion", + ), + ) + + assert await limiter.dual_cache.async_get_cache("provider_spend:openai:1d") == 0.25 + + +@pytest.mark.asyncio +async def test_budget_of_other_provider_is_untouched(disable_budget_sync): + """A provider without its own budget must not bleed into a configured one.""" + limiter = RouterBudgetLimiting( + dual_cache=DualCache(), + provider_budget_config={"openai": {"budget_limit": 10.0, "time_period": "1d"}}, + ) + + await _log_success( + limiter, + _success_kwargs(provider_in_litellm_params=None, provider_in_payload="anthropic"), + ) + + assert await limiter.dual_cache.async_get_cache("provider_spend:openai:1d") in (None, 0.0) + + +@pytest.mark.asyncio +async def test_deployment_budget_tracked_when_provider_is_unresolvable(disable_budget_sync): + """An unresolvable provider must not abort the deployment and tag budgets that follow it.""" + limiter = RouterBudgetLimiting( + dual_cache=DualCache(), + provider_budget_config=None, + model_list=[ + { + "model_name": "some-model", + "litellm_params": { + "model": "openai/gpt-4o", + "max_budget": 10.0, + "budget_duration": "1d", + }, + "model_info": {"id": "deployment-1"}, + } + ], + ) + + await _log_success( + limiter, + _success_kwargs(provider_in_litellm_params=None, provider_in_payload=None), + ) + + assert await limiter.dual_cache.async_get_cache("deployment_spend:deployment-1:1d") == 0.25 From 1fd04abb9278ef3b33379009874f75fe32e27315 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 25 Sep 2026 14:32:41 -0700 Subject: [PATCH 4/5] fix(responses): fall back on pre-output stream drops, fail truncated streams, honor request_timeout (#43133) * fix(responses): fall back on pre-output stream drops, fail truncated streams, honor request_timeout A native /v1/responses stream that drops before any output item now raises the router's fallback-eligible MidStreamFallbackError, so configured fallbacks retry the original input. A stream that ends with a clean EOF or a [DONE] marker but no response.completed, response.incomplete or response.failed event now raises litellm.APIConnectionError instead of ending as if it had completed: fallback-eligible before any output, an explicit error after partial output. The sync iterator mirrors every branch. resolve_llm_passthrough_timeout now consults an explicitly set litellm_settings.request_timeout right after the router timeout and before general_settings.pass_through_request_timeout, so the router's native responses path honors it. * test(responses): give the normal-completion stream tests a terminal event --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/passthrough/timeout_utils.py | 7 +- litellm/responses/streaming_iterator.py | 68 +++++-- ...t_base_responses_api_streaming_iterator.py | 54 ++++-- .../test_pass_through_endpoints.py | 24 +++ .../responses/test_streaming_iterator.py | 168 +++++++++++++++++- tests/unit/test_router/test_router.py | 127 +++++++++++++ 6 files changed, 420 insertions(+), 28 deletions(-) diff --git a/litellm/passthrough/timeout_utils.py b/litellm/passthrough/timeout_utils.py index fc67aa8c553..f600c7817f2 100644 --- a/litellm/passthrough/timeout_utils.py +++ b/litellm/passthrough/timeout_utils.py @@ -5,6 +5,8 @@ from typing import Final from pydantic import TypeAdapter +from litellm.litellm_core_utils.request_timeout_resolver import get_configured_request_timeout + DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS: Final = 600.0 _SECONDS: Final = TypeAdapter(float) @@ -48,8 +50,8 @@ def resolve_llm_passthrough_timeout( Anthropic /v1/messages). Non-streaming precedence: kwargs timeout/request_timeout -> litellm_params - timeout/request_timeout -> router_timeout -> general_settings.pass_through_request_timeout - -> 600s default. + timeout/request_timeout -> router_timeout -> litellm.request_timeout (litellm_settings.request_timeout, + when explicitly set) -> general_settings.pass_through_request_timeout -> 600s default. Streaming (``kwargs["stream"]`` truthy) resolves ``stream_timeout`` at every level before any generic timeout, matching ``Router._get_stream_timeout`` on the completion route: @@ -73,6 +75,7 @@ def resolve_llm_passthrough_timeout( deployment.get("timeout"), deployment.get("request_timeout"), router_timeout, + get_configured_request_timeout(), ) winner: Final = next((val for val in candidates if val is not None), None) return resolve_pass_through_request_timeout() if winner is None else _SECONDS.validate_python(winner) diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index fdc702af005..1ef39775bd3 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -10,7 +10,7 @@ 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, overload, runtime_checkable +from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn, Protocol, overload, runtime_checkable import httpx from openai._streaming import SSEDecoder @@ -265,6 +265,9 @@ def _mid_stream_fallback_eligible(mapped_exception: Exception) -> bool: return not isinstance(status_code, int) or status_code >= 500 or status_code == 429 +_PRE_OUTPUT_LIFECYCLE_EVENT_TYPES: Final = frozenset({"response.created", "response.in_progress", "response.queued"}) + + class BaseResponsesAPIStreamingIterator: """ Base class for streaming iterators that process responses from the Responses API. @@ -292,6 +295,7 @@ class BaseResponsesAPIStreamingIterator: self.start_time = getattr(logging_obj, "start_time", datetime.now()) self._failure_handled = False # Track if failure handler has been called self._yielded_first_chunk = False + self._output_started = False self._generated_content = "" self._generated_tool_arguments = "" self._completed_response_cached = False @@ -879,6 +883,46 @@ class BaseResponsesAPIStreamingIterator: except Exception: pass + def _note_yielded_event(self, event: ResponsesAPIStreamingResponse) -> None: + self._yielded_first_chunk = True + if event.type not in _PRE_OUTPUT_LIFECYCLE_EVENT_TYPES: + self._output_started = True + + def _fallback_error(self, original: Exception) -> MidStreamFallbackError: + return MidStreamFallbackError( + message=str(original), + model=self.model or "", + llm_provider=self.custom_llm_provider or "", + original_exception=original, + generated_content="", + is_pre_first_chunk=not self._yielded_first_chunk, + ) + + def _stream_ended_early_error(self) -> litellm.APIConnectionError: + return litellm.APIConnectionError( + message=( + f"{self.custom_llm_provider or 'provider'} closed the responses stream before any terminal event " + "(response.completed, response.incomplete or response.failed)" + ), + llm_provider=self.custom_llm_provider or "", + model=self.model or "", + ) + + def _raise_if_ended_without_terminal_event(self) -> None: + if self.completed_response is not None: + return + error: Final = self._stream_ended_early_error() + self._handle_failure(error) + if self._output_started: + raise error + raise self._fallback_error(error) from error + + def _raise_for_transport_error(self, error: httpx.ReadError | httpx.RemoteProtocolError) -> NoReturn: + self._handle_failure(error) + if self._output_started: + raise error + raise self._fallback_error(error) from error + async def call_post_streaming_hooks_for_testing( iterator: object, chunk: ResponsesAPIStreamingResponse @@ -934,12 +978,14 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): sse = await self.stream_iterator.__anext__() except StopAsyncIteration: self.finished = True + self._raise_if_ended_without_terminal_event() raise StopAsyncIteration self._check_max_streaming_duration() result = self._process_chunk(sse.data) if self.finished: + self._raise_if_ended_without_terminal_event() raise StopAsyncIteration elif result is not None: self._maybe_raise_for_error_event(result) @@ -948,7 +994,7 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): result = await self._call_post_streaming_deployment_hook( chunk=result, ) - self._yielded_first_chunk = True + self._note_yielded_event(result) return result # If result is None, continue the loop to get the next chunk @@ -957,10 +1003,9 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): raise except (httpx.ReadError, httpx.RemoteProtocolError) as e: self.finished = True - if self.completed_response is None: - self._handle_failure(e) - raise - raise StopAsyncIteration from e + if self.completed_response is not None: + raise StopAsyncIteration from e + self._raise_for_transport_error(e) except httpx.HTTPError as e: # Handle HTTP errors self.finished = True @@ -1016,12 +1061,14 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): sse = next(self.stream_iterator) except StopIteration: self.finished = True + self._raise_if_ended_without_terminal_event() raise StopIteration self._check_max_streaming_duration() result = self._process_chunk(sse.data) if self.finished: + self._raise_if_ended_without_terminal_event() raise StopIteration elif result is not None: self._maybe_raise_for_error_event(result) @@ -1030,7 +1077,7 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): async_function=self._call_post_streaming_deployment_hook, chunk=result, ) - self._yielded_first_chunk = True + self._note_yielded_event(result) return result # If result is None, continue the loop to get the next chunk @@ -1039,10 +1086,9 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): raise except (httpx.ReadError, httpx.RemoteProtocolError) as e: self.finished = True - if self.completed_response is None: - self._handle_failure(e) - raise - raise StopIteration from e + if self.completed_response is not None: + raise StopIteration from e + self._raise_for_transport_error(e) except httpx.HTTPError as e: # Handle HTTP errors self.finished = True diff --git a/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py b/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py index 47b377dc9a4..da37803b64a 100644 --- a/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py +++ b/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py @@ -381,6 +381,36 @@ class TestBaseResponsesAPIStreamingIterator: ) raise + @staticmethod + def _config_completing_after_one_delta() -> Mock: + mock_config = Mock(spec=BaseResponsesAPIConfig) + completed_response = ResponsesAPIResponse( + id="resp_123", + created_at=0, + status="completed", + model="gpt-5.5", + object="response", + output=[], + usage=ResponseAPIUsage(input_tokens=1, output_tokens=1, total_tokens=2), + ) + + def _transform(model, parsed_chunk, logging_obj): + if parsed_chunk.get("type") == "response.completed": + return ResponseCompletedEvent( + type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED, + response=completed_response, + ) + return OutputTextDeltaEvent( + type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, + item_id="msg_123", + output_index=0, + content_index=0, + delta=parsed_chunk["delta"], + ) + + mock_config.transform_streaming_response.side_effect = _transform + return mock_config + @pytest.mark.asyncio async def test_stop_async_iteration_not_logged_as_failure(self): """ @@ -399,6 +429,7 @@ class TestBaseResponsesAPIStreamingIterator: async def mock_aiter_bytes(): yield b'data: {"type": "response.output_text.delta", "delta": "test"}\n\n' + yield b'data: {"type": "response.completed", "response": {"id": "resp_123"}}\n\n' mock_response.aiter_bytes = mock_aiter_bytes @@ -408,11 +439,7 @@ class TestBaseResponsesAPIStreamingIterator: mock_logging_obj.async_failure_handler = Mock() mock_logging_obj.failure_handler = Mock() - mock_config = Mock(spec=BaseResponsesAPIConfig) - mock_delta_event = Mock() - mock_delta_event.type = ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA - mock_delta_event.delta = "test" - mock_config.transform_streaming_response.return_value = mock_delta_event + mock_config = self._config_completing_after_one_delta() # Create the iterator instance iterator = ResponsesAPIStreamingIterator( @@ -432,8 +459,9 @@ class TestBaseResponsesAPIStreamingIterator: except StopAsyncIteration: pass # This is expected - # Verify we got the chunk - assert len(chunks_received) == 1 + # Verify we got the delta and the terminal event + assert len(chunks_received) == 2 + assert iterator.completed_response is not None # CRITICAL: Verify that failure handlers were NOT called # StopAsyncIteration is a normal end of stream, not a failure @@ -460,6 +488,7 @@ class TestBaseResponsesAPIStreamingIterator: def mock_iter_bytes(): yield b'data: {"type": "response.output_text.delta", "delta": "test"}\n\n' + yield b'data: {"type": "response.completed", "response": {"id": "resp_123"}}\n\n' mock_response.iter_bytes = mock_iter_bytes @@ -469,11 +498,7 @@ class TestBaseResponsesAPIStreamingIterator: mock_logging_obj.async_failure_handler = Mock() mock_logging_obj.failure_handler = Mock() - mock_config = Mock(spec=BaseResponsesAPIConfig) - mock_delta_event = Mock() - mock_delta_event.type = ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA - mock_delta_event.delta = "test" - mock_config.transform_streaming_response.return_value = mock_delta_event + mock_config = self._config_completing_after_one_delta() # Create the iterator instance iterator = SyncResponsesAPIStreamingIterator( @@ -493,8 +518,9 @@ class TestBaseResponsesAPIStreamingIterator: except StopIteration: pass # This is expected - # Verify we got the chunk - assert len(chunks_received) == 1 + # Verify we got the delta and the terminal event + assert len(chunks_received) == 2 + assert iterator.completed_response is not None # CRITICAL: Verify that failure handlers were NOT called # StopIteration is a normal end of stream, not a failure diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index a40741c8fdb..3469df082e0 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -23,6 +23,7 @@ from starlette.datastructures import UploadFile as StarletteUploadFile import litellm from litellm._logging import verbose_proxy_logger +from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy._types import ProxyException, UserAPIKeyAuth @@ -1165,6 +1166,29 @@ def test_resolve_llm_passthrough_timeout_precedence(): assert resolve_llm_passthrough_timeout() == 6.0 +def test_resolve_llm_passthrough_timeout_honors_explicit_global_request_timeout(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr("litellm.request_timeout", 44.0, raising=False) + monkeypatch.setattr("litellm.request_timeout_explicitly_set", True, raising=False) + + with patch("litellm.proxy.proxy_server.general_settings", {"pass_through_request_timeout": 6}): + assert resolve_llm_passthrough_timeout() == 44.0 + assert resolve_llm_passthrough_timeout(kwargs={"stream": True}) == 44.0 + assert resolve_llm_passthrough_timeout(router_timeout=120) == 120.0 + assert resolve_llm_passthrough_timeout(kwargs={"stream": True}, router_stream_timeout=900) == 900.0 + assert resolve_llm_passthrough_timeout(litellm_params={"timeout": 90}) == 90.0 + assert resolve_llm_passthrough_timeout(kwargs={"timeout": 45}) == 45.0 + + +def test_resolve_llm_passthrough_timeout_skips_unset_global_request_timeout(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr("litellm.request_timeout", float(DEFAULT_REQUEST_TIMEOUT_SECONDS), raising=False) + monkeypatch.setattr("litellm.request_timeout_explicitly_set", False, raising=False) + + with patch("litellm.proxy.proxy_server.general_settings", {"pass_through_request_timeout": 6}): + assert resolve_llm_passthrough_timeout() == 6.0 + with patch("litellm.proxy.proxy_server.general_settings", {}): + assert resolve_llm_passthrough_timeout() == DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS + + def test_resolve_llm_passthrough_timeout_stream_timeout_precedence(): assert ( resolve_llm_passthrough_timeout( diff --git a/tests/test_litellm/responses/test_streaming_iterator.py b/tests/test_litellm/responses/test_streaming_iterator.py index dbf54ec3b9b..9dbbc20591e 100644 --- a/tests/test_litellm/responses/test_streaming_iterator.py +++ b/tests/test_litellm/responses/test_streaming_iterator.py @@ -6,13 +6,14 @@ completion_start_time = end_time.""" import json from datetime import datetime from typing import Final, Optional -from unittest.mock import Mock, patch +from unittest.mock import AsyncMock, Mock, patch import httpx import pytest from pydantic_core import PydanticSerializationError import litellm +from litellm.exceptions import MidStreamFallbackError from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig from litellm.responses.streaming_iterator import ( @@ -251,6 +252,171 @@ def test_sync_transport_error_before_completed_event_raises(): pass +_DONE_MARKER: Final = b"data: [DONE]\n\n" +_CREATED_EVENT: Final = _sse_event({"type": "response.created"}) +_IN_PROGRESS_EVENT: Final = _sse_event({"type": "response.in_progress"}) +_PARTIAL_OUTPUT_EVENTS: Final = _COMPLETE_STREAM_EVENTS[:-1] +_PRE_OUTPUT_PREFIXES: Final = [ + pytest.param([], True, id="nothing-yielded"), + pytest.param([_CREATED_EVENT], False, id="created"), + pytest.param([_CREATED_EVENT, _IN_PROGRESS_EVENT], False, id="created-and-in-progress"), +] + + +def _failure_tracking_logging_obj() -> Mock: + logging_obj: Final = _logging_obj_stub() + logging_obj.async_failure_handler = AsyncMock() + return logging_obj + + +def _assert_failure_logged_once(logging_obj: Mock, exception: Exception) -> None: + assert logging_obj.async_failure_handler.await_count == 1 + assert logging_obj.async_failure_handler.await_args.kwargs["exception"] is exception + + +@pytest.mark.asyncio +@pytest.mark.parametrize("prefix, pre_first_chunk", _PRE_OUTPUT_PREFIXES) +@pytest.mark.parametrize("trailing_error", _TRAILING_ERRORS, ids=type) +async def test_transport_error_before_any_output_raises_fallback_error(prefix, pre_first_chunk, trailing_error): + """A connection lost while only lifecycle events (response.created / response.in_progress) + have streamed is fallback-eligible, so it must surface as the MidStreamFallbackError the + router re-routes, carrying the raw transport error and no generated content.""" + logging_obj: Final = _failure_tracking_logging_obj() + iterator: Final = _make_iterator(sse_events=prefix, logging_obj=logging_obj, trailing_error=trailing_error) + + with pytest.raises(MidStreamFallbackError) as exc_info: + async for _ in iterator: + pass + + assert exc_info.value.original_exception is trailing_error + assert exc_info.value.is_pre_first_chunk is pre_first_chunk + assert exc_info.value.generated_content == "" + _assert_failure_logged_once(logging_obj, trailing_error) + + +@pytest.mark.asyncio +async def test_transport_error_after_output_started_is_not_fallback_eligible(): + logging_obj: Final = _failure_tracking_logging_obj() + trailing_error: Final = httpx.ReadError("Response payload is not completed") + iterator: Final = _make_iterator( + sse_events=_PARTIAL_OUTPUT_EVENTS, logging_obj=logging_obj, trailing_error=trailing_error + ) + + with pytest.raises(httpx.ReadError) as exc_info: + async for _ in iterator: + pass + + assert exc_info.value is trailing_error + _assert_failure_logged_once(logging_obj, trailing_error) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("trailer", [[], [_DONE_MARKER]], ids=["eof", "done-marker"]) +async def test_stream_ending_after_partial_output_without_terminal_event_raises(trailer): + """A clean EOF or `[DONE]` after output text but with no response.completed / + response.incomplete / response.failed is a truncated answer: the partial events still + reach the caller, then an explicit error follows instead of a normal end of stream.""" + logging_obj: Final = _failure_tracking_logging_obj() + iterator: Final = _make_iterator(sse_events=[*_PARTIAL_OUTPUT_EVENTS, *trailer], logging_obj=logging_obj) + + created: Final = await iterator.__anext__() + delta: Final = await iterator.__anext__() + with pytest.raises(litellm.APIConnectionError) as exc_info: + await iterator.__anext__() + + assert (created.type, delta.type) == ("response.created", "response.output_text.delta") + assert not isinstance(exc_info.value, MidStreamFallbackError) + assert exc_info.value.llm_provider == "openai" + _assert_failure_logged_once(logging_obj, exc_info.value) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("prefix, pre_first_chunk", _PRE_OUTPUT_PREFIXES) +@pytest.mark.parametrize("trailer", [[], [_DONE_MARKER]], ids=["eof", "done-marker"]) +async def test_stream_ending_before_any_output_raises_fallback_error(prefix, pre_first_chunk, trailer): + logging_obj: Final = _failure_tracking_logging_obj() + iterator: Final = _make_iterator(sse_events=[*prefix, *trailer], logging_obj=logging_obj) + + with pytest.raises(MidStreamFallbackError) as exc_info: + async for _ in iterator: + pass + + assert isinstance(exc_info.value.original_exception, litellm.APIConnectionError) + assert exc_info.value.is_pre_first_chunk is pre_first_chunk + assert exc_info.value.generated_content == "" + _assert_failure_logged_once(logging_obj, exc_info.value.original_exception) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("trailer", [[], [_DONE_MARKER]], ids=["eof", "done-marker"]) +async def test_complete_stream_still_ends_normally(trailer): + logging_obj: Final = _failure_tracking_logging_obj() + iterator: Final = _make_iterator(sse_events=[*_COMPLETE_STREAM_EVENTS, *trailer], logging_obj=logging_obj) + + seen: Final = [event.type async for event in iterator] + + assert seen[-1] == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + assert logging_obj.async_failure_handler.await_count == 0 + + +@pytest.mark.parametrize("trailing_error", _TRAILING_ERRORS, ids=type) +def test_sync_transport_error_before_any_output_raises_fallback_error(trailing_error): + logging_obj: Final = _failure_tracking_logging_obj() + iterator: Final = _make_sync_iterator( + sse_events=[_CREATED_EVENT, _IN_PROGRESS_EVENT], + logging_obj=logging_obj, + trailing_error=trailing_error, + ) + + with pytest.raises(MidStreamFallbackError) as exc_info: + for _ in iterator: + pass + + assert exc_info.value.original_exception is trailing_error + assert exc_info.value.is_pre_first_chunk is False + assert exc_info.value.generated_content == "" + _assert_failure_logged_once(logging_obj, trailing_error) + + +@pytest.mark.parametrize("trailer", [[], [_DONE_MARKER]], ids=["eof", "done-marker"]) +def test_sync_stream_ending_after_partial_output_without_terminal_event_raises(trailer): + logging_obj: Final = _failure_tracking_logging_obj() + iterator: Final = _make_sync_iterator(sse_events=[*_PARTIAL_OUTPUT_EVENTS, *trailer], logging_obj=logging_obj) + + created: Final = next(iterator) + delta: Final = next(iterator) + with pytest.raises(litellm.APIConnectionError) as exc_info: + next(iterator) + + assert (created.type, delta.type) == ("response.created", "response.output_text.delta") + assert not isinstance(exc_info.value, MidStreamFallbackError) + _assert_failure_logged_once(logging_obj, exc_info.value) + + +def test_sync_stream_ending_before_any_output_raises_fallback_error(): + logging_obj: Final = _failure_tracking_logging_obj() + iterator: Final = _make_sync_iterator(sse_events=[_CREATED_EVENT], logging_obj=logging_obj) + + with pytest.raises(MidStreamFallbackError) as exc_info: + for _ in iterator: + pass + + assert isinstance(exc_info.value.original_exception, litellm.APIConnectionError) + assert exc_info.value.is_pre_first_chunk is False + _assert_failure_logged_once(logging_obj, exc_info.value.original_exception) + + +@pytest.mark.parametrize("trailer", [[], [_DONE_MARKER]], ids=["eof", "done-marker"]) +def test_sync_complete_stream_still_ends_normally(trailer): + logging_obj: Final = _failure_tracking_logging_obj() + iterator: Final = _make_sync_iterator(sse_events=[*_COMPLETE_STREAM_EVENTS, *trailer], logging_obj=logging_obj) + + seen: Final = [event.type for event in iterator] + + assert seen[-1] == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + assert logging_obj.async_failure_handler.await_count == 0 + + def test_stream_cache_write_completes_when_asyncio_run_closes_the_loop(monkeypatch): """ Regression test for LIT-6184 on the /v1/responses streaming surface: the diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index 80131534183..10669de9cc8 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -4486,6 +4486,107 @@ async def test_aresponses_streaming_iterator_pre_first_chunk_skips_continuation( assert fbk["input"] == "Hello" # original input, no continuation messages +def _make_native_responses_iterator(*, sse_payloads: tuple[dict[str, str], ...], trailing_error: Exception | None): + """A real ResponsesAPIStreamingIterator over canned SSE bytes, so the router test covers the + iterator's own transport-error classification instead of a hand-built MidStreamFallbackError.""" + from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig + from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator + + async def aiter_bytes(): + for payload in sse_payloads: + yield f"data: {json.dumps(payload)}\n\n".encode() + if trailing_error is not None: + raise trailing_error + + def transform(model, parsed_chunk, logging_obj): + return MagicMock(type=parsed_chunk["type"]) + + response: Final = MagicMock() + response.headers = {} + response.aiter_bytes = aiter_bytes + config: Final = MagicMock(spec=BaseResponsesAPIConfig) + config.transform_streaming_response.side_effect = transform + logging_obj: Final = MagicMock(spec=LiteLLMLogging) + logging_obj.completion_start_time = None + logging_obj.model_call_details = {"litellm_params": {}} + return ResponsesAPIStreamingIterator( + response=response, + model="gpt-4", + responses_api_provider_config=config, + logging_obj=logging_obj, + litellm_metadata={}, + custom_llm_provider="openai", + ) + + +_RESPONSES_LIFECYCLE_PAYLOADS: Final = ({"type": "response.created"}, {"type": "response.in_progress"}) + + +@pytest.mark.asyncio +async def test_aresponses_streaming_iterator_falls_back_on_transport_drop_before_output(): + """A connection lost after response.created but before any output item is re-routed to the + fallback with the original input, the same as a provider error event would be.""" + router: Final = _make_router_with_fallback() + src: Final = _make_native_responses_iterator( + sse_payloads=_RESPONSES_LIFECYCLE_PAYLOADS, + trailing_error=httpx.ReadError("Response payload is not completed"), + ) + + with patch.object( + router, + "async_function_with_fallbacks_common_utils", + return_value=_AsyncList([MagicMock(type="response.completed")]), + ) as mock_fallback_utils: + wrapped: Final = await router._aresponses_streaming_iterator( + response=src, + initial_kwargs={ + "model": "gpt-4", + "stream": True, + "input": "Hello", + "original_generic_function": litellm.aresponses, + }, + ) + seen: Final = [chunk.type async for chunk in wrapped] + + assert seen == ["response.created", "response.in_progress", "response.completed"] + assert isinstance(mock_fallback_utils.call_args.kwargs["e"], MidStreamFallbackError) + assert mock_fallback_utils.call_args.kwargs["kwargs"]["input"] == "Hello" + + +@pytest.mark.asyncio +async def test_aresponses_streaming_iterator_surfaces_transport_drop_when_no_fallback_lands(): + transport_error: Final = httpx.ReadError("Response payload is not completed") + router: Final = _make_router_with_fallback() + src: Final = _make_native_responses_iterator( + sse_payloads=_RESPONSES_LIFECYCLE_PAYLOADS, trailing_error=transport_error + ) + + async def reraise_trigger(**kwargs): + raise kwargs["e"] + + with patch.object( + router, "async_function_with_fallbacks_common_utils", new=AsyncMock(side_effect=reraise_trigger) + ) as mock_fallback_utils: + wrapped: Final = await router._aresponses_streaming_iterator( + response=src, + initial_kwargs={ + "model": "gpt-4", + "stream": True, + "input": "Hello", + "original_generic_function": litellm.aresponses, + }, + ) + with pytest.raises(httpx.ReadError) as exc_info: + async for _ in wrapped: + pass + + assert exc_info.value is transport_error + assert mock_fallback_utils.await_count == 1 + trigger: Final = mock_fallback_utils.await_args.kwargs["e"] + assert isinstance(trigger, MidStreamFallbackError) + assert trigger.original_exception is transport_error + + @pytest.mark.asyncio async def test_aresponses_streaming_iterator_partial_content_injects_continuation(): """Mid-stream error: input is rewritten to include user prompt + @@ -6090,6 +6191,32 @@ def test_update_kwargs_with_deployment_passthrough_router_stream_timeout_sources assert _passthrough_timeout(default_router, default_router.model_list[0], stream=False) == 120.0 +def test_update_kwargs_with_deployment_passthrough_honors_global_request_timeout(monkeypatch: pytest.MonkeyPatch): + """litellm_settings.request_timeout must bound the native responses route when neither the + deployment nor the router carries a timeout, while a deployment timeout keeps winning.""" + monkeypatch.setattr("litellm.request_timeout", 44.0, raising=False) + monkeypatch.setattr("litellm.request_timeout_explicitly_set", True, raising=False) + router: Final = litellm.Router( + model_list=[ + { + "model_name": "responses-global-timeout", + "litellm_params": {"model": "openai/gpt-5.4-mini", "api_key": "fake-key"}, + }, + { + "model_name": "responses-deployment-timeout", + "litellm_params": {"model": "openai/gpt-5.4-mini", "api_key": "fake-key", "timeout": 3}, + }, + ], + ) + global_only, per_deployment = router.model_list + + with patch("litellm.proxy.proxy_server.general_settings", {"pass_through_request_timeout": 6}): + assert _passthrough_timeout(router, global_only, stream=True) == 44.0 + assert _passthrough_timeout(router, global_only, stream=False) == 44.0 + assert _passthrough_timeout(router, per_deployment, stream=True) == 3.0 + assert _passthrough_timeout(router, per_deployment, stream=False) == 3.0 + + @pytest.mark.asyncio async def test_router_acompletion_with_unknown_model_and_default_fallback(): """ From a09f8b84a4fc97c57caaf8fb0a446f170c706f10 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 25 Sep 2026 14:52:35 -0700 Subject: [PATCH 5/5] fix(sentry): scrub PII and secrets inside object reprs and nested locals, add SENTRY_SEND_DEFAULT_PII opt-in (#43123) * fix(sentry): scrub PII and secrets inside object reprs and nested locals, add SENTRY_SEND_DEFAULT_PII opt-in * fix(sentry): keep the SDK denylist and filter the request headers a virtual key arrives in * fix(sentry): leave source context lines unscrubbed * fix(sentry): filter bracketed secret values and cap the JSON walk depth * ci(deps): install sentry-sdk in the proxy-dev group so the unit shards import it * fix(sentry): scrub source-context names outside real stack frames * fix(sentry): tie the key pattern floor to the custom key minimum * fix(sentry): keep the key pattern floor at or below a generated key's length --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/constants.py | 15 + litellm/litellm_core_utils/litellm_logging.py | 17 +- .../litellm_core_utils/sentry_scrubbing.py | 152 ++++++++++ pyproject.toml | 1 + .../code_coverage_tests/recursive_detector.py | 1 + .../test_litellm_logging.py | 148 ++-------- .../test_sentry_scrubbing.py | 278 ++++++++++++++++++ uv.lock | 2 + 8 files changed, 481 insertions(+), 133 deletions(-) create mode 100644 litellm/litellm_core_utils/sentry_scrubbing.py create mode 100644 tests/test_litellm/litellm_core_utils/test_sentry_scrubbing.py diff --git a/litellm/constants.py b/litellm/constants.py index e7ba1f6b07f..8316761c95b 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1952,6 +1952,15 @@ SENTRY_DENYLIST: Final = [ "auth_token", "jwt_token", "private_key", + "authorization", + "api-key", + "x-api-key", + "x-goog-api-key", + "ocp-apim-subscription-key", + "x-litellm-api-key", + "x-mcp-auth", + "cookie", + "set-cookie", "SLACK_WEBHOOK_URL", "ALERTING_WEBHOOK_URL", "webhook_url", @@ -1974,6 +1983,12 @@ SENTRY_DENYLIST: Final = [ ] SENTRY_PII_DENYLIST: Final = [ "user_id", + "user_email", + "end_user_id", + "user_api_key_hash", + "user_api_key_user_id", + "user_api_key_user_email", + "user_api_key_end_user_id", "email", "phone", "address", diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 83ab2bc11a2..d8182140a17 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -43,8 +43,6 @@ from litellm.constants import ( DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT, EMPTY_MAPPING, PROVIDER_REQUEST_ID_HEADERS, - SENTRY_DENYLIST, - SENTRY_PII_DENYLIST, ) from litellm.cost_calculator import ( RealtimeAPITokenUsageProcessor, @@ -4423,21 +4421,10 @@ def set_callbacks(callback_list, function_id=None): print_verbose("Package 'sentry_sdk' is missing. Installing it...") subprocess.check_call([sys.executable, "-m", "pip", "install", "sentry_sdk"]) import sentry_sdk - from sentry_sdk.scrubber import EventScrubber + from litellm.litellm_core_utils.sentry_scrubbing import build_sentry_init_options sentry_sdk_instance = sentry_sdk - sentry_trace_rate = os.environ.get("SENTRY_API_TRACE_RATE", "1.0") - sentry_sample_rate = ( - os.environ.get("SENTRY_API_SAMPLE_RATE") if "SENTRY_API_SAMPLE_RATE" in os.environ else "1.0" - ) - sentry_sdk_instance.init( - dsn=os.environ.get("SENTRY_DSN"), - traces_sample_rate=float(sentry_trace_rate), - sample_rate=float(sentry_sample_rate if sentry_sample_rate else 1.0), - send_default_pii=False, # Prevent sending Personal Identifiable Information - event_scrubber=EventScrubber(denylist=SENTRY_DENYLIST, pii_denylist=SENTRY_PII_DENYLIST), - environment=os.environ.get("SENTRY_ENVIRONMENT", "production"), - ) + sentry_sdk_instance.init(**build_sentry_init_options(os.environ)) capture_exception = sentry_sdk_instance.capture_exception add_breadcrumb = sentry_sdk_instance.add_breadcrumb elif callback == "slack": diff --git a/litellm/litellm_core_utils/sentry_scrubbing.py b/litellm/litellm_core_utils/sentry_scrubbing.py new file mode 100644 index 00000000000..4c14cabc2ab --- /dev/null +++ b/litellm/litellm_core_utils/sentry_scrubbing.py @@ -0,0 +1,152 @@ +from __future__ import annotations + +import re +from collections.abc import Callable, Mapping, Sequence +from functools import reduce +from typing import TYPE_CHECKING, Final, TypeAlias, cast + +from pydantic import JsonValue +from sentry_sdk.scrubber import DEFAULT_DENYLIST, DEFAULT_PII_DENYLIST, EventScrubber +from typing_extensions import ReadOnly, TypedDict + +from litellm.constants import ( + LENGTH_OF_LITELLM_GENERATED_KEY, + MINIMUM_CUSTOM_KEY_LENGTH, + SENTRY_DENYLIST, + SENTRY_PII_DENYLIST, +) +from litellm.secret_managers.main import str_to_bool + +if TYPE_CHECKING: + from sentry_sdk.types import Event, Hint + +EventScrubFn: TypeAlias = "Callable[[Event, Hint], Event]" +JsonPath: TypeAlias = tuple[str, ...] + +FILTERED: Final = "[Filtered]" +SEND_DEFAULT_PII_ENV: Final = "SENTRY_SEND_DEFAULT_PII" +SECRET_FIELD_NAMES: Final = tuple(DEFAULT_DENYLIST) + tuple(SENTRY_DENYLIST) +PII_FIELD_NAMES: Final = tuple(DEFAULT_PII_DENYLIST) + tuple(SENTRY_PII_DENYLIST) + +KEY_PREFIX: Final = "sk-" + + +def build_key_pattern(custom_key_minimum: int, generated_key_bytes: int) -> re.Pattern[str]: + generated_suffix_length: Final = (generated_key_bytes * 4 + 2) // 3 + floor: Final = min(custom_key_minimum - len(KEY_PREFIX), generated_suffix_length) + return re.compile(rf"{KEY_PREFIX}[A-Za-z0-9_-]{{{floor},}}") + + +LITELLM_KEY_PATTERN: Final = build_key_pattern(MINIMUM_CUSTOM_KEY_LENGTH, LENGTH_OF_LITELLM_GENERATED_KEY) +SOURCE_CONTEXT_KEYS: Final = frozenset({"pre_context", "context_line", "post_context"}) +STACK_FRAME_PATHS: Final = frozenset( + { + ("exception", "values", "*", "stacktrace", "frames", "*"), + ("threads", "values", "*", "stacktrace", "frames", "*"), + ("stacktrace", "frames", "*"), + } +) +MAX_SCRUB_DEPTH: Final = 64 +EMAIL_PATTERN: Final = re.compile(r"[A-Za-z0-9._%+-]+@[A-Za-z0-9-]+(?:\.[A-Za-z0-9-]+)*\.[A-Za-z]{2,}") +SHA256_HEX_PATTERN: Final = re.compile(r"(? re.Pattern[str]: + names: Final = "|".join(re.escape(name) for name in field_names) + return re.compile( + rf"(?P(?{QUOTED_VALUE}|{BRACKETED_VALUE}|{BARE_VALUE})", + re.IGNORECASE, + ) + + +def build_string_scrubber(send_default_pii: bool) -> Callable[[str], str]: + field_names: Final = SECRET_FIELD_NAMES if send_default_pii else SECRET_FIELD_NAMES + PII_FIELD_NAMES + field_pattern: Final = build_repr_field_pattern(field_names) + value_patterns: Final = ( + (LITELLM_KEY_PATTERN,) if send_default_pii else (LITELLM_KEY_PATTERN, EMAIL_PATTERN, SHA256_HEX_PATTERN) + ) + + def scrub(text: str) -> str: + fields_scrubbed: Final = field_pattern.sub(_filtered_field, text) + return _substitute_all(value_patterns, fields_scrubbed) + + return scrub + + +def _filtered_field(match: re.Match[str]) -> str: + quote: Final = '"' if match.group("value").startswith('"') else "'" + return f"{match.group('field')}{quote}{FILTERED}{quote}" + + +def _substitute_all(patterns: Sequence[re.Pattern[str]], text: str) -> str: + return reduce(lambda scrubbed, pattern: pattern.sub(FILTERED, scrubbed), patterns, text) + + +def scrub_json_strings(value: JsonValue, scrub: Callable[[str], str], path: JsonPath = ()) -> JsonValue: + if len(path) > MAX_SCRUB_DEPTH: + return FILTERED + if isinstance(value, str): + return scrub(value) + if isinstance(value, dict): + unscrubbed_keys: Final = SOURCE_CONTEXT_KEYS if path in STACK_FRAME_PATHS else frozenset[str]() + return { # mutable-ok: JSON object + key: item if key in unscrubbed_keys else scrub_json_strings(item, scrub, (*path, key)) + for key, item in value.items() + } + if isinstance(value, list): + return [scrub_json_strings(item, scrub, (*path, "*")) for item in value] # mutable-ok: JSON array + return value + + +def build_event_scrubber(send_default_pii: bool) -> EventScrubFn: + scrub: Final = build_string_scrubber(send_default_pii) + + def scrub_event(event: Event, _hint: Hint) -> Event: + json_event: Final = cast("JsonValue", event) # cast-ok: [LIT006] the SDK serialized the event to JSON already + return cast("Event", scrub_json_strings(json_event, scrub)) # cast-ok: [LIT006] same JSON shape going back + + return scrub_event + + +def send_default_pii_from_env(env: Mapping[str, str]) -> bool: + return str_to_bool(env.get(SEND_DEFAULT_PII_ENV)) is True + + +def build_sentry_init_options(env: Mapping[str, str]) -> SentryInitOptions: + send_default_pii: Final = send_default_pii_from_env(env) + scrub_event: Final = build_event_scrubber(send_default_pii) + return SentryInitOptions( + dsn=env.get("SENTRY_DSN"), + traces_sample_rate=float(env.get("SENTRY_API_TRACE_RATE") or "1.0"), + sample_rate=float(env.get("SENTRY_API_SAMPLE_RATE") or "1.0"), + send_default_pii=send_default_pii, + event_scrubber=EventScrubber( + denylist=list(SECRET_FIELD_NAMES), # mutable-ok: EventScrubber appends pii_denylist onto denylist in place + pii_denylist=list(PII_FIELD_NAMES), # mutable-ok: EventScrubber takes List[str] + recursive=True, + send_default_pii=send_default_pii, + ), + before_send=scrub_event, + before_send_transaction=scrub_event, + environment=env.get("SENTRY_ENVIRONMENT", "production"), + ) diff --git a/pyproject.toml b/pyproject.toml index f2364b5e77b..28b00379cc7 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -249,6 +249,7 @@ proxy-dev = [ "prisma==0.11.0", "hypercorn==0.17.3", "prometheus-client==0.20.0", + "sentry-sdk==2.21.0", "opentelemetry-api==1.33.1", "opentelemetry-sdk==1.33.1", "opentelemetry-exporter-otlp==1.33.1", diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py index dc1f8592612..659dc438f2d 100644 --- a/tests/code_coverage_tests/recursive_detector.py +++ b/tests/code_coverage_tests/recursive_detector.py @@ -72,6 +72,7 @@ IGNORE_FUNCTIONS = [ "_string_leaves", # bounded by the nesting depth of a safe_json_structure output (a finite JSON tree, no cycles possible). "_replace_string_leaves", # bounded by the nesting depth of a safe_json_structure output (a finite JSON tree, no cycles possible). "_sort_processed_sets", # bounded by the nesting depth of the log-record extra it walks (a finite JSON tree, no cycles possible). + "scrub_json_strings", # max depth set (MAX_SCRUB_DEPTH); fails closed by returning "[Filtered]" for anything nested past the cap. ] diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 3afa31cc801..c4829ced9f3 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -20,7 +20,7 @@ from openai._legacy_response import HttpxBinaryResponseContent import litellm from litellm._logging import session_id_var, trace_id_var -from litellm.constants import SENTRY_DENYLIST, SENTRY_PII_DENYLIST +from litellm.constants import SENTRY_PII_DENYLIST from litellm.cost_calculator import ocr_batch_cost from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging @@ -357,108 +357,23 @@ def test_post_call_serializes_dict_with_datetime(logging_obj): assert "2026-05-11" in serialized -def test_sentry_sample_rate(monkeypatch): - existing_sample_rate = os.getenv("SENTRY_API_SAMPLE_RATE") - try: - # test with default value by removing the environment variable - if existing_sample_rate: - del os.environ["SENTRY_API_SAMPLE_RATE"] - - set_callbacks(["sentry"]) - # Check if the default sample rate is set to 1.0 - assert os.environ.get("SENTRY_API_SAMPLE_RATE") == "1.0" - - # test with custom value - monkeypatch.setenv("SENTRY_API_SAMPLE_RATE", "0.5") - - set_callbacks(["sentry"]) - # Check if the custom sample rate is set correctly - assert os.environ.get("SENTRY_API_SAMPLE_RATE") == "0.5" - except Exception as e: - print(f"Error: {e}") - finally: - # Restore the original environment variable - if existing_sample_rate: - monkeypatch.setenv("SENTRY_API_SAMPLE_RATE", existing_sample_rate) - else: - if "SENTRY_API_SAMPLE_RATE" in os.environ: - del os.environ["SENTRY_API_SAMPLE_RATE"] - - def test_sentry_environment(monkeypatch): - """Test that SENTRY_ENVIRONMENT is properly handled during Sentry initialization""" - existing_environment = os.getenv("SENTRY_ENVIRONMENT") - existing_dsn = os.getenv("SENTRY_DSN") + import sentry_sdk - # Create mock sentry_sdk module - mock_event_scrubber_instance = MagicMock() - mock_event_scrubber_cls = MagicMock(return_value=mock_event_scrubber_instance) - - mock_scrubber_module = MagicMock() - mock_scrubber_module.EventScrubber = mock_event_scrubber_cls - - mock_sentry_sdk = MagicMock() - mock_sentry_sdk.scrubber = mock_scrubber_module mock_init = MagicMock() - mock_sentry_sdk.init = mock_init + monkeypatch.setattr(sentry_sdk, "init", mock_init) + monkeypatch.setenv("SENTRY_DSN", "https://test@sentry.io/123456") + monkeypatch.delenv("SENTRY_ENVIRONMENT", raising=False) - # Inject mocks into sys.modules - sys.modules["sentry_sdk"] = mock_sentry_sdk - sys.modules["sentry_sdk.scrubber"] = mock_scrubber_module - - try: - # Set a mock DSN to allow Sentry initialization - monkeypatch.setenv("SENTRY_DSN", "https://test@sentry.io/123456") - - # Test with default value (no environment set) - if existing_environment: - del os.environ["SENTRY_ENVIRONMENT"] + set_callbacks(["sentry"]) + assert mock_init.call_args[1]["environment"] == "production" + for environment in ("development", "staging"): + monkeypatch.setenv("SENTRY_ENVIRONMENT", environment) mock_init.reset_mock() set_callbacks(["sentry"]) - # Check that init was called with default environment "production" mock_init.assert_called_once() - call_kwargs = mock_init.call_args[1] - assert call_kwargs["environment"] == "production" - - # Test with custom environment value - monkeypatch.setenv("SENTRY_ENVIRONMENT", "development") - - mock_init.reset_mock() - set_callbacks(["sentry"]) - # Check that init was called with custom environment "development" - mock_init.assert_called_once() - call_kwargs = mock_init.call_args[1] - assert call_kwargs["environment"] == "development" - - # Test with staging environment - monkeypatch.setenv("SENTRY_ENVIRONMENT", "staging") - - mock_init.reset_mock() - set_callbacks(["sentry"]) - # Check that init was called with custom environment "staging" - mock_init.assert_called_once() - call_kwargs = mock_init.call_args[1] - assert call_kwargs["environment"] == "staging" - - except Exception as e: - print(f"Error: {e}") - raise - finally: - # Restore the original environment variables - if existing_environment: - monkeypatch.setenv("SENTRY_ENVIRONMENT", existing_environment) - else: - if "SENTRY_ENVIRONMENT" in os.environ: - del os.environ["SENTRY_ENVIRONMENT"] - - if existing_dsn: - monkeypatch.setenv("SENTRY_DSN", existing_dsn) - else: - if "SENTRY_DSN" in os.environ: - del os.environ["SENTRY_DSN"] - - + assert mock_init.call_args[1]["environment"] == environment def test_use_custom_pricing_for_model(): from litellm.litellm_core_utils.litellm_logging import use_custom_pricing_for_model @@ -3100,37 +3015,34 @@ def test_speech_call_is_still_priced_from_input_characters(call_type): def test_sentry_event_scrubber_initialization(monkeypatch): - # Step 1: Create a fake sentry_sdk.scrubber module - mock_event_scrubber_instance = MagicMock() - mock_event_scrubber_cls = MagicMock(return_value=mock_event_scrubber_instance) + import sentry_sdk - mock_scrubber_module = MagicMock() - mock_scrubber_module.EventScrubber = mock_event_scrubber_cls - - # Step 2: Create a fake sentry_sdk module and insert into sys.modules - mock_sentry_sdk = MagicMock() - mock_sentry_sdk.scrubber = mock_scrubber_module mock_init = MagicMock() - mock_sentry_sdk.init = mock_init + monkeypatch.setattr(sentry_sdk, "init", mock_init) + monkeypatch.delenv("SENTRY_SEND_DEFAULT_PII", raising=False) - # Step 3: Inject both into sys.modules BEFORE import occurs - sys.modules["sentry_sdk"] = mock_sentry_sdk - sys.modules["sentry_sdk.scrubber"] = mock_scrubber_module - - # Step 4: Run the actual sentry setup code set_callbacks(["sentry"]) - # Step 5: Assert the EventScrubber was constructed correctly - mock_event_scrubber_cls.assert_called_once_with( - denylist=SENTRY_DENYLIST, - pii_denylist=SENTRY_PII_DENYLIST, - ) - - # Step 6: Assert the event_scrubber and PII args were passed mock_init.assert_called_once() call_args = mock_init.call_args[1] - assert call_args["event_scrubber"] == mock_event_scrubber_instance assert call_args["send_default_pii"] is False + assert call_args["event_scrubber"].recursive is True + assert {name.lower() for name in SENTRY_PII_DENYLIST} <= {name.lower() for name in call_args["event_scrubber"].denylist} + assert call_args["before_send"] is call_args["before_send_transaction"] + + +def test_sentry_send_default_pii_opt_in(monkeypatch): + import sentry_sdk + + mock_init = MagicMock() + monkeypatch.setattr(sentry_sdk, "init", mock_init) + monkeypatch.setenv("SENTRY_SEND_DEFAULT_PII", "true") + + set_callbacks(["sentry"]) + + call_args = mock_init.call_args[1] + assert call_args["send_default_pii"] is True + assert not {name.lower() for name in SENTRY_PII_DENYLIST} & {name.lower() for name in call_args["event_scrubber"].denylist} def test_get_masked_values(): diff --git a/tests/test_litellm/litellm_core_utils/test_sentry_scrubbing.py b/tests/test_litellm/litellm_core_utils/test_sentry_scrubbing.py new file mode 100644 index 00000000000..9aae3999129 --- /dev/null +++ b/tests/test_litellm/litellm_core_utils/test_sentry_scrubbing.py @@ -0,0 +1,278 @@ +import hashlib +import json +import secrets +from collections.abc import Callable, Mapping +from functools import reduce +from typing import Final, cast + +import pytest +import sentry_sdk +from pydantic import JsonValue +from sentry_sdk.envelope import Envelope +from sentry_sdk.transport import Transport +from sentry_sdk.utils import event_from_exception + +from litellm.constants import LENGTH_OF_LITELLM_GENERATED_KEY, MINIMUM_CUSTOM_KEY_LENGTH +from litellm.litellm_core_utils.sentry_scrubbing import ( + FILTERED, + MAX_SCRUB_DEPTH, + build_key_pattern, + build_sentry_init_options, + build_string_scrubber, + scrub_json_strings, +) +from litellm.proxy._types import LiteLLM_UserTable, UserAPIKeyAuth + +EMAIL: Final = "qa.user@example.com" +VIRTUAL_KEY: Final = "sk-virtual-key-under-test" +KEY_HASH: Final = hashlib.sha256(VIRTUAL_KEY.encode()).hexdigest() +MASTER_KEY: Final = "sk-master-key-under-test" +DATABASE_URL: Final = "postgresql://litellm:db-password-under-test@db.internal:5432/litellm" +PII_ON: Final = {"SENTRY_DSN": "https://key@sentry.example/1", "SENTRY_SEND_DEFAULT_PII": "true"} +PII_OFF: Final = {"SENTRY_DSN": "https://key@sentry.example/1"} + + +class RecordingTransport(Transport): + def __init__(self) -> None: + super().__init__() + self.last_envelope: Envelope | None = None + + def capture_envelope(self, envelope: Envelope) -> None: + self.last_envelope = envelope + + +def reject_request( + valid_token: UserAPIKeyAuth, + user_obj: LiteLLM_UserTable, + general_settings: Mapping[str, str], + data: Mapping[str, Mapping[str, str]], + raw_headers: Mapping[str, str], +) -> None: + raise RuntimeError(f"key {valid_token.token} owned by {user_obj.user_email} was rejected") + + +def raise_with_identity_locals() -> None: + reject_request( + valid_token=UserAPIKeyAuth(token=KEY_HASH, key_name="sk-...test", user_id=EMAIL, user_email=EMAIL), + user_obj=LiteLLM_UserTable(user_id=EMAIL, user_email=EMAIL, user_role="internal_user"), + general_settings={"master_key": MASTER_KEY, "database_url": DATABASE_URL}, + data={"metadata": {"user_api_key_hash": KEY_HASH, "user_api_key_user_email": EMAIL}}, + raw_headers={"authorization": f"Bearer {VIRTUAL_KEY}", "x-api-key": VIRTUAL_KEY, "content-type": "application/json"}, + ) + + +def raise_with_source_context_named_locals() -> None: + metadata: Final = {"context_line": f"Bearer {VIRTUAL_KEY}", "pre_context": [EMAIL], "post_context": [KEY_HASH]} + stacktrace: Final = {"frames": [{"context_line": MASTER_KEY, "pre_context": [EMAIL]}]} + raise RuntimeError(f"rejected with {len(metadata)} metadata fields and {len(stacktrace)} stack fields") + + +def capture_serialized_event(env: Mapping[str, str], raiser: Callable[[], None] = raise_with_identity_locals) -> str: + transport: Final = RecordingTransport() + client: Final = sentry_sdk.Client(transport=transport, **build_sentry_init_options(env)) + try: + raiser() + except RuntimeError as error: + event, hint = event_from_exception(error, client_options=client.options) + client.capture_event(event, hint=hint) + assert transport.last_envelope is not None + return json.dumps(transport.last_envelope.items[0].payload.json) + + +def innermost_frame_vars(serialized: str) -> dict[str, JsonValue]: + event: Final = json.loads(serialized) + frames: Final = event["exception"]["values"][0]["stacktrace"]["frames"] + return frames[-1]["vars"] + + +def test_default_event_carries_no_email_hash_or_secret_anywhere() -> None: + serialized: Final = capture_serialized_event(PII_OFF) + assert EMAIL not in serialized + assert KEY_HASH not in serialized + assert MASTER_KEY not in serialized + assert VIRTUAL_KEY not in serialized + assert "db-password-under-test" not in serialized + frame_vars: Final = innermost_frame_vars(serialized) + assert frame_vars["raw_headers"] == {"authorization": FILTERED, "x-api-key": FILTERED, "content-type": "'application/json'"} + assert f"token='{FILTERED}'" in frame_vars["valid_token"] + assert f"user_id='{FILTERED}'" in frame_vars["valid_token"] + assert f"user_email='{FILTERED}'" in frame_vars["user_obj"] + assert frame_vars["general_settings"] == {"master_key": FILTERED, "database_url": FILTERED} + assert frame_vars["data"] == {"metadata": {"user_api_key_hash": FILTERED, "user_api_key_user_email": FILTERED}} + assert "key_name='sk-...test'" in frame_vars["valid_token"] + assert "user_role='internal_user'" in frame_vars["user_obj"] + + +def test_source_context_lines_are_left_readable() -> None: + frames: Final = json.loads(capture_serialized_event(PII_OFF))["exception"]["values"][0]["stacktrace"]["frames"] + source_lines: Final = tuple( + line + for frame in frames + for line in (*frame.get("pre_context", []), frame.get("context_line", ""), *frame.get("post_context", [])) + ) + assert any("token=KEY_HASH" in line for line in source_lines) + assert not any(FILTERED in line for line in source_lines) + + +def test_source_context_names_outside_stack_frames_are_scrubbed() -> None: + serialized: Final = capture_serialized_event(PII_OFF, raise_with_source_context_named_locals) + assert VIRTUAL_KEY not in serialized + assert MASTER_KEY not in serialized + assert EMAIL not in serialized + assert KEY_HASH not in serialized + frame_vars: Final = innermost_frame_vars(serialized) + assert frame_vars["metadata"] == { + "context_line": f"'Bearer {FILTERED}'", + "pre_context": [f"'{FILTERED}'"], + "post_context": [f"'{FILTERED}'"], + } + assert frame_vars["stacktrace"] == {"frames": [{"context_line": f"'{FILTERED}'", "pre_context": [f"'{FILTERED}'"]}]} + innermost_frame: Final = json.loads(serialized)["exception"]["values"][0]["stacktrace"]["frames"][-1] + assert "raise RuntimeError" in innermost_frame["context_line"] + assert FILTERED not in json.dumps(innermost_frame["pre_context"]) + + +def test_default_event_keeps_the_exception_message_shape() -> None: + serialized: Final = capture_serialized_event(PII_OFF) + message: Final = json.loads(serialized)["exception"]["values"][0]["value"] + assert message == f"key {FILTERED} owned by {FILTERED} was rejected" + + +def test_pii_opt_in_keeps_identifiers_and_still_scrubs_secrets() -> None: + serialized: Final = capture_serialized_event(PII_ON) + frame_vars: Final = innermost_frame_vars(serialized) + assert f"user_id='{EMAIL}'" in frame_vars["valid_token"] + assert f"user_email='{EMAIL}'" in frame_vars["user_obj"] + assert frame_vars["data"] == { + "metadata": {"user_api_key_hash": f"'{KEY_HASH}'", "user_api_key_user_email": f"'{EMAIL}'"} + } + assert f"token='{FILTERED}'" in frame_vars["valid_token"] + assert frame_vars["general_settings"] == {"master_key": FILTERED, "database_url": FILTERED} + assert frame_vars["raw_headers"] == {"authorization": FILTERED, "x-api-key": FILTERED, "content-type": "'application/json'"} + assert MASTER_KEY not in serialized + assert VIRTUAL_KEY not in serialized + assert "db-password-under-test" not in serialized + + +def test_transaction_events_are_scrubbed_too() -> None: + transport: Final = RecordingTransport() + client: Final = sentry_sdk.Client(transport=transport, **build_sentry_init_options(PII_OFF)) + client.capture_event( + { + "type": "transaction", + "transaction": "/user/info", + "contexts": {"trace": {"trace_id": "a" * 32, "span_id": "b" * 16}}, + "spans": [{"description": f"lookup {EMAIL} by {KEY_HASH}", "span_id": "c" * 16, "trace_id": "a" * 32}], + } + ) + assert transport.last_envelope is not None + serialized: Final = json.dumps(transport.last_envelope.items[0].payload.json) + assert EMAIL not in serialized + assert KEY_HASH not in serialized + assert f"lookup {FILTERED} by {FILTERED}" in serialized + + +@pytest.mark.parametrize( + ("text", "expected"), + [ + ( + "UserAPIKeyAuth(token='abc', key_alias='team-a', user_id=None)", + f"UserAPIKeyAuth(token='{FILTERED}', key_alias='team-a', user_id=None)", + ), + ('{"api_key": "sk-1", "model": "gpt-5"}', f'{{"api_key": "{FILTERED}", "model": "gpt-5"}}'), + ("{'user_id': 'u-1', 'max_budget': 5}", f"{{'user_id': '{FILTERED}', 'max_budget': 5}}"), + ("Config(OPENAI_API_KEY=sk-live, timeout=10)", f"Config(OPENAI_API_KEY='{FILTERED}', timeout=10)"), + ("lookup for somebody@example.com failed", f"lookup for {FILTERED} failed"), + (f"hashed key {KEY_HASH} not found", f"hashed key {FILTERED} not found"), + ("request id 0123456789abcdef0123456789abcdef stays", "request id 0123456789abcdef0123456789abcdef stays"), + ("monkey=banana", "monkey=banana"), + ( + "{'x-api-key': 'k-1', 'cookie': 'session=abc', 'content-type': 'application/json'}", + f"{{'x-api-key': '{FILTERED}', 'cookie': '{FILTERED}', 'content-type': 'application/json'}}", + ), + ( + "headers={'x-tenant-key': 'sk-custom-header-key-0123456789'} key_name='sk-...6789'", + f"headers={{'x-tenant-key': '{FILTERED}'}} key_name='sk-...6789'", + ), + ( + "master_key={'value': 'not-a-litellm-key'} timeout=10", + f"master_key='{FILTERED}' timeout=10", + ), + ( + "credentials=[{'value': ('deep', 'secret')}], model='gpt-5'", + f"credentials='{FILTERED}', model='gpt-5'", + ), + ], +) +def test_string_scrubber_rewrites_field_and_value_forms(text: str, expected: str) -> None: + assert build_string_scrubber(send_default_pii=False)(text) == expected + + +def test_bare_key_floor_follows_the_custom_key_minimum() -> None: + scrub: Final = build_string_scrubber(send_default_pii=False) + shortest_key: Final = "sk-" + "a" * (MINIMUM_CUSTOM_KEY_LENGTH - len("sk-")) + assert scrub(f"label={shortest_key} model=gpt-5") == f"label={FILTERED} model=gpt-5" + assert scrub(f"label={shortest_key[:-1]} model=gpt-5") == f"label={shortest_key[:-1]} model=gpt-5" + + +def test_key_pattern_floor_never_exceeds_a_generated_key() -> None: + generated_key: Final = "sk-" + secrets.token_urlsafe(LENGTH_OF_LITELLM_GENERATED_KEY) + stricter_custom_minimum: Final = len(generated_key) + 10 + assert build_key_pattern(stricter_custom_minimum, LENGTH_OF_LITELLM_GENERATED_KEY).fullmatch(generated_key) + assert build_key_pattern(stricter_custom_minimum, LENGTH_OF_LITELLM_GENERATED_KEY).fullmatch(generated_key[:-1]) is None + + +def test_json_walk_fails_closed_past_the_depth_cap() -> None: + scrub: Final = build_string_scrubber(send_default_pii=False) + nested: Final = reduce(lambda inner, _: [inner], range(MAX_SCRUB_DEPTH + 1), cast("JsonValue", "api_key=sk-1")) + assert FILTERED in json.dumps(scrub_json_strings(nested, scrub)) + assert "sk-1" not in json.dumps(scrub_json_strings(nested, scrub)) + assert scrub_json_strings([["api_key=sk-1"]], scrub) == [[f"api_key='{FILTERED}'"]] + + +def test_string_scrubber_with_pii_on_only_scrubs_secrets() -> None: + scrub: Final = build_string_scrubber(send_default_pii=True) + assert scrub(f"user_id='{EMAIL}', token='{KEY_HASH}', email {EMAIL} hash {KEY_HASH}") == ( + f"user_id='{EMAIL}', token='{FILTERED}', email {EMAIL} hash {KEY_HASH}" + ) + assert scrub(f"headers={{'authorization': 'Bearer {VIRTUAL_KEY}'}} sent {VIRTUAL_KEY}") == ( + f"headers={{'authorization': '{FILTERED}'}} sent {FILTERED}" + ) + + +@pytest.mark.parametrize( + ("env", "expected"), + [ + ({}, False), + ({"SENTRY_SEND_DEFAULT_PII": "true"}, True), + ({"SENTRY_SEND_DEFAULT_PII": "True"}, True), + ({"SENTRY_SEND_DEFAULT_PII": "false"}, False), + ({"SENTRY_SEND_DEFAULT_PII": "yes please"}, False), + ], +) +def test_send_default_pii_comes_from_the_environment(env: Mapping[str, str], expected: bool) -> None: + assert build_sentry_init_options(env)["send_default_pii"] is expected + + +def test_init_options_read_dsn_rates_and_environment() -> None: + options: Final = build_sentry_init_options( + { + "SENTRY_DSN": "https://key@sentry.example/7", + "SENTRY_API_TRACE_RATE": "0.25", + "SENTRY_API_SAMPLE_RATE": "0.5", + "SENTRY_ENVIRONMENT": "staging", + } + ) + assert options["dsn"] == "https://key@sentry.example/7" + assert options["traces_sample_rate"] == 0.25 + assert options["sample_rate"] == 0.5 + assert options["environment"] == "staging" + assert options["event_scrubber"].recursive is True + + +def test_init_options_defaults() -> None: + options: Final = build_sentry_init_options({}) + assert options["dsn"] is None + assert options["traces_sample_rate"] == 1.0 + assert options["sample_rate"] == 1.0 + assert options["environment"] == "production" diff --git a/uv.lock b/uv.lock index c235171ecb2..527f53bd372 100644 --- a/uv.lock +++ b/uv.lock @@ -4743,6 +4743,7 @@ proxy-dev = [ { name = "opentelemetry-sdk" }, { name = "prisma" }, { name = "prometheus-client" }, + { name = "sentry-sdk" }, ] [package.metadata] @@ -4956,6 +4957,7 @@ proxy-dev = [ { name = "opentelemetry-sdk", specifier = "==1.33.1" }, { name = "prisma", specifier = "==0.11.0" }, { name = "prometheus-client", specifier = "==0.20.0" }, + { name = "sentry-sdk", specifier = "==2.21.0" }, ] [[package]]