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/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/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/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/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/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/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/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/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/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/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() 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/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 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(): """ 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, 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]]