mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
Merge remote-tracking branch 'origin/main' into litellm_agent365_fail_open_default
This commit is contained in:
commit
8838391f7e
22 changed files with 1187 additions and 221 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
152
litellm/litellm_core_utils/sentry_scrubbing.py
Normal file
152
litellm/litellm_core_utils/sentry_scrubbing.py
Normal 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"),
|
||||
)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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/<id>/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/<id>/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/<id>/files calls will 403 for "
|
||||
"non-admin keys.",
|
||||
type(original_stream_response).__name__,
|
||||
)
|
||||
|
||||
async def base_passthrough_process_llm_request(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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`
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
278
tests/test_litellm/litellm_core_utils/test_sentry_scrubbing.py
Normal file
278
tests/test_litellm/litellm_core_utils/test_sentry_scrubbing.py
Normal 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"
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
137
tests/test_litellm/router_strategy/test_budget_limiter.py
Normal file
137
tests/test_litellm/router_strategy/test_budget_limiter.py
Normal 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
|
||||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
2
uv.lock
generated
|
|
@ -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]]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue