Merge branch 'main' of https://github.com/BerriAI/litellm into include-cost-in-usage

This commit is contained in:
Ben Langfeld 2026-09-25 19:17:07 -03:00
commit 6e670e4a82
No known key found for this signature in database
17 changed files with 1043 additions and 167 deletions

View file

@ -1953,6 +1953,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",
@ -1975,6 +1984,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",

View file

@ -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":

View file

@ -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"(?<![0-9A-Za-z])[0-9a-f]{64}(?![0-9A-Za-z])")
QUOTED_VALUE: Final = r"'(?:[^'\\]|\\.)*'|\"(?:[^\"\\]|\\.)*\""
BRACKET_ATOM: Final = rf"(?:{QUOTED_VALUE})|[^\[\]{{}}()'\"]"
NESTED_BRACKET_LEVELS: Final = 3
BRACKETED_VALUE: Final = reduce(
lambda inner, _: rf"[\[{{(](?:{BRACKET_ATOM}|{inner})*[\]}})]",
range(NESTED_BRACKET_LEVELS),
rf"[\[{{(](?:{BRACKET_ATOM})*[\]}})]",
)
BARE_VALUE: Final = r"(?!None(?![0-9A-Za-z_]))[^,)\]}\s]+"
class SentryInitOptions(TypedDict):
dsn: ReadOnly[str | None]
traces_sample_rate: ReadOnly[float]
sample_rate: ReadOnly[float]
send_default_pii: ReadOnly[bool]
event_scrubber: ReadOnly[EventScrubber]
before_send: ReadOnly[EventScrubFn]
before_send_transaction: ReadOnly[EventScrubFn]
environment: ReadOnly[str]
def build_repr_field_pattern(field_names: Sequence[str]) -> re.Pattern[str]:
names: Final = "|".join(re.escape(name) for name in field_names)
return re.compile(
rf"(?P<field>(?<![0-9A-Za-z_])(?:{names})=|['\"](?:{names})['\"]:\s*)(?P<value>{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"),
)

View file

@ -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)

View file

@ -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

View file

@ -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}"

View file

@ -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",

View file

@ -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.
]

View file

@ -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

View file

@ -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():

View file

@ -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"

View file

@ -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(

View file

@ -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

View file

@ -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

View file

@ -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():
"""

View file

@ -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,

2
uv.lock generated
View file

@ -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]]