mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
Merge pull request #42618 from BerriAI/litellm_backport_stable_1_102_x_realtime_otel
chore(release): backport #42388 and #41462 to stable/1.102.x
This commit is contained in:
commit
d09bbae1c6
21 changed files with 522 additions and 103 deletions
|
|
@ -5,6 +5,7 @@ from collections.abc import Callable, Iterable, Mapping
|
|||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, TypedDict, cast
|
||||
|
||||
import litellm
|
||||
|
|
@ -20,7 +21,9 @@ from litellm.integrations.opentelemetry_utils.gen_ai_semconv import (
|
|||
OTELSemconvCategory,
|
||||
parse_semconv_opt_in,
|
||||
)
|
||||
from litellm.integrations.otel.model.baggage import promoted_metadata
|
||||
from litellm.integrations.otel.model.db_endpoint import db_span_attributes
|
||||
from litellm.integrations.otel.model.metadata import flatten_metadata
|
||||
from litellm.integrations.otel.model.semconv import Metric
|
||||
from litellm.litellm_core_utils.internal_call_metadata import is_unbilled_non_inference_call_from_params
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
|
@ -288,6 +291,7 @@ class OpenTelemetryConfig:
|
|||
# under ``litellm.team.metadata``. Empty by default so none of a team's
|
||||
# metadata leaves the process until explicitly allowlisted.
|
||||
baggage_team_metadata_keys: list[str] = field(default_factory=list)
|
||||
baggage_metadata_keys: list[str] = field(default_factory=list)
|
||||
# Prometheus-style include/exclude control over which attributes are stamped
|
||||
# on emitted metrics, to cap metric cardinality.
|
||||
attributes: OTELMetricAttributeFilter | None = None
|
||||
|
|
@ -314,6 +318,9 @@ class OpenTelemetryConfig:
|
|||
self.baggage_team_metadata_keys = _normalize_team_metadata_keys(
|
||||
self.baggage_team_metadata_keys
|
||||
) or _normalize_team_metadata_keys(os.getenv("LITELLM_OTEL_BAGGAGE_TEAM_METADATA_KEYS"))
|
||||
self.baggage_metadata_keys = _normalize_team_metadata_keys(
|
||||
self.baggage_metadata_keys
|
||||
) or _normalize_team_metadata_keys(os.getenv("LITELLM_OTEL_BAGGAGE_METADATA_KEYS"))
|
||||
|
||||
@classmethod
|
||||
def from_env(cls):
|
||||
|
|
@ -366,11 +373,14 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
**kwargs,
|
||||
):
|
||||
team_metadata_keys_override: Final = kwargs.pop("baggage_team_metadata_keys", None)
|
||||
metadata_keys_override: Final = kwargs.pop("baggage_metadata_keys", None)
|
||||
metric_attributes_override: Final = kwargs.pop("attributes", None)
|
||||
if config is None:
|
||||
config = OpenTelemetryConfig.from_env()
|
||||
if team_metadata_keys_override is not None:
|
||||
config.baggage_team_metadata_keys = _normalize_team_metadata_keys(team_metadata_keys_override)
|
||||
if metadata_keys_override is not None:
|
||||
config.baggage_metadata_keys = _normalize_team_metadata_keys(metadata_keys_override)
|
||||
if metric_attributes_override is not None:
|
||||
config.attributes = _build_metric_attribute_filter(metric_attributes_override)
|
||||
|
||||
|
|
@ -1542,6 +1552,11 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
if team_metadata:
|
||||
self.safe_set_attribute(span=span, key=TEAM_METADATA_ATTRIBUTE, value=team_metadata)
|
||||
|
||||
if self.config.baggage_metadata_keys:
|
||||
flat_metadata: Final = MappingProxyType(dict(flatten_metadata(metadata)))
|
||||
for key, value in promoted_metadata(flat_metadata, tuple(self.config.baggage_metadata_keys)).items():
|
||||
self.safe_set_attribute(span=span, key=key, value=value)
|
||||
|
||||
model_group: Final = standard_logging_payload.get("model_group")
|
||||
if model_group:
|
||||
self.safe_set_attribute(span=span, key=MODEL_GROUP_ATTRIBUTE, value=model_group)
|
||||
|
|
|
|||
|
|
@ -33,6 +33,7 @@ from litellm.integrations.otel.model.metadata import (
|
|||
LLMCallEvent,
|
||||
RequestIdentity,
|
||||
auth_metadata,
|
||||
metadata_from_request_data,
|
||||
model_from_request_data,
|
||||
)
|
||||
from litellm.integrations.otel.model.payloads import (
|
||||
|
|
@ -679,7 +680,12 @@ class OpenTelemetryV2(CustomLogger):
|
|||
# / errors are the FastAPI instrumentor's job, so we don't touch it here.
|
||||
# ====================================================================== #
|
||||
|
||||
def seed_request_identity(self, user_api_key_dict: object, model: str | None = None) -> None:
|
||||
def seed_request_identity(
|
||||
self,
|
||||
user_api_key_dict: object,
|
||||
model: str | None = None,
|
||||
request_metadata: Mapping[str, object] | None = None,
|
||||
) -> None:
|
||||
"""Attach request-identity Baggage to the current context + server span.
|
||||
|
||||
Seeding identity into Baggage makes **every** span emitted afterwards for
|
||||
|
|
@ -691,7 +697,7 @@ class OpenTelemetryV2(CustomLogger):
|
|||
isn't determined yet, which is correct.
|
||||
"""
|
||||
try:
|
||||
identity: Final = RequestIdentity.from_user_api_key_auth(user_api_key_dict)
|
||||
identity: Final = RequestIdentity.from_user_api_key_auth(user_api_key_dict, request_metadata)
|
||||
bag: Final = promoted_baggage(
|
||||
identity,
|
||||
model,
|
||||
|
|
@ -743,6 +749,7 @@ class OpenTelemetryV2(CustomLogger):
|
|||
self.seed_request_identity(
|
||||
user_api_key_dict,
|
||||
model=model_from_request_data(data),
|
||||
request_metadata=metadata_from_request_data(data),
|
||||
)
|
||||
return data
|
||||
|
||||
|
|
|
|||
|
|
@ -15,9 +15,10 @@ never promoted whole.
|
|||
|
||||
import json
|
||||
from collections.abc import Callable, Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from litellm.integrations.otel.model.metadata import RequestIdentity
|
||||
from litellm.integrations.otel.model.metadata import REQUESTER_METADATA_PATH, RequestIdentity
|
||||
from litellm.integrations.otel.model.semconv import GenAI, LiteLLM
|
||||
|
||||
# Attribute key -> value extractor over (identity, request_model,
|
||||
|
|
@ -79,17 +80,23 @@ def promoted_baggage(
|
|||
``team_metadata_keys`` selects sub-keys of the team's metadata to promote
|
||||
under ``litellm.team.metadata``. Empty values are dropped.
|
||||
"""
|
||||
out: Final[dict[str, str]] = {}
|
||||
for key, extract in _PROMOTABLE.items():
|
||||
if key in promoted_keys:
|
||||
value = extract(identity, request_model, team_metadata_keys)
|
||||
if value:
|
||||
out[key] = value
|
||||
for meta_key in metadata_keys:
|
||||
value = identity.metadata.get(meta_key)
|
||||
if value:
|
||||
out[f"{LiteLLM.METADATA_PREFIX}{meta_key}"] = value
|
||||
return out
|
||||
identity_values: Final = {
|
||||
key: value
|
||||
for key, extract in _PROMOTABLE.items()
|
||||
if key in promoted_keys and (value := extract(identity, request_model, team_metadata_keys))
|
||||
}
|
||||
return {**identity_values, **promoted_metadata(identity.metadata, metadata_keys)}
|
||||
|
||||
|
||||
def promoted_metadata(metadata: Mapping[str, str], metadata_keys: tuple[str, ...]) -> Mapping[str, str]:
|
||||
"""Allowlisted entries of a flattened metadata mapping under ``litellm.metadata.*``."""
|
||||
return MappingProxyType(
|
||||
{
|
||||
f"{LiteLLM.METADATA_PREFIX}{meta_key.removeprefix(REQUESTER_METADATA_PATH)}": value
|
||||
for meta_key in metadata_keys
|
||||
if (value := metadata.get(meta_key))
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _filtered_team_metadata_json(
|
||||
|
|
|
|||
|
|
@ -210,7 +210,10 @@ class OpenTelemetryV2Config(BaseSettings):
|
|||
validation_alias=AliasChoices("baggage_metadata_keys", "LITELLM_OTEL_BAGGAGE_METADATA_KEYS"),
|
||||
description=(
|
||||
"Metadata sub-keys promoted under the ``litellm.metadata.*`` "
|
||||
"namespace. Configure via the ``LITELLM_OTEL_BAGGAGE_METADATA_KEYS`` "
|
||||
"namespace. A dotted path such as ``requester_metadata.trace_id`` "
|
||||
"reads the caller's nested ``metadata.trace_id`` and is promoted as "
|
||||
"``litellm.metadata.trace_id``; other dotted keys keep their full path. "
|
||||
"Configure via the ``LITELLM_OTEL_BAGGAGE_METADATA_KEYS`` "
|
||||
"env var (comma-separated) or "
|
||||
"``callback_settings.otel.baggage_metadata_keys`` in config.yaml."
|
||||
),
|
||||
|
|
|
|||
|
|
@ -49,6 +49,8 @@ if TYPE_CHECKING:
|
|||
from litellm.types.utils import StandardLoggingPayload
|
||||
|
||||
LANGFUSE_TRACE_NAME_HEADER: Final = "langfuse_trace_name"
|
||||
REQUESTER_METADATA_KEY: Final = "requester_metadata"
|
||||
REQUESTER_METADATA_PATH: Final = f"{REQUESTER_METADATA_KEY}."
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
|
|
@ -78,7 +80,7 @@ class RequestIdentity:
|
|||
model, not just the user-facing one.
|
||||
"""
|
||||
raw_meta: Final = cast(Mapping[str, object], payload.get("metadata") or {})
|
||||
metadata = {key: str(value) for key, value in raw_meta.items() if isinstance(value, (str, bool, int, float))}
|
||||
metadata: Final = MappingProxyType(dict(flatten_metadata(raw_meta)))
|
||||
return cls(
|
||||
call_id=as_str(payload.get("litellm_call_id")) or as_str(payload.get("id")),
|
||||
# StandardLoggingMetadata's canonical key is ``user_api_key_team_id``;
|
||||
|
|
@ -95,7 +97,9 @@ class RequestIdentity:
|
|||
)
|
||||
|
||||
@classmethod
|
||||
def from_user_api_key_auth(cls, auth: object) -> RequestIdentity:
|
||||
def from_user_api_key_auth(
|
||||
cls, auth: object, request_metadata: Mapping[str, object] | None = None
|
||||
) -> RequestIdentity:
|
||||
"""Identity from a ``UserAPIKeyAuth`` (duck-typed to keep this module
|
||||
free of a proxy import).
|
||||
|
||||
|
|
@ -103,11 +107,13 @@ class RequestIdentity:
|
|||
guardrail, or service span is created — so the whole request's spans
|
||||
inherit identity, not just the LLM-call span. Metadata sub-keys use the
|
||||
``user_api_key_*`` names that ``baggage.DEFAULT_BAGGAGE_METADATA_KEYS``
|
||||
promotes.
|
||||
promotes; ``request_metadata`` (the caller's ``requester_metadata``
|
||||
snapshot) is flattened to dotted keys so ``requester_metadata.<key>``
|
||||
resolves too.
|
||||
"""
|
||||
get: Final = lambda name: getattr(auth, name, None) # noqa: E731
|
||||
metadata: Final = {
|
||||
meta_key: str(value)
|
||||
auth_meta: Final = tuple(
|
||||
(meta_key, str(value))
|
||||
for meta_key, attr in (
|
||||
("user_api_key_user_id", "user_id"),
|
||||
("user_api_key_org_id", "org_id"),
|
||||
|
|
@ -115,7 +121,9 @@ class RequestIdentity:
|
|||
("user_api_key_end_user_id", "end_user_id"),
|
||||
)
|
||||
if (value := get(attr))
|
||||
}
|
||||
)
|
||||
request_meta: Final = flatten_metadata(request_metadata) if request_metadata is not None else ()
|
||||
metadata: Final = MappingProxyType(dict((*request_meta, *auth_meta)))
|
||||
return cls(
|
||||
team_id=as_str(get("team_id")),
|
||||
team_alias=as_str(get("team_alias")),
|
||||
|
|
@ -351,6 +359,35 @@ def model_from_request_data(data: object) -> str | None:
|
|||
return None
|
||||
|
||||
|
||||
def metadata_from_request_data(data: object) -> Mapping[str, object] | None:
|
||||
"""The caller's ``requester_metadata`` snapshot from a pre-call ``data`` dict, keyed under its wrapper.
|
||||
|
||||
The proxy stores it under ``metadata`` or ``litellm_metadata`` depending on the route;
|
||||
the proxy-owned siblings (``user_api_key_*``, ``requester_ip_address``) are not read.
|
||||
"""
|
||||
top: Final = _as_str_mapping(data)
|
||||
if top is None:
|
||||
return None
|
||||
snapshots: Final = tuple(
|
||||
snapshot
|
||||
for name in ("metadata", "litellm_metadata")
|
||||
if (nested := _as_str_mapping(top.get(name))) is not None
|
||||
and (snapshot := _as_str_mapping(nested.get(REQUESTER_METADATA_KEY))) is not None
|
||||
)
|
||||
return MappingProxyType({REQUESTER_METADATA_KEY: snapshots[0]}) if snapshots else None
|
||||
|
||||
|
||||
def flatten_metadata(raw: Mapping[str, object]) -> Iterator[tuple[str, str]]:
|
||||
"""Scalar leaves of a nested metadata mapping, keyed by their dotted path."""
|
||||
stack: Final = list(tuple(raw.items())[::-1]) # mutable-ok: iterative worklist keeps the walk off the call stack
|
||||
while stack:
|
||||
key, value = stack.pop()
|
||||
if (nested := _as_str_mapping(value)) is not None:
|
||||
stack.extend(tuple((f"{key}.{sub_key}", sub_value) for sub_key, sub_value in nested.items())[::-1])
|
||||
elif isinstance(value, (str, bool, int, float)):
|
||||
yield key, str(value)
|
||||
|
||||
|
||||
def resolve_provider_model(payload: StandardLoggingPayload) -> str | None:
|
||||
"""The model litellm dispatched to the provider, from the payload.
|
||||
|
||||
|
|
|
|||
|
|
@ -9,13 +9,20 @@ frame itself fail, which is how a loud failure turns back into a silent one.
|
|||
"""
|
||||
|
||||
import json
|
||||
from typing import Final
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Protocol
|
||||
|
||||
from litellm.types.realtime import RealtimeErrorDetail, RealtimeErrorEvent
|
||||
|
||||
WEBSOCKET_CLOSE_REASON_MAX_BYTES: Final = 123
|
||||
|
||||
|
||||
class _ClientWebSocket(Protocol):
|
||||
async def send_text(self, data: str) -> None: ...
|
||||
|
||||
async def close(self, code: int = ..., reason: str | None = ...) -> None: ...
|
||||
|
||||
|
||||
def realtime_error_event(message: str, error_type: str) -> str:
|
||||
detail: Final[RealtimeErrorDetail] = {"type": error_type, "message": message}
|
||||
event: Final[RealtimeErrorEvent] = {"type": "error", "error": detail}
|
||||
|
|
@ -37,3 +44,28 @@ def client_close_code(upstream_code: int) -> int:
|
|||
if upstream_code in EXTERNAL_CLOSE_CODES or 3000 <= upstream_code < 5000:
|
||||
return upstream_code
|
||||
return int(CloseCode.INTERNAL_ERROR)
|
||||
|
||||
|
||||
def upstream_handshake_close_code(status_code: int) -> int:
|
||||
from websockets.frames import CloseCode
|
||||
|
||||
refusal_codes: Final = MappingProxyType(
|
||||
{
|
||||
401: int(CloseCode.POLICY_VIOLATION),
|
||||
403: int(CloseCode.POLICY_VIOLATION),
|
||||
429: int(CloseCode.TRY_AGAIN_LATER),
|
||||
}
|
||||
)
|
||||
return refusal_codes.get(status_code, int(CloseCode.INTERNAL_ERROR))
|
||||
|
||||
|
||||
async def close_after_upstream_handshake_refusal(websocket: _ClientWebSocket, status_code: int) -> None:
|
||||
message: Final = f"Upstream realtime handshake rejected with HTTP {status_code}"
|
||||
try:
|
||||
await websocket.send_text(realtime_error_event(message, error_type="server_error"))
|
||||
except Exception: # noqa: BLE001 # best-effort notice: a dead client socket must not skip the close below
|
||||
pass
|
||||
await websocket.close(
|
||||
code=upstream_handshake_close_code(status_code),
|
||||
reason=websocket_close_reason(message, fallback="Upstream handshake rejected"),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -8,11 +8,15 @@ from collections.abc import Mapping
|
|||
from types import MappingProxyType
|
||||
from typing import Any, Final, Protocol, cast
|
||||
|
||||
from litellm._logging import _redact_string, verbose_proxy_logger
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
|
||||
from litellm.types.realtime import RealtimeQueryParams
|
||||
|
||||
from ....litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
from ....litellm_core_utils.realtime_errors import (
|
||||
close_after_upstream_handshake_refusal,
|
||||
realtime_error_event,
|
||||
)
|
||||
from ....litellm_core_utils.realtime_streaming import (
|
||||
RealTimeStreaming,
|
||||
ScopedWebSocket,
|
||||
|
|
@ -49,7 +53,9 @@ def azure_realtime_protocol_for_client(
|
|||
|
||||
|
||||
class _ProxyClientWebSocket(Protocol):
|
||||
"""Client-facing websocket handle: this path only closes it after a failed handshake."""
|
||||
"""Client-facing websocket handle: this path only writes to it after a failed handshake."""
|
||||
|
||||
async def send_text(self, data: str) -> None: ...
|
||||
|
||||
async def close(self, code: int = ..., reason: str | None = ...) -> None: ...
|
||||
|
||||
|
|
@ -181,7 +187,16 @@ class AzureOpenAIRealtime(AzureChatCompletion):
|
|||
)
|
||||
await realtime_streaming.bidirectional_forward()
|
||||
|
||||
except websockets.exceptions.InvalidStatusCode as e:
|
||||
await websocket.close(code=e.status_code, reason=_redact_string(str(e)))
|
||||
except websockets.exceptions.InvalidStatus as e:
|
||||
verbose_proxy_logger.exception("Error in AzureOpenAIRealtime.async_realtime")
|
||||
await close_after_upstream_handshake_refusal(websocket, e.response.status_code)
|
||||
except Exception:
|
||||
verbose_proxy_logger.exception("Error in AzureOpenAIRealtime.async_realtime")
|
||||
try:
|
||||
await websocket.send_text(realtime_error_event("Internal server error", error_type="server_error"))
|
||||
except Exception: # noqa: BLE001 # best-effort notice: a dead client socket must not skip the close below
|
||||
pass
|
||||
try:
|
||||
await websocket.close(code=1011, reason="Internal server error")
|
||||
except Exception: # noqa: BLE001 # the lower layer may have closed the socket already; closing twice is not an error
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -30,7 +30,11 @@ from litellm.litellm_core_utils.audio_utils.subtitle_utils import (
|
|||
)
|
||||
from litellm.litellm_core_utils.get_litellm_params import AWS_CREDENTIAL_KWARGS_KEYS
|
||||
from litellm.litellm_core_utils.llm_request_utils import serialize_multipart_form_fields
|
||||
from litellm.litellm_core_utils.realtime_errors import realtime_error_event, websocket_close_reason
|
||||
from litellm.litellm_core_utils.realtime_errors import (
|
||||
close_after_upstream_handshake_refusal,
|
||||
realtime_error_event,
|
||||
websocket_close_reason,
|
||||
)
|
||||
from litellm.litellm_core_utils.realtime_streaming import RealTimeStreaming
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
from litellm.llms.base_llm.anthropic_messages.transformation import (
|
||||
|
|
@ -6228,9 +6232,9 @@ class BaseLLMHTTPHandler:
|
|||
|
||||
await realtime_streaming.bidirectional_forward()
|
||||
|
||||
except websockets.exceptions.InvalidStatusCode as e:
|
||||
except websockets.exceptions.InvalidStatus as e:
|
||||
verbose_logger.exception("Error connecting to backend: %s", e)
|
||||
await websocket.close(code=e.status_code, reason=_redact_string(str(e)))
|
||||
await close_after_upstream_handshake_refusal(websocket, e.response.status_code)
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Error connecting to backend: %s", e)
|
||||
redacted_error: Final = _redact_string(str(e))
|
||||
|
|
@ -6640,9 +6644,9 @@ class BaseLLMHTTPHandler:
|
|||
)
|
||||
await streaming.bidirectional_forward()
|
||||
|
||||
except websockets.exceptions.InvalidStatusCode as e:
|
||||
except websockets.exceptions.InvalidStatus as e:
|
||||
verbose_logger.exception("Error connecting to responses WS backend: %s", e)
|
||||
await websocket.close(code=e.status_code, reason=_redact_string(str(e)))
|
||||
await close_after_upstream_handshake_refusal(websocket, e.response.status_code)
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Error in responses WS: %s", e)
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
|
|||
from litellm.types.realtime import RealtimeQueryParams
|
||||
|
||||
from ....litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
from ....litellm_core_utils.realtime_errors import close_after_upstream_handshake_refusal
|
||||
from ....litellm_core_utils.realtime_streaming import (
|
||||
RealtimeEventNormalizer,
|
||||
RealTimeStreaming,
|
||||
|
|
@ -175,8 +176,8 @@ class OpenAIRealtime(OpenAIChatCompletion):
|
|||
)
|
||||
await realtime_streaming.bidirectional_forward()
|
||||
|
||||
except websockets.exceptions.InvalidStatusCode as e:
|
||||
await websocket.close(code=e.status_code, reason=_redact_string(str(e)))
|
||||
except websockets.exceptions.InvalidStatus as e:
|
||||
await close_after_upstream_handshake_refusal(websocket, e.response.status_code)
|
||||
except Exception as e:
|
||||
try:
|
||||
await websocket.close(code=1011, reason=_redact_string(f"Internal server error: {e}"))
|
||||
|
|
|
|||
|
|
@ -294,6 +294,7 @@ from litellm.litellm_core_utils.core_helpers import (
|
|||
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.realtime_errors import (
|
||||
close_after_upstream_handshake_refusal,
|
||||
realtime_error_event,
|
||||
websocket_close_reason,
|
||||
)
|
||||
|
|
@ -12026,9 +12027,9 @@ async def realtime_websocket_endpoint(
|
|||
user_model=user_model,
|
||||
)
|
||||
await llm_call
|
||||
except websockets.exceptions.InvalidStatusCode as e:
|
||||
except websockets.exceptions.InvalidStatus as e:
|
||||
verbose_proxy_logger.exception("Invalid status code")
|
||||
await websocket.close(code=e.status_code, reason="Invalid status code")
|
||||
await close_after_upstream_handshake_refusal(websocket, e.response.status_code)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Internal server error")
|
||||
redacted_error: Final = _redact_string(str(e))
|
||||
|
|
|
|||
|
|
@ -286,7 +286,7 @@ def function_call_item(events: tuple[ReceivedEvent, ...]) -> OutputItem | None:
|
|||
# ---- session + client --------------------------------------------------
|
||||
|
||||
|
||||
def _as_text(message: str | bytes) -> str:
|
||||
def as_text(message: str | bytes) -> str:
|
||||
return message.decode("utf-8") if isinstance(message, bytes) else message
|
||||
|
||||
|
||||
|
|
@ -304,7 +304,7 @@ class RealtimeSession:
|
|||
collected: list[ReceivedEvent] = []
|
||||
while time.monotonic() < deadline:
|
||||
try:
|
||||
text = _as_text(
|
||||
text = as_text(
|
||||
self.connection.recv(timeout=deadline - time.monotonic())
|
||||
)
|
||||
except TimeoutError:
|
||||
|
|
|
|||
|
|
@ -13,8 +13,9 @@ failure. See REALTIME_COVERAGE_MATRIX.md.
|
|||
"""
|
||||
|
||||
import pytest
|
||||
from lifecycle import ResourceManager
|
||||
from models import LiteLLMParamsBody
|
||||
from pydantic import BaseModel
|
||||
|
||||
from realtime_client import (
|
||||
PROVIDERS,
|
||||
ConversationItemCreate,
|
||||
|
|
@ -27,14 +28,17 @@ from realtime_client import (
|
|||
RealtimeProvider,
|
||||
ResponseCreate,
|
||||
ResponseDone,
|
||||
ServerEnvelope,
|
||||
SessionConfig,
|
||||
SessionUpdate,
|
||||
as_text,
|
||||
function_call_item,
|
||||
parse_last,
|
||||
realtime_model,
|
||||
transcript,
|
||||
user_message,
|
||||
)
|
||||
from websockets.exceptions import ConnectionClosedError
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
|
|
@ -147,3 +151,39 @@ def test_tool_call_round_trip(
|
|||
second = session.collect_until("response.done", timeout=60)
|
||||
|
||||
assert "72" in transcript(second), "follow-up did not use the tool result"
|
||||
|
||||
|
||||
_REFUSED_UPSTREAMS = (
|
||||
RealtimeProvider(
|
||||
"azure-bad-key",
|
||||
"azure-realtime-refused",
|
||||
LiteLLMParamsBody(
|
||||
model="azure/gpt-realtime",
|
||||
api_key="invalid-e2e-key",
|
||||
api_version="2025-08-28",
|
||||
realtime_protocol="GA",
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("provider", _REFUSED_UPSTREAMS, ids=[p.id for p in _REFUSED_UPSTREAMS])
|
||||
def test_upstream_handshake_refusal_is_an_error_event_and_policy_close(
|
||||
client: RealtimeClient,
|
||||
resources: ResourceManager,
|
||||
scoped_key: str,
|
||||
provider: RealtimeProvider,
|
||||
) -> None:
|
||||
model_name, model_id = client.provision(provider)
|
||||
resources.defer(lambda: client.proxy.delete_model(model_id))
|
||||
|
||||
with client.connect(key=scoped_key, model=model_name) as session:
|
||||
first = ServerEnvelope.model_validate_json(
|
||||
as_text(session.connection.recv(timeout=15))
|
||||
)
|
||||
assert first.type == "error", first
|
||||
with pytest.raises(ConnectionClosedError) as closed:
|
||||
session.connection.recv(timeout=15)
|
||||
|
||||
assert closed.value.rcvd is not None
|
||||
assert closed.value.rcvd.code == 1008, closed.value
|
||||
|
|
|
|||
|
|
@ -168,6 +168,46 @@ def test_allowlisted_metadata_subkey_promoted_blob_excluded():
|
|||
assert all("private_note" not in k for k in span.attributes)
|
||||
|
||||
|
||||
def test_nested_metadata_key_promoted_under_caller_path():
|
||||
"""A dotted allowlist entry reads the nested caller metadata the proxy stores
|
||||
under ``requester_metadata`` and lands on the LLM-call span under the caller's
|
||||
own path (``litellm.metadata.trace_id``, ``litellm.metadata.nested.deep``);
|
||||
a pre-existing flat dotted key keeps its full name, and unlisted siblings and
|
||||
the blob stay out."""
|
||||
engine, exporter = _engine_and_exporter()
|
||||
payload = _payload()
|
||||
payload["metadata"]["a.b"] = "flat"
|
||||
payload["metadata"]["requester_metadata"] = {
|
||||
"trace_id": "abc",
|
||||
"attempt": 0,
|
||||
"empty": "",
|
||||
"nested": {"deep": "x", "skipped": "y"},
|
||||
}
|
||||
data = LLMCallSpanData.from_standard_logging_payload(payload)
|
||||
bag = promoted_baggage(
|
||||
data.identity,
|
||||
data.request_model,
|
||||
BAGGAGE_PROMOTED_KEYS,
|
||||
metadata_keys=(
|
||||
"requester_metadata.trace_id",
|
||||
"requester_metadata.attempt",
|
||||
"requester_metadata.empty",
|
||||
"requester_metadata.nested.deep",
|
||||
"a.b",
|
||||
),
|
||||
)
|
||||
engine.emit(SpanRole.LLM_CALL, data, ctx_mod.set_request_baggage(bag))
|
||||
(span,) = exporter.get_finished_spans()
|
||||
assert span.attributes[f"{LiteLLM.METADATA_PREFIX}trace_id"] == "abc"
|
||||
assert span.attributes[f"{LiteLLM.METADATA_PREFIX}attempt"] == "0"
|
||||
assert span.attributes[f"{LiteLLM.METADATA_PREFIX}nested.deep"] == "x"
|
||||
assert span.attributes[f"{LiteLLM.METADATA_PREFIX}a.b"] == "flat"
|
||||
assert f"{LiteLLM.METADATA_PREFIX}empty" not in span.attributes
|
||||
assert f"{LiteLLM.METADATA_PREFIX}deep" not in span.attributes
|
||||
assert f"{LiteLLM.METADATA_PREFIX}nested.skipped" not in span.attributes
|
||||
assert not any(k.startswith(f"{LiteLLM.METADATA_PREFIX}requester_metadata") for k in span.attributes)
|
||||
|
||||
|
||||
def test_http_attributes_never_promoted():
|
||||
"""Even if http.* is present in baggage, the processor must not stamp it on
|
||||
child spans (it belongs on the SERVER span only)."""
|
||||
|
|
|
|||
|
|
@ -1623,17 +1623,22 @@ def test_provider_model_and_team_metadata_on_real_boundary_flow():
|
|||
def test_pre_call_hook_seeds_baggage_onto_server_and_child_spans():
|
||||
"""The pre-call hook seeds identity Baggage in the request context so the
|
||||
server span (stamped directly) AND later child spans (service here, via the
|
||||
Baggage processor) carry identity — not just the LLM-call span."""
|
||||
Baggage processor) carry identity — not just the LLM-call span. Only the
|
||||
caller's ``requester_metadata`` is read from the request dict, so a proxy-owned
|
||||
sibling such as ``requester_ip_address`` is not stamped from here even though
|
||||
the default allowlist names it, and an unlisted caller key is not promoted."""
|
||||
logger, exporter = _logger()
|
||||
server = logger._emitter.start_span(
|
||||
SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME
|
||||
)
|
||||
data = {
|
||||
"model": "gpt-4o",
|
||||
"metadata": {"requester_ip_address": "127.0.0.1", "requester_metadata": {"trace_id": "abc"}},
|
||||
}
|
||||
|
||||
async def _flow():
|
||||
# pre-call seeds baggage + stamps the active server span
|
||||
await logger.async_pre_call_hook(
|
||||
_Auth(), None, {"model": "gpt-4o"}, "completion"
|
||||
)
|
||||
await logger.async_pre_call_hook(_Auth(), None, data, "completion")
|
||||
# a later service call (same task) must inherit the identity
|
||||
await logger.async_service_success_hook(
|
||||
payload=_ServicePayload("redis", "set"), parent_otel_span=server
|
||||
|
|
@ -1653,6 +1658,46 @@ def test_pre_call_hook_seeds_baggage_onto_server_and_child_spans():
|
|||
srv.attributes[LiteLLM.TEAM_ID] == "t1"
|
||||
) # stamped directly on the server span
|
||||
assert srv.attributes[f"{LiteLLM.METADATA_PREFIX}user_api_key_user_id"] == "u1"
|
||||
assert not any(
|
||||
k in (f"{LiteLLM.METADATA_PREFIX}requester_ip_address", f"{LiteLLM.METADATA_PREFIX}trace_id")
|
||||
for s in (redis, srv)
|
||||
for k in s.attributes
|
||||
)
|
||||
|
||||
|
||||
def test_pre_call_hook_promotes_nested_request_metadata_key():
|
||||
"""``baggage_metadata_keys: [requester_metadata.trace_id]`` reads the caller's
|
||||
``metadata.trace_id`` (snapshotted by the proxy under ``requester_metadata``)
|
||||
and stamps ``litellm.metadata.trace_id`` on the server, LLM-call and service
|
||||
spans of the request; unlisted siblings are not promoted."""
|
||||
cfg = OpenTelemetryV2Config(exporter="in_memory", baggage_metadata_keys=["requester_metadata.trace_id"])
|
||||
exporter = InMemorySpanExporter()
|
||||
logger = OpenTelemetryV2(config=cfg, tracer_provider=providers.build_tracer_provider(cfg, exporter=exporter))
|
||||
server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME)
|
||||
data = {"model": "gpt-4o", "metadata": {"requester_metadata": {"trace_id": "abc", "nested": {"deep": "x"}}}}
|
||||
kwargs = _kwargs()
|
||||
|
||||
async def _flow():
|
||||
await logger.async_pre_call_hook(_Auth(), None, data, "completion")
|
||||
logger.log_pre_api_call(model="gpt-4o", messages=[], kwargs=kwargs)
|
||||
await logger.async_log_success_event(kwargs, None, None, None)
|
||||
await logger.async_service_success_hook(payload=_ServicePayload("redis", "set"), parent_otel_span=server)
|
||||
|
||||
with trace.use_span(server, end_on_exit=False):
|
||||
asyncio.run(_flow())
|
||||
server.end()
|
||||
|
||||
spans = {s.name: s for s in exporter.get_finished_spans()}
|
||||
key = f"{LiteLLM.METADATA_PREFIX}trace_id"
|
||||
assert spans[LITELLM_PROXY_REQUEST_SPAN_NAME].attributes[key] == "abc"
|
||||
assert spans["chat gpt-4o"].attributes[key] == "abc"
|
||||
assert spans["redis set"].attributes[key] == "abc"
|
||||
assert data == {"model": "gpt-4o", "metadata": {"requester_metadata": {"trace_id": "abc", "nested": {"deep": "x"}}}}
|
||||
assert not any(
|
||||
k.startswith(f"{LiteLLM.METADATA_PREFIX}requester_metadata") or k.endswith("deep")
|
||||
for s in spans.values()
|
||||
for k in s.attributes
|
||||
)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
|
|
|||
|
|
@ -5581,6 +5581,38 @@ class TestOpenTelemetryInferenceIdentityAttributes(unittest.TestCase):
|
|||
otel.set_attributes(span, kwargs, {"model": "azure/gpt-4o"})
|
||||
assert "http.route" not in self._attr(span, exp)
|
||||
|
||||
def test_nested_metadata_key_promoted_under_caller_path(self):
|
||||
"""``baggage_metadata_keys: [requester_metadata.trace_id]`` stamps the
|
||||
caller's nested metadata value as ``litellm.metadata.trace_id`` and a deeper
|
||||
path keeps its dotted name; unlisted siblings stay inside the
|
||||
``metadata.requester_metadata`` blob."""
|
||||
otel = OpenTelemetry(
|
||||
config=OpenTelemetryConfig(
|
||||
baggage_metadata_keys=["requester_metadata.trace_id", "requester_metadata.nested.deep"]
|
||||
)
|
||||
)
|
||||
kwargs = self._kwargs()
|
||||
kwargs["standard_logging_object"]["metadata"]["requester_metadata"] = {
|
||||
"trace_id": "abc",
|
||||
"nested": {"deep": "x", "skipped": "y"},
|
||||
}
|
||||
span, exp = self._span()
|
||||
otel.set_attributes(span, kwargs, {"model": "azure/gpt-4o"})
|
||||
attrs = self._attr(span, exp)
|
||||
assert attrs["litellm.metadata.trace_id"] == "abc"
|
||||
assert attrs["litellm.metadata.nested.deep"] == "x"
|
||||
assert "litellm.metadata.deep" not in attrs
|
||||
assert "litellm.metadata.nested.skipped" not in attrs
|
||||
assert not any(k.startswith("litellm.metadata.requester_metadata") for k in attrs)
|
||||
|
||||
def test_metadata_keys_default_to_none_promoted(self):
|
||||
otel = OpenTelemetry()
|
||||
kwargs = self._kwargs()
|
||||
kwargs["standard_logging_object"]["metadata"]["requester_metadata"] = {"trace_id": "abc"}
|
||||
span, exp = self._span()
|
||||
otel.set_attributes(span, kwargs, {"model": "azure/gpt-4o"})
|
||||
assert not any(k.startswith("litellm.metadata.") for k in self._attr(span, exp))
|
||||
|
||||
def test_team_metadata_json_helper(self):
|
||||
keys = ["a", "b"]
|
||||
assert OpenTelemetry._team_metadata_json(None, keys) is None
|
||||
|
|
@ -5631,6 +5663,11 @@ class TestOpenTelemetryTeamMetadataKeysConfig(unittest.TestCase):
|
|||
cfg = OpenTelemetryConfig(baggage_team_metadata_keys=["from_arg"])
|
||||
assert cfg.baggage_team_metadata_keys == ["from_arg"]
|
||||
|
||||
def test_metadata_keys_from_kwargs_and_env(self):
|
||||
with patch.dict("os.environ", {"LITELLM_OTEL_BAGGAGE_METADATA_KEYS": "requester_metadata.trace_id, a.b"}):
|
||||
assert OpenTelemetryConfig().baggage_metadata_keys == ["requester_metadata.trace_id", "a.b"]
|
||||
assert OpenTelemetry(baggage_metadata_keys="x.y").config.baggage_metadata_keys == ["x.y"]
|
||||
|
||||
|
||||
class TestOpenTelemetryMetricAttributeFiltering(unittest.TestCase):
|
||||
"""LIT-3600: include/exclude control over which attributes are stamped on
|
||||
|
|
|
|||
|
|
@ -1,13 +1,17 @@
|
|||
import json
|
||||
from typing import cast
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.litellm_core_utils.realtime_errors import (
|
||||
WEBSOCKET_CLOSE_REASON_MAX_BYTES,
|
||||
client_close_code,
|
||||
close_after_upstream_handshake_refusal,
|
||||
realtime_error_event,
|
||||
upstream_handshake_close_code,
|
||||
websocket_close_reason,
|
||||
)
|
||||
from litellm.types.realtime import RealtimeErrorEvent
|
||||
|
||||
|
||||
def test_realtime_error_event_shape():
|
||||
|
|
@ -52,3 +56,52 @@ def test_websocket_close_reason_truncates_multibyte_message_by_bytes():
|
|||
)
|
||||
def test_client_close_code_only_forwards_codes_a_server_may_send(upstream_code, expected):
|
||||
assert client_close_code(upstream_code) == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("status_code", "expected"),
|
||||
[(401, 1008), (403, 1008), (429, 1013), (500, 1011)],
|
||||
)
|
||||
def test_upstream_handshake_close_code_maps_http_status_to_close_code(status_code: int, expected: int):
|
||||
assert upstream_handshake_close_code(status_code) == expected
|
||||
|
||||
|
||||
class _RecordingWebSocket:
|
||||
def __init__(self) -> None:
|
||||
self.sent: list[str] = []
|
||||
self.closed: tuple[int, str | None] | None = None
|
||||
|
||||
async def send_text(self, data: str) -> None:
|
||||
self.sent.append(data)
|
||||
|
||||
async def close(self, code: int = 1000, reason: str | None = None) -> None:
|
||||
self.closed = (code, reason)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_close_after_upstream_handshake_refusal_sends_error_event_then_policy_close():
|
||||
websocket = _RecordingWebSocket()
|
||||
|
||||
await close_after_upstream_handshake_refusal(websocket, 401)
|
||||
|
||||
assert len(websocket.sent) == 1
|
||||
event = cast(RealtimeErrorEvent, json.loads(websocket.sent[0]))
|
||||
assert event["type"] == "error"
|
||||
assert event["error"]["type"] == "server_error"
|
||||
assert "401" in event["error"]["message"]
|
||||
assert websocket.closed is not None
|
||||
assert websocket.closed[0] == 1008
|
||||
assert websocket.closed[1]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_close_after_upstream_handshake_refusal_still_closes_when_send_fails():
|
||||
class _DeadWebSocket(_RecordingWebSocket):
|
||||
async def send_text(self, data: str) -> None:
|
||||
raise RuntimeError("socket gone")
|
||||
|
||||
websocket = _DeadWebSocket()
|
||||
|
||||
await close_after_upstream_handshake_refusal(websocket, 500)
|
||||
|
||||
assert websocket.closed == (1011, "Upstream realtime handshake rejected with HTTP 500")
|
||||
|
|
|
|||
1
tests/test_litellm/llms/azure/realtime/__init__.py
Normal file
1
tests/test_litellm/llms/azure/realtime/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
|
||||
86
tests/test_litellm/llms/azure/realtime/test_handler.py
Normal file
86
tests/test_litellm/llms/azure/realtime/test_handler.py
Normal file
|
|
@ -0,0 +1,86 @@
|
|||
import json
|
||||
from typing import cast
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
class _RecordingClientWebSocket:
|
||||
scope: dict[str, list[tuple[bytes, bytes]]] = {"headers": []}
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.sent: list[str] = []
|
||||
self.closed: list[tuple[int, str | None]] = []
|
||||
|
||||
async def send_text(self, data: str) -> None:
|
||||
self.sent.append(data)
|
||||
|
||||
async def close(self, code: int = 1000, reason: str | None = None) -> None:
|
||||
self.closed.append((code, reason))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_realtime_upstream_handshake_refusal_sends_error_event_then_policy_close():
|
||||
from websockets.datastructures import Headers
|
||||
from websockets.exceptions import InvalidStatus
|
||||
from websockets.http11 import Response
|
||||
|
||||
from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime
|
||||
from litellm.types.realtime import RealtimeErrorEvent
|
||||
|
||||
handler = AzureOpenAIRealtime()
|
||||
model = "gpt-realtime"
|
||||
|
||||
dummy_websocket = _RecordingClientWebSocket()
|
||||
dummy_logging_obj = MagicMock()
|
||||
|
||||
refused = InvalidStatus(Response(401, "Unauthorized", Headers()))
|
||||
|
||||
with patch("websockets.connect", side_effect=refused):
|
||||
await handler.async_realtime( # pyright: ignore[reportUnknownMemberType] # handler's websocket param is a Protocol here but the mock connect type is incomplete
|
||||
model=model,
|
||||
websocket=dummy_websocket,
|
||||
logging_obj=dummy_logging_obj,
|
||||
api_base="https://example.openai.azure.com",
|
||||
api_key="bad-key",
|
||||
api_version="2025-08-28",
|
||||
query_params={"model": model},
|
||||
)
|
||||
|
||||
assert len(dummy_websocket.sent) == 1
|
||||
event = cast(RealtimeErrorEvent, json.loads(dummy_websocket.sent[0]))
|
||||
assert event["type"] == "error"
|
||||
assert event["error"]["type"] == "server_error"
|
||||
assert "401" in event["error"]["message"]
|
||||
assert dummy_websocket.closed and dummy_websocket.closed[0][0] == 1008
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_realtime_unexpected_error_sends_error_event_then_internal_close():
|
||||
from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime
|
||||
from litellm.types.realtime import RealtimeErrorEvent
|
||||
|
||||
handler = AzureOpenAIRealtime()
|
||||
model = "gpt-realtime"
|
||||
|
||||
dummy_websocket = _RecordingClientWebSocket()
|
||||
dummy_logging_obj = MagicMock()
|
||||
|
||||
with patch("websockets.connect", side_effect=OSError("connection reset")):
|
||||
await handler.async_realtime( # pyright: ignore[reportUnknownMemberType] # same as above
|
||||
model=model,
|
||||
websocket=dummy_websocket,
|
||||
logging_obj=dummy_logging_obj,
|
||||
api_base="https://example.openai.azure.com",
|
||||
api_key="bad-key",
|
||||
api_version="2025-08-28",
|
||||
query_params={"model": model},
|
||||
)
|
||||
|
||||
assert len(dummy_websocket.sent) == 1
|
||||
event = cast(RealtimeErrorEvent, json.loads(dummy_websocket.sent[0]))
|
||||
assert event["type"] == "error"
|
||||
assert event["error"]["type"] == "server_error"
|
||||
assert event["error"]["message"] == "Internal server error"
|
||||
assert "connection reset" not in dummy_websocket.sent[0]
|
||||
assert dummy_websocket.closed and dummy_websocket.closed[0] == (1011, "Internal server error")
|
||||
|
|
@ -416,3 +416,52 @@ async def test_async_realtime_ws_url_has_no_ssl():
|
|||
|
||||
# Verify ssl is None for ws:// URLs (the fix for issue #19222)
|
||||
assert called_kwargs["ssl"] is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_realtime_upstream_handshake_refusal_sends_error_event_then_policy_close():
|
||||
from typing import cast
|
||||
|
||||
from websockets.datastructures import Headers
|
||||
from websockets.exceptions import InvalidStatus
|
||||
from websockets.http11 import Response
|
||||
|
||||
from litellm.llms.openai.realtime.handler import OpenAIRealtime
|
||||
from litellm.types.realtime import RealtimeErrorEvent
|
||||
|
||||
handler = OpenAIRealtime()
|
||||
model = "gpt-realtime"
|
||||
|
||||
sent: list[str] = []
|
||||
closed: list[tuple[int, str | None]] = []
|
||||
|
||||
class RecordingClientWebSocket:
|
||||
scope: dict[str, list[tuple[bytes, bytes]]] = {"headers": []}
|
||||
|
||||
async def send_text(self, data: str) -> None:
|
||||
sent.append(data)
|
||||
|
||||
async def close(self, code: int = 1000, reason: str | None = None) -> None:
|
||||
closed.append((code, reason))
|
||||
|
||||
dummy_websocket = RecordingClientWebSocket()
|
||||
dummy_logging_obj = MagicMock()
|
||||
|
||||
refused = InvalidStatus(Response(401, "Unauthorized", Headers()))
|
||||
|
||||
with patch("websockets.connect", side_effect=refused):
|
||||
await handler.async_realtime( # pyright: ignore[reportUnknownMemberType] # handler's websocket param is Any
|
||||
model=model,
|
||||
websocket=dummy_websocket,
|
||||
logging_obj=dummy_logging_obj,
|
||||
api_base="https://api.openai.com/",
|
||||
api_key="bad-key",
|
||||
query_params={"model": model},
|
||||
)
|
||||
|
||||
assert len(sent) == 1
|
||||
event = cast(RealtimeErrorEvent, json.loads(sent[0]))
|
||||
assert event["type"] == "error"
|
||||
assert event["error"]["type"] == "server_error"
|
||||
assert "401" in event["error"]["message"]
|
||||
assert closed and closed[0][0] == 1008
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
Tests for _redact_string usage in error/logging paths.
|
||||
|
||||
Covers actual execution of redaction in:
|
||||
- WebSocket close reasons in realtime handlers (openai, azure, bedrock)
|
||||
- WebSocket close reasons in realtime handlers (openai, bedrock)
|
||||
- Gemini RAG ingestion x-goog-api-key header usage
|
||||
- Traceback redaction pattern used in proxy streaming
|
||||
- Router fallback-failure traceback redaction
|
||||
|
|
@ -72,25 +72,6 @@ class TestOpenAIRealtimeRedaction:
|
|||
api_key="test-key",
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalid_status_code_redacts_reason(self):
|
||||
import websockets.exceptions
|
||||
|
||||
from litellm.llms.openai.realtime.handler import OpenAIRealtime
|
||||
|
||||
handler = OpenAIRealtime()
|
||||
exc = websockets.exceptions.InvalidStatusCode(403, None)
|
||||
exc.status_code = 403
|
||||
|
||||
kwargs = self._call_kwargs()
|
||||
mock_ws = kwargs["websocket"]
|
||||
p1, p2, p3 = self._make_patches(handler)
|
||||
with p1, p2, p3, patch("websockets.connect", side_effect=exc):
|
||||
await handler.async_realtime(**kwargs)
|
||||
|
||||
mock_ws.close.assert_called_once()
|
||||
assert mock_ws.close.call_args[1]["code"] == 403
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generic_exception_redacts_reason(self):
|
||||
from litellm.llms.openai.realtime.handler import OpenAIRealtime
|
||||
|
|
@ -111,41 +92,6 @@ class TestOpenAIRealtimeRedaction:
|
|||
assert "sk-1234567890abcdefghij" not in mock_ws.close.call_args[1]["reason"]
|
||||
|
||||
|
||||
class TestAzureRealtimeRedaction:
|
||||
"""Test that Azure realtime handler redacts secrets in websocket close reasons."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalid_status_code_redacts_reason(self):
|
||||
import websockets.exceptions
|
||||
|
||||
from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime
|
||||
|
||||
handler = AzureOpenAIRealtime()
|
||||
mock_ws = AsyncMock()
|
||||
exc = websockets.exceptions.InvalidStatusCode(403, None)
|
||||
exc.status_code = 403
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
handler,
|
||||
"_construct_url",
|
||||
return_value="wss://test.openai.azure.com/openai/realtime",
|
||||
),
|
||||
patch("websockets.connect", side_effect=exc),
|
||||
):
|
||||
await handler.async_realtime(
|
||||
model="gpt-4",
|
||||
websocket=mock_ws,
|
||||
logging_obj=MagicMock(),
|
||||
api_base="https://test.openai.azure.com/",
|
||||
api_key="test-key",
|
||||
api_version="2024-10-01-preview",
|
||||
)
|
||||
|
||||
mock_ws.close.assert_called_once()
|
||||
assert mock_ws.close.call_args[1]["code"] == 403
|
||||
|
||||
|
||||
class TestBedrockRealtimeRedaction:
|
||||
"""Test that _redact_string produces safe close reasons for Bedrock-style errors."""
|
||||
|
||||
|
|
|
|||
12
uv.lock
generated
12
uv.lock
generated
|
|
@ -315,16 +315,16 @@ vertex = [
|
|||
|
||||
[[package]]
|
||||
name = "anyio"
|
||||
version = "4.13.0"
|
||||
version = "4.14.2"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "exceptiongroup", marker = "python_full_version < '3.11'" },
|
||||
{ name = "idna" },
|
||||
{ name = "typing-extensions", marker = "python_full_version < '3.13'" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/19/14/2c5dd9f512b66549ae92767a9c7b330ae88e1932ca57876909410251fe13/anyio-4.13.0.tar.gz", hash = "sha256:334b70e641fd2221c1505b3890c69882fe4a2df910cba14d97019b90b24439dc", size = 231622, upload-time = "2026-03-24T12:59:09.671Z" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/61/cc/a381afa6efea9f496eff839d4a6a1aed3bfafc7b3ab4b0d1b243a12573dd/anyio-4.14.2.tar.gz", hash = "sha256:cfa139f3ed1a23ee8f88a145ddb5ac7605b8bbfd8592baacd7ce3d8bb4313c7f", size = 260176, upload-time = "2026-07-12T20:29:07.082Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/da/42/e921fccf5015463e32a3cf6ee7f980a6ed0f395ceeaa45060b61d86486c2/anyio-4.13.0-py3-none-any.whl", hash = "sha256:08b310f9e24a9594186fd75b4f73f4a4152069e3853f1ed8bfbf58369f4ad708", size = 114353, upload-time = "2026-03-24T12:59:08.246Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/da/35/f2287558c17e29fafc8ef3daf819bb9834061cfa43bff8014f7df7f63bdc/anyio-4.14.2-py3-none-any.whl", hash = "sha256:9f505dda5ac9f0c8309b5e8bd445a8c2bf7246f3ce950121e45ea15bc41d1494", size = 125813, upload-time = "2026-07-12T20:29:05.763Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -9080,11 +9080,11 @@ wheels = [
|
|||
|
||||
[[package]]
|
||||
name = "soupsieve"
|
||||
version = "2.8.4"
|
||||
version = "2.9"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/47/2c/0a5f6f8ee0d5589e48c7640213ed5175d52cf540a06725b628cc1a45d6ce/soupsieve-2.8.4.tar.gz", hash = "sha256:e121fd02e975c695e4e9e8774a5ee35d74714b59307868dcc5319ad2d9e3328e", size = 121110, upload-time = "2026-05-24T13:55:57.154Z" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/80/f1/93422647dd7e461f23d254e6b2bfa687a85b53aeb4903fcdbb74474d4584/soupsieve-2.9.tar.gz", hash = "sha256:acee8417325c5653e1377dc31eccad59eb82cbc65942afe6174c53b3aaad63fc", size = 122122, upload-time = "2026-07-19T01:35:18.425Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/5e/f5/0c41cb68dcae6b7de4fac4188a3a9589e21fb31df21ea3a2e888db95e6c9/soupsieve-2.8.4-py3-none-any.whl", hash = "sha256:e7e6b0769c8f51ed59acab6e994b00621096cfb1c640a7509295987388fbaf65", size = 37304, upload-time = "2026-05-24T13:55:55.406Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/7b/d6/3185ab5ad1280319b31986898f3206dd7227cd75e293d4dba2a5e6bf27a0/soupsieve-2.9-py3-none-any.whl", hash = "sha256:a2b2c76d67df2382d245409fd71e321a571717e58463efa32ace87dcadac2c12", size = 37387, upload-time = "2026-07-19T01:35:17.106Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue