From 2453936a82d300af31cb9ed9c4b6dae0208627b0 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Thu, 4 Jun 2026 00:18:35 +0530 Subject: [PATCH 01/11] Litellm websocket improvements (#29563) * Add support for websocket via codex * Add model alias and creds support * fix: skip cost tracking for WS session wrapper call types The @client decorator on _aresponses_websocket fires async_success_handler with result=None after the session ends. This triggered cost tracking errors because standard_logging_object is never built for None results. Per-turn costs are correctly tracked by individual litellm.aresponses calls inside the session. The outer session-level logging obj should not attempt cost tracking. Fix: skip _aresponses_websocket and _arealtime call types in deployment_callback_on_success, RouterBudgetLimiting.async_log_success_event, and _PROXY_track_cost_callback. * fix: address Greptile review comments Fix JSON injection: use json.dumps instead of f-string interpolation for model name in WS body. Add 30s timeout for first WS frame to prevent unbounded connection resource tie-up. Restore per-event model override in streaming_iterator; fall back to connection-level model when event omits it. Strengthen regression test: inject alias into kwargs via _update_kwargs_with_deployment mock so the test would fail on un-fixed code. * fix: handle nested response.create format in first-frame model extraction When ?model= is omitted, the first WS frame can carry the model in either flat format (first_event["model"]) or nested format (first_event["response"]["model"]). The flat-only check would silently reject clients using the nested wire format. Mirrors the same two-format logic in _build_base_call_kwargs. * fix: don't force connection-level custom_llm_provider on per-event model overrides If a client sends a different model per response.create turn, litellm needs to re-resolve the provider from that model string. Forcing the connection-level custom_llm_provider would silently route the request to the wrong backend. Only inject custom_llm_provider when the per-event model matches the connection-level model. * refactor: extract WS model extraction into testable function Pull the flat/nested model extraction into _extract_model_from_first_ws_event so tests import and exercise the real function rather than a copy. * fix: compare providers not full model strings in _inject_credentials The model == self.model guard was too strict: same-provider model variants (e.g., vertex_ai/gemini-2.0 -> vertex_ai/gemini-1.5 on one connection) would lose custom_llm_provider, breaking routing when a custom api_base is in use. Compare the provider extracted by get_llm_provider instead, so same-provider variants still inherit the connection-level provider while cross-provider overrides let litellm re-resolve. * style: black formatting * refactor: extract first-frame model resolution to fix PLR0915 (too many statements) * Fix responses WebSocket first-frame validation * fix: classify WS first-frame read errors and clarify cost-skip log Distinguish client disconnects from server errors when reading the responses WebSocket first frame, make the cost-tracking skip log message accurate for session wrappers (which do carry a model), and resolve the connection-level provider once per session instead of on every response.create event. * test: cover WS first-frame read errors and same-provider credential injection Adds regression tests for the still-uncovered responses WebSocket paths: the timeout, invalid-JSON and missing-model branches of _read_ws_model_from_first_frame, plus the provider comparison in ManagedResponsesWebSocketHandler._same_provider and _inject_credentials (same-provider model variants keep the connection provider; cross-provider models re-resolve). * fix(responses-ws): fall back to explicit custom_llm_provider when connection model is unresolvable When a WebSocket session is opened with a custom deployment alias that litellm cannot resolve to a provider, _connection_provider was None, so _same_provider returned False for every resolvable per-event model and the connection-level custom_llm_provider was dropped. Use the explicitly-set custom_llm_provider as the connection provider in that case so same-provider per-event models still inherit it while genuinely cross-provider models continue to re-resolve. --------- Co-authored-by: Cursor Agent Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/llms/custom_httpx/llm_http_handler.py | 3 + .../proxy/hooks/proxy_track_cost_callback.py | 10 +- .../proxy/response_api_endpoints/endpoints.py | 171 +++++- litellm/responses/streaming_iterator.py | 52 +- litellm/router.py | 5 +- litellm/router_strategy/budget_limiter.py | 3 + .../response_api_endpoints/test_endpoints.py | 515 ++++++++++++++++++ tests/test_litellm/test_router.py | 55 ++ 8 files changed, 796 insertions(+), 18 deletions(-) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 8b3add5398b..5c502c56ffe 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -5528,6 +5528,7 @@ class BaseLLMHTTPHandler: user_api_key_dict: Optional[Any] = None, litellm_metadata: Optional[Dict[str, Any]] = None, custom_llm_provider: Optional[str] = None, + first_message: Optional[str] = None, **kwargs: Any, ): """ @@ -5559,6 +5560,7 @@ class BaseLLMHTTPHandler: api_base=api_base, timeout=timeout, custom_llm_provider=custom_llm_provider, + first_message=first_message, **kwargs, ) await handler.run() @@ -5624,6 +5626,7 @@ class BaseLLMHTTPHandler: logging_obj=logging_obj, user_api_key_dict=user_api_key_dict, request_data=_request_data, + first_message=first_message, ) await streaming.bidirectional_forward() diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 3688f25ac44..b4a4fd571d0 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -285,9 +285,15 @@ class _ProxyDBLogger(CustomLogger): await _release_budget_reservation(budget_reservation=budget_reservation) # Non-model call types (health checks, afile_delete) have no model or standard_logging_object. # Use .get() for "stream" to avoid KeyError on health checks. - if sl_object is None and not kwargs.get("model"): + # WS session wrappers (_aresponses_websocket, _arealtime) also reach here with + # result=None; their per-turn costs are tracked on the inner aresponses/realtime calls. + if sl_object is None and ( + not kwargs.get("model") + or kwargs.get("call_type") + in ("_aresponses_websocket", "_arealtime") + ): verbose_proxy_logger.warning( - "Cost tracking - skipping, no standard_logging_object and no model for call_type=%s", + "Cost tracking - skipping, no standard_logging_object for call_type=%s", kwargs.get("call_type", "unknown"), ) return diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 8023853e263..023f903194b 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -6,7 +6,7 @@ from uuid import uuid4 import fastapi from fastapi import APIRouter, Depends, HTTPException, Request, Response -from starlette.websockets import WebSocket +from starlette.websockets import WebSocket, WebSocketDisconnect from litellm._logging import verbose_proxy_logger from litellm.integrations.custom_guardrail import ModifyResponseException @@ -935,12 +935,146 @@ async def cancel_response( ) +async def _read_ws_model_from_first_frame( + websocket: WebSocket, +) -> Optional[tuple]: + """Read the first WS frame and return (model, raw_message), or None on error. + + Sends an appropriate error frame and closes the socket before returning None. + """ + try: + first_message = await asyncio.wait_for(websocket.receive_text(), timeout=30) + except asyncio.TimeoutError: + await websocket.close(code=1008, reason="Timed out waiting for first message") + return None + except WebSocketDisconnect: + return None + except Exception: + verbose_proxy_logger.exception( + "Responses WebSocket error reading first message" + ) + await websocket.close(code=1011, reason="Internal server error") + return None + + try: + first_event = json.loads(first_message) + except json.JSONDecodeError: + await websocket.send_text( + json.dumps( + { + "type": "error", + "error": { + "type": "invalid_request_error", + "message": "First message is not valid JSON.", + }, + } + ) + ) + await websocket.close(code=1008, reason="Invalid JSON in first message") + return None + + if ( + not isinstance(first_event, dict) + or first_event.get("type") != "response.create" + ): + await websocket.send_text( + json.dumps( + { + "type": "error", + "error": { + "type": "invalid_request_error", + "message": "First message must be a response.create JSON object.", + }, + } + ) + ) + await websocket.close(code=1008, reason="Invalid first message") + return None + + model = _extract_model_from_first_ws_event(first_event) + if not model: + await websocket.send_text( + json.dumps( + { + "type": "error", + "error": { + "type": "invalid_request_error", + "message": "No model provided. Supply ?model= in the URL or include 'model' in the first response.create event.", + }, + } + ) + ) + await websocket.close(code=1008, reason="No model provided") + return None + + return model, first_message + + +def _extract_model_from_first_ws_event(first_event: Any) -> Optional[str]: + """Extract model from a response.create WS event, handling flat and nested formats. + + Flat: {"type": "response.create", "model": "gpt-4o", ...} + Nested: {"type": "response.create", "response": {"model": "gpt-4o", ...}} + """ + if not isinstance(first_event, dict): + return None + nested = first_event.get("response") + return ( + nested.get("model") if isinstance(nested, dict) else None + ) or first_event.get("model") + + +async def _enforce_responses_ws_first_frame_model_auth( + request: Request, + model: str, + user_api_key_dict: UserAPIKeyAuth, + llm_router: Optional[Any], +) -> None: + from litellm.proxy.auth.user_api_key_auth import ( + _enforce_key_and_fallback_model_access, + _run_centralized_common_checks, + ) + from litellm.proxy.proxy_server import ( + general_settings, + llm_model_list, + master_key, + user_custom_auth, + ) + + request_data = {"model": model} + route = request.scope.get("path") or "/v1/responses" + if master_key is None and not ( + general_settings.get("enable_jwt_auth", False) + or general_settings.get("enable_oauth2_auth", False) + or general_settings.get("enable_oauth2_proxy_auth", False) + ): + return + if user_custom_auth is not None and not general_settings.get( + "custom_auth_run_common_checks", False + ): + return + await _enforce_key_and_fallback_model_access( + valid_token=user_api_key_dict, + request_data=request_data, + route=route, + request=request, + llm_model_list=llm_model_list, + llm_router=llm_router, + ) + await _run_centralized_common_checks( + user_api_key_auth_obj=user_api_key_dict, + request=request, + request_data=request_data, + route=route, + ) + + @router.websocket("/v1/responses") @router.websocket("/responses") async def responses_websocket_endpoint( websocket: WebSocket, - model: str = fastapi.Query( - ..., description="The model to use for the responses WebSocket session." + model: Optional[str] = fastapi.Query( + None, description="The model to use for the responses WebSocket session." ), user_api_key_dict=Depends(user_api_key_auth_websocket), ): @@ -950,6 +1084,10 @@ async def responses_websocket_endpoint( Keeps a persistent WebSocket connection for response.create events, enabling lower-latency agentic workflows with many tool-call round trips. + Follows the OpenAI split: the bearer token is validated at connection time + (before accept); the model is resolved either from the ?model= query param + or from the first response.create frame, whichever is present. + See: https://developers.openai.com/api/docs/guides/websocket-mode/ """ from litellm.proxy.proxy_server import ( @@ -966,7 +1104,8 @@ async def responses_websocket_endpoint( ) from litellm.proxy.route_llm_request import route_request - # Accept the WebSocket handshake + # Accept the WebSocket handshake. Key was already validated by the Depends + # above; we can safely accept regardless of whether ?model= was supplied. requested_protocols = [ p.strip() for p in (websocket.headers.get("sec-websocket-protocol") or "").split(",") @@ -977,10 +1116,19 @@ async def responses_websocket_endpoint( accept_kwargs["subprotocol"] = requested_protocols[0] await websocket.accept(**accept_kwargs) + first_message: Optional[str] = None + if not model: + result = await _read_ws_model_from_first_frame(websocket) + if result is None: + return + model, first_message = result + data: Dict[str, Any] = { "model": model, "websocket": websocket, } + if first_message is not None: + data["first_message"] = first_message # Construct a synthetic Request for pre-call processing headers_list = list(websocket.scope.get("headers") or []) @@ -993,14 +1141,23 @@ async def responses_websocket_endpoint( request = Request(scope=scope) request._url = websocket.url + _body_bytes = json.dumps({"model": model}).encode() + async def return_body(): - return f'{{"model": "{model}"}}'.encode() + return _body_bytes request.body = return_body # type: ignore # Phase 1: pre-call processing (auth, guardrails, rate limits) base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data) try: + if first_message is not None: + await _enforce_responses_ws_first_frame_model_auth( + request=request, + model=model, + user_api_key_dict=user_api_key_dict, + llm_router=llm_router, + ) ( data, litellm_logging_obj, @@ -1027,7 +1184,7 @@ async def responses_websocket_endpoint( { "type": "error", "error": { - "type": "pre_call_error", + "type": "invalid_request_error", "message": str(e), }, } @@ -1035,7 +1192,7 @@ async def responses_websocket_endpoint( ) except Exception: pass - await websocket.close(code=1011, reason="Pre-call error") + await websocket.close(code=1008, reason="Pre-call error") return # Phase 2: route to upstream provider diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index c4e72cb7dc5..dfc43bc29b5 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -1251,6 +1251,7 @@ class ResponsesWebSocketStreaming: logging_obj: LiteLLMLoggingObj, user_api_key_dict: Optional[Any] = None, request_data: Optional[Dict] = None, + first_message: Optional[str] = None, ): self.websocket = websocket self.backend_ws = backend_ws @@ -1259,6 +1260,7 @@ class ResponsesWebSocketStreaming: self.request_data: Dict = request_data or {} self.messages: list[Dict] = [] self.input_messages: list[Dict[str, str]] = [] + self.first_message = first_message def _should_store_event(self, event_obj: dict) -> bool: return event_obj.get("type") in RESPONSES_WS_LOGGED_EVENT_TYPES @@ -1362,6 +1364,11 @@ class ResponsesWebSocketStreaming: async def client_to_backend(self) -> None: """Forward response.create events from client to backend.""" try: + if self.first_message is not None: + self._store_input(self.first_message) + self._store_event(self.first_message) + await self.backend_ws.send(self.first_message) # type: ignore[union-attr] + while True: message = await self.websocket.receive_text() @@ -1440,6 +1447,7 @@ class ManagedResponsesWebSocketHandler: api_base: Optional[str] = None, timeout: Optional[float] = None, custom_llm_provider: Optional[str] = None, + first_message: Optional[str] = None, **kwargs: Any, ) -> None: self.websocket = websocket @@ -1451,6 +1459,8 @@ class ManagedResponsesWebSocketHandler: self.api_base = api_base self.timeout = timeout self.custom_llm_provider = custom_llm_provider + self._connection_provider = self._resolve_provider(model) or custom_llm_provider + self.first_message = first_message # Carry through safe pass-through kwargs (e.g. extra_headers) self.extra_kwargs: Dict[str, Any] = { k: v for k, v in kwargs.items() if k not in _MANAGED_WS_SKIP_KWARGS @@ -1648,8 +1658,30 @@ class ManagedResponsesWebSocketHandler: # cross-connection multi-turn when spend logs are committed) call_kwargs["previous_response_id"] = previous_response_id + @staticmethod + def _resolve_provider(model: Optional[str]) -> Optional[str]: + """Resolve the LLM provider for a model string, or None if unresolvable.""" + if not model: + return None + try: + from litellm import get_llm_provider + + _, provider, _, _ = get_llm_provider(model=model) + return provider + except Exception: + return None + + def _same_provider(self, model: Optional[str]) -> bool: + """Return True if model uses the same LLM provider as the connection model.""" + if model is None or model == self.model: + return True + event_provider = self._resolve_provider(model) + if event_provider is None: + return False + return event_provider == self._connection_provider + def _inject_credentials( - self, call_kwargs: Dict[str, Any], event_model: Optional[str] + self, call_kwargs: Dict[str, Any], model: Optional[str] = None ) -> None: """Inject connection-level credentials and metadata into call_kwargs.""" if self.api_key is not None: @@ -1658,10 +1690,12 @@ class ManagedResponsesWebSocketHandler: call_kwargs["api_base"] = self.api_base if self.timeout is not None: call_kwargs["timeout"] = self.timeout - # Only propagate custom_llm_provider when no per-request model override exists. - # If the payload specifies a different model, let litellm re-resolve the - # provider so we don't accidentally force the wrong backend. - if self.custom_llm_provider is not None and not event_model: + # Only force connection-level custom_llm_provider when the per-event model + # uses the same provider as the connection model. If the provider differs + # (e.g., connection is vertex_ai but event says openai/gpt-4), let litellm + # re-resolve from the model string. Same-provider model variants (e.g., + # vertex_ai/gemini-2.0 -> vertex_ai/gemini-1.5) still inherit the provider. + if self.custom_llm_provider is not None and self._same_provider(model): call_kwargs["custom_llm_provider"] = self.custom_llm_provider if self.litellm_metadata: call_kwargs["litellm_metadata"] = dict(self.litellm_metadata) @@ -1776,8 +1810,7 @@ class ManagedResponsesWebSocketHandler: call_kwargs = self._build_base_call_kwargs(msg_obj) call_kwargs["stream"] = True - event_model: Optional[str] = call_kwargs.pop("model", None) - model = event_model or self.model + model = call_kwargs.pop("model", None) or self.model previous_response_id: Optional[str] = call_kwargs.pop( "previous_response_id", None @@ -1794,7 +1827,7 @@ class ManagedResponsesWebSocketHandler: self._apply_history( call_kwargs, previous_response_id, current_messages, prior_history ) - self._inject_credentials(call_kwargs, event_model) + self._inject_credentials(call_kwargs, model=model) self._update_proxy_request(call_kwargs, model) call_kwargs.update(self.extra_kwargs) @@ -1819,6 +1852,9 @@ class ManagedResponsesWebSocketHandler: each one before waiting for the next message. """ try: + if self.first_message is not None: + await self._process_response_create(self.first_message) + while True: try: message = await self.websocket.receive_text() diff --git a/litellm/router.py b/litellm/router.py index 7aaf989919c..a92590d3dba 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -4636,11 +4636,11 @@ class Router: except Exception: custom_llm_provider = None - # Build response kwargs response_kwargs = { **data, "caching": self.cache_responses, **kwargs, + "model": model_name, } # Only set custom_llm_provider if it's not None if custom_llm_provider is not None: @@ -7126,6 +7126,9 @@ class Router: from litellm.types.caching import RedisPipelineIncrementOperation try: + # WS session wrappers fire with result=None; per-turn costs tracked by inner calls. + if kwargs.get("call_type") in ("_aresponses_websocket", "_arealtime"): + return standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get( "standard_logging_object", None ) diff --git a/litellm/router_strategy/budget_limiter.py b/litellm/router_strategy/budget_limiter.py index f677e40b934..0bb69ca0319 100644 --- a/litellm/router_strategy/budget_limiter.py +++ b/litellm/router_strategy/budget_limiter.py @@ -430,6 +430,9 @@ class RouterBudgetLimiting(CustomLogger): async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): """Original method now uses helper functions""" verbose_router_logger.debug("in RouterBudgetLimiting.async_log_success_event") + # WS session wrappers fire with result=None; per-turn costs tracked by inner calls. + if kwargs.get("call_type") in ("_aresponses_websocket", "_arealtime"): + return standard_logging_payload: Optional[StandardLoggingPayload] = kwargs.get( "standard_logging_object", None ) diff --git a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py index 1929c443720..07d1a9d14f9 100644 --- a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py @@ -196,3 +196,518 @@ class TestResponsesAPIEndpoints(unittest.TestCase): assert "x-litellm-response-cost" in response.headers response_cost_value = float(response.headers["x-litellm-response-cost"]) assert response_cost_value == pytest.approx(0.0005, abs=1e-10) + + +import json + + +class TestManagedResponsesWSFirstMessage: + @pytest.mark.asyncio + async def test_first_message_processed_before_loop(self): + """ + ManagedResponsesWebSocketHandler must process first_message before + entering its receive loop. Regression for clients that connect without + ?model= (e.g. Codex) and send model inside the first response.create event. + """ + from litellm.responses.streaming_iterator import ManagedResponsesWebSocketHandler + + first = json.dumps( + { + "type": "response.create", + "model": "gpt-4o-mini", + "store": False, + "input": [ + { + "type": "message", + "role": "user", + "content": [{"type": "input_text", "text": "hi"}], + } + ], + } + ) + + ws = MagicMock() + ws.receive_text = AsyncMock(side_effect=Exception("disconnect")) + ws.send_text = AsyncMock() + + processed: list = [] + + async def fake_process(msg: str) -> None: + processed.append(msg) + + handler = ManagedResponsesWebSocketHandler( + websocket=ws, + model="gpt-4o-mini", + logging_obj=MagicMock(), + first_message=first, + ) + handler._process_response_create = fake_process # type: ignore[method-assign] + + await handler.run() + + assert processed == [first] + + @pytest.mark.asyncio + async def test_no_first_message_falls_through_to_loop(self): + """When first_message is None, run() goes straight to receive_text().""" + from litellm.responses.streaming_iterator import ManagedResponsesWebSocketHandler + + subsequent = json.dumps({"type": "response.create", "model": "gpt-4o-mini"}) + + ws = MagicMock() + ws.receive_text = AsyncMock(side_effect=[subsequent, Exception("disconnect")]) + ws.send_text = AsyncMock() + + processed: list = [] + + async def fake_process(msg: str) -> None: + processed.append(msg) + + handler = ManagedResponsesWebSocketHandler( + websocket=ws, + model="gpt-4o-mini", + logging_obj=MagicMock(), + first_message=None, + ) + handler._process_response_create = fake_process # type: ignore[method-assign] + + await handler.run() + + assert processed == [subsequent] + + +class TestResponsesWSStreamingFirstMessage: + @pytest.mark.asyncio + async def test_client_to_backend_replays_first_message(self): + """ + ResponsesWebSocketStreaming.client_to_backend must send first_message to + the backend before entering the receive loop. + """ + from litellm.responses.streaming_iterator import ResponsesWebSocketStreaming + + first = json.dumps({"type": "response.create", "model": "gpt-4o-mini", "input": []}) + + ws = MagicMock() + ws.receive_text = AsyncMock(side_effect=Exception("disconnect")) + + backend_ws = MagicMock() + backend_ws.send = AsyncMock() + + streaming = ResponsesWebSocketStreaming( + websocket=ws, + backend_ws=backend_ws, + logging_obj=MagicMock(), + first_message=first, + ) + + await streaming.client_to_backend() + + backend_ws.send.assert_awaited_once_with(first) + + +class TestWSSessionCostTracking: + @pytest.mark.asyncio + async def test_router_budget_limiter_skips_aresponses_websocket_call_type(self): + """ + RouterBudgetLimiting.async_log_success_event must not raise when + call_type='_aresponses_websocket', even when standard_logging_object is None. + Per-turn costs are tracked by individual aresponses calls inside the session; + the outer session wrapper fires with result=None. + """ + from litellm.router_strategy.budget_limiter import RouterBudgetLimiting + + limiter = RouterBudgetLimiting.__new__(RouterBudgetLimiting) + kwargs = { + "call_type": "_aresponses_websocket", + "standard_logging_object": None, + "litellm_params": {"custom_llm_provider": "vertex_ai"}, + } + await limiter.async_log_success_event( + kwargs=kwargs, + response_obj=None, + start_time=None, + end_time=None, + ) + + @pytest.mark.asyncio + async def test_router_budget_limiter_skips_arealtime_call_type(self): + """Same guard applies to _arealtime WS session wrappers.""" + from litellm.router_strategy.budget_limiter import RouterBudgetLimiting + + limiter = RouterBudgetLimiting.__new__(RouterBudgetLimiting) + kwargs = { + "call_type": "_arealtime", + "standard_logging_object": None, + "litellm_params": {"custom_llm_provider": "openai"}, + } + await limiter.async_log_success_event( + kwargs=kwargs, + response_obj=None, + start_time=None, + end_time=None, + ) + + +class TestWSModelExtraction: + """Test _extract_model_from_first_ws_event for flat and nested frame formats.""" + + def test_flat_format_extracts_model(self): + from litellm.proxy.response_api_endpoints.endpoints import ( + _extract_model_from_first_ws_event, + ) + event = {"type": "response.create", "model": "gpt-4o", "input": "hello"} + assert _extract_model_from_first_ws_event(event) == "gpt-4o" + + def test_nested_format_extracts_model(self): + from litellm.proxy.response_api_endpoints.endpoints import ( + _extract_model_from_first_ws_event, + ) + event = {"type": "response.create", "response": {"model": "gpt-4o", "input": "hello"}} + assert _extract_model_from_first_ws_event(event) == "gpt-4o" + + def test_nested_format_takes_precedence_over_flat(self): + from litellm.proxy.response_api_endpoints.endpoints import ( + _extract_model_from_first_ws_event, + ) + event = { + "type": "response.create", + "model": "flat-model", + "response": {"model": "nested-model"}, + } + assert _extract_model_from_first_ws_event(event) == "nested-model" + + def test_no_model_returns_none(self): + from litellm.proxy.response_api_endpoints.endpoints import ( + _extract_model_from_first_ws_event, + ) + event = {"type": "response.create", "input": "hello"} + assert _extract_model_from_first_ws_event(event) is None + + def test_non_object_returns_none(self): + from litellm.proxy.response_api_endpoints.endpoints import ( + _extract_model_from_first_ws_event, + ) + + assert _extract_model_from_first_ws_event([]) is None + + +class TestResponsesWSFirstFrameValidation: + @pytest.mark.asyncio + async def test_rejects_non_response_create_first_frame(self): + from litellm.proxy.response_api_endpoints.endpoints import ( + _read_ws_model_from_first_frame, + ) + + ws = MagicMock() + ws.receive_text = AsyncMock( + return_value=json.dumps({"type": "session.update", "model": "gpt-4o"}) + ) + ws.send_text = AsyncMock() + ws.close = AsyncMock() + + result = await _read_ws_model_from_first_frame(ws) + + assert result is None + ws.send_text.assert_awaited_once() + ws.close.assert_awaited_once_with(code=1008, reason="Invalid first message") + error_payload = json.loads(ws.send_text.await_args.args[0]) + assert ( + error_payload["error"]["message"] + == "First message must be a response.create JSON object." + ) + + @pytest.mark.asyncio + async def test_rejects_non_object_json_first_frame(self): + from litellm.proxy.response_api_endpoints.endpoints import ( + _read_ws_model_from_first_frame, + ) + + ws = MagicMock() + ws.receive_text = AsyncMock(return_value=json.dumps(["gpt-4o"])) + ws.send_text = AsyncMock() + ws.close = AsyncMock() + + result = await _read_ws_model_from_first_frame(ws) + + assert result is None + ws.send_text.assert_awaited_once() + ws.close.assert_awaited_once_with(code=1008, reason="Invalid first message") + + @pytest.mark.asyncio + async def test_client_disconnect_first_frame_does_not_close(self): + from fastapi import WebSocketDisconnect + + from litellm.proxy.response_api_endpoints.endpoints import ( + _read_ws_model_from_first_frame, + ) + + ws = MagicMock() + ws.receive_text = AsyncMock(side_effect=WebSocketDisconnect(code=1006)) + ws.send_text = AsyncMock() + ws.close = AsyncMock() + + result = await _read_ws_model_from_first_frame(ws) + + assert result is None + ws.close.assert_not_awaited() + ws.send_text.assert_not_awaited() + + @pytest.mark.asyncio + async def test_server_error_first_frame_closes_with_internal_error(self): + from litellm.proxy.response_api_endpoints.endpoints import ( + _read_ws_model_from_first_frame, + ) + + ws = MagicMock() + ws.receive_text = AsyncMock(side_effect=RuntimeError("boom")) + ws.send_text = AsyncMock() + ws.close = AsyncMock() + + result = await _read_ws_model_from_first_frame(ws) + + assert result is None + ws.close.assert_awaited_once_with(code=1011, reason="Internal server error") + + +class TestResponsesWSFirstFrameModelAuth: + @pytest.mark.asyncio + async def test_endpoint_enforces_auth_after_model_from_first_frame(self): + from litellm.proxy.response_api_endpoints.endpoints import ( + responses_websocket_endpoint, + ) + + ws = MagicMock() + ws.headers = {} + ws.query_params = {} + ws.scope = {"headers": []} + ws.url = "ws://testserver/v1/responses" + ws.accept = AsyncMock() + ws.receive_text = AsyncMock( + return_value=json.dumps( + {"type": "response.create", "model": "gpt-4o-mini", "input": []} + ) + ) + ws.close = AsyncMock() + + processor = MagicMock() + processor.common_processing_pre_call_logic = AsyncMock( + return_value=({"model": "gpt-4o-mini"}, MagicMock()) + ) + + async def fake_llm_call(): + return None + + with ( + patch( + "litellm.proxy.response_api_endpoints.endpoints._enforce_responses_ws_first_frame_model_auth", + new_callable=AsyncMock, + ) as mock_model_auth, + patch( + "litellm.proxy.response_api_endpoints.endpoints.ProxyBaseLLMRequestProcessing", + return_value=processor, + ), + patch( + "litellm.proxy.route_llm_request.route_request", + new_callable=AsyncMock, + return_value=fake_llm_call(), + ), + ): + await responses_websocket_endpoint( + websocket=ws, + model=None, + user_api_key_dict=MagicMock(), + ) + + mock_model_auth.assert_awaited_once() + + @pytest.mark.asyncio + async def test_reruns_model_auth_for_first_frame_model(self): + from starlette.requests import Request + + from litellm.proxy.response_api_endpoints.endpoints import ( + _enforce_responses_ws_first_frame_model_auth, + ) + + request = Request( + {"type": "http", "method": "POST", "path": "/v1/responses", "headers": []} + ) + user_api_key_dict = MagicMock() + llm_router = MagicMock() + + with ( + patch( + "litellm.proxy.auth.user_api_key_auth._enforce_key_and_fallback_model_access", + new_callable=AsyncMock, + ) as mock_key_check, + patch( + "litellm.proxy.auth.user_api_key_auth._run_centralized_common_checks", + new_callable=AsyncMock, + ) as mock_common_checks, + patch( + "litellm.proxy.proxy_server.llm_model_list", + [], + ), + patch("litellm.proxy.proxy_server.master_key", "sk-test"), + patch("litellm.proxy.proxy_server.user_custom_auth", None), + patch("litellm.proxy.proxy_server.general_settings", {}), + ): + await _enforce_responses_ws_first_frame_model_auth( + request=request, + model="gpt-4o-mini", + user_api_key_dict=user_api_key_dict, + llm_router=llm_router, + ) + + mock_key_check.assert_awaited_once_with( + valid_token=user_api_key_dict, + request_data={"model": "gpt-4o-mini"}, + route="/v1/responses", + request=request, + llm_model_list=[], + llm_router=llm_router, + ) + mock_common_checks.assert_awaited_once_with( + user_api_key_auth_obj=user_api_key_dict, + request=request, + request_data={"model": "gpt-4o-mini"}, + route="/v1/responses", + ) + + +class TestReadWSModelFromFirstFrameErrors: + @pytest.mark.asyncio + async def test_timeout_closes_without_error_frame(self): + import asyncio + + from litellm.proxy.response_api_endpoints.endpoints import ( + _read_ws_model_from_first_frame, + ) + + ws = MagicMock() + ws.receive_text = AsyncMock(side_effect=asyncio.TimeoutError()) + ws.send_text = AsyncMock() + ws.close = AsyncMock() + + result = await _read_ws_model_from_first_frame(ws) + + assert result is None + ws.send_text.assert_not_awaited() + ws.close.assert_awaited_once_with( + code=1008, reason="Timed out waiting for first message" + ) + + @pytest.mark.asyncio + async def test_invalid_json_sends_error_and_closes(self): + from litellm.proxy.response_api_endpoints.endpoints import ( + _read_ws_model_from_first_frame, + ) + + ws = MagicMock() + ws.receive_text = AsyncMock(return_value="this is not json") + ws.send_text = AsyncMock() + ws.close = AsyncMock() + + result = await _read_ws_model_from_first_frame(ws) + + assert result is None + payload = json.loads(ws.send_text.await_args.args[0]) + assert payload["error"]["message"] == "First message is not valid JSON." + ws.close.assert_awaited_once_with( + code=1008, reason="Invalid JSON in first message" + ) + + @pytest.mark.asyncio + async def test_missing_model_sends_error_and_closes(self): + from litellm.proxy.response_api_endpoints.endpoints import ( + _read_ws_model_from_first_frame, + ) + + ws = MagicMock() + ws.receive_text = AsyncMock( + return_value=json.dumps({"type": "response.create", "input": []}) + ) + ws.send_text = AsyncMock() + ws.close = AsyncMock() + + result = await _read_ws_model_from_first_frame(ws) + + assert result is None + payload = json.loads(ws.send_text.await_args.args[0]) + assert "No model provided" in payload["error"]["message"] + ws.close.assert_awaited_once_with(code=1008, reason="No model provided") + + @pytest.mark.asyncio + async def test_valid_first_frame_returns_model_and_raw(self): + from litellm.proxy.response_api_endpoints.endpoints import ( + _read_ws_model_from_first_frame, + ) + + raw = json.dumps({"type": "response.create", "model": "gpt-4o", "input": []}) + ws = MagicMock() + ws.receive_text = AsyncMock(return_value=raw) + ws.send_text = AsyncMock() + ws.close = AsyncMock() + + result = await _read_ws_model_from_first_frame(ws) + + assert result == ("gpt-4o", raw) + ws.send_text.assert_not_awaited() + ws.close.assert_not_awaited() + + +class TestManagedResponsesSameProvider: + def _handler(self, model, custom_llm_provider=None): + from litellm.responses.streaming_iterator import ( + ManagedResponsesWebSocketHandler, + ) + + return ManagedResponsesWebSocketHandler( + websocket=MagicMock(), + model=model, + logging_obj=MagicMock(), + custom_llm_provider=custom_llm_provider, + ) + + def test_none_model_treated_as_same_provider(self): + assert self._handler("openai/gpt-4o")._same_provider(None) is True + + def test_identical_model_is_same_provider(self): + assert self._handler("openai/gpt-4o")._same_provider("openai/gpt-4o") is True + + def test_same_provider_different_model(self): + assert self._handler("gpt-4o")._same_provider("gpt-4o-mini") is True + + def test_different_provider_is_not_same(self): + assert ( + self._handler("gpt-4o")._same_provider("vertex_ai/gemini-2.0-flash") + is False + ) + + def test_inject_credentials_keeps_provider_for_same_provider_model(self): + handler = self._handler("gpt-4o", custom_llm_provider="openai") + call_kwargs: dict = {} + handler._inject_credentials(call_kwargs, model="gpt-4o-mini") + assert call_kwargs["custom_llm_provider"] == "openai" + + def test_inject_credentials_drops_provider_for_cross_provider_model(self): + handler = self._handler("gpt-4o", custom_llm_provider="openai") + call_kwargs: dict = {} + handler._inject_credentials(call_kwargs, model="vertex_ai/gemini-2.0-flash") + assert "custom_llm_provider" not in call_kwargs + + def test_unresolvable_connection_model_falls_back_to_custom_provider(self): + handler = self._handler( + "my-custom-deployment", custom_llm_provider="openai" + ) + assert handler._same_provider("gpt-4o-mini") is True + call_kwargs: dict = {} + handler._inject_credentials(call_kwargs, model="gpt-4o-mini") + assert call_kwargs["custom_llm_provider"] == "openai" + + def test_unresolvable_connection_model_still_drops_cross_provider(self): + handler = self._handler( + "my-custom-deployment", custom_llm_provider="openai" + ) + call_kwargs: dict = {} + handler._inject_credentials(call_kwargs, model="vertex_ai/gemini-2.0-flash") + assert "custom_llm_provider" not in call_kwargs diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index e9287a95438..cd235d8de67 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -982,6 +982,61 @@ async def test_router_ageneric_api_call_with_fallbacks_helper(): assert router.fail_calls["gpt-3.5-turbo"] == initial_fail_count + 1 +@pytest.mark.asyncio +async def test_ageneric_api_call_deployment_model_overrides_alias(): + """ + Regression: when a model alias (e.g. "not-gemini-2.5-flash") maps to a deployment + with model="vertex_ai/gemini-2.5-flash", the underlying litellm function must receive + the deployment model, not the alias. Before the fix, **kwargs overwrote data["model"]. + """ + from unittest.mock import patch + + captured: dict = {} + + async def capture_model(**kwargs): + captured["model"] = kwargs.get("model") + return {"result": "ok"} + + router = litellm.Router( + model_list=[ + { + "model_name": "not-gemini-2.5-flash", + "litellm_params": { + "model": "vertex_ai/gemini-2.5-flash", + "api_key": "fake-key", + }, + } + ] + ) + + def inject_alias_into_kwargs(deployment, kwargs, function_name=None): + # Simulate the alias leaking into kwargs (as happens when + # _ageneric_api_call_with_fallbacks sets kwargs["model"] = alias before + # calling the helper through async_function_with_fallbacks). + kwargs["model"] = "not-gemini-2.5-flash" + + with patch.object(router, "async_get_available_deployment") as mock_dep, \ + patch.object(router, "_update_kwargs_with_deployment", side_effect=inject_alias_into_kwargs), \ + patch.object(router, "async_routing_strategy_pre_call_checks"), \ + patch.object(router, "_get_client", return_value=None): + mock_dep.return_value = { + "model_name": "not-gemini-2.5-flash", + "litellm_params": { + "model": "vertex_ai/gemini-2.5-flash", + "api_key": "fake-key", + }, + } + + await router._ageneric_api_call_with_fallbacks_helper( + model="not-gemini-2.5-flash", + original_generic_function=capture_model, + ) + + assert captured["model"] == "vertex_ai/gemini-2.5-flash", ( + f"Expected deployment model 'vertex_ai/gemini-2.5-flash', got '{captured['model']}'" + ) + + def test_router_get_model_access_groups_team_only_models(): """ Test that Router.get_model_access_groups returns the correct response for team-only models From 5119b9462f96482a10572bd0d6ade3c552b0556f Mon Sep 17 00:00:00 2001 From: milan-berri Date: Wed, 3 Jun 2026 22:09:50 +0300 Subject: [PATCH 02/11] =?UTF-8?q?feat(arize/phoenix):=20OpenInference=20re?= =?UTF-8?q?ndering=20parity=20=E2=80=94=20tool=5Fcalls,=20cost,=20passthro?= =?UTF-8?q?ugh=20I/O,=20session/user,=20multimodal,=20cache=20tokens=20(#2?= =?UTF-8?q?8800)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat(arize): enrich OpenInference attributes for better span rendering Pure rendering enhancements to the Arize / Arize Phoenix integration. No existing attribute keys or values are removed or overwritten; every new emit is independently try/except-wrapped and fires only when its source data is present so existing behavior is preserved. What this adds - Coerce non-dict response objects (e.g. httpx.Response from passthrough routes) via JSON decode so id/model/usage extraction stops crashing with "'Response' object has no attribute 'get'". Dicts and Pydantic objects with .get pass through unchanged. - Set OPENINFERENCE_SPAN_KIND defensively early so a downstream failure can't blank the kind; the original late write (incl. TOOL upgrade) is preserved. - Add "passthrough" keyword to _infer_open_inference_span_kind so allm_passthrough_route / llm_passthrough_route resolve to LLM instead of UNKNOWN. - Emit cache token breakdown: LLM_TOKEN_COUNT_PROMPT_DETAILS_CACHE_READ / _CACHE_WRITE / _AUDIO. Sources covered: OpenAI prompt_tokens_details and Anthropic / Bedrock cache_{read,creation}_input_tokens. - Render assistant tool_calls on both input and output messages via MESSAGE_TOOL_CALLS.* (Pydantic-aware, handles ModelResponse choices). Tool-result input messages also get MESSAGE_TOOL_CALL_ID and MESSAGE_NAME. - Render multimodal list-shaped content via MESSAGE_CONTENTS.* (OpenAI image_url, Anthropic source.{media_type,data} as data: URI). Legacy MESSAGE_CONTENT write is unchanged. - Emit SESSION_ID (end_user_id / trace_id), USER_ID (only when not already set by optional_params.user or model_params.user), and litellm.{team_id,team_alias,key_alias} from StandardLoggingPayload metadata. - Emit llm.response.cost as float from StandardLoggingPayload.response_cost. - Bedrock / Anthropic passthrough normalization: extract input from additional_args.complete_input_dict and output from the coerced provider response so INPUT_VALUE / OUTPUT_VALUE / LLM_INPUT_MESSAGES / LLM_OUTPUT_MESSAGES are populated. Only runs when call_type contains "passthrough" / "pass_through". Tests - 15 new unit tests covering each addition plus explicit regression guards (USER_ID overwrite protection, passthrough normalizer scope, coerce identity for dicts/.get-bearing objects, no spurious cache emits). - Existing test_arize_set_attributes count bumped from 26 to 27 to account for the additional defensive span.kind write (same value, written twice). - tests/test_litellm/integrations/arize/: 70 passed (55 baseline + 15 new). tests/test_litellm/integrations/test_opentelemetry.py: 221 passed. Co-authored-by: Cursor * refactor(arize): collapse additive try/except blocks into _safe_emit helper The additive attribute emitters all share the same shape: run a callable, swallow any exception to debug log so it cannot blank the span. Hoisting that pattern into a single _safe_emit(label, fn, *args, **kwargs) helper removes 5 repeated try/except blocks. Behavior unchanged; arize test suite still passes (70/70). Co-authored-by: Cursor * fix(arize): emit cost under canonical llm.cost.total key Arize's "Total Cost" column reads the OpenInference-standard `llm.cost.total` attribute. The previous custom `llm.response.cost` key never surfaced in the trace list. Now emits both keys (canonical + legacy) so renderers + any existing consumers both work. Co-authored-by: Cursor * fix(arize): keep span.kind=LLM for tool-using completions + render tool_calls in Output A chat completion that passes `tools=[...]` or returns `tool_calls` is still an LLM call per the OpenInference spec — TOOL is reserved for actual tool execution. The previous override demoted these to TOOL, breaking Arize's LLM-scoped dashboards/evals and skewing token/cost analytics for any tool-using traffic. Additionally, when an assistant response had no text content but did request tool calls, `output.value` was set to the empty string so Arize's "Output" pane rendered blank. Now serializes the tool_calls into a compact JSON summary in `output.value` (the structured `MESSAGE_TOOL_CALLS.*` attributes are still emitted unchanged). Cleanups: - extract `_get_tool_calls` and `_normalize_tool_call` helpers, deduplicating the dict-vs-Pydantic + function-dict logic across `_set_choice_outputs`, `_emit_message_tool_calls`, and the new `_summarize_tool_calls_for_output`. - drop redundant late `OPENINFERENCE_SPAN_KIND` write — the defensive early write is now the single source of truth. - remove a dead local re-import of `MessageAttributes`/`SpanAttributes`. Tests: 73 pass (added regression guard asserting span.kind stays LLM for completions that pass tools AND return tool_calls; existing call_count assertion restored to 26). Co-authored-by: Cursor * chore(arize): tighten cleanup — fold _get_tool_calls into _safe_get Two tiny cleanups, no behavior change: - collapse `_get_tool_calls` to use `_safe_get`, removing a 7-line hand-rolled dict-vs-attribute fallback that duplicated existing logic. - trim the `_set_choice_outputs` tool-call summary comment from 4 lines to 2 (was over-explaining). Co-authored-by: Cursor * fix(arize): address Greptile review — drop session_id=trace_id fallback, remove dead code, fix Black Three Greptile-flagged issues + the Black formatting CI failure. 1. SESSION_ID no longer falls back to trace_id. Previously every span without an explicit `user_api_key_end_user_id` would have its session.id set to the per-request trace_id, which creates one distinct "session" per request and breaks Arize's Session-grouping analytics. Now SESSION_ID is emitted only when an explicit end-user identifier exists, and the trace_id is emitted under its own `litellm.trace_id` key so spans remain filterable by trace. 2. Removed dead `ArizeOTELAttributes.set_response_output_messages` override. Confirmed zero callers in the entire repo (the live path is `_set_choice_outputs` via `_set_response_attributes`). The override was preexisting dead code, but the expansion of `_set_choice_outputs` in this PR made the divergence misleading. 3. Removed permanently-dead first branch in cache_write detection. `_safe_get(prompt_token_details, "cache_creation_tokens")` looks for a key that neither OpenAI's `prompt_tokens_details` nor Anthropic's payload ever exposes. Now reads straight off `usage` for `cache_creation_input_tokens`. 4. Reformatted both files under Black 26.3.1 (the version CI uses via `uv sync --frozen`). Local previously used 24.10.0. Tests: 74/74 pass in the arize suite (added `test_arize_does_not_use_trace_id_as_session_id_fallback`). Combined arize + opentelemetry suite: 295/295 pass. End-to-end verified live: tool-call still emits `span.kind=LLM` and JSON tool_calls in `output.value`; `session.id` is now correctly unset when no end_user_id is provided; `litellm.trace_id` is populated; Bedrock passthrough input/output unchanged. Co-authored-by: Cursor * fix(arize): gate passthrough prompt export on message redaction - Skip the complete_input_dict bridge in _maybe_normalize_passthrough when should_redact_message_logging() is true, so enabling redaction no longer leaks raw passthrough prompts into Arize (Veria security finding). - Split passthrough input/output rendering into helpers to satisfy PLR0915. - Remove dead call_type assignment (F841). Validated live against a Bedrock passthrough proxy exporting to Arize: non-redacted renders the real prompt on litellm_request; global turn_off_message_logging yields input.value=redacted-by-litellm with the raw_gen_ai_request child span suppressed and no SSN/marker leakage. Co-authored-by: Cursor --------- Co-authored-by: Cursor --- litellm/integrations/arize/_utils.py | 702 ++++++++++++++-- .../integrations/arize/test_arize_utils.py | 748 +++++++++++++++++- 2 files changed, 1395 insertions(+), 55 deletions(-) diff --git a/litellm/integrations/arize/_utils.py b/litellm/integrations/arize/_utils.py index a1bf65141c9..75710e10498 100644 --- a/litellm/integrations/arize/_utils.py +++ b/litellm/integrations/arize/_utils.py @@ -8,18 +8,23 @@ from litellm.integrations.opentelemetry_utils.base_otel_llm_obs_attributes impor BaseLLMObsOTELAttributes, safe_set_attribute, ) +from litellm.litellm_core_utils.redact_messages import ( + should_redact_message_logging, +) from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.types.utils import StandardLoggingPayload if TYPE_CHECKING: from opentelemetry.trace import Span from litellm.integrations._types.open_inference import ( - MessageAttributes, - ImageAttributes, - SpanAttributes, AudioAttributes, EmbeddingAttributes, + ImageAttributes, + MessageAttributes, + MessageContentAttributes, OpenInferenceSpanKindValues, + SpanAttributes, + ToolCallAttributes, ) @@ -53,40 +58,24 @@ class ArizeOTELAttributes(BaseLLMObsOTELAttributes): msg.get("content", ""), ) - @staticmethod - @override - def set_response_output_messages(span: "Span", response_obj): - """ - Sets output message attributes on the span from the LLM response. - Args: - span: The OpenTelemetry span to set attributes on - response_obj: The response object containing choices with messages - """ - from litellm.integrations._types.open_inference import ( - MessageAttributes, - SpanAttributes, - ) + # Additive: emit structured tool_calls / multimodal content + # so Arize/Phoenix can render tool-using and image-bearing + # turns. These set NEW attribute keys (MESSAGE_TOOL_CALLS / + # MESSAGE_NAME / MESSAGE_TOOL_CALL_ID / MESSAGE_CONTENTS.*) — + # never replace the MESSAGE_CONTENT write above. + _safe_emit( + f"input message extras (idx={idx})", + _emit_input_message_extras, + span, + prefix, + msg, + ) - for idx, choice in enumerate(response_obj.get("choices", [])): - response_message = choice.get("message", {}) - safe_set_attribute( - span, - SpanAttributes.OUTPUT_VALUE, - response_message.get("content", ""), - ) - - # This shows up under `output_messages` tab on the span page. - prefix = f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.{idx}" - safe_set_attribute( - span, - f"{prefix}.{MessageAttributes.MESSAGE_ROLE}", - response_message.get("role"), - ) - safe_set_attribute( - span, - f"{prefix}.{MessageAttributes.MESSAGE_CONTENT}", - response_message.get("content", ""), - ) + # Note: `BaseLLMObsOTELAttributes.set_response_output_messages` is not + # overridden here. The live code path uses `_set_choice_outputs` (called + # via `_set_response_attributes` from `set_attributes`) which handles + # tool_calls, multimodal output, embeddings, audio, images, and structured + # outputs in a single place. def _set_response_attributes(span: "Span", response_obj): @@ -106,11 +95,17 @@ def _set_response_attributes(span: "Span", response_obj): def _set_choice_outputs(span: "Span", response_obj, msg_attrs, span_attrs): for idx, choice in enumerate(response_obj.get("choices", [])): response_message = choice.get("message", {}) - safe_set_attribute( - span, - span_attrs.OUTPUT_VALUE, - response_message.get("content", ""), - ) + content = response_message.get("content", "") + + # Tool-only assistant responses have empty content; serialize the + # tool_calls into OUTPUT_VALUE so Arize's "Output" pane isn't blank. + output_value = content + if not output_value: + tool_calls = _get_tool_calls(response_message) + if tool_calls: + output_value = _summarize_tool_calls_for_output(tool_calls) + + safe_set_attribute(span, span_attrs.OUTPUT_VALUE, output_value) prefix = f"{span_attrs.LLM_OUTPUT_MESSAGES}.{idx}" safe_set_attribute( span, @@ -120,7 +115,18 @@ def _set_choice_outputs(span: "Span", response_obj, msg_attrs, span_attrs): safe_set_attribute( span, f"{prefix}.{msg_attrs.MESSAGE_CONTENT}", - response_message.get("content", ""), + content, + ) + + # Additive: emit assistant tool_calls so tool-using turns render in + # Arize/Phoenix. Sets new MESSAGE_TOOL_CALLS keys only — does not + # change MESSAGE_CONTENT/MESSAGE_ROLE writes above. + _safe_emit( + f"output tool_calls (idx={idx})", + _emit_message_tool_calls, + span, + prefix, + response_message, ) @@ -278,6 +284,43 @@ def _set_usage_outputs(span: "Span", response_obj, span_attrs): reasoning_tokens, ) + # Additive: cache token breakdown so prompt-caching savings render in + # Arize. Sources covered: + # - OpenAI Chat Completions: `prompt_tokens_details.cached_tokens` + # - Anthropic / Bedrock-Anthropic: `cache_read_input_tokens`, + # `cache_creation_input_tokens` + # All emits are conditional, so when none of these fields exist (the + # situation in the existing test fixtures) no extra attributes are set. + prompt_token_details = _safe_get(usage, "prompt_tokens_details") or _safe_get( + usage, "input_tokens_details" + ) + cache_read = _safe_get(prompt_token_details, "cached_tokens") or _safe_get( + usage, "cache_read_input_tokens" + ) + if cache_read: + safe_set_attribute( + span, + span_attrs.LLM_TOKEN_COUNT_PROMPT_DETAILS_CACHE_READ, + cache_read, + ) + # Anthropic / Bedrock-Anthropic only — OpenAI's `prompt_tokens_details` + # does not expose a cache-write count, so we read straight off `usage`. + cache_write = _safe_get(usage, "cache_creation_input_tokens") + if cache_write: + safe_set_attribute( + span, + span_attrs.LLM_TOKEN_COUNT_PROMPT_DETAILS_CACHE_WRITE, + cache_write, + ) + + audio_prompt_tokens = _safe_get(prompt_token_details, "audio_tokens") + if audio_prompt_tokens: + safe_set_attribute( + span, + span_attrs.LLM_TOKEN_COUNT_PROMPT_DETAILS_AUDIO, + audio_prompt_tokens, + ) + def _infer_open_inference_span_kind(call_type: Optional[str]) -> str: """ @@ -321,6 +364,10 @@ def _infer_open_inference_span_kind(call_type: Optional[str]) -> str: "videos", "realtime", "pass_through", + # `passthrough` (no underscore) is what real call_types use: + # `allm_passthrough_route`, `llm_passthrough_route`. Without + # this they fell through to UNKNOWN, blanking span.kind. + "passthrough", "anthropic_messages", "ocr", ) @@ -396,6 +443,18 @@ def set_attributes( """ Populates span with OpenInference-compliant LLM attributes for Arize and Phoenix tracing. """ + # Coerce non-dict response objects (e.g. httpx.Response from passthrough + # routes) into a dict so downstream `.get()` calls don't crash. Existing + # dict / `.get()`-bearing objects (incl. Pydantic OpenAI Responses API + # models) are returned unchanged, preserving the existing test behavior. + response_obj_for_attrs = _coerce_response_obj_for_attrs(response_obj) + + # Set span.kind defensively before anything else. If a downstream step + # throws, the span still has a kind so Arize can render it correctly + # (an LLM call instead of UNKNOWN). This is the single source of truth + # for span.kind — no late re-write happens below. + _safe_emit("early span kind", _set_early_span_kind, span, kwargs) + try: optional_params = _sanitize_optional_params(kwargs.get("optional_params")) litellm_params = kwargs.get("litellm_params", {}) or {} @@ -415,25 +474,22 @@ def set_attributes( metadata_tools = _extract_metadata_tools(metadata) optional_tools = _extract_optional_tools(optional_params) - call_type = standard_logging_payload.get("call_type") _set_request_attributes( span=span, kwargs=kwargs, standard_logging_payload=standard_logging_payload, optional_params=optional_params, litellm_params=litellm_params, - response_obj=response_obj, + response_obj=response_obj_for_attrs, span_attrs=SpanAttributes, ) - span_kind = _infer_open_inference_span_kind(call_type=call_type) + # span.kind was already set above by `_set_early_span_kind`. We do + # NOT re-write it here based on tool presence: a chat completion + # that passes `tools=[...]` (or returns `tool_calls`) is still an + # LLM call per the OpenInference spec — TOOL is reserved for actual + # tool execution spans, not LLM calls that request tools. _set_tool_attributes(span, optional_tools, metadata_tools) - if ( - optional_tools or metadata_tools - ) and span_kind != OpenInferenceSpanKindValues.TOOL.value: - span_kind = OpenInferenceSpanKindValues.TOOL.value - - safe_set_attribute(span, SpanAttributes.OPENINFERENCE_SPAN_KIND, span_kind) attributes.set_messages(span, kwargs) model_params = ( @@ -443,7 +499,7 @@ def set_attributes( ) _set_model_params(span, model_params, SpanAttributes) - _set_response_attributes(span=span, response_obj=response_obj) + _set_response_attributes(span=span, response_obj=response_obj_for_attrs) except Exception as e: verbose_logger.error( @@ -452,6 +508,22 @@ def set_attributes( if hasattr(span, "record_exception"): span.record_exception(e) + # Additive emitters. Each is independently guarded so a failure can never + # blank the attributes set by the main try-block above. New attributes are + # written under new keys; existing attributes are not overwritten. + slp = kwargs.get("standard_logging_object") + _safe_emit("session/user attrs", _set_session_and_user_attrs, span, kwargs, slp) + _safe_emit("response cost", _set_response_cost_attr, span, slp) + _safe_emit( + "passthrough normalization", + _maybe_normalize_passthrough, + span, + kwargs, + response_obj, + response_obj_for_attrs, + slp, + ) + def _sanitize_optional_params(optional_params: Optional[dict]) -> dict: if not isinstance(optional_params, dict): @@ -534,3 +606,529 @@ def _set_model_params(span: "Span", model_params: Optional[dict], span_attrs) -> user_id = model_params.get("user") if user_id is not None: safe_set_attribute(span, span_attrs.USER_ID, user_id) + + +# --------------------------------------------------------------------------- +# Additive rendering helpers (introduced to enhance Arize/Phoenix rendering +# without changing any previously-emitted attribute keys or values). +# --------------------------------------------------------------------------- + + +def _safe_emit(label: str, fn, *args, **kwargs) -> None: + """Run an additive attribute emitter, swallowing any error so it cannot + blank attributes set elsewhere on the span. Failures are logged at debug. + """ + try: + fn(*args, **kwargs) + except Exception as e: + verbose_logger.debug("[Arize] %s skipped: %s", label, e) + + +def _set_early_span_kind(span: "Span", kwargs: dict) -> None: + """Defensively set OPENINFERENCE_SPAN_KIND before any other logic runs.""" + slp = kwargs.get("standard_logging_object") + call_type = slp.get("call_type") if isinstance(slp, dict) else None + safe_set_attribute( + span, + SpanAttributes.OPENINFERENCE_SPAN_KIND, + _infer_open_inference_span_kind(call_type=call_type), + ) + + +def _coerce_response_obj_for_attrs(response_obj): + """Return a `.get`-compatible view of `response_obj` when possible. + + - dicts and Pydantic models that already expose `.get` are returned + unchanged (preserves all current behavior, including the Responses API + flow which relies on Pydantic attribute access). + - `httpx.Response` and other text-only responses (passthrough routes) + are JSON-decoded so the standard extraction paths can read fields like + `id`, `model`, and `usage`. On failure the original object is returned + so behavior is no worse than today. + """ + if response_obj is None or hasattr(response_obj, "get"): + return response_obj + text = getattr(response_obj, "text", None) + if isinstance(text, str) and text: + try: + parsed = json.loads(text) + if isinstance(parsed, dict): + return parsed + except Exception: + pass + return response_obj + + +def _coerce_text(value) -> Optional[str]: + """Best-effort text extraction from a message-content value. + + Returns None when no textual portion can be derived. Handles: + - plain strings + - lists of OpenAI-style content parts (`{"type": "text", "text": ...}`) + - lists of Anthropic-style content parts (`{"type": "text", "text": ...}` + or `{"type": "input_text", "text": ...}`) + """ + if value is None: + return None + if isinstance(value, str): + return value + if isinstance(value, list): + parts = [] + for part in value: + if isinstance(part, str): + parts.append(part) + elif isinstance(part, dict): + text = part.get("text") or part.get("input_text") + if isinstance(text, str): + parts.append(text) + if parts: + return "\n".join(parts) + return None + + +def _to_plain_dict(value): + """Best-effort: coerce a value (Pydantic model / dict / None) to a dict. + + Returns the original value when no safe conversion exists. Used to bridge + OpenAI Pydantic message/tool_call objects into the dict-based helpers. + """ + if value is None or isinstance(value, dict): + return value + model_dump = getattr(value, "model_dump", None) + if callable(model_dump): + try: + return model_dump() + except Exception: + pass + return value + + +def _get_tool_calls(message) -> Optional[list]: + """Return ``message.tool_calls`` only when it's a non-empty list. + + Works for dicts and Pydantic message objects via ``_safe_get``. + """ + tool_calls = _safe_get(message, "tool_calls") + return tool_calls if isinstance(tool_calls, list) and tool_calls else None + + +def _normalize_tool_call(raw_tc) -> Optional[Dict[str, Any]]: + """Normalize a single tool_call (dict or Pydantic) into a stable shape: + + {"id": str|None, "type": str, "function": {"name": str|None, "arguments": str|None}} + + Arguments are coerced to a JSON string per OpenInference convention. + Returns ``None`` when ``raw_tc`` cannot be coerced to a dict. + """ + tc = _to_plain_dict(raw_tc) + if not isinstance(tc, dict): + return None + function = _to_plain_dict(tc.get("function")) + name = function.get("name") if isinstance(function, dict) else None + args = function.get("arguments") if isinstance(function, dict) else None + if args is not None and not isinstance(args, str): + try: + args = json.dumps(args) + except Exception: + args = str(args) + return { + "id": tc.get("id"), + "type": tc.get("type", "function"), + "function": {"name": name, "arguments": args}, + } + + +def _summarize_tool_calls_for_output(tool_calls) -> str: + """Render a tool_calls list as a compact JSON string for OUTPUT_VALUE. + + Best-effort: returns ``str(tool_calls)`` if anything unexpected happens + so OUTPUT_VALUE is never blanked on a malformed payload. + """ + try: + normalized = [n for n in (_normalize_tool_call(tc) for tc in tool_calls) if n] + return json.dumps({"tool_calls": normalized}) + except Exception: + return str(tool_calls) + + +def _emit_message_tool_calls(span: "Span", prefix: str, message) -> None: + """Emit ``MESSAGE_TOOL_CALLS.*`` for an assistant message that requested + tool calls. Pure addition: only writes when ``tool_calls`` is non-empty. + + Accepts dicts or Pydantic message objects (e.g. ``litellm.Message``); the + same applies to each tool_call entry. + """ + tool_calls = _get_tool_calls(message) + if not tool_calls: + return + for tc_idx, raw_tc in enumerate(tool_calls): + tc = _normalize_tool_call(raw_tc) + if tc is None: + continue + tc_prefix = f"{prefix}.{MessageAttributes.MESSAGE_TOOL_CALLS}.{tc_idx}" + if tc["id"]: + safe_set_attribute( + span, f"{tc_prefix}.{ToolCallAttributes.TOOL_CALL_ID}", tc["id"] + ) + fn = tc["function"] + if fn["name"]: + safe_set_attribute( + span, + f"{tc_prefix}.{ToolCallAttributes.TOOL_CALL_FUNCTION_NAME}", + fn["name"], + ) + if fn["arguments"] is not None: + safe_set_attribute( + span, + f"{tc_prefix}.{ToolCallAttributes.TOOL_CALL_FUNCTION_ARGUMENTS_JSON}", + fn["arguments"], + ) + + +def _emit_input_message_extras(span: "Span", prefix: str, message: dict) -> None: + """Emit additive attributes for an input message: + + - `MESSAGE_NAME` and `MESSAGE_TOOL_CALL_ID` (commonly set on tool-result + messages so traces show which tool produced which result). + - `MESSAGE_TOOL_CALLS.*` when an assistant message requested tools. + - `MESSAGE_CONTENTS.*` structured content for list-shaped content + (multimodal text + image parts). The plain `MESSAGE_CONTENT` write is + still performed by the caller, so renderers that only read the legacy + key continue to work. + """ + if not isinstance(message, dict): + return + + name = message.get("name") + if name: + safe_set_attribute(span, f"{prefix}.{MessageAttributes.MESSAGE_NAME}", name) + + tool_call_id = message.get("tool_call_id") + if tool_call_id: + safe_set_attribute( + span, + f"{prefix}.{MessageAttributes.MESSAGE_TOOL_CALL_ID}", + tool_call_id, + ) + + _emit_message_tool_calls(span, prefix, message) + + content = message.get("content") + if isinstance(content, list): + contents_prefix = f"{prefix}.{MessageAttributes.MESSAGE_CONTENTS}" + for part_idx, part in enumerate(content): + if not isinstance(part, dict): + continue + part_prefix = f"{contents_prefix}.{part_idx}" + part_type = part.get("type") + if part_type in ("text", "input_text"): + text = part.get("text") + if isinstance(text, str): + safe_set_attribute( + span, + f"{part_prefix}.{MessageContentAttributes.MESSAGE_CONTENT_TYPE}", + "text", + ) + safe_set_attribute( + span, + f"{part_prefix}.{MessageContentAttributes.MESSAGE_CONTENT_TEXT}", + text, + ) + elif part_type in ("image_url", "image", "input_image"): + url = None + image = part.get("image_url") + if isinstance(image, dict): + url = image.get("url") + elif isinstance(image, str): + url = image + if not url: + # Anthropic-style source.{type=base64,media_type,data} + source = part.get("source") + if isinstance(source, dict) and source.get("data"): + media_type = source.get("media_type", "image/jpeg") + url = f"data:{media_type};base64,{source['data']}" + elif isinstance(part.get("url"), str): + url = part["url"] + if url: + safe_set_attribute( + span, + f"{part_prefix}.{MessageContentAttributes.MESSAGE_CONTENT_TYPE}", + "image", + ) + safe_set_attribute( + span, + f"{part_prefix}.message_content.image.image.url", + url, + ) + + +def _set_session_and_user_attrs( + span: "Span", kwargs: dict, standard_logging_payload +) -> None: + """Emit `SESSION_ID` / `USER_ID` / team metadata when source data exists. + + `SESSION_ID` is emitted only when an explicit end-user identifier exists + (`metadata.user_api_key_end_user_id`). We deliberately do NOT fall back + to `trace_id`, because that would create a distinct "session" for every + single request and distort Arize's Session-grouping analytics. The + `trace_id` is still emitted under its own `litellm.trace_id` key so + spans remain filterable by trace. + + USER_ID is *only* emitted when no upstream path (model_params.user or + optional_params.user) has already set it, to avoid overwriting an + existing value with a possibly-different one from API-key metadata. + """ + if not isinstance(standard_logging_payload, dict): + return + metadata = standard_logging_payload.get("metadata") or {} + if not isinstance(metadata, dict): + return + + session_id = metadata.get("user_api_key_end_user_id") + if session_id: + safe_set_attribute(span, SpanAttributes.SESSION_ID, str(session_id)) + + trace_id = standard_logging_payload.get("trace_id") + if trace_id: + safe_set_attribute(span, "litellm.trace_id", str(trace_id)) + + optional_params = kwargs.get("optional_params") or {} + model_params = standard_logging_payload.get("model_parameters") or {} + has_user_already = bool( + (isinstance(optional_params, dict) and optional_params.get("user")) + or (isinstance(model_params, dict) and model_params.get("user")) + ) + if not has_user_already: + user_id = metadata.get("user_api_key_user_id") + if user_id: + safe_set_attribute(span, SpanAttributes.USER_ID, str(user_id)) + + team_id = metadata.get("user_api_key_team_id") + if team_id: + safe_set_attribute(span, "litellm.team_id", str(team_id)) + team_alias = metadata.get("user_api_key_team_alias") + if team_alias: + safe_set_attribute(span, "litellm.team_alias", str(team_alias)) + key_alias = metadata.get("user_api_key_alias") + if key_alias: + safe_set_attribute(span, "litellm.key_alias", str(key_alias)) + + +def _set_response_cost_attr(span: "Span", standard_logging_payload) -> None: + """Emit cost attributes from the StandardLoggingPayload when present. + + Uses the OpenInference `llm.cost.total` key so Arize / Phoenix can + surface the cost in their "Total Cost" column. LiteLLM only tracks a + single total in `StandardLoggingPayload.response_cost`, so we cannot + split it into prompt/completion. We also keep the legacy + `llm.response.cost` key for back-compat with any consumer querying it. + """ + if not isinstance(standard_logging_payload, dict): + return + cost = standard_logging_payload.get("response_cost") + if cost is None: + return + try: + cost_value = float(cost) + except (TypeError, ValueError): + return + safe_set_attribute(span, "llm.cost.total", cost_value) + safe_set_attribute(span, "llm.response.cost", cost_value) + + +def _is_passthrough_call_type(call_type: Optional[str]) -> bool: + if not call_type: + return False + lowered = str(call_type).lower() + return "passthrough" in lowered or "pass_through" in lowered + + +def _maybe_normalize_passthrough( + span: "Span", + kwargs: dict, + raw_response_obj, + coerced_response_obj, + standard_logging_payload, +) -> None: + """Surface input/output text for passthrough routes (e.g. Bedrock + InvokeModel) so the parent span renders as more than `usage` numbers. + + Only runs when `call_type` is a passthrough variant. Reads from: + - `kwargs["additional_args"]["complete_input_dict"]` for input + - the coerced response (or `kwargs["original_response"]`) for output + + All emits are best-effort: if the provider shape isn't recognized the + helper exits silently. Existing chat/completion paths never enter this + helper because their call_type doesn't contain "passthrough". + + TEMPORARY BRIDGE: passthrough handlers don't populate the + StandardLoggingPayload `messages` field today (they call + `transform_response(messages=[])`), so the input is only available via + `additional_args.complete_input_dict`. The proper fix is upstream in + `base_passthrough_logging_handler._create_response_logging_payload()`: + once that populates SLP `messages`/`response`, every callback gets + passthrough I/O (with central redaction) for free and this helper's + `complete_input_dict` fallback can be deleted. See follow-up issue. + """ + call_type = ( + standard_logging_payload.get("call_type") + if isinstance(standard_logging_payload, dict) + else None + ) + if not _is_passthrough_call_type(call_type): + return + + # Respect LiteLLM's central message-redaction contract. The normal + # chat/completion path is redacted by `perform_redaction` before + # callbacks run, but `complete_input_dict` (read below) is NOT covered by + # that layer — so without this gate, an operator who enabled redaction + # would still see raw passthrough prompts in Arize. Skip entirely when + # redaction is on so neither input nor output leaks through this bridge. + if should_redact_message_logging(kwargs): + return + + # --- INPUT -------------------------------------------------------------- + additional_args = kwargs.get("additional_args") or {} + complete_input_dict = ( + additional_args.get("complete_input_dict") + if isinstance(additional_args, dict) + else None + ) + if isinstance(complete_input_dict, dict): + _set_passthrough_input_attributes(span, complete_input_dict.get("messages")) + + # --- OUTPUT ------------------------------------------------------------- + parsed_response = _parse_passthrough_response( + raw_response_obj, coerced_response_obj, kwargs + ) + if not isinstance(parsed_response, dict): + return + + _set_passthrough_output_attributes(span, parsed_response) + + +def _set_passthrough_input_attributes(span: "Span", messages) -> None: + """Render passthrough request messages into INPUT_VALUE + LLM_INPUT_MESSAGES.""" + if not (isinstance(messages, list) and messages): + return + # Set INPUT_VALUE from the last user message text if discoverable. + last_text = None + for msg in reversed(messages): + if isinstance(msg, dict): + last_text = _coerce_text(msg.get("content")) + if last_text: + break + if last_text: + safe_set_attribute(span, SpanAttributes.INPUT_VALUE, last_text) + # Mirror messages into LLM_INPUT_MESSAGES so the input pane renders. + for idx, msg in enumerate(messages): + if not isinstance(msg, dict): + continue + prefix = f"{SpanAttributes.LLM_INPUT_MESSAGES}.{idx}" + role = msg.get("role") + if role: + safe_set_attribute( + span, + f"{prefix}.{MessageAttributes.MESSAGE_ROLE}", + role, + ) + text = _coerce_text(msg.get("content")) + if text is not None: + safe_set_attribute( + span, + f"{prefix}.{MessageAttributes.MESSAGE_CONTENT}", + text, + ) + + +def _set_passthrough_output_attributes(span: "Span", parsed_response: dict) -> None: + """Render passthrough response into OUTPUT_VALUE + LLM_OUTPUT_MESSAGES.""" + # Anthropic / Bedrock-Anthropic: `content` is a list of typed parts. + content_list = parsed_response.get("content") + if isinstance(content_list, list) and content_list: + texts = [] + for part in content_list: + if isinstance(part, dict) and isinstance(part.get("text"), str): + texts.append(part["text"]) + joined = "\n\n".join(t for t in texts if t) + if joined: + safe_set_attribute(span, SpanAttributes.OUTPUT_VALUE, joined) + prefix = f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.0" + safe_set_attribute( + span, + f"{prefix}.{MessageAttributes.MESSAGE_ROLE}", + parsed_response.get("role", "assistant"), + ) + safe_set_attribute( + span, + f"{prefix}.{MessageAttributes.MESSAGE_CONTENT}", + joined, + ) + + # OpenAI-style passthrough: `choices[0].message.content` + choices = parsed_response.get("choices") + if isinstance(choices, list) and choices: + first = choices[0] + if isinstance(first, dict): + msg = first.get("message") + if isinstance(msg, dict): + text = _coerce_text(msg.get("content")) + if text: + safe_set_attribute(span, SpanAttributes.OUTPUT_VALUE, text) + prefix = f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.0" + safe_set_attribute( + span, + f"{prefix}.{MessageAttributes.MESSAGE_ROLE}", + msg.get("role", "assistant"), + ) + safe_set_attribute( + span, + f"{prefix}.{MessageAttributes.MESSAGE_CONTENT}", + text, + ) + + +def _parse_passthrough_response(raw_response_obj, coerced_response_obj, kwargs): + """Return a dict view of the provider response for passthrough routes.""" + # Prefer the coerced view (already JSON-parsed for httpx.Response). + candidates = [] + if isinstance(coerced_response_obj, dict): + candidates.append(coerced_response_obj) + if ( + isinstance(raw_response_obj, dict) + and raw_response_obj is not coerced_response_obj + ): + candidates.append(raw_response_obj) + + for candidate in candidates: + # StandardPassThroughResponseObject wrapper: {"response": "..."}. + if ( + "response" in candidate + and "content" not in candidate + and "choices" not in candidate + ): + inner = candidate.get("response") + if isinstance(inner, str): + try: + parsed = json.loads(inner) + if isinstance(parsed, dict): + return parsed + except Exception: + continue + if isinstance(inner, dict): + return inner + else: + return candidate + + # Fallback: kwargs["original_response"] from the OTel base path. + original = kwargs.get("original_response") if isinstance(kwargs, dict) else None + if isinstance(original, dict): + return original + if isinstance(original, str): + try: + parsed = json.loads(original) + if isinstance(parsed, dict): + return parsed + except Exception: + return None + return None diff --git a/tests/test_litellm/integrations/arize/test_arize_utils.py b/tests/test_litellm/integrations/arize/test_arize_utils.py index 86c5448d468..83c3351319a 100644 --- a/tests/test_litellm/integrations/arize/test_arize_utils.py +++ b/tests/test_litellm/integrations/arize/test_arize_utils.py @@ -83,7 +83,12 @@ def test_arize_set_attributes(): # Apply attribute setting via ArizeLogger ArizeLogger.set_arize_attributes(span, kwargs, response_obj) - # Validate that the expected number of attributes were set + # Validate that the expected number of attributes were set. + # OPENINFERENCE_SPAN_KIND is written exactly once (defensively, before + # the main attribute pipeline) so a partial failure cannot blank it. + # Per the OpenInference spec, a chat completion that passes `tools=[...]` + # is still an LLM span — not TOOL (TOOL is reserved for actual tool + # execution by application code). assert span.set_attribute.call_count == 26 # Metadata attached to the span @@ -108,8 +113,15 @@ def test_arize_set_attributes(): # Response metadata span.set_attribute.assert_any_call("llm.response.id", "chatcmpl-ID") span.set_attribute.assert_any_call("llm.response.model", "gpt-4o") - # Span kind is set to TOOL when tools are present - span.set_attribute.assert_any_call(SpanAttributes.OPENINFERENCE_SPAN_KIND, "TOOL") + # Span kind stays LLM even when tools are passed (OpenInference spec). + span.set_attribute.assert_any_call(SpanAttributes.OPENINFERENCE_SPAN_KIND, "LLM") + # And TOOL must never be written for an LLM chat completion call. + span_kind_writes = [ + c.args[1] + for c in span.set_attribute.call_args_list + if c.args[0] == SpanAttributes.OPENINFERENCE_SPAN_KIND + ] + assert "TOOL" not in span_kind_writes # Request message content and metadata span.set_attribute.assert_any_call( @@ -451,3 +463,733 @@ def test_construct_dynamic_arize_headers(): dynamic_params_space_key_and_api_key ) expected_headers = {"arize-space-id": "test_space_key", "api_key": "test_api_key"} + + +# --------------------------------------------------------------------------- +# Additive rendering-enhancement tests. None of these assert that previously +# emitted attributes were removed or changed — they only assert that the new +# attributes appear in their respective scenarios. +# --------------------------------------------------------------------------- + + +def _collect_calls(span): + """Helper: return dict[attr_name] = value of all set_attribute calls.""" + out = {} + for call in span.set_attribute.call_args_list: + args = call.args + if len(args) >= 2: + out[args[0]] = args[1] + return out + + +def test_arize_emits_cache_tokens_openai_style(): + """OpenAI prompt_tokens_details.cached_tokens → cache_read attr.""" + from unittest.mock import MagicMock + + from litellm.integrations.arize._utils import _set_usage_outputs + + span = MagicMock() + response_obj = { + "usage": { + "total_tokens": 100, + "completion_tokens": 60, + "prompt_tokens": 40, + "prompt_tokens_details": {"cached_tokens": 32, "audio_tokens": 8}, + } + } + _set_usage_outputs(span, response_obj, SpanAttributes) + attrs = _collect_calls(span) + assert attrs[SpanAttributes.LLM_TOKEN_COUNT_PROMPT_DETAILS_CACHE_READ] == 32 + assert attrs[SpanAttributes.LLM_TOKEN_COUNT_PROMPT_DETAILS_AUDIO] == 8 + + +def test_arize_emits_cache_tokens_anthropic_style(): + """Anthropic/Bedrock cache_read_input_tokens / cache_creation_input_tokens.""" + from unittest.mock import MagicMock + + from litellm.integrations.arize._utils import _set_usage_outputs + + span = MagicMock() + response_obj = { + "usage": { + "input_tokens": 100, + "output_tokens": 50, + "cache_read_input_tokens": 80, + "cache_creation_input_tokens": 20, + } + } + _set_usage_outputs(span, response_obj, SpanAttributes) + attrs = _collect_calls(span) + assert attrs[SpanAttributes.LLM_TOKEN_COUNT_PROMPT_DETAILS_CACHE_READ] == 80 + assert attrs[SpanAttributes.LLM_TOKEN_COUNT_PROMPT_DETAILS_CACHE_WRITE] == 20 + + +def test_arize_emits_no_cache_tokens_when_absent(): + """Regression guard: when no cache fields exist, no cache attrs emitted.""" + from unittest.mock import MagicMock + + from litellm.integrations.arize._utils import _set_usage_outputs + + span = MagicMock() + response_obj = { + "usage": {"total_tokens": 10, "completion_tokens": 4, "prompt_tokens": 6} + } + _set_usage_outputs(span, response_obj, SpanAttributes) + attrs = _collect_calls(span) + assert SpanAttributes.LLM_TOKEN_COUNT_PROMPT_DETAILS_CACHE_READ not in attrs + assert SpanAttributes.LLM_TOKEN_COUNT_PROMPT_DETAILS_CACHE_WRITE not in attrs + + +def test_passthrough_call_type_resolves_to_llm_span_kind(): + """`allm_passthrough_route` should map to LLM (was UNKNOWN before fix).""" + from litellm.integrations._types.open_inference import OpenInferenceSpanKindValues + from litellm.integrations.arize._utils import _infer_open_inference_span_kind + + assert ( + _infer_open_inference_span_kind("allm_passthrough_route") + == OpenInferenceSpanKindValues.LLM.value + ) + assert ( + _infer_open_inference_span_kind("llm_passthrough_route") + == OpenInferenceSpanKindValues.LLM.value + ) + + +def test_arize_chat_completion_with_tools_stays_llm_span_kind(): + """Regression guard against the old `TOOL` override: a chat completion + that passes `tools=[...]` AND returns `tool_calls` must remain LLM.""" + from unittest.mock import MagicMock + + from litellm.types.utils import Choices, ModelResponse + + span = MagicMock() + kwargs = { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "weather?"}], + "standard_logging_object": { + "model_parameters": {}, + "metadata": {}, + "call_type": "completion", + }, + "optional_params": { + "tools": [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "weather", + "parameters": {"type": "object", "properties": {}}, + }, + } + ] + }, + "litellm_params": {"custom_llm_provider": "openai"}, + } + response_obj = ModelResponse( + usage={"total_tokens": 10, "completion_tokens": 4, "prompt_tokens": 6}, + choices=[ + Choices( + message={ + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "call_x", + "type": "function", + "function": {"name": "get_weather", "arguments": "{}"}, + } + ], + } + ) + ], + model="gpt-4o", + id="r-toolkind", + ) + + ArizeLogger.set_arize_attributes(span, kwargs, response_obj) + span_kind_writes = [ + c.args[1] + for c in span.set_attribute.call_args_list + if c.args[0] == SpanAttributes.OPENINFERENCE_SPAN_KIND + ] + assert span_kind_writes, "span.kind must be written" + assert all(v == "LLM" for v in span_kind_writes) + assert "TOOL" not in span_kind_writes + + +def test_arize_emits_assistant_tool_calls_on_output_message(): + """Assistant tool_calls should surface as MESSAGE_TOOL_CALLS.* attrs.""" + from unittest.mock import MagicMock + + from litellm.types.utils import Choices, ModelResponse + + span = MagicMock() + kwargs = { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "weather?"}], + "standard_logging_object": { + "model_parameters": {}, + "metadata": {}, + "call_type": "completion", + }, + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "openai"}, + } + response_obj = ModelResponse( + usage={"total_tokens": 10, "completion_tokens": 4, "prompt_tokens": 6}, + choices=[ + Choices( + message={ + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "call_abc", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{"location": "SF"}', + }, + } + ], + } + ) + ], + model="gpt-4o", + id="chatcmpl-1", + ) + ArizeLogger.set_arize_attributes(span, kwargs, response_obj) + attrs = _collect_calls(span) + base = f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.0.{MessageAttributes.MESSAGE_TOOL_CALLS}.0" + assert attrs[f"{base}.{ToolCallAttributes.TOOL_CALL_ID}"] == "call_abc" + assert ( + attrs[f"{base}.{ToolCallAttributes.TOOL_CALL_FUNCTION_NAME}"] == "get_weather" + ) + assert ( + attrs[f"{base}.{ToolCallAttributes.TOOL_CALL_FUNCTION_ARGUMENTS_JSON}"] + == '{"location": "SF"}' + ) + + +def test_arize_output_value_falls_back_to_tool_calls_summary(): + """When the assistant returns no text content but did request tool + calls, OUTPUT_VALUE should contain a JSON summary so Arize's Output + pane shows something.""" + from unittest.mock import MagicMock + + from litellm.types.utils import Choices, ModelResponse + + span = MagicMock() + kwargs = { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "weather?"}], + "standard_logging_object": { + "model_parameters": {}, + "metadata": {}, + "call_type": "completion", + }, + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "openai"}, + } + response_obj = ModelResponse( + usage={"total_tokens": 10, "completion_tokens": 4, "prompt_tokens": 6}, + choices=[ + Choices( + message={ + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "call_abc", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{"location": "SF"}', + }, + } + ], + } + ) + ], + model="gpt-4o", + id="r-tc-out", + ) + ArizeLogger.set_arize_attributes(span, kwargs, response_obj) + attrs = _collect_calls(span) + + # OUTPUT_VALUE should contain the tool_call name + arguments JSON + out = attrs[SpanAttributes.OUTPUT_VALUE] + assert "tool_calls" in out + assert "get_weather" in out + assert "SF" in out + + +def test_arize_output_value_unchanged_when_content_present(): + """Regression guard: when content is non-empty, OUTPUT_VALUE must be + exactly that content (no summary written).""" + from unittest.mock import MagicMock + + from litellm.types.utils import Choices, ModelResponse + + span = MagicMock() + kwargs = { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hi"}], + "standard_logging_object": { + "model_parameters": {}, + "metadata": {}, + "call_type": "completion", + }, + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "openai"}, + } + response_obj = ModelResponse( + usage={"total_tokens": 4, "completion_tokens": 2, "prompt_tokens": 2}, + choices=[ + Choices( + message={ + "role": "assistant", + "content": "hello world", + "tool_calls": [ + { + "id": "call_x", + "type": "function", + "function": {"name": "n", "arguments": "{}"}, + } + ], + } + ) + ], + model="gpt-4o", + id="r-content", + ) + ArizeLogger.set_arize_attributes(span, kwargs, response_obj) + attrs = _collect_calls(span) + assert attrs[SpanAttributes.OUTPUT_VALUE] == "hello world" + + +def test_arize_emits_tool_call_id_and_name_on_input_tool_message(): + """A tool-result input message should expose tool_call_id + name.""" + from unittest.mock import MagicMock + + from litellm.types.utils import Choices, ModelResponse + + span = MagicMock() + kwargs = { + "model": "gpt-4o", + "messages": [ + {"role": "user", "content": "weather?"}, + { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "call_abc", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{"location": "SF"}', + }, + } + ], + }, + { + "role": "tool", + "tool_call_id": "call_abc", + "name": "get_weather", + "content": "sunny, 72F", + }, + ], + "standard_logging_object": { + "model_parameters": {}, + "metadata": {}, + "call_type": "completion", + }, + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "openai"}, + } + response_obj = ModelResponse( + usage={"total_tokens": 10, "completion_tokens": 4, "prompt_tokens": 6}, + choices=[Choices(message={"role": "assistant", "content": "It's sunny."})], + model="gpt-4o", + id="chatcmpl-2", + ) + ArizeLogger.set_arize_attributes(span, kwargs, response_obj) + attrs = _collect_calls(span) + # Assistant tool_call surfaces on input msg index 1 + assistant_base = f"{SpanAttributes.LLM_INPUT_MESSAGES}.1.{MessageAttributes.MESSAGE_TOOL_CALLS}.0" + assert attrs[f"{assistant_base}.{ToolCallAttributes.TOOL_CALL_ID}"] == "call_abc" + # Tool message at index 2 + tool_prefix = f"{SpanAttributes.LLM_INPUT_MESSAGES}.2" + assert ( + attrs[f"{tool_prefix}.{MessageAttributes.MESSAGE_TOOL_CALL_ID}"] == "call_abc" + ) + assert attrs[f"{tool_prefix}.{MessageAttributes.MESSAGE_NAME}"] == "get_weather" + + +def test_arize_emits_multimodal_input_contents(): + """List-shaped content should populate MESSAGE_CONTENTS.* alongside the + legacy MESSAGE_CONTENT (which stays for back-compat).""" + from unittest.mock import MagicMock + + from litellm.types.utils import Choices, ModelResponse + + span = MagicMock() + kwargs = { + "model": "gpt-4o", + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is in this image?"}, + { + "type": "image_url", + "image_url": {"url": "https://example.com/cat.png"}, + }, + ], + } + ], + "standard_logging_object": { + "model_parameters": {}, + "metadata": {}, + "call_type": "completion", + }, + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "openai"}, + } + response_obj = ModelResponse( + usage={"total_tokens": 10, "completion_tokens": 4, "prompt_tokens": 6}, + choices=[Choices(message={"role": "assistant", "content": "A cat."})], + model="gpt-4o", + id="chatcmpl-img", + ) + ArizeLogger.set_arize_attributes(span, kwargs, response_obj) + attrs = _collect_calls(span) + base = f"{SpanAttributes.LLM_INPUT_MESSAGES}.0.{MessageAttributes.MESSAGE_CONTENTS}" + assert attrs[f"{base}.0.message_content.type"] == "text" + assert attrs[f"{base}.0.message_content.text"] == "What is in this image?" + assert attrs[f"{base}.1.message_content.type"] == "image" + assert ( + attrs[f"{base}.1.message_content.image.image.url"] + == "https://example.com/cat.png" + ) + + +def test_arize_emits_session_and_user_attrs_from_metadata(): + """end_user_id → SESSION_ID; user_api_key_user_id → USER_ID (only when + optional_params.user/model_params.user absent).""" + from unittest.mock import MagicMock + + from litellm.types.utils import Choices, ModelResponse + + span = MagicMock() + kwargs = { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hi"}], + "standard_logging_object": { + "model_parameters": {}, + "metadata": { + "user_api_key_user_id": "user_42", + "user_api_key_end_user_id": "session_99", + "user_api_key_team_id": "team_7", + "user_api_key_team_alias": "alpha", + "user_api_key_alias": "key_alpha", + }, + "call_type": "completion", + }, + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "openai"}, + } + response_obj = ModelResponse( + usage={"total_tokens": 4, "completion_tokens": 2, "prompt_tokens": 2}, + choices=[Choices(message={"role": "assistant", "content": "hello"})], + model="gpt-4o", + id="r1", + ) + ArizeLogger.set_arize_attributes(span, kwargs, response_obj) + attrs = _collect_calls(span) + assert attrs[SpanAttributes.SESSION_ID] == "session_99" + assert attrs[SpanAttributes.USER_ID] == "user_42" + assert attrs["litellm.team_id"] == "team_7" + assert attrs["litellm.team_alias"] == "alpha" + assert attrs["litellm.key_alias"] == "key_alpha" + + +def test_arize_does_not_use_trace_id_as_session_id_fallback(): + """SESSION_ID must NOT fall back to trace_id (one session-per-request + would distort Arize Session analytics). trace_id is emitted under its + own `litellm.trace_id` key instead. + """ + from unittest.mock import MagicMock + + from litellm.types.utils import Choices, ModelResponse + + span = MagicMock() + kwargs = { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hi"}], + "standard_logging_object": { + "model_parameters": {}, + "metadata": {}, + "call_type": "completion", + "trace_id": "trace-xyz-123", + }, + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "openai"}, + } + response_obj = ModelResponse( + usage={"total_tokens": 4, "completion_tokens": 2, "prompt_tokens": 2}, + choices=[Choices(message={"role": "assistant", "content": "hi"})], + model="gpt-4o", + id="r-trace", + ) + ArizeLogger.set_arize_attributes(span, kwargs, response_obj) + attrs = _collect_calls(span) + + # SESSION_ID must NOT be derived from trace_id. + assert SpanAttributes.SESSION_ID not in attrs + # trace_id surfaces under its own key. + assert attrs["litellm.trace_id"] == "trace-xyz-123" + + +def test_arize_does_not_overwrite_user_id_from_optional_params(): + """If optional_params.user is set, metadata USER_ID must NOT overwrite.""" + from unittest.mock import MagicMock + + from litellm.types.utils import Choices, ModelResponse + + span = MagicMock() + kwargs = { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hi"}], + "standard_logging_object": { + "model_parameters": {"user": "from_model_params"}, + "metadata": {"user_api_key_user_id": "from_metadata"}, + "call_type": "completion", + }, + "optional_params": {"user": "from_optional_params"}, + "litellm_params": {"custom_llm_provider": "openai"}, + } + response_obj = ModelResponse( + usage={"total_tokens": 4, "completion_tokens": 2, "prompt_tokens": 2}, + choices=[Choices(message={"role": "assistant", "content": "hello"})], + model="gpt-4o", + id="r2", + ) + ArizeLogger.set_arize_attributes(span, kwargs, response_obj) + user_id_writes = [ + c.args[1] + for c in span.set_attribute.call_args_list + if c.args[0] == SpanAttributes.USER_ID + ] + assert "from_metadata" not in user_id_writes + + +def test_arize_emits_response_cost(): + """StandardLoggingPayload.response_cost → llm.cost.total (+ legacy llm.response.cost).""" + from unittest.mock import MagicMock + + from litellm.types.utils import Choices, ModelResponse + + span = MagicMock() + kwargs = { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hi"}], + "standard_logging_object": { + "model_parameters": {}, + "metadata": {}, + "call_type": "completion", + "response_cost": 0.0012345, + }, + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "openai"}, + } + response_obj = ModelResponse( + usage={"total_tokens": 4, "completion_tokens": 2, "prompt_tokens": 2}, + choices=[Choices(message={"role": "assistant", "content": "hello"})], + model="gpt-4o", + id="r3", + ) + ArizeLogger.set_arize_attributes(span, kwargs, response_obj) + attrs = _collect_calls(span) + assert attrs["llm.cost.total"] == 0.0012345 + assert attrs["llm.response.cost"] == 0.0012345 # legacy key still emitted + + +def test_arize_passthrough_bedrock_anthropic_normalization(): + """Bedrock-Anthropic passthrough: input/output text must be set so the + span renders something other than raw provider attrs.""" + from unittest.mock import MagicMock + + span = MagicMock() + bedrock_response_body = { + "id": "msg_01", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "The capital of France is Paris."}], + "model": "anthropic.claude-sonnet-4-v1:0", + "stop_reason": "end_turn", + "usage": {"input_tokens": 18, "output_tokens": 12}, + } + + class FakeHttpxResponse: + """Minimal httpx.Response stand-in: has `.text` and no `.get`.""" + + def __init__(self, body): + self.text = json.dumps(body) + + response_obj = FakeHttpxResponse(bedrock_response_body) + kwargs = { + "model": "anthropic.claude-sonnet-4-v1:0", + "messages": [ + { + "role": "user", + "content": json.dumps({"messages": [{"role": "user", "content": "?"}]}), + } + ], + "additional_args": { + "complete_input_dict": { + "anthropic_version": "bedrock-2023-05-31", + "max_tokens": 64, + "messages": [ + {"role": "user", "content": "What is the capital of France?"} + ], + } + }, + "standard_logging_object": { + "model_parameters": {}, + "metadata": {}, + "call_type": "allm_passthrough_route", + }, + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "bedrock"}, + } + ArizeLogger.set_arize_attributes(span, kwargs, response_obj) + attrs = _collect_calls(span) + + # Input rendering + assert attrs[SpanAttributes.INPUT_VALUE] == "What is the capital of France?" + msg0 = f"{SpanAttributes.LLM_INPUT_MESSAGES}.0" + assert attrs[f"{msg0}.{MessageAttributes.MESSAGE_ROLE}"] == "user" + assert ( + attrs[f"{msg0}.{MessageAttributes.MESSAGE_CONTENT}"] + == "What is the capital of France?" + ) + + # Output rendering (Anthropic content[].text) + assert attrs[SpanAttributes.OUTPUT_VALUE] == "The capital of France is Paris." + out0 = f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.0" + assert attrs[f"{out0}.{MessageAttributes.MESSAGE_ROLE}"] == "assistant" + assert ( + attrs[f"{out0}.{MessageAttributes.MESSAGE_CONTENT}"] + == "The capital of France is Paris." + ) + + # Token counts (Bedrock input_tokens/output_tokens) — extracted via + # coercion of the non-dict response. + assert attrs[SpanAttributes.LLM_TOKEN_COUNT_PROMPT] == 18 + assert attrs[SpanAttributes.LLM_TOKEN_COUNT_COMPLETION] == 12 + + # Span kind defended even though the call_type is a passthrough variant. + span_kind_writes = [ + c.args[1] + for c in span.set_attribute.call_args_list + if c.args[0] == SpanAttributes.OPENINFERENCE_SPAN_KIND + ] + assert span_kind_writes # at least one + assert all(v == "LLM" for v in span_kind_writes) + + +def test_arize_passthrough_call_type_does_not_run_on_chat_completion(): + """Guard: passthrough normalizer must not fire for normal chat calls. + + If it did, it could double-write input/output for ordinary completions. + """ + from unittest.mock import MagicMock + + from litellm.integrations.arize._utils import _maybe_normalize_passthrough + + span = MagicMock() + _maybe_normalize_passthrough( + span, + { + "additional_args": { + "complete_input_dict": {"messages": [{"role": "user", "content": "x"}]} + } + }, + {"choices": [{"message": {"role": "assistant", "content": "y"}}]}, + {"choices": [{"message": {"role": "assistant", "content": "y"}}]}, + {"call_type": "completion"}, + ) + assert span.set_attribute.call_count == 0 + + +def test_arize_passthrough_skipped_when_message_redaction_enabled(): + """Security guard: when message-logging redaction is enabled, the + passthrough normalizer must NOT export the raw prompt (read from + `complete_input_dict`, which bypasses central redaction) to the span. + """ + from unittest.mock import MagicMock + + from litellm.integrations.arize._utils import _maybe_normalize_passthrough + + span = MagicMock() + kwargs = { + "additional_args": { + "complete_input_dict": { + "messages": [ + {"role": "user", "content": "Patient John Doe, SSN 123-45-6789"} + ] + } + }, + # Enables redaction via the dynamic-param path inside + # should_redact_message_logging(), without touching globals. + "standard_callback_dynamic_params": {"turn_off_message_logging": True}, + } + _maybe_normalize_passthrough( + span, + kwargs, + {"content": [{"type": "text", "text": "secret response"}]}, + {"content": [{"type": "text", "text": "secret response"}]}, + {"call_type": "allm_passthrough_route"}, + ) + # Nothing — neither input nor output — should be written to the span. + assert span.set_attribute.call_count == 0 + + +def test_arize_coerce_response_obj_passes_dicts_through_untouched(): + """Regression guard for the BaseModel/dict path.""" + from litellm.integrations.arize._utils import _coerce_response_obj_for_attrs + + d = {"id": "x", "model": "m"} + assert _coerce_response_obj_for_attrs(d) is d + + class HasGet: + def get(self, *a, **k): # noqa: D401 + return None + + obj = HasGet() + assert _coerce_response_obj_for_attrs(obj) is obj + + assert _coerce_response_obj_for_attrs(None) is None + + +def test_arize_coerce_response_obj_parses_httpx_like(): + """httpx.Response-like objects without `.get` should JSON-decode.""" + from litellm.integrations.arize._utils import _coerce_response_obj_for_attrs + + class FakeHttpxResponse: + text = '{"id": "msg_1", "model": "claude"}' + + parsed = _coerce_response_obj_for_attrs(FakeHttpxResponse()) + assert parsed == {"id": "msg_1", "model": "claude"} + + +def test_arize_coerce_response_obj_returns_original_on_bad_json(): + from litellm.integrations.arize._utils import _coerce_response_obj_for_attrs + + class BadJson: + text = "not-json" + + obj = BadJson() + assert _coerce_response_obj_for_attrs(obj) is obj From 2bbdbfa5c348e198eb21461731c784e12897f01f Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 3 Jun 2026 12:13:02 -0700 Subject: [PATCH 03/11] fix: passthrough endpoints duplicate logs (#29598) * fix duplicate cost callbacks for anthropic streaming pass-through Two bugs caused _PROXY_track_cost_callback to see stream=True + complete_streaming_response=None on every streaming pass-through request, making the dedup guard in dispatch_success_handlers permanently inactive: 1. pass_through_endpoints.py created the Logging object with stream=False for all requests. _is_assembled_stream_success short-circuits on self.stream is not True, so has_dispatched_final_stream_success was never set and any second dispatch went through unchecked. Fix: set logging_obj.stream = True after stream detection. 2. _create_anthropic_response_logging_payload set complete_streaming_response inside the try block after litellm.completion_cost(), so a pricing error caused an early return without setting it on model_call_details. Fix: set complete_streaming_response before the try block. Co-Authored-By: Claude Sonnet 4.6 * fix stream * add stream to logging obj * test(pass_through): give mock logging object a real model_call_details dict The anthropic passthrough logging payload now records the assembled response on model_call_details before cost calculation, which requires model_call_details to support item assignment. In production it is always a dict; the existing unit test stubbed the logging object with a bare Mock whose attribute is not subscriptable, so the new assignment raised TypeError. Use a real dict to match the production logging object. * test(pass_through): cover streaming logging-obj stream flag The streaming branch of pass_through_request that marks the logging object as streaming (logging_obj.stream and model_call_details["stream"]) had no unit coverage, so the patch coverage gate flagged it. Add a regression test that drives a streaming pass-through request through pass_through_request and asserts the logging object is flagged as a stream before dispatch. * test(pass_through): cover SSE-response stream flag fallback branch The auto-detected streaming branch of pass_through_request (when a request that was not flagged as streaming returns a text/event-stream response) sets logging_obj.stream and model_call_details["stream"] but had no unit coverage, so the codecov patch gate failed at 60%. Drive a non-streaming pass-through request whose upstream response is SSE through pass_through_request and assert the logging object is flagged as a stream before dispatch. * fix(pass_through): gate complete_streaming_response on stream flag perform_redaction only scrubs complete_streaming_response when model_call_details["stream"] is True. Setting it unconditionally for non-streaming Anthropic pass-through responses left the assembled response unredacted in model_call_details, which is handed to logging callbacks as kwargs when message logging is disabled. Only record it for actual streaming responses so redaction always applies. --------- Co-authored-by: mubashir1osmani Co-authored-by: Claude Sonnet 4.6 --- .../anthropic_passthrough_logging_handler.py | 7 + .../pass_through_endpoints.py | 6 + .../test_unit_test_anthropic_pass_through.py | 1 + ...t_anthropic_passthrough_logging_handler.py | 332 ++++++++++++++++++ .../test_pass_through_endpoints.py | 125 +++++++ 5 files changed, 471 insertions(+) diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py index 3be26eb572d..a94672f9487 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py @@ -114,6 +114,13 @@ class AnthropicPassthroughLoggingHandler: handles streaming and non-streaming responses """ + # Only record complete_streaming_response for actual streaming responses. + # perform_redaction scrubs this field only when stream is True, so setting + # it on a non-streaming response would bypass message redaction. + if logging_obj.model_call_details.get("stream") is True: + logging_obj.model_call_details["complete_streaming_response"] = ( + litellm_model_response + ) try: # Get custom_llm_provider from logging object if available (e.g., azure_ai for Azure Anthropic) custom_llm_provider = logging_obj.model_call_details.get( diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 985785ad77e..e49e1302ab5 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -1061,6 +1061,9 @@ async def pass_through_request( # noqa: PLR0915 ) if stream: + logging_obj.stream = True + logging_obj.model_call_details["stream"] = True + if is_multipart: response = ( await HttpPassThroughEndpointHelpers.make_multipart_http_request( @@ -1139,6 +1142,9 @@ async def pass_through_request( # noqa: PLR0915 verbose_proxy_logger.debug("response.headers= %s", response.headers) if _is_streaming_response(response) is True: + logging_obj.stream = True + logging_obj.model_call_details["stream"] = True + try: response.raise_for_status() except httpx.HTTPStatusError as e: diff --git a/tests/pass_through_unit_tests/test_unit_test_anthropic_pass_through.py b/tests/pass_through_unit_tests/test_unit_test_anthropic_pass_through.py index 455c72ff636..5ab0319da47 100644 --- a/tests/pass_through_unit_tests/test_unit_test_anthropic_pass_through.py +++ b/tests/pass_through_unit_tests/test_unit_test_anthropic_pass_through.py @@ -318,6 +318,7 @@ def test_handle_logging_anthropic_collected_chunks(all_chunks): from litellm.types.utils import ModelResponse litellm_logging_obj = Mock() + litellm_logging_obj.model_call_details = {} pass_through_logging_obj = Mock() sent_args = { diff --git a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py index 0a9e3031030..f8b6fbde3dc 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py @@ -1043,3 +1043,335 @@ class TestPureTextFastPathParity: AnthropicPassthroughLoggingHandler._collapse_pure_text_chunks(all_chunks) is None ) + + +class TestStreamFalseDeduplication: + """ + Regression tests for the duplicate-callback bug where a streaming pass-through + request had stream=False hardcoded on its Logging object. + + Before the fix: + - logging_obj.stream was always False for pass-through requests + - _is_assembled_stream_success() checked `self.stream is not True` and returned + False immediately, so has_dispatched_final_stream_success was never set + - Any second dispatch_success_handlers call went through unchecked + + After the fix: + - pass_through_endpoints.py sets logging_obj.stream = True after detecting stream + - _create_anthropic_response_logging_payload sets complete_streaming_response on + model_call_details so callbacks see the correct assembled response state + - _is_assembled_stream_success returns True, dedup guard fires on first dispatch + """ + + @staticmethod + def _sse(event, data): + return f"event: {event}\ndata: {json.dumps(data)}\n\n".encode() + + @staticmethod + def _make_logging_obj(stream: bool = False) -> LiteLLMLoggingObj: + logging_obj = LiteLLMLoggingObj( + model="claude-3-5-sonnet-20241022", + messages=[{"role": "user", "content": "hello"}], + stream=stream, + call_type="pass_through_endpoint", + start_time=datetime.now(), + litellm_call_id="test-call-id", + function_id="1245", + ) + return logging_obj + + @staticmethod + def _build_chunks(): + frames = [ + TestStreamFalseDeduplication._sse( + "message_start", + { + "type": "message_start", + "message": { + "id": "msg_abc", + "type": "message", + "role": "assistant", + "model": "claude-3-5-sonnet-20241022", + "content": [], + "stop_reason": None, + "stop_sequence": None, + "usage": {"input_tokens": 10, "output_tokens": 0}, + }, + }, + ), + TestStreamFalseDeduplication._sse( + "content_block_start", + { + "type": "content_block_start", + "index": 0, + "content_block": {"type": "text", "text": ""}, + }, + ), + TestStreamFalseDeduplication._sse( + "content_block_delta", + { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "text_delta", "text": "Hello"}, + }, + ), + TestStreamFalseDeduplication._sse( + "content_block_stop", {"type": "content_block_stop", "index": 0} + ), + TestStreamFalseDeduplication._sse( + "message_delta", + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn", "stop_sequence": None}, + "usage": {"output_tokens": 5}, + }, + ), + TestStreamFalseDeduplication._sse("message_stop", {"type": "message_stop"}), + ] + from litellm.proxy.pass_through_endpoints.streaming_handler import ( + PassThroughStreamingHandler, + ) + + return PassThroughStreamingHandler._convert_raw_bytes_to_str_lines(frames) + + def test_complete_streaming_response_set_on_model_call_details(self): + """ + After the fix, _create_anthropic_response_logging_payload must set + complete_streaming_response on logging_obj.model_call_details so that + callbacks like _PROXY_track_cost_callback see the assembled response + instead of None. + + Before the fix: model_call_details had no complete_streaming_response key. + The log showed: "kwargs stream: True + complete streaming response: None" + """ + from litellm.types.passthrough_endpoints.pass_through_endpoints import ( + EndpointType, + ) + + # pass_through_request sets the stream flag before the streaming handler + # reconstructs the response; mirror that here. + logging_obj = self._make_logging_obj(stream=True) + logging_obj.model_call_details["stream"] = True + all_chunks = list(self._build_chunks()) + + result = AnthropicPassthroughLoggingHandler._handle_logging_anthropic_collected_chunks( + litellm_logging_obj=logging_obj, + passthrough_success_handler_obj=MagicMock(), + url_route="/anthropic/v1/messages", + request_body={"model": "claude-3-5-sonnet-20241022", "stream": True}, + endpoint_type=EndpointType.ANTHROPIC, + start_time=datetime.now(), + all_chunks=all_chunks, + end_time=datetime.now(), + ) + + # The assembled response must be stored on model_call_details so callbacks + # can identify this as a completed streaming call, not an in-progress one. + assert ( + logging_obj.model_call_details.get("complete_streaming_response") + is not None + ), "complete_streaming_response must be set on model_call_details after assembly" + + # The returned result must match what was stored + assert result["result"] is logging_obj.model_call_details.get( + "complete_streaming_response" + ) + + def test_dedup_guard_fires_when_stream_true_on_logging_obj(self): + """ + When logging_obj.stream is True (set by pass_through_endpoints.py after + detecting a streaming request), dispatch_success_handlers must set + has_dispatched_final_stream_success=True on the first call so that any + second call is a no-op. + + This is the _is_assembled_stream_success gate: with stream=False it + always returned False and the guard was permanently disabled. + """ + from litellm.types.passthrough_endpoints.pass_through_endpoints import ( + EndpointType, + ) + from litellm.types.utils import ModelResponse + + # Simulate what pass_through_endpoints.py now does after stream detection + logging_obj = self._make_logging_obj(stream=False) + logging_obj.stream = True # fix applied + logging_obj.model_call_details["stream"] = True + + # Simulate what _create_anthropic_response_logging_payload now does + mock_response = ModelResponse(model="claude-3-5-sonnet-20241022") + logging_obj.model_call_details["complete_streaming_response"] = mock_response + + assert logging_obj._is_assembled_stream_success(result=mock_response) is True + + # First dispatch sets the flag + assert not logging_obj.model_call_details.get( + "has_dispatched_final_stream_success" + ) + logging_obj.model_call_details["has_dispatched_final_stream_success"] = True + + # Second dispatch would be blocked — simulate the guard check + would_skip = bool( + logging_obj._is_assembled_stream_success(result=mock_response) + and logging_obj.model_call_details.get( + "has_dispatched_final_stream_success" + ) + ) + assert would_skip is True, ( + "Dedup guard must block a second dispatch_success_handlers call for the " + "same assembled streaming response" + ) + + def test_sse_fallback_path_sets_stream_true_for_dedup(self): + """ + When a nominally non-streaming request receives an SSE response + (_is_streaming_response returns True), the fallback branch in + pass_through_endpoints.py must set logging_obj.stream = True so the + dedup guard activates. + + Before the fix the fallback path never set stream=True, so + _is_assembled_stream_success always returned False and duplicate + callback dispatches were never blocked. + """ + from litellm.types.utils import ModelResponse + + # logging_obj starts with stream=False, as created before the request + logging_obj = self._make_logging_obj(stream=False) + assert logging_obj._is_assembled_stream_success(result=MagicMock()) is False + + # Simulate what the SSE fallback branch in pass_through_endpoints.py now does + logging_obj.stream = True + logging_obj.model_call_details["stream"] = True + + mock_response = ModelResponse(model="claude-3-5-sonnet-20241022") + logging_obj.model_call_details["complete_streaming_response"] = mock_response + + # With stream=True the dedup guard must be active + assert logging_obj._is_assembled_stream_success(result=mock_response) is True + + logging_obj.model_call_details["has_dispatched_final_stream_success"] = True + + would_skip = bool( + logging_obj._is_assembled_stream_success(result=mock_response) + and logging_obj.model_call_details.get( + "has_dispatched_final_stream_success" + ) + ) + assert would_skip is True + + def test_stream_false_logging_obj_bypasses_dedup_guard(self): + """ + Demonstrates the pre-fix state: with stream=False on the logging object, + _is_assembled_stream_success always returns False regardless of whether + complete_streaming_response is set. This means the dedup guard can never + fire, so duplicate dispatches go through unchecked. + + This test documents the old broken behavior so the fix is clearly justified. + """ + from litellm.types.utils import ModelResponse + + logging_obj = self._make_logging_obj(stream=False) + mock_response = ModelResponse(model="claude-3-5-sonnet-20241022") + logging_obj.model_call_details["complete_streaming_response"] = mock_response + + # With stream=False, _is_assembled_stream_success returns False even though + # complete_streaming_response is present — the guard is permanently disabled. + assert logging_obj._is_assembled_stream_success(result=mock_response) is False + + +class TestNonStreamingResponseRedaction: + """ + Regression tests ensuring _create_anthropic_response_logging_payload only sets + complete_streaming_response for streaming responses. perform_redaction scrubs + that field exclusively when model_call_details["stream"] is True, so storing it + on a non-streaming response would deliver the unredacted response to logging + callbacks when message logging is disabled. + """ + + @staticmethod + def _make_logging_obj(stream: bool) -> LiteLLMLoggingObj: + logging_obj = LiteLLMLoggingObj( + model="claude-3-5-sonnet-20241022", + messages=[{"role": "user", "content": "hello"}], + stream=stream, + call_type="pass_through_endpoint", + start_time=datetime.now(), + litellm_call_id="test-call-id", + function_id="1245", + ) + # pass_through_request mirrors the stream flag onto model_call_details, + # which is the key perform_redaction inspects. + logging_obj.model_call_details["stream"] = stream + return logging_obj + + def test_non_streaming_does_not_set_complete_streaming_response(self): + from litellm.types.utils import ModelResponse + + logging_obj = self._make_logging_obj(stream=False) + response = ModelResponse(model="claude-3-5-sonnet-20241022") + + AnthropicPassthroughLoggingHandler._create_anthropic_response_logging_payload( + litellm_model_response=response, + model="claude-3-5-sonnet-20241022", + kwargs={}, + start_time=datetime.now(), + end_time=datetime.now(), + logging_obj=logging_obj, + ) + + assert ( + "complete_streaming_response" not in logging_obj.model_call_details + ), "non-streaming responses must not populate complete_streaming_response" + + def test_streaming_sets_complete_streaming_response(self): + from litellm.types.utils import ModelResponse + + logging_obj = self._make_logging_obj(stream=True) + response = ModelResponse(model="claude-3-5-sonnet-20241022") + + AnthropicPassthroughLoggingHandler._create_anthropic_response_logging_payload( + litellm_model_response=response, + model="claude-3-5-sonnet-20241022", + kwargs={}, + start_time=datetime.now(), + end_time=datetime.now(), + logging_obj=logging_obj, + ) + + assert ( + logging_obj.model_call_details.get("complete_streaming_response") + is response + ) + + def test_non_streaming_response_is_redacted_when_message_logging_off(self): + from litellm.litellm_core_utils.redact_messages import ( + redact_message_input_output_from_logging, + ) + from litellm.types.utils import Choices, Message, ModelResponse + + logging_obj = self._make_logging_obj(stream=False) + response = ModelResponse( + model="claude-3-5-sonnet-20241022", + choices=[Choices(message=Message(role="assistant", content="secret"))], + ) + + AnthropicPassthroughLoggingHandler._create_anthropic_response_logging_payload( + litellm_model_response=response, + model="claude-3-5-sonnet-20241022", + kwargs={}, + start_time=datetime.now(), + end_time=datetime.now(), + logging_obj=logging_obj, + ) + + logging_obj.model_call_details["litellm_params"] = { + "metadata": {"headers": {"x-litellm-enable-message-redaction": True}} + } + + redacted = redact_message_input_output_from_logging( + model_call_details=logging_obj.model_call_details, + result=response, + ) + + leaked = logging_obj.model_call_details.get("complete_streaming_response") + assert leaked is None + assert redacted.choices[0].message.content == "redacted-by-litellm" diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 61299e2662a..89c57bcced6 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -1050,6 +1050,131 @@ async def test_pass_through_request_contains_proxy_server_request_in_kwargs(): assert metadata["user_api_key_user_id"] == "test-user-id" +@pytest.mark.asyncio +async def test_pass_through_request_streaming_marks_logging_obj_as_stream(): + """ + Regression: a streaming pass-through request must flag its logging object as + streaming (logging_obj.stream and model_call_details["stream"]) before the + response is dispatched, so cost/success callbacks treat it as a stream and the + streaming dedup guard fires instead of double-logging. + """ + with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging: + with patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client" + ) as mock_get_client: + with patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.PassThroughStreamingHandler.chunk_processor" + ) as mock_chunk_processor: + mock_proxy_logging.pre_call_hook = AsyncMock( + return_value={"model": "claude-3", "stream": True} + ) + mock_proxy_logging.post_call_failure_hook = AsyncMock() + + upstream_response = MagicMock() + upstream_response.status_code = 200 + upstream_response.headers = {} + upstream_response.raise_for_status = MagicMock() + + async_client = MagicMock() + async_client.build_request = MagicMock(return_value=MagicMock()) + async_client.send = AsyncMock(return_value=upstream_response) + mock_get_client.return_value = MagicMock(client=async_client) + + async def _empty_chunks(*args, **kwargs): + return + yield # pragma: no cover + + mock_chunk_processor.return_value = _empty_chunks() + + mock_request = MagicMock(spec=Request) + mock_request.method = "POST" + mock_request.url = "http://test-proxy.com/v1/messages" + mock_request.body = AsyncMock( + return_value=b'{"model": "claude-3", "stream": true}' + ) + mock_request.headers = Headers({}) + mock_request.query_params = QueryParams({}) + + await pass_through_request( + request=mock_request, + target="http://target-api.com/v1/messages", + custom_headers={}, + user_api_key_dict=MagicMock(), + stream=True, + ) + + async_client.send.assert_awaited_once() + assert async_client.send.call_args.kwargs["stream"] is True + + mock_chunk_processor.assert_called_once() + logging_obj = mock_chunk_processor.call_args.kwargs[ + "litellm_logging_obj" + ] + assert logging_obj.stream is True + assert logging_obj.model_call_details["stream"] is True + + +@pytest.mark.asyncio +async def test_pass_through_request_sse_response_marks_logging_obj_as_stream(): + """ + Regression: a request that is not flagged as streaming up front but whose + upstream response comes back as an SSE stream (content-type text/event-stream) + must still flag its logging object as streaming before dispatch. Otherwise the + cost/success callbacks treat the assembled stream as a non-stream and the dedup + guard never fires, double-logging the request. + """ + with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging: + with patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client" + ) as mock_get_client: + with patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.PassThroughStreamingHandler.chunk_processor" + ) as mock_chunk_processor: + mock_proxy_logging.pre_call_hook = AsyncMock( + return_value={"model": "claude-3"} + ) + mock_proxy_logging.post_call_failure_hook = AsyncMock() + + upstream_response = MagicMock() + upstream_response.status_code = 200 + upstream_response.headers = {"content-type": "text/event-stream"} + upstream_response.raise_for_status = MagicMock() + + async_client = MagicMock() + async_client.request = AsyncMock(return_value=upstream_response) + mock_get_client.return_value = MagicMock(client=async_client) + + async def _empty_chunks(*args, **kwargs): + return + yield # pragma: no cover + + mock_chunk_processor.return_value = _empty_chunks() + + mock_request = MagicMock(spec=Request) + mock_request.method = "POST" + mock_request.url = "http://test-proxy.com/v1/messages" + mock_request.body = AsyncMock(return_value=b'{"model": "claude-3"}') + mock_request.headers = Headers({}) + mock_request.query_params = QueryParams({}) + + await pass_through_request( + request=mock_request, + target="http://target-api.com/v1/messages", + custom_headers={}, + user_api_key_dict=MagicMock(), + stream=False, + ) + + async_client.request.assert_awaited_once() + + mock_chunk_processor.assert_called_once() + logging_obj = mock_chunk_processor.call_args.kwargs[ + "litellm_logging_obj" + ] + assert logging_obj.stream is True + assert logging_obj.model_call_details["stream"] is True + + @pytest.mark.asyncio async def test_create_pass_through_endpoint(): """ From 84969aaf15fa4c510279dde68b511a82cfdd49f8 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 3 Jun 2026 13:37:53 -0700 Subject: [PATCH 04/11] fix(ci): keep coverage rename green when a parallel node runs no tests (#29608) * fix(ci): keep coverage rename green when a parallel node runs no tests local_testing_part1 and local_testing_part2 run with parallelism 4. When CircleCI reruns only the failed tests, the failed test lands on a single node and the other nodes receive an empty bucket, so pytest never writes coverage.xml or .coverage. The unguarded "mv coverage.xml ..." then exits 1 and turns the whole job red even though the rerun passed; the next persist_to_workspace step would fail the same way on the missing paths. Guard the rename so a node with no coverage emits empty placeholders instead. coverage combine tolerates the empty files, so the downstream upload-coverage job keeps the real nodes' data intact. * fix(ci): pre-create test-results in litellm_router_testing for empty-bucket reruns litellm_router_testing also runs with parallelism 4. On a rerun of only the failed tests, a node can receive no tests, so the test command never creates test-results and the final store_test_results step can fail on the missing path. Pre-create the directory up front, matching what local_testing_part1 and part2 already do and CircleCI's own guidance for parallel reruns. * test(openai): retry wildcard chat completion on transient OpenAI 500 build_and_test reddened on test_openai_wildcard_chat_completion when the real gpt-3.5-turbo-0125 call returned an OpenAI 500 ("The server had an error while processing your request"). The base branch passed the same call concurrently, so the 500 is an intermittent OpenAI server error, not a regression. Add the same pytest-retry marker the sibling real-call tests in this file already use so a transient upstream 500 no longer fails CI. --- .circleci/config.yml | 27 +++++++++++++++++++++++---- tests/test_openai_endpoints.py | 1 + 2 files changed, 24 insertions(+), 4 deletions(-) diff --git a/.circleci/config.yml b/.circleci/config.yml index f5a8fe77a25..6ee5634f54f 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -249,8 +249,15 @@ jobs: - run: name: Rename the coverage files command: | - mv coverage.xml local_testing_part1_coverage.xml - mv .coverage local_testing_part1_coverage + # When CI reruns only the failed tests, a parallel node can receive + # zero tests and pytest never writes coverage. Emit empty placeholders + # so persist_to_workspace and the downstream coverage combine stay green. + if [ -f coverage.xml ]; then + mv coverage.xml local_testing_part1_coverage.xml + mv .coverage local_testing_part1_coverage + else + touch local_testing_part1_coverage.xml local_testing_part1_coverage + fi # Store test results - store_test_results: @@ -314,8 +321,15 @@ jobs: - run: name: Rename the coverage files command: | - mv coverage.xml local_testing_part2_coverage.xml - mv .coverage local_testing_part2_coverage + # When CI reruns only the failed tests, a parallel node can receive + # zero tests and pytest never writes coverage. Emit empty placeholders + # so persist_to_workspace and the downstream coverage combine stay green. + if [ -f coverage.xml ]; then + mv coverage.xml local_testing_part2_coverage.xml + mv .coverage local_testing_part2_coverage + else + touch local_testing_part2_coverage.xml local_testing_part2_coverage + fi # Store test results - store_test_results: @@ -464,6 +478,11 @@ jobs: - run: name: Run tests command: | + # On a "rerun failed tests" build a parallel node can receive no + # tests, so the test command never creates test-results. Pre-create it + # so store_test_results doesn't fail the node on a missing path. + mkdir -p test-results + TEST_FILES=$(circleci tests glob "tests/local_testing/**/test_*.py") echo "$TEST_FILES" | circleci tests run \ diff --git a/tests/test_openai_endpoints.py b/tests/test_openai_endpoints.py index e898b88a556..880cc1ebbef 100644 --- a/tests/test_openai_endpoints.py +++ b/tests/test_openai_endpoints.py @@ -522,6 +522,7 @@ async def test_image_generation(): await image_generation(session=session, key=key_2) +@pytest.mark.flaky(retries=5, delay=1) @pytest.mark.asyncio async def test_openai_wildcard_chat_completion(): """ From b4aee2c7ddb4d9f01c6359065d89ed928f71ee6b Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 3 Jun 2026 13:46:43 -0700 Subject: [PATCH 05/11] test(vcr): close out the remaining VCR live-call leaks (#29603) * Fix remaining VCR live-call leaks * test(vcr): dedupe live-test helpers and drop spurious kwargs Extract the duplicated isVertexQuotaError/runVertexRequestOrSkip Vertex quota-skip helpers into tests/pass_through_tests/vertex_test_helpers.js and the duplicated _skip_live_prompt_caching_test guard into tests/_live_test_helpers.py so each lives in one place. In test_aarun_thread_litellm, build a separate message_data carrying role/content for add_message and a thread_data without them for run_thread/run_thread_stream/get_messages, which no longer receive the spurious message fields. * test(overhead): assert mock transport is exercised in non-streaming and stream tests --- tests/_live_test_helpers.py | 10 + tests/_vcr_conftest_common.py | 19 + tests/litellm_utils_tests/conftest.py | 27 +- .../test_aws_secret_manager.py | 11 +- .../test_litellm_overhead.py | 342 +++++----- tests/llm_translation/base_llm_unit_tests.py | 5 + tests/llm_translation/conftest.py | 8 +- .../test_bedrock_completion.py | 5 + .../test_bedrock_invoke_tests.py | 10 +- tests/local_testing/conftest.py | 3 - tests/local_testing/test_assistants.py | 601 ++++++++++-------- tests/logging_callback_tests/conftest.py | 9 +- .../test_amazing_s3_logs.py | 58 +- tests/ocr_tests/conftest.py | 22 +- tests/ocr_tests/test_ocr_vertex_ai.py | 8 + tests/pass_through_tests/test_vertex.test.js | 18 +- tests/pass_through_tests/test_vertex_ai.py | 17 +- .../test_vertex_with_spend.test.js | 18 +- .../pass_through_tests/vertex_test_helpers.js | 27 + ..._anthropic_messages_prompt_caching_test.py | 15 +- tests/pass_through_unit_tests/conftest.py | 11 +- 21 files changed, 702 insertions(+), 542 deletions(-) create mode 100644 tests/_live_test_helpers.py create mode 100644 tests/pass_through_tests/vertex_test_helpers.js diff --git a/tests/_live_test_helpers.py b/tests/_live_test_helpers.py new file mode 100644 index 00000000000..a79b81e82c1 --- /dev/null +++ b/tests/_live_test_helpers.py @@ -0,0 +1,10 @@ +import os + +import pytest + + +def _skip_live_prompt_caching_test(): + if os.environ.get("LITELLM_RUN_LIVE_PROMPT_CACHING_TESTS") != "1": + pytest.skip("Live prompt-caching E2E tests are opt-in") + if os.environ.get("CASSETTE_REDIS_URL"): + pytest.skip("Live prompt-caching E2E tests cannot run under VCR replay") diff --git a/tests/_vcr_conftest_common.py b/tests/_vcr_conftest_common.py index 995c333ad0e..4d5a73779ea 100644 --- a/tests/_vcr_conftest_common.py +++ b/tests/_vcr_conftest_common.py @@ -1930,6 +1930,25 @@ def emit_vcr_classification_summary(terminalreporter) -> None: continue terminalreporter.write_line(f" [{verdict}] {n}") + leak_verdicts = ( + VERDICT_PARTIAL, + VERDICT_MISS_OVERFLOW, + VERDICT_MISS_NOT_PERSISTED, + VERDICT_UNMARKED_LIVE_CALL, + ) + leak_counts = {verdict: counts.get(verdict, 0) for verdict in leak_verdicts} + total_leaks = sum(leak_counts.values()) + terminalreporter.write_sep("-", "VCR COST LEAK CHECK", bold=True) + if total_leaks: + rendered = ", ".join( + f"{verdict}={count}" for verdict, count in leak_counts.items() if count + ) + terminalreporter.write_line(f" FAIL: {rendered}") + else: + terminalreporter.write_line( + " PASS: no overflow, partial, not-persisted, or unmarked live-call verdicts" + ) + overflow = snapshot["overflow_tests"] if overflow: terminalreporter.write_sep( diff --git a/tests/litellm_utils_tests/conftest.py b/tests/litellm_utils_tests/conftest.py index d20203da3a7..68c281a045f 100644 --- a/tests/litellm_utils_tests/conftest.py +++ b/tests/litellm_utils_tests/conftest.py @@ -28,32 +28,9 @@ from tests._vcr_conftest_common import ( # noqa: E402,F401 _verbose_state = VerboseReporterState() +_VCR_INCOMPATIBLE_FILES = frozenset() -# Files where VCR replay breaks the test: -# - ``test_litellm_overhead.py``: asserts overhead/total < 40%, which -# inverts when cached replay collapses the upstream time to microseconds. -_VCR_INCOMPATIBLE_FILES = frozenset( - { - "test_litellm_overhead.py", - } -) - -# AWS Secrets Manager resource-lifecycle tests. Each run creates a secret -# under a per-run unique name (``litellm_test_``) and either asserts the -# API response echoes that exact unique name or reads it straight back. The -# name *must* be unique per run because AWS enforces a >=7-day deletion -# recovery window — a fixed name can't be re-created on the daily VCR -# re-record. Deterministic replay returns the previously-recorded (different) -# name, so the unique-name round-trip cannot be reproduced offline. The -# config-parsing tests in the same file (settings / STS endpoint) make no such -# unique-resource calls and stay VCR-cached. -_VCR_INCOMPATIBLE_NODEID_SUFFIXES: tuple[str, ...] = ( - "::test_write_and_read_simple_secret", - "::test_write_and_read_json_secret", - "::test_read_nonexistent_secret", - "::test_primary_secret_functionality", - "::test_write_secret_with_description_and_tags", -) +_VCR_INCOMPATIBLE_NODEID_SUFFIXES: tuple[str, ...] = () @pytest.fixture(scope="function", autouse=True) diff --git a/tests/litellm_utils_tests/test_aws_secret_manager.py b/tests/litellm_utils_tests/test_aws_secret_manager.py index 674f9b3ca82..46e8d004534 100644 --- a/tests/litellm_utils_tests/test_aws_secret_manager.py +++ b/tests/litellm_utils_tests/test_aws_secret_manager.py @@ -10,7 +10,6 @@ from dotenv import load_dotenv import litellm.types import litellm.types.utils - load_dotenv() import io @@ -52,6 +51,11 @@ def skip_on_throttling(func): def check_aws_credentials(): """Helper function to check if AWS credentials are set""" + if os.getenv("LITELLM_RUN_LIVE_AWS_SECRET_MANAGER_TESTS") != "1": + pytest.skip("Live AWS Secrets Manager E2E tests are opt-in") + if os.getenv("CASSETTE_REDIS_URL"): + pytest.skip("Live AWS Secrets Manager E2E tests cannot run under VCR replay") + required_vars = ["AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY", "AWS_REGION_NAME"] missing_vars = [var for var in required_vars if not os.getenv(var)] if missing_vars: @@ -444,6 +448,11 @@ async def test_end_to_end_iam_role_secret_write(): - TEST_IAM_ROLE_ARN environment variable with ARN of a role that can be assumed - Proper AWS credentials configured (via instance profile, IAM role, or environment) """ + if os.getenv("LITELLM_RUN_LIVE_AWS_SECRET_MANAGER_TESTS") != "1": + pytest.skip("Live AWS Secrets Manager E2E tests are opt-in") + if os.getenv("CASSETTE_REDIS_URL"): + pytest.skip("Live AWS Secrets Manager E2E tests cannot run under VCR replay") + # Skip if TEST_IAM_ROLE_ARN is not set test_role_arn = os.getenv("TEST_IAM_ROLE_ARN") if not test_role_arn: diff --git a/tests/litellm_utils_tests/test_litellm_overhead.py b/tests/litellm_utils_tests/test_litellm_overhead.py index 3a428e9d588..95c376c24ff 100644 --- a/tests/litellm_utils_tests/test_litellm_overhead.py +++ b/tests/litellm_utils_tests/test_litellm_overhead.py @@ -1,237 +1,185 @@ +import asyncio import json -import os -import sys import time -from contextlib import asynccontextmanager, contextmanager -from datetime import datetime -from unittest.mock import AsyncMock, patch, MagicMock + import httpx import pytest -import asyncio -sys.path.insert( - 0, os.path.abspath("../..") -) # Adds the parent directory to the system path import litellm +OPENAI_API_BASE = "https://example.openai.test/v1" -# Fake Vertex AI Gemini response for mocking -FAKE_VERTEX_GEMINI_RESPONSE = { - "candidates": [ + +def _completion_payload(response_id="chatcmpl-test"): + return { + "id": response_id, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "Hello"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, + } + + +def _stream_payload(response_id="chatcmpl-stream"): + chunks = [ { - "content": { - "parts": [{"text": "Hello! How can I help you today?"}], - "role": "model", - }, - "finishReason": "STOP", - } - ], - "usageMetadata": { - "promptTokenCount": 5, - "candidatesTokenCount": 8, - "totalTokenCount": 13, - }, -} + "id": response_id, + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o", + "choices": [ + { + "index": 0, + "delta": {"role": "assistant", "content": "Hello"}, + "finish_reason": None, + } + ], + }, + { + "id": response_id, + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o", + "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, + }, + ] + return ( + "".join(f"data: {json.dumps(chunk)}\n\n" for chunk in chunks) + + "data: [DONE]\n\n" + ).encode() -def _make_fake_httpx_response(url: str) -> httpx.Response: - """Create a fake httpx.Response that looks like a Vertex AI Gemini response.""" - response = httpx.Response( - status_code=200, - json=FAKE_VERTEX_GEMINI_RESPONSE, - request=httpx.Request("POST", url), +def _mock_openai_completion_transport( + monkeypatch, *, stream=False, response_id="chatcmpl-test" +): + from litellm.llms.custom_httpx.aiohttp_transport import LiteLLMAiohttpTransport + + calls = {"count": 0} + + async def delayed_response(_transport, request): + calls["count"] += 1 + await asyncio.sleep(0.2) + if stream: + return httpx.Response( + 200, + content=_stream_payload(response_id), + headers={"content-type": "text/event-stream"}, + request=request, + ) + return httpx.Response( + 200, json=_completion_payload(response_id), request=request + ) + + monkeypatch.setattr( + LiteLLMAiohttpTransport, + "handle_async_request", + delayed_response, ) - return response + return calls -@asynccontextmanager -async def _vertex_ai_mocks(): - """Context manager that mocks Vertex AI auth and HTTP calls. - - Mocks at the httpx.AsyncClient.send level so that the - @track_llm_api_timing decorator on AsyncHTTPHandler.post still runs, - preserving the overhead measurement. - """ - fake_response = _make_fake_httpx_response( - "https://fake-vertex-endpoint/v1/models/gemini-1.5-flash:generateContent" - ) - - async def fake_send(self, request, **kwargs): - await asyncio.sleep(0.2) # simulate ~200ms network latency - return fake_response - - with ( - patch( - "litellm.llms.vertex_ai.vertex_llm_base.VertexBase._ensure_access_token_async", - new_callable=AsyncMock, - return_value=("Bearer fake-token", "fake-project"), - ), - patch.object( - httpx.AsyncClient, - "send", - new=fake_send, - ), - ): - yield - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - "model", - [ - "bedrock/mistral.mistral-7b-instruct-v0:2", - "openai/gpt-4o", - "openai/self_hosted", - "bedrock/anthropic.claude-3-5-haiku-20241022-v1:0", - "vertex_ai/gemini-1.5-flash", - ], -) -async def test_litellm_overhead_non_streaming(model): - """ - - Test we can see the litellm overhead and that it is less than 40% of the total request time - """ - - litellm._turn_on_debug() - start_time = datetime.now() - kwargs = { - "messages": [{"role": "user", "content": "Hello, world!"}], - "model": model, - } - ######################################################### - # Specific cases for models - ######################################################### - if model == "vertex_ai/gemini-1.5-flash": - kwargs["vertex_project"] = "fake-project" - kwargs["vertex_location"] = "us-central1" - if model == "openai/self_hosted": - kwargs["api_base"] = os.environ.get("FAKE_OPENAI_API_BASE") - - async def _run(): - return await litellm.acompletion(**kwargs) - - if model == "vertex_ai/gemini-1.5-flash": - async with _vertex_ai_mocks(): - response = await _run() - else: - response = await _run() - ######################################################### - # End of specific cases for models - ######################################################### - end_time = datetime.now() - total_time_ms = (end_time - start_time).total_seconds() * 1000 - print(response) - print(response._hidden_params) +def _assert_overhead_is_smaller_than_total(response, total_time_ms): litellm_overhead_ms = response._hidden_params["litellm_overhead_time_ms"] - # calculate percent of overhead caused by litellm overhead_percent = litellm_overhead_ms * 100 / total_time_ms - print("##########################\n") - print("total_time_ms", total_time_ms) - print("response litellm_overhead_ms", litellm_overhead_ms) - print("litellm overhead_percent {}%".format(overhead_percent)) - print("##########################\n") + assert litellm_overhead_ms > 0 assert litellm_overhead_ms < 1000 - - # latency overhead should be less than total request time - assert litellm_overhead_ms < (end_time - start_time).total_seconds() * 1000 - - # latency overhead should be under 40% of total request time + assert litellm_overhead_ms < total_time_ms assert overhead_percent < 40 - pass + +@pytest.fixture(autouse=True) +def reset_litellm_state(): + litellm.cache = None + litellm.success_callback = [] + litellm._async_success_callback = [] + litellm.failure_callback = [] + litellm.callbacks = [] + yield + litellm.cache = None + litellm.callbacks = [] @pytest.mark.asyncio -@pytest.mark.parametrize( - "model", - [ - "bedrock/mistral.mistral-7b-instruct-v0:2", - "openai/gpt-4o", - "bedrock/anthropic.claude-3-5-haiku-20241022-v1:0", - "openai/self_hosted", - ], -) -async def test_litellm_overhead_stream(model): +async def test_litellm_overhead_non_streaming(monkeypatch): + calls = _mock_openai_completion_transport( + monkeypatch, response_id="chatcmpl-non-stream" + ) - litellm._turn_on_debug() - start_time = datetime.now() - kwargs = { - "messages": [{"role": "user", "content": "Hello, world!"}], - "model": model, - "stream": True, - } - ######################################################### - # Specific cases for models - ######################################################### - if model == "openai/self_hosted": - kwargs["api_base"] = "https://exampleopenaiendpoint-production.up.railway.app/" - # warmup call for auth validation on vertex_ai models - await litellm.acompletion(**kwargs) + start_time = time.perf_counter() + response = await litellm.acompletion( + model="gpt-4o", + api_key="test-key", + api_base=OPENAI_API_BASE, + messages=[{"role": "user", "content": "Hello, world!"}], + ) + total_time_ms = (time.perf_counter() - start_time) * 1000 - response = await litellm.acompletion(**kwargs) - - async for chunk in response: - print() - - end_time = datetime.now() - total_time_ms = (end_time - start_time).total_seconds() * 1000 - print(response) - print(response._hidden_params) - litellm_overhead_ms = response._hidden_params["litellm_overhead_time_ms"] - # calculate percent of overhead caused by litellm - overhead_percent = litellm_overhead_ms * 100 / total_time_ms - print("##########################\n") - print("total_time_ms", total_time_ms) - print("response litellm_overhead_ms", litellm_overhead_ms) - print("litellm overhead_percent {}%".format(overhead_percent)) - print("##########################\n") - assert litellm_overhead_ms > 0 - assert litellm_overhead_ms < 1000 - - # latency overhead should be less than total request time - assert litellm_overhead_ms < (end_time - start_time).total_seconds() * 1000 - - # latency overhead should be under 40% of total request time - assert overhead_percent < 40 - - pass + assert calls["count"] == 1 + _assert_overhead_is_smaller_than_total(response, total_time_ms) @pytest.mark.asyncio -async def test_litellm_overhead_cache_hit(): - """ - Test that litellm overhead is tracked on cache hits. - Makes two identical requests and checks that the second one (cache hit) has overhead in hidden params. - """ +async def test_litellm_overhead_stream(monkeypatch): + calls = _mock_openai_completion_transport( + monkeypatch, stream=True, response_id="chatcmpl-stream" + ) + + start_time = time.perf_counter() + response = await litellm.acompletion( + model="gpt-4o", + api_key="test-key", + api_base=OPENAI_API_BASE, + messages=[{"role": "user", "content": "Hello, world!"}], + stream=True, + ) + + async for _chunk in response: + pass + + total_time_ms = (time.perf_counter() - start_time) * 1000 + + assert calls["count"] == 1 + _assert_overhead_is_smaller_than_total(response, total_time_ms) + + +@pytest.mark.asyncio +async def test_litellm_overhead_cache_hit(monkeypatch): from litellm.caching.caching import Cache - litellm._turn_on_debug() + calls = _mock_openai_completion_transport(monkeypatch, response_id="chatcmpl-cache") litellm.cache = Cache() - print("test2 for caching") - litellm.set_verbose = True + messages = [{"role": "user", "content": "Hello, world! Cache test"}] response1 = await litellm.acompletion( - model="gpt-4.1-nano", messages=messages, caching=True + model="gpt-4o", + api_key="test-key", + api_base=OPENAI_API_BASE, + messages=messages, + caching=True, ) - await asyncio.sleep(2) - # Wait for any pending background tasks to complete - pending_tasks = [task for task in asyncio.all_tasks() if not task.done()] - print("all pending tasks", pending_tasks) - if pending_tasks: - await asyncio.wait(pending_tasks, timeout=1.0) - + await asyncio.sleep(0.5) response2 = await litellm.acompletion( - model="gpt-4.1-nano", messages=messages, caching=True + model="gpt-4o", + api_key="test-key", + api_base=OPENAI_API_BASE, + messages=messages, + caching=True, ) - print("RESPONSE 1", response1) - print("RESPONSE 2", response2) + + assert calls["count"] == 1 assert response1.id == response2.id - - print("response 2 hidden params", response2._hidden_params) - assert "_response_ms" in response2._hidden_params - total_time_ms = response2._hidden_params["_response_ms"] + assert response2._hidden_params["litellm_overhead_time_ms"] > 0 assert ( - response2._hidden_params["litellm_overhead_time_ms"] > 0 - and response2._hidden_params["litellm_overhead_time_ms"] < total_time_ms + response2._hidden_params["litellm_overhead_time_ms"] + < response2._hidden_params["_response_ms"] ) diff --git a/tests/llm_translation/base_llm_unit_tests.py b/tests/llm_translation/base_llm_unit_tests.py index 77850dac457..fef1d23d867 100644 --- a/tests/llm_translation/base_llm_unit_tests.py +++ b/tests/llm_translation/base_llm_unit_tests.py @@ -30,6 +30,10 @@ from litellm.types.utils import Usage, ModelResponse from abc import ABC, abstractmethod from openai import OpenAI +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", ".."))) + +from tests._live_test_helpers import _skip_live_prompt_caching_test # noqa: E402 + def _usage_format_tests(usage: litellm.Usage): """ @@ -960,6 +964,7 @@ class BaseLLMChatTest(ABC): @pytest.mark.flaky(retries=4, delay=1) def test_prompt_caching(self): + _skip_live_prompt_caching_test() print("test_prompt_caching") litellm.set_verbose = True from litellm.utils import supports_prompt_caching diff --git a/tests/llm_translation/conftest.py b/tests/llm_translation/conftest.py index d346dae4308..dba3812ee1c 100644 --- a/tests/llm_translation/conftest.py +++ b/tests/llm_translation/conftest.py @@ -39,13 +39,7 @@ from tests._vcr_conftest_common import ( # noqa: E402,F401 # itself run under a live cassette context. _VCR_AUTO_MARKER_SKIP_FILES = frozenset({"test_vcr_redis_persister.py"}) -# Tests that observe live cross-call provider state (e.g. prompt-cache -# warm-up between two consecutive calls); replay can't reproduce that state. -_VCR_INCOMPATIBLE_NODEID_SUFFIXES = ( - "::test_prompt_caching", - "TestBedrockInvokeNovaJson::test_json_response_pydantic_obj", - "::test_bedrock_converse__streaming_passthrough", -) +_VCR_INCOMPATIBLE_NODEID_SUFFIXES: tuple[str, ...] = () _verbose_state = VerboseReporterState() diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py index aecd7bc699a..9cf253c379d 100644 --- a/tests/llm_translation/test_bedrock_completion.py +++ b/tests/llm_translation/test_bedrock_completion.py @@ -3220,6 +3220,11 @@ async def test_bedrock_converse__streaming_passthrough(monkeypatch): from litellm.integrations.custom_logger import CustomLogger import asyncio + if os.environ.get("LITELLM_RUN_LIVE_BEDROCK_PASSTHROUGH_TESTS") != "1": + pytest.skip("Live Bedrock passthrough E2E tests are opt-in") + if os.environ.get("CASSETTE_REDIS_URL"): + pytest.skip("Live Bedrock passthrough E2E tests cannot run under VCR replay") + class MockCustomLogger(CustomLogger): pass diff --git a/tests/llm_translation/test_bedrock_invoke_tests.py b/tests/llm_translation/test_bedrock_invoke_tests.py index 23f436d5b28..901b43542f7 100644 --- a/tests/llm_translation/test_bedrock_invoke_tests.py +++ b/tests/llm_translation/test_bedrock_invoke_tests.py @@ -3,7 +3,6 @@ import pytest import sys import os - sys.path.insert( 0, os.path.abspath("../..") ) # Adds the parent directory to the system path @@ -41,6 +40,15 @@ class TestBedrockInvokeNovaJson(BaseLLMChatTest): f"Skipping non-JSON test: {request.function.__name__} does not contain 'json'" ) + def test_json_response_pydantic_obj(self): + if os.environ.get("LITELLM_RUN_LIVE_BEDROCK_NOVA_JSON_TESTS") != "1": + pytest.skip("Live Bedrock Nova response-schema E2E tests are opt-in") + if os.environ.get("CASSETTE_REDIS_URL"): + pytest.skip( + "Live Bedrock Nova response-schema E2E tests cannot run under VCR replay" + ) + super().test_json_response_pydantic_obj() + def test_nova_invoke_remove_empty_system_messages(): """Test that _remove_empty_system_messages removes empty system list.""" diff --git a/tests/local_testing/conftest.py b/tests/local_testing/conftest.py index 0831313c136..d45caec22d8 100644 --- a/tests/local_testing/conftest.py +++ b/tests/local_testing/conftest.py @@ -57,13 +57,10 @@ from tests._vcr_conftest_common import ( # noqa: E402,F401 # blacklisting was masking valid cache opportunities. # Files where VCR replay breaks the test: -# - ``test_assistants.py``: polls fresh per-session run IDs that no cassette -# can match, so every CI run re-records and the suite times out. # - ``test_router_caching.py``: asserts upstream returns a *new* id per call, # which a deterministic cassette replay violates. _VCR_INCOMPATIBLE_FILES = frozenset( { - "test_assistants.py", "test_router_caching.py", } ) diff --git a/tests/local_testing/test_assistants.py b/tests/local_testing/test_assistants.py index ee1c8fb6518..8dc4f9e48e1 100644 --- a/tests/local_testing/test_assistants.py +++ b/tests/local_testing/test_assistants.py @@ -1,22 +1,13 @@ -# What is this? -## Unit Tests for OpenAI Assistants API -import json import os import sys -import traceback - -from dotenv import load_dotenv - -load_dotenv() -sys.path.insert( - 0, os.path.abspath("../..") -) # Adds the parent directory to the system path -import asyncio -import logging import pytest +from dotenv import load_dotenv from openai.types.beta.assistant import Assistant -from typing_extensions import override +from openai.types.beta.assistant_deleted import AssistantDeleted + +load_dotenv() +sys.path.insert(0, os.path.abspath("../..")) import litellm from litellm import create_thread, get_thread @@ -25,40 +16,264 @@ from litellm.llms.openai.openai import ( AsyncAssistantEventHandler, AsyncCursorPage, MessageData, - OpenAIAssistantsAPI, + OpenAIMessage as Message, + Run, + SyncCursorPage, + Thread, ) -from litellm.llms.openai.openai import OpenAIMessage as Message -from litellm.llms.openai.openai import SyncCursorPage, Thread -""" -V0 Scope: - -- Add Message -> `/v1/threads/{thread_id}/messages` -- Run Thread -> `/v1/threads/{thread_id}/run` -""" +ASSISTANT_INSTRUCTIONS = ( + "You are a personal math tutor. When asked a question, write and run Python " + "code to answer the question." +) +ASSISTANT_ID = "asst_test" +THREAD_ID = "thread_test" +MESSAGE_ID = "msg_test" +RUN_ID = "run_test" -def _add_azure_related_dynamic_params(data: dict) -> dict: - data["api_version"] = "2024-02-15-preview" - data["api_base"] = os.getenv("AZURE_AI_API_BASE") - data["api_key"] = os.getenv("AZURE_AI_API_KEY") +def _assistant(**overrides): + data = { + "id": ASSISTANT_ID, + "object": "assistant", + "created_at": 1, + "name": "Math Tutor", + "description": None, + "model": "gpt-4.1", + "instructions": ASSISTANT_INSTRUCTIONS, + "tools": [], + "metadata": {}, + "top_p": 1.0, + "temperature": 1.0, + "response_format": "auto", + } + data.update(overrides) + return Assistant(**data) + + +def _thread(thread_id=THREAD_ID): + return Thread(id=thread_id, object="thread", created_at=1, metadata={}) + + +def _message(thread_id=THREAD_ID): + return Message( + id=MESSAGE_ID, + object="thread.message", + created_at=1, + thread_id=thread_id, + role="user", + content=[ + { + "type": "text", + "text": {"value": "Hey, how's it going?", "annotations": []}, + } + ], + assistant_id=None, + run_id=None, + attachments=[], + metadata={}, + status="completed", + ) + + +def _run(thread_id=THREAD_ID, assistant_id=ASSISTANT_ID): + return Run( + id=RUN_ID, + object="thread.run", + created_at=1, + assistant_id=assistant_id, + thread_id=thread_id, + status="completed", + started_at=1, + expires_at=None, + cancelled_at=None, + failed_at=None, + completed_at=1, + last_error=None, + model="gpt-4.1", + instructions=ASSISTANT_INSTRUCTIONS, + tools=[], + metadata={}, + usage={"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + required_action=None, + incomplete_details=None, + temperature=1.0, + top_p=1.0, + max_prompt_tokens=None, + max_completion_tokens=None, + truncation_strategy={"type": "auto", "last_messages": None}, + response_format="auto", + tool_choice="auto", + parallel_tool_calls=True, + ) + + +def _sync_page(data): + first_id = data[0].id if data else None + return SyncCursorPage( + data=data, + object="list", + first_id=first_id, + last_id=first_id, + has_more=False, + ) + + +def _async_page(data): + first_id = data[0].id if data else None + return AsyncCursorPage( + data=data, + object="list", + first_id=first_id, + last_id=first_id, + has_more=False, + ) + + +class _FakeAssistantEventHandler(AssistantEventHandler): + def until_done(self): + return None + + +class _FakeAsyncAssistantEventHandler(AsyncAssistantEventHandler): + async def until_done(self): + return None + + +class _FakeAssistantStream: + def __enter__(self): + return _FakeAssistantEventHandler() + + def __exit__(self, exc_type, exc, tb): + return False + + +class _FakeAsyncAssistantStream: + async def __aenter__(self): + return _FakeAsyncAssistantEventHandler() + + async def __aexit__(self, exc_type, exc, tb): + return False + + +class _SyncAssistants: + def list(self, **_kwargs): + return _sync_page([_assistant()]) + + def create(self, **kwargs): + return _assistant(**kwargs) + + def delete(self, assistant_id): + return AssistantDeleted( + id=assistant_id, object="assistant.deleted", deleted=True + ) + + +class _AsyncAssistants: + async def list(self, **_kwargs): + return _async_page([_assistant()]) + + async def create(self, **kwargs): + return _assistant(**kwargs) + + async def delete(self, assistant_id): + return AssistantDeleted( + id=assistant_id, object="assistant.deleted", deleted=True + ) + + +class _SyncMessages: + def create(self, thread_id, **_kwargs): + return _message(thread_id) + + def list(self, thread_id): + return _sync_page([_message(thread_id)]) + + +class _AsyncMessages: + async def create(self, thread_id, **_kwargs): + return _message(thread_id) + + async def list(self, thread_id): + return _async_page([_message(thread_id)]) + + +class _SyncRuns: + def create_and_poll(self, thread_id, assistant_id, **_kwargs): + return _run(thread_id=thread_id, assistant_id=assistant_id) + + def stream(self, **_kwargs): + return _FakeAssistantStream() + + +class _AsyncRuns: + async def create_and_poll(self, thread_id, assistant_id, **_kwargs): + return _run(thread_id=thread_id, assistant_id=assistant_id) + + def stream(self, **_kwargs): + return _FakeAsyncAssistantStream() + + +class _SyncThreads: + def __init__(self): + self.messages = _SyncMessages() + self.runs = _SyncRuns() + + def create(self, **_kwargs): + return _thread() + + def retrieve(self, thread_id): + return _thread(thread_id) + + +class _AsyncThreads: + def __init__(self): + self.messages = _AsyncMessages() + self.runs = _AsyncRuns() + + async def create(self, **_kwargs): + return _thread() + + async def retrieve(self, thread_id): + return _thread(thread_id) + + +class _FakeBeta: + def __init__(self, *, async_mode): + self.assistants = _AsyncAssistants() if async_mode else _SyncAssistants() + self.threads = _AsyncThreads() if async_mode else _SyncThreads() + + +class _FakeAssistantClient: + def __init__(self, *, async_mode): + self.beta = _FakeBeta(async_mode=async_mode) + + +@pytest.fixture +def assistant_client(sync_mode): + return _FakeAssistantClient(async_mode=not sync_mode) + + +def _request_data(provider, assistant_client, **kwargs): + data = {"custom_llm_provider": provider, "client": assistant_client, **kwargs} + if provider == "azure": + data.update( + { + "api_version": "2024-02-15-preview", + "api_base": "https://example.azure.test", + "api_key": "test-key", + } + ) return data @pytest.mark.parametrize("provider", ["openai", "azure"]) -@pytest.mark.parametrize( - "sync_mode", - [True, False], -) +@pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.asyncio -async def test_get_assistants(provider, sync_mode): - data = { - "custom_llm_provider": provider, - } - if provider == "azure": - data = _add_azure_related_dynamic_params(data) +async def test_get_assistants(provider, sync_mode, assistant_client): + data = _request_data(provider, assistant_client) - if sync_mode == True: + if sync_mode: assistants = litellm.get_assistants(**data) assert isinstance(assistants, SyncCursorPage) else: @@ -67,276 +282,152 @@ async def test_get_assistants(provider, sync_mode): @pytest.mark.parametrize("provider", ["azure", "openai"]) -@pytest.mark.parametrize( - "sync_mode", - [True, False], -) +@pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.asyncio() -@pytest.mark.flaky(retries=3, delay=1) -async def test_create_delete_assistants(provider, sync_mode): - litellm.ssl_verify = False - litellm._turn_on_debug() - data = { - "custom_llm_provider": provider, - "model": "gpt-4.1", - "instructions": "You are a personal math tutor. When asked a question, write and run Python code to answer the question.", - "name": "Math Tutor", - "tools": [{"type": "code_interpreter"}], - } - if provider == "azure": - data = _add_azure_related_dynamic_params(data) +async def test_create_delete_assistants(provider, sync_mode, assistant_client): + data = _request_data( + provider, + assistant_client, + model="gpt-4.1", + instructions=ASSISTANT_INSTRUCTIONS, + name="Math Tutor", + tools=[{"type": "code_interpreter"}], + ) - if sync_mode == True: + if sync_mode: assistant = litellm.create_assistants(**data) - - print("New assistants", assistant) assert isinstance(assistant, Assistant) - assert ( - assistant.instructions - == "You are a personal math tutor. When asked a question, write and run Python code to answer the question." - ) + assert assistant.instructions == ASSISTANT_INSTRUCTIONS assert assistant.id is not None - # delete the created assistant - delete_data = { - "custom_llm_provider": provider, - "assistant_id": assistant.id, - } - if provider == "azure": - delete_data = _add_azure_related_dynamic_params(delete_data) - response = litellm.delete_assistant(**delete_data) - print("Response deleting assistant", response) + response = litellm.delete_assistant( + **_request_data( + provider, + assistant_client, + assistant_id=assistant.id, + ) + ) assert response.id == assistant.id else: assistant = await litellm.acreate_assistants(**data) - print("New assistants", assistant) assert isinstance(assistant, Assistant) - assert ( - assistant.instructions - == "You are a personal math tutor. When asked a question, write and run Python code to answer the question." - ) + assert assistant.instructions == ASSISTANT_INSTRUCTIONS assert assistant.id is not None - # delete the created assistant - delete_data = { - "custom_llm_provider": provider, - "assistant_id": assistant.id, - } - if provider == "azure": - delete_data = _add_azure_related_dynamic_params(delete_data) - response = await litellm.adelete_assistant(**delete_data) - print("Response deleting assistant", response) + response = await litellm.adelete_assistant( + **_request_data( + provider, + assistant_client, + assistant_id=assistant.id, + ) + ) assert response.id == assistant.id -@pytest.mark.parametrize("provider", ["openai", "azure"]) -@pytest.mark.parametrize("sync_mode", [True, False]) -@pytest.mark.asyncio -async def test_create_thread_litellm(sync_mode, provider) -> Thread: +async def _create_thread_litellm(sync_mode, provider, assistant_client) -> Thread: message: MessageData = {"role": "user", "content": "Hey, how's it going?"} # type: ignore - data = { - "custom_llm_provider": provider, - "message": [message], - } - if provider == "azure": - data = _add_azure_related_dynamic_params(data) + data = _request_data(provider, assistant_client, message=[message]) if sync_mode: new_thread = create_thread(**data) else: new_thread = await litellm.acreate_thread(**data) - assert isinstance( - new_thread, Thread - ), f"type of thread={type(new_thread)}. Expected Thread-type" - + assert isinstance(new_thread, Thread) return new_thread @pytest.mark.parametrize("provider", ["openai", "azure"]) @pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.asyncio -async def test_get_thread_litellm(provider, sync_mode): - new_thread = test_create_thread_litellm(sync_mode, provider) +async def test_create_thread_litellm(sync_mode, provider, assistant_client): + await _create_thread_litellm(sync_mode, provider, assistant_client) - if asyncio.iscoroutine(new_thread): - _new_thread = await new_thread - else: - _new_thread = new_thread - data = { - "custom_llm_provider": provider, - "thread_id": _new_thread.id, - } - if provider == "azure": - data = _add_azure_related_dynamic_params(data) +@pytest.mark.parametrize("provider", ["openai", "azure"]) +@pytest.mark.parametrize("sync_mode", [True, False]) +@pytest.mark.asyncio +async def test_get_thread_litellm(provider, sync_mode, assistant_client): + new_thread = await _create_thread_litellm(sync_mode, provider, assistant_client) + data = _request_data(provider, assistant_client, thread_id=new_thread.id) if sync_mode: received_thread = get_thread(**data) else: received_thread = await litellm.aget_thread(**data) - assert isinstance( - received_thread, Thread - ), f"type of thread={type(received_thread)}. Expected Thread-type" - return new_thread + assert isinstance(received_thread, Thread) @pytest.mark.parametrize("provider", ["openai", "azure"]) @pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.asyncio -async def test_add_message_litellm(sync_mode, provider): +async def test_add_message_litellm(sync_mode, provider, assistant_client): + new_thread = await _create_thread_litellm(sync_mode, provider, assistant_client) message: MessageData = {"role": "user", "content": "Hey, how's it going?"} # type: ignore - new_thread = test_create_thread_litellm(sync_mode, provider) + data = _request_data(provider, assistant_client, thread_id=new_thread.id, **message) - if asyncio.iscoroutine(new_thread): - _new_thread = await new_thread - else: - _new_thread = new_thread - # add message to thread - message: MessageData = {"role": "user", "content": "Hey, how's it going?"} # type: ignore - - data = {"custom_llm_provider": provider, "thread_id": _new_thread.id, **message} - if provider == "azure": - data = _add_azure_related_dynamic_params(data) if sync_mode: added_message = litellm.add_message(**data) else: added_message = await litellm.a_add_message(**data) - print(f"added message: {added_message}") - assert isinstance(added_message, Message) -@pytest.mark.parametrize( - "provider", - [ - "azure", - "openai", - ], -) # -@pytest.mark.parametrize( - "sync_mode", - [ - True, - False, - ], -) -@pytest.mark.parametrize( - "is_streaming", - [True, False], -) # +@pytest.mark.parametrize("provider", ["azure", "openai"]) +@pytest.mark.parametrize("sync_mode", [True, False]) +@pytest.mark.parametrize("is_streaming", [True, False]) @pytest.mark.asyncio -@pytest.mark.flaky(retries=3, delay=1) -async def test_aarun_thread_litellm(sync_mode, provider, is_streaming): - """ - - Get Assistants - - Create thread - - Create run w/ Assistants + Thread - """ - import openai +async def test_aarun_thread_litellm( + sync_mode, provider, is_streaming, assistant_client +): + get_assistants_data = _request_data(provider, assistant_client) + if sync_mode: + assistants = litellm.get_assistants(**get_assistants_data) + else: + assistants = await litellm.aget_assistants(**get_assistants_data) - try: - get_assistants_data = { - "custom_llm_provider": provider, - } - if provider == "azure": - get_assistants_data = _add_azure_related_dynamic_params(get_assistants_data) - if sync_mode: - assistants = litellm.get_assistants(**get_assistants_data) + assistant_id = assistants.data[0].id + new_thread = await _create_thread_litellm(sync_mode, provider, assistant_client) + message: MessageData = {"role": "user", "content": "Hey, how's it going?"} # type: ignore + thread_data = _request_data(provider, assistant_client, thread_id=new_thread.id) + message_data = _request_data( + provider, assistant_client, thread_id=new_thread.id, **message + ) + + if sync_mode: + added_message = litellm.add_message(**message_data) + assert isinstance(added_message, Message) + + if is_streaming: + run = litellm.run_thread_stream(assistant_id=assistant_id, **thread_data) + with run as run: + assert isinstance(run, AssistantEventHandler) + run.until_done() else: - assistants = await litellm.aget_assistants(**get_assistants_data) + run = litellm.run_thread( + assistant_id=assistant_id, stream=is_streaming, **thread_data + ) + assert run.status == "completed" + messages = litellm.get_messages(**thread_data) + assert isinstance(messages.data[0], Message) + else: + added_message = await litellm.a_add_message(**message_data) + assert isinstance(added_message, Message) - ## get the first assistant ### - try: - assistant_id = assistants.data[0].id - except IndexError: - pytest.skip("No assistants found") - - new_thread = test_create_thread_litellm(sync_mode=sync_mode, provider=provider) - - if asyncio.iscoroutine(new_thread): - _new_thread = await new_thread + if is_streaming: + run = litellm.arun_thread_stream(assistant_id=assistant_id, **thread_data) + async with run as run: + assert isinstance(run, AsyncAssistantEventHandler) + await run.until_done() else: - _new_thread = new_thread - - thread_id = _new_thread.id - - # add message to thread - message: MessageData = {"role": "user", "content": "Hey, how's it going?"} # type: ignore - - data = {"custom_llm_provider": provider, "thread_id": _new_thread.id, **message} - if provider == "azure": - data = _add_azure_related_dynamic_params(data) - - if sync_mode: - added_message = litellm.add_message(**data) - - if is_streaming: - run = litellm.run_thread_stream(assistant_id=assistant_id, **data) - with run as run: - assert isinstance(run, AssistantEventHandler) - print(run) - run.until_done() - else: - run = litellm.run_thread( - assistant_id=assistant_id, stream=is_streaming, **data - ) - if run.status == "completed": - messages = litellm.get_messages( - thread_id=_new_thread.id, custom_llm_provider=provider - ) - assert isinstance(messages.data[0], Message) - elif ( - run.status == "failed" - and run.last_error - and "No connection matching model" in run.last_error.message - ): - pytest.skip(f"Azure deployment not found: {run.last_error.message}") - else: - pytest.fail( - "An unexpected error occurred when running the thread, {}".format( - run - ) - ) - - else: - added_message = await litellm.a_add_message(**data) - - if is_streaming: - run = litellm.arun_thread_stream(assistant_id=assistant_id, **data) - async with run as run: - print(f"run: {run}") - assert isinstance( - run, - AsyncAssistantEventHandler, - ) - print(run) - await run.until_done() - else: - run = await litellm.arun_thread( - custom_llm_provider=provider, - thread_id=thread_id, - assistant_id=assistant_id, - ) - - if run.status == "completed": - messages = await litellm.aget_messages( - thread_id=_new_thread.id, custom_llm_provider=provider - ) - assert isinstance(messages.data[0], Message) - elif ( - run.status == "failed" - and run.last_error - and "No connection matching model" in run.last_error.message - ): - pytest.skip(f"Azure deployment not found: {run.last_error.message}") - else: - pytest.fail( - "An unexpected error occurred when running the thread, {}".format( - run - ) - ) - except openai.APIError as e: - pass + run = await litellm.arun_thread( + custom_llm_provider=provider, + thread_id=new_thread.id, + assistant_id=assistant_id, + client=assistant_client, + ) + assert run.status == "completed" + messages = await litellm.aget_messages(**thread_data) + assert isinstance(messages.data[0], Message) diff --git a/tests/logging_callback_tests/conftest.py b/tests/logging_callback_tests/conftest.py index 6dde85f2ca7..dedff9a5aee 100644 --- a/tests/logging_callback_tests/conftest.py +++ b/tests/logging_callback_tests/conftest.py @@ -42,14 +42,7 @@ _RESPX_CONFLICTING_FILES = frozenset( } ) -# Files where VCR replay breaks the test: -# - ``test_amazing_s3_logs.py``: vcrpy's boto3 stub intercepts a real S3 -# PUT/LIST round-trip the test asserts on, so the per-run id is never found. -_VCR_INCOMPATIBLE_FILES = frozenset( - { - "test_amazing_s3_logs.py", - } -) +_VCR_INCOMPATIBLE_FILES = frozenset() _VCR_INCOMPATIBLE_NODEID_SUFFIXES: tuple[str, ...] = () diff --git a/tests/logging_callback_tests/test_amazing_s3_logs.py b/tests/logging_callback_tests/test_amazing_s3_logs.py index dab2a0cc0b9..08b9ac7d01a 100644 --- a/tests/logging_callback_tests/test_amazing_s3_logs.py +++ b/tests/logging_callback_tests/test_amazing_s3_logs.py @@ -1,6 +1,7 @@ import sys import os import io, asyncio +from collections import defaultdict # import logging # logging.basicConfig(level=logging.DEBUG) @@ -18,6 +19,60 @@ from litellm._logging import verbose_logger import logging +class _FakeS3Paginator: + def __init__(self, objects): + self.objects = objects + + def paginate(self, Bucket): + keys = sorted(self.objects[Bucket]) + if not keys: + return [{}] + return [{"Contents": [{"Key": key} for key in keys]}] + + +class _FakeS3Client: + def __init__(self): + self.objects = defaultdict(dict) + + def clear(self): + self.objects.clear() + + def put_object(self, Bucket, Key, Body, **_kwargs): + self.objects[Bucket][Key] = Body + return {"ResponseMetadata": {"HTTPStatusCode": 200}} + + def delete_object(self, Bucket, Key): + self.objects[Bucket].pop(Key, None) + return {"ResponseMetadata": {"HTTPStatusCode": 204}} + + def get_paginator(self, name): + assert name == "list_objects_v2" + return _FakeS3Paginator(self.objects) + + def list_objects(self, Bucket): + keys = sorted(self.objects[Bucket]) + return {"Contents": [{"Key": key, "LastModified": 0} for key in keys]} + + +_FAKE_S3_CLIENT = _FakeS3Client() + + +@pytest.fixture(autouse=True) +def fake_s3_client(monkeypatch): + _FAKE_S3_CLIENT.clear() + + def fake_boto3_client(service_name, *args, **kwargs): + assert service_name == "s3" + return _FAKE_S3_CLIENT + + monkeypatch.setattr(boto3, "client", fake_boto3_client) + litellm.success_callback = [] + litellm.callbacks = [] + yield _FAKE_S3_CLIENT + litellm.success_callback = [] + litellm.callbacks = [] + + @pytest.mark.asyncio @pytest.mark.parametrize( "sync_mode,streaming", [(True, True), (True, False), (False, True), (False, False)] @@ -172,6 +227,7 @@ async def test_basic_s3_v2_logging_failure(): model="gpt-5-mini", api_key="invalid-api-key", messages=[{"role": "user", "content": "This is a test"}], + mock_response=Exception("forced failure for S3 logging test"), ) except Exception as e: print(f"Expected error: {e}") @@ -407,7 +463,7 @@ from litellm.integrations.s3_v2 import S3Logger class TestS3Logger(S3Logger): def __init__(self, *args, **kwargs): self.recorded_requests = {} - self.logged_standard_logging_payload: Optional[StandardLoggingPayload] = None + self.logged_standard_logging_payload = None super().__init__(*args, **kwargs) async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): diff --git a/tests/ocr_tests/conftest.py b/tests/ocr_tests/conftest.py index 7a74dde3e41..09d535dee4b 100644 --- a/tests/ocr_tests/conftest.py +++ b/tests/ocr_tests/conftest.py @@ -26,27 +26,7 @@ from tests._vcr_conftest_common import ( # noqa: E402,F401 vcr_config_dict, ) -# Vertex AI MaaS Mistral OCR tests that cannot be VCR-cached in CI. -# -# ``vertex_ai/mistral-ocr-2505`` is a Model-as-a-Service partner model that -# must be explicitly enabled in the GCP project's Model Garden. It is not -# provisioned in the CI project (``litellm-ci-cd``), so the live -# ``:rawPredict`` call fails on every run and ``BaseOCRTest`` catches the -# provider error and skips. Because the doomed live call is recorded but the -# test then skips, the persister refuses to save it (skipped tests don't -# persist) and the cassette is never seeded — so the test re-records live and -# is classified MISS:NOT_PERSISTED on every single run, forever. No cassette -# can be recorded until the model is provisioned. Mark the tests VCR- -# incompatible so they are honestly accounted as live calls (UNMARKED:LIVE_CALL) -# rather than phantom cache misses; behaviour is unchanged (they still run and -# still skip on the provider error). The sibling direct-Mistral and Azure OCR -# tests replay from cache normally and are unaffected. Remove these entries if -# the MaaS model is enabled in the CI project. -_VCR_INCOMPATIBLE_NODEID_SUFFIXES: tuple[str, ...] = ( - "test_ocr_vertex_ai.py::TestVertexAIMistralOCR::test_ocr_response_structure", - "test_ocr_vertex_ai.py::TestVertexAIMistralOCR::test_basic_ocr_with_url[True]", - "test_ocr_vertex_ai.py::TestVertexAIMistralOCR::test_basic_ocr_with_url[False]", -) +_VCR_INCOMPATIBLE_NODEID_SUFFIXES: tuple[str, ...] = () _verbose_state = VerboseReporterState() diff --git a/tests/ocr_tests/test_ocr_vertex_ai.py b/tests/ocr_tests/test_ocr_vertex_ai.py index 1b58b955de6..1ba5b9d0883 100644 --- a/tests/ocr_tests/test_ocr_vertex_ai.py +++ b/tests/ocr_tests/test_ocr_vertex_ai.py @@ -62,6 +62,14 @@ class TestVertexAIMistralOCR(BaseOCRTest): sending to the API, since Vertex AI OCR endpoint doesn't have internet access. """ + def setup_method(self): + if os.environ.get("LITELLM_RUN_LIVE_VERTEX_MISTRAL_OCR_TESTS") != "1": + pytest.skip("Live Vertex AI Mistral OCR E2E tests are opt-in") + if os.environ.get("CASSETTE_REDIS_URL"): + pytest.skip( + "Live Vertex AI Mistral OCR E2E tests cannot run under VCR replay" + ) + def get_base_ocr_call_args(self) -> dict: """ Return the base OCR call args for Vertex AI Mistral OCR. diff --git a/tests/pass_through_tests/test_vertex.test.js b/tests/pass_through_tests/test_vertex.test.js index e0e879c2897..3663d35d192 100644 --- a/tests/pass_through_tests/test_vertex.test.js +++ b/tests/pass_through_tests/test_vertex.test.js @@ -8,6 +8,8 @@ const { writeFileSync } = require('fs'); // Import fetch if the SDK uses it const originalFetch = global.fetch || require('node-fetch'); +const { runVertexRequestOrSkip } = require('./vertex_test_helpers'); + // Monkey-patch the fetch used internally global.fetch = async function patchedFetch(url, options) { // Modify the URL to use HTTP instead of HTTPS @@ -89,7 +91,12 @@ describe('Vertex AI Tests', () => { contents: [{role: 'user', parts: [{text: 'How are you doing today tell me your name?'}]}], }; - const streamingResult = await generativeModel.generateContentStream(request); + const streamingResult = await runVertexRequestOrSkip(() => + generativeModel.generateContentStream(request) + ); + if (streamingResult === null) { + return; + } // Add some assertions expect(streamingResult).toBeDefined(); @@ -122,11 +129,16 @@ describe('Vertex AI Tests', () => { ); const request = {contents: [{role: 'user', parts: [{text: 'What is 2+2?'}]}]}; - const result = await generativeModel.generateContent(request); + const result = await runVertexRequestOrSkip(() => + generativeModel.generateContent(request) + ); + if (result === null) { + return; + } expect(result).toBeDefined(); expect(result.response).toBeDefined(); console.log('non-streaming response:', JSON.stringify(result.response)); }, VERTEX_TEST_TIMEOUT_MS ); -}); \ No newline at end of file +}); diff --git a/tests/pass_through_tests/test_vertex_ai.py b/tests/pass_through_tests/test_vertex_ai.py index bf1200489aa..0ac66b470c6 100644 --- a/tests/pass_through_tests/test_vertex_ai.py +++ b/tests/pass_through_tests/test_vertex_ai.py @@ -12,7 +12,6 @@ import os import pytest import asyncio - # Path to your service account JSON file SERVICE_ACCOUNT_FILE = "path/to/your/service-account.json" @@ -95,6 +94,15 @@ async def call_spend_logs_endpoint(): LITE_LLM_ENDPOINT = "http://localhost:4000" +def _is_vertex_quota_error(exc: Exception) -> bool: + message = str(exc) + return ( + "429" in message + or "Too Many Requests" in message + or "RESOURCE_EXHAUSTED" in message + ) + + @pytest.mark.asyncio() async def test_basic_vertex_ai_pass_through_with_spendlog(): @@ -109,7 +117,12 @@ async def test_basic_vertex_ai_pass_through_with_spendlog(): ) model = GenerativeModel(model_name="gemini-3.1-flash-lite") - response = model.generate_content("hi") + try: + response = model.generate_content("hi") + except Exception as exc: + if _is_vertex_quota_error(exc): + pytest.skip("Vertex AI quota exhausted") + raise print("response", response) diff --git a/tests/pass_through_tests/test_vertex_with_spend.test.js b/tests/pass_through_tests/test_vertex_with_spend.test.js index 4dee890dc78..5914908e66a 100644 --- a/tests/pass_through_tests/test_vertex_with_spend.test.js +++ b/tests/pass_through_tests/test_vertex_with_spend.test.js @@ -10,6 +10,8 @@ const originalFetch = global.fetch || require('node-fetch'); let lastCallId; +const { runVertexRequestOrSkip } = require('./vertex_test_helpers'); + // Monkey-patch the fetch used internally global.fetch = async function patchedFetch(url, options) { // Modify the URL to use HTTP instead of HTTPS @@ -93,7 +95,12 @@ describe('Vertex AI Tests', () => { contents: [{role: 'user', parts: [{text: 'Say "hello test" and nothing else'}]}] }; - const result = await generativeModel.generateContent(request); + const result = await runVertexRequestOrSkip(() => + generativeModel.generateContent(request) + ); + if (result === null) { + return; + } expect(result).toBeDefined(); // Use the captured callId @@ -152,7 +159,12 @@ describe('Vertex AI Tests', () => { contents: [{role: 'user', parts: [{text: 'Say "hello test" and nothing else'}]}] }; - const streamingResult = await generativeModel.generateContentStream(request); + const streamingResult = await runVertexRequestOrSkip(() => + generativeModel.generateContentStream(request) + ); + if (streamingResult === null) { + return; + } expect(streamingResult).toBeDefined(); @@ -198,4 +210,4 @@ describe('Vertex AI Tests', () => { expect(spendData[0].spend).toBeGreaterThan(0); expect(spendData[0].custom_llm_provider).toBe('vertex_ai'); }, 90000); -}); \ No newline at end of file +}); diff --git a/tests/pass_through_tests/vertex_test_helpers.js b/tests/pass_through_tests/vertex_test_helpers.js new file mode 100644 index 00000000000..d637f20f916 --- /dev/null +++ b/tests/pass_through_tests/vertex_test_helpers.js @@ -0,0 +1,27 @@ +function isVertexQuotaError(error) { + const message = [ + error && error.message, + error && error.stack, + error && error.cause && JSON.stringify(error.cause), + ].filter(Boolean).join('\n'); + + return ( + message.includes('429') || + message.includes('Too Many Requests') || + message.includes('RESOURCE_EXHAUSTED') + ); +} + +async function runVertexRequestOrSkip(requestFn) { + try { + return await requestFn(); + } catch (error) { + if (isVertexQuotaError(error)) { + console.warn('Vertex AI quota exhausted; skipping live provider assertions for this run'); + return null; + } + throw error; + } +} + +module.exports = { isVertexQuotaError, runVertexRequestOrSkip }; diff --git a/tests/pass_through_unit_tests/base_anthropic_messages_prompt_caching_test.py b/tests/pass_through_unit_tests/base_anthropic_messages_prompt_caching_test.py index 5fc4ecefb33..e8d14b00681 100644 --- a/tests/pass_through_unit_tests/base_anthropic_messages_prompt_caching_test.py +++ b/tests/pass_through_unit_tests/base_anthropic_messages_prompt_caching_test.py @@ -18,14 +18,15 @@ from abc import ABC, abstractmethod from typing import Any, Dict, List sys.path.insert(0, os.path.abspath("../../..")) +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", ".."))) import pytest import litellm +from tests._live_test_helpers import _skip_live_prompt_caching_test # Large document for caching tests (needs 1024+ tokens for Claude models) -LARGE_DOCUMENT_FOR_CACHING = ( - """ +LARGE_DOCUMENT_FOR_CACHING = """ This is a comprehensive legal agreement between Party A and Party B. ARTICLE 1: DEFINITIONS @@ -77,9 +78,7 @@ ARTICLE 9: GENERAL PROVISIONS 9.5 Waiver of any provision shall not constitute ongoing waiver. IN WITNESS WHEREOF, the parties have executed this Agreement. -""" - * 8 -) # Repeat to ensure we have enough tokens (need 1024+ for Claude models) +""" * 8 # Repeat to ensure we have enough tokens (need 1024+ for Claude models) class BaseAnthropicMessagesPromptCachingTest(ABC): @@ -130,6 +129,7 @@ class BaseAnthropicMessagesPromptCachingTest(ABC): This validates that the cache_control field is being passed through correctly and the provider is creating a cache. """ + _skip_live_prompt_caching_test() litellm._turn_on_debug() messages = self.get_messages_with_cache_control() @@ -167,6 +167,7 @@ class BaseAnthropicMessagesPromptCachingTest(ABC): This validates that caching is working end-to-end. """ + _skip_live_prompt_caching_test() litellm._turn_on_debug() messages = self.get_messages_with_cache_control() @@ -207,6 +208,7 @@ class BaseAnthropicMessagesPromptCachingTest(ABC): """ E2E test: Prompt caching with system message should work. """ + _skip_live_prompt_caching_test() litellm._turn_on_debug() messages = [ @@ -268,6 +270,7 @@ class BaseAnthropicMessagesPromptCachingTest(ABC): This validates that cache_creation_input_tokens and cache_read_input_tokens are correctly returned in the streaming response's message_delta event. """ + _skip_live_prompt_caching_test() litellm._turn_on_debug() messages = self.get_messages_with_cache_control() @@ -365,6 +368,7 @@ class BaseAnthropicMessagesPromptCachingTest(ABC): """ E2E test: Second streaming call should return cache_read_input_tokens > 0. """ + _skip_live_prompt_caching_test() litellm._turn_on_debug() messages = self.get_messages_with_cache_control() @@ -443,6 +447,7 @@ class BaseAnthropicMessagesPromptCachingTest(ABC): didn't include cache fields in message_start, causing clients to think caching wasn't supported. """ + _skip_live_prompt_caching_test() litellm._turn_on_debug() messages = self.get_messages_with_cache_control() diff --git a/tests/pass_through_unit_tests/conftest.py b/tests/pass_through_unit_tests/conftest.py index 390e14b7f11..10615ddcb73 100644 --- a/tests/pass_through_unit_tests/conftest.py +++ b/tests/pass_through_unit_tests/conftest.py @@ -19,16 +19,7 @@ from tests._vcr_conftest_common import ( # noqa: E402,F401 vcr_config_dict, ) -# Tests that observe live cross-call provider state — typically a -# warm-up call followed by an assertion that the *second* call sees the -# upstream's prompt-cache (Anthropic / Bedrock prompt-caching). VCR's -# deterministic replay can't model this: both calls match the same -# cassette episode, so the second call returns the first call's -# pre-warmup response. Opt these out so they run live (no caching). -_VCR_INCOMPATIBLE_NODEID_SUFFIXES = ( - "::test_prompt_caching_returns_cache_read_tokens_on_second_call", - "::test_prompt_caching_streaming_second_call_returns_cache_read", -) +_VCR_INCOMPATIBLE_NODEID_SUFFIXES: tuple[str, ...] = () _verbose_state = VerboseReporterState() From 97ba7e1a30588b9bce87c9efc96b34bf2d1de375 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Wed, 3 Jun 2026 14:07:59 -0700 Subject: [PATCH 06/11] fix(key_generate): exempt UI/CLI session tokens from the budget ceiling for team keys (#29612) Non-admin users creating a team key through the UI were rejected with "max_budget cannot exceed the caller's own max_budget (0.25)". The request is authenticated by a UI/CLI session token whose max_budget is the per-session chat spend cap (max_ui_session_budget, default $0.25), and the delegated-authority budget ceiling (GHSA-q775-qw9r-2r4g) treated that cap as a delegation limit. Skip the ceiling only when a session token creates a team key (data.team_id set); that key's spend is bounded by the team budget at request time. Personal keys and every other non-admin caller keep the ceiling, so a session token cannot mint an arbitrary-budget personal key. --- .../key_management_endpoints.py | 9 ++ .../test_key_management_endpoints.py | 84 +++++++++++++++++++ 2 files changed, 93 insertions(+) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index cf90f0661b3..771f8287b6e 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -722,8 +722,17 @@ async def _common_key_generation_helper( # noqa: PLR0915 # Delegated-authority ceiling (GHSA-q775-qw9r-2r4g): a non-admin caller # with an explicit budget cannot grant a key a higher budget than their own. # Callers with max_budget=None (unlimited) can delegate any budget. + # A UI/CLI session token's max_budget is a per-session chat spend cap + # (max_ui_session_budget), not a delegation authority, so it is exempt only + # when creating a team key - that key's spend is bounded by the team budget + # at request time. Personal keys keep the ceiling; nothing else bounds them. + is_ui_session_team_key = ( + user_api_key_dict.team_id == UI_SESSION_TOKEN_TEAM_ID + and data.team_id is not None + ) if ( user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value + and not is_ui_session_team_key and _requested_max_budget is not None and user_api_key_dict.max_budget is not None and _requested_max_budget > user_api_key_dict.max_budget diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index cda22da6ebd..ca4a3f4fea0 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -11539,3 +11539,87 @@ async def test_ghsa_q775_admin_bypasses_budget_ceiling(): litellm_changed_by=None, ) assert result is not None + + +@pytest.mark.asyncio +async def test_ghsa_q775_ui_session_token_team_key_exempt_from_budget_ceiling(): + """ + Regression: a UI/CLI session token (team_id=litellm-dashboard) creating a + TEAM key (data.team_id set) is exempt from the delegated-authority ceiling. + The session max_budget is a per-session chat spend cap (max_ui_session_budget, + default $0.25), not a delegation authority, and the team key's spend is bounded + by the team budget at request time. This is the team-admin key-creation flow + blocked since v1.86.x. Calls the helper directly so the ceiling runs (mocking + out _common_key_generation_helper would mock out the check under test). + """ + from litellm.constants import UI_SESSION_TOKEN_TEAM_ID + + data = GenerateKeyRequest(max_budget=500, team_id="team-abc") + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-ui-session", + user_id="user-1", + team_id=UI_SESSION_TOKEN_TEAM_ID, + max_budget=0.25, + ) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", AsyncMock()), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + patch("litellm.proxy.proxy_server.llm_router", None), + patch("litellm.proxy.proxy_server.premium_user", False), + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "default_user_id"), + ): + try: + await _common_key_generation_helper( + data=data, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + team_table=MagicMock(), + ) + except (HTTPException, ProxyException) as err: + msg = str(getattr(err, "detail", "")) + str(getattr(err, "message", "")) + assert ( + "cannot exceed" not in msg.lower() + ), "UI/CLI session token creating a team key must be exempt from the ceiling" + + +@pytest.mark.asyncio +async def test_ghsa_q775_ui_session_token_personal_key_still_capped(): + """ + Security regression for GHSA-q775: the session-token exemption must NOT extend + to personal keys. A UI/CLI session token (team_id=litellm-dashboard) creating a + key with no data.team_id is still bound by the ceiling; otherwise a session + token - or a leaked one, whose blast radius is the $0.25 chat cap - could mint + an arbitrary-budget personal key, the exact escalation GHSA-q775 closed. Unlike + a team key, nothing else bounds a personal key's spend. + """ + from litellm.constants import UI_SESSION_TOKEN_TEAM_ID + + data = GenerateKeyRequest(max_budget=500) + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-ui-session", + user_id="user-1", + team_id=UI_SESSION_TOKEN_TEAM_ID, + max_budget=0.25, + ) + + mock_prisma_client = AsyncMock() + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + patch("litellm.proxy.proxy_server.user_custom_key_generate", None), + ): + with pytest.raises((HTTPException, ProxyException)) as exc_info: + await generate_key_fn( + data=data, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + ) + err = exc_info.value + code = getattr(err, "status_code", None) or getattr(err, "code", None) + msg = str(getattr(err, "detail", "")) + str(getattr(err, "message", "")) + assert str(code) == "400" + assert "cannot exceed" in msg.lower() From 5ee526d78ee34a084d6c9ac7f83b98e4aaaef0c5 Mon Sep 17 00:00:00 2001 From: milan-berri Date: Thu, 4 Jun 2026 00:59:18 +0300 Subject: [PATCH 07/11] fix(realtime): allow null transcripts in stream logging payloads (#29625) Allow realtime event transcript fields to be nullable so GA conversation.item payloads with transcript=null don't fail logging normalization and suppress success callbacks. Co-authored-by: Cursor --- litellm/types/llms/openai.py | 4 +-- tests/test_litellm/test_cost_calculator.py | 38 ++++++++++++++++++++++ 2 files changed, 40 insertions(+), 2 deletions(-) diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 51e408f8c63..0c854d89bb1 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -1892,7 +1892,7 @@ class OpenAIRealtimeStreamResponseOutputItemContent(TypedDict, total=False): """The ID of the previous conversation item for reference""" text: str """The text content, used for 'input_text' / 'text' / 'output_text' content types""" - transcript: str + transcript: Optional[str] """The transcript content, used for 'input_audio' / 'audio' content types""" type: Literal[ "input_audio", @@ -1998,7 +1998,7 @@ class OpenAIRealtimeResponseContentPart(TypedDict, total=False): text: str """The text content, if type is 'text' or 'output_text'""" - transcript: str + transcript: Optional[str] """The transcript content, if type is 'audio' or 'output_audio'""" type: Union[ diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index d973f8b4542..3d45a3409d8 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -12,6 +12,7 @@ from pydantic import BaseModel import litellm from litellm.cost_calculator import ( + RealtimeAPITokenUsageProcessor, completion_cost, cost_per_token, handle_realtime_stream_cost_calculation, @@ -385,6 +386,43 @@ def test_handle_realtime_stream_cost_calculation(): assert cost == 0.0 # No usage, no cost +def test_realtime_logging_object_allows_null_transcript_in_conversation_item_added(): + results: OpenAIRealtimeStreamList = [ + { + "type": "conversation.item.added", + "event_id": "event_added", + "item": { + "id": "item_123", + "type": "message", + "role": "assistant", + "status": "in_progress", + "content": [{"type": "audio", "transcript": None}], + }, + }, + { + "type": "response.done", + "event_id": "event_done", + "response": { + "id": "resp_123", + "object": "realtime.response", + "status": "completed", + "usage": {"input_tokens": 11, "output_tokens": 7, "total_tokens": 18}, + }, + }, + ] + + usage = RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results( + results=results + ) + logging_result = RealtimeAPITokenUsageProcessor.create_logging_realtime_object( + usage=usage, + results=results, + ) + + assert logging_result.usage.total_tokens == 18 + assert logging_result.results[0]["item"]["content"][0]["transcript"] is None + + def test_custom_pricing_with_router_model_id(): from litellm import Router From c7f1bcfd0d25f261384fb930b0026dbdaa23e01d Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 3 Jun 2026 15:50:20 -0700 Subject: [PATCH 08/11] build(ui): migrate eslint to flat config and bump eslint-config-next to 16 (#29626) ESLint 9 defaults to flat config and eslint-config-next was pinned at 15 while Next is on 16, so eslint only ran with ESLINT_USE_FLAT_CONFIG=false and next lint is gone on Next 16. Replace .eslintrc.json with a native flat eslint.config.mjs (config-next 16 ships flat configs, so no FlatCompat shim is needed), bump eslint-config-next to 16.2.6, add @eslint/js and typescript-eslint as explicit devDeps for the recommended rule sets, and point the lint script at eslint directly. This only makes eslint runnable on modern tooling; it does not wire it into CI. The same rules carry over (next/core-web-vitals, eslint and typescript-eslint recommended, prettier, unused-imports) --- ui/litellm-dashboard/.eslintrc.json | 17 - ui/litellm-dashboard/eslint.config.mjs | 33 ++ ui/litellm-dashboard/package-lock.json | 546 ++++++++++++++++++++----- ui/litellm-dashboard/package.json | 6 +- 4 files changed, 482 insertions(+), 120 deletions(-) delete mode 100644 ui/litellm-dashboard/.eslintrc.json create mode 100644 ui/litellm-dashboard/eslint.config.mjs diff --git a/ui/litellm-dashboard/.eslintrc.json b/ui/litellm-dashboard/.eslintrc.json deleted file mode 100644 index 90edda434cc..00000000000 --- a/ui/litellm-dashboard/.eslintrc.json +++ /dev/null @@ -1,17 +0,0 @@ -{ - "extends": ["next/core-web-vitals", "eslint:recommended", "plugin:@typescript-eslint/recommended", "prettier"], - "plugins": ["unused-imports"], - "rules": { - "unused-imports/no-unused-imports": "error", - "@typescript-eslint/no-explicit-any": "off", - "@typescript-eslint/no-unused-vars": "off", - "@typescript-eslint/no-unused-expressions": "off", - "@typescript-eslint/ban-ts-comment": "off", - "prefer-const": "off", - "no-empty": "off", - "no-prototype-builtins": "off", - "no-useless-catch": "off", - "no-useless-escape": "off", - "no-self-assign": "off" - } -} diff --git a/ui/litellm-dashboard/eslint.config.mjs b/ui/litellm-dashboard/eslint.config.mjs new file mode 100644 index 00000000000..838cd4a069d --- /dev/null +++ b/ui/litellm-dashboard/eslint.config.mjs @@ -0,0 +1,33 @@ +import js from "@eslint/js"; +import tseslint from "typescript-eslint"; +import nextCoreWebVitals from "eslint-config-next/core-web-vitals"; +import prettier from "eslint-config-prettier/flat"; +import unusedImports from "eslint-plugin-unused-imports"; + +const eslintConfig = [ + { + ignores: [".next/**", "out/**", "build/**", "coverage/**", "next-env.d.ts"], + }, + js.configs.recommended, + ...tseslint.configs.recommended, + ...nextCoreWebVitals, + prettier, + { + plugins: { "unused-imports": unusedImports }, + rules: { + "unused-imports/no-unused-imports": "error", + "@typescript-eslint/no-explicit-any": "off", + "@typescript-eslint/no-unused-vars": "off", + "@typescript-eslint/no-unused-expressions": "off", + "@typescript-eslint/ban-ts-comment": "off", + "prefer-const": "off", + "no-empty": "off", + "no-prototype-builtins": "off", + "no-useless-catch": "off", + "no-useless-escape": "off", + "no-self-assign": "off", + }, + }, +]; + +export default eslintConfig; diff --git a/ui/litellm-dashboard/package-lock.json b/ui/litellm-dashboard/package-lock.json index 97bc797fd54..858844d7e5c 100644 --- a/ui/litellm-dashboard/package-lock.json +++ b/ui/litellm-dashboard/package-lock.json @@ -37,6 +37,7 @@ "uuid": "14.0.0" }, "devDependencies": { + "@eslint/js": "9.39.2", "@playwright/test": "1.58.1", "@tailwindcss/forms": "0.5.11", "@testing-library/dom": "10.4.1", @@ -56,7 +57,7 @@ "autoprefixer": "10.4.24", "dotenv": "17.2.3", "eslint": "9.39.2", - "eslint-config-next": "15.5.10", + "eslint-config-next": "16.2.6", "eslint-config-prettier": "10.1.8", "eslint-plugin-unused-imports": "4.3.0", "jsdom": "27.4.0", @@ -65,6 +66,7 @@ "prettier": "3.2.5", "tailwindcss": "3.4.19", "typescript": "5.9.3", + "typescript-eslint": "8.60.1", "vite": "7.3.2", "vitest": "3.2.4" }, @@ -266,13 +268,13 @@ "license": "MIT" }, "node_modules/@babel/code-frame": { - "version": "7.29.0", - "resolved": "https://registry.npmjs.org/@babel/code-frame/-/code-frame-7.29.0.tgz", - "integrity": "sha512-9NhCeYjq9+3uxgdtp20LSiJXJvN0FeCtNGpJxuMFZ1Kv3cWUNb6DOhJwUvcVCzKGR66cw4njwM6hrJLqgOwbcw==", + "version": "7.29.7", + "resolved": "https://registry.npmjs.org/@babel/code-frame/-/code-frame-7.29.7.tgz", + "integrity": "sha512-Aup7aUOfpbAUg2ROOJN6Iw5f9DMBlzu0mIkm/malLQFN/YQgO48wCj0Kxa3sEHJvPVFg7siR+qRInwXd2qhQKw==", "dev": true, "license": "MIT", "dependencies": { - "@babel/helper-validator-identifier": "^7.28.5", + "@babel/helper-validator-identifier": "^7.29.7", "js-tokens": "^4.0.0", "picocolors": "^1.1.1" }, @@ -280,10 +282,170 @@ "node": ">=6.9.0" } }, + "node_modules/@babel/compat-data": { + "version": "7.29.7", + "resolved": "https://registry.npmjs.org/@babel/compat-data/-/compat-data-7.29.7.tgz", + "integrity": "sha512-locTkQyKvwIEgBzVrn8693ebc97F2U8ZHjbXwDXJ5Fn2TCpNwTlKcaKLkdHop5c/icOFE7qt7Q9JC5hnKNa6Gg==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=6.9.0" + } + }, + "node_modules/@babel/core": { + "version": "7.29.7", + "resolved": "https://registry.npmjs.org/@babel/core/-/core-7.29.7.tgz", + "integrity": "sha512-RgHBCvtjbOK2gXSNBNIkNoEc9qoVEtau3hj8gEqKQuL3HZAibKarWFEI3Lfm6EYKkLalOh8eSrj9b+ch9H/VBA==", + "dev": true, + "license": "MIT", + "dependencies": { + "@babel/code-frame": "^7.29.7", + "@babel/generator": "^7.29.7", + "@babel/helper-compilation-targets": "^7.29.7", + "@babel/helper-module-transforms": "^7.29.7", + "@babel/helpers": "^7.29.7", + "@babel/parser": "^7.29.7", + "@babel/template": "^7.29.7", + "@babel/traverse": "^7.29.7", + "@babel/types": "^7.29.7", + "@jridgewell/remapping": "^2.3.5", + "convert-source-map": "^2.0.0", + "debug": "^4.1.0", + "gensync": "^1.0.0-beta.2", + "json5": "^2.2.3", + "semver": "^6.3.1" + }, + "engines": { + "node": ">=6.9.0" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/babel" + } + }, + "node_modules/@babel/core/node_modules/json5": { + "version": "2.2.3", + "resolved": "https://registry.npmjs.org/json5/-/json5-2.2.3.tgz", + "integrity": "sha512-XmOWe7eyHYH14cLdVPoyg+GOH3rYX++KpzrylJwSW98t3Nk+U8XOl8FWKOgwtzdb8lXGf6zYwDUzeHMWfxasyg==", + "dev": true, + "license": "MIT", + "bin": { + "json5": "lib/cli.js" + }, + "engines": { + "node": ">=6" + } + }, + "node_modules/@babel/core/node_modules/semver": { + "version": "6.3.1", + "resolved": "https://registry.npmjs.org/semver/-/semver-6.3.1.tgz", + "integrity": "sha512-BR7VvDCVHO+q2xBEWskxS6DJE1qRnb7DxzUrogb71CWoSficBxYsiAGd+Kl0mmq/MprG9yArRkyrQxTO6XjMzA==", + "dev": true, + "license": "ISC", + "bin": { + "semver": "bin/semver.js" + } + }, + "node_modules/@babel/generator": { + "version": "7.29.7", + "resolved": "https://registry.npmjs.org/@babel/generator/-/generator-7.29.7.tgz", + "integrity": "sha512-DkXD5OJQaAQIdZ1bt3UZdEnHAn9Imd3IVBdX03UFe+ony9Ojw5pzr9YVKGDY1jt+Gcn/FnGkNf8r+Vj5NOJWtQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "@babel/parser": "^7.29.7", + "@babel/types": "^7.29.7", + "@jridgewell/gen-mapping": "^0.3.12", + "@jridgewell/trace-mapping": "^0.3.28", + "jsesc": "^3.0.2" + }, + "engines": { + "node": ">=6.9.0" + } + }, + "node_modules/@babel/helper-compilation-targets": { + "version": "7.29.7", + "resolved": "https://registry.npmjs.org/@babel/helper-compilation-targets/-/helper-compilation-targets-7.29.7.tgz", + "integrity": "sha512-wem6WaBj4NaVYVdNhLPPVacES6ZJ+KBBfSkTMD3YZxbP3rm3Di85tJU5ljaUNhaOynt+Aj0xruhYuzQBt8n71g==", + "dev": true, + "license": "MIT", + "dependencies": { + "@babel/compat-data": "^7.29.7", + "@babel/helper-validator-option": "^7.29.7", + "browserslist": "^4.24.0", + "lru-cache": "^5.1.1", + "semver": "^6.3.1" + }, + "engines": { + "node": ">=6.9.0" + } + }, + "node_modules/@babel/helper-compilation-targets/node_modules/lru-cache": { + "version": "5.1.1", + "resolved": "https://registry.npmjs.org/lru-cache/-/lru-cache-5.1.1.tgz", + "integrity": "sha512-KpNARQA3Iwv+jTA0utUVVbrh+Jlrr1Fv0e56GGzAFOXN7dk/FviaDW8LHmK52DlcH4WP2n6gI8vN1aesBFgo9w==", + "dev": true, + "license": "ISC", + "dependencies": { + "yallist": "^3.0.2" + } + }, + "node_modules/@babel/helper-compilation-targets/node_modules/semver": { + "version": "6.3.1", + "resolved": "https://registry.npmjs.org/semver/-/semver-6.3.1.tgz", + "integrity": "sha512-BR7VvDCVHO+q2xBEWskxS6DJE1qRnb7DxzUrogb71CWoSficBxYsiAGd+Kl0mmq/MprG9yArRkyrQxTO6XjMzA==", + "dev": true, + "license": "ISC", + "bin": { + "semver": "bin/semver.js" + } + }, + "node_modules/@babel/helper-globals": { + "version": "7.29.7", + "resolved": "https://registry.npmjs.org/@babel/helper-globals/-/helper-globals-7.29.7.tgz", + "integrity": "sha512-3nQVUAtvkKH9zahfWgw96Jc/uFOmjACE1kQz82E2lqWmHBgjzbNlsC22nuQTfahmWeQtTq5nQ/4Nnd2A1wj4zA==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=6.9.0" + } + }, + "node_modules/@babel/helper-module-imports": { + "version": "7.29.7", + "resolved": "https://registry.npmjs.org/@babel/helper-module-imports/-/helper-module-imports-7.29.7.tgz", + "integrity": "sha512-ejHwrQQYcm9xnTivShn2IDOlIzInN34AXskvq9QicvCtEzq1Vzclu/tKF8Jq1Cg8JG2GL6/EmjgsCT7lXepE3g==", + "dev": true, + "license": "MIT", + "dependencies": { + "@babel/traverse": "^7.29.7", + "@babel/types": "^7.29.7" + }, + "engines": { + "node": ">=6.9.0" + } + }, + "node_modules/@babel/helper-module-transforms": { + "version": "7.29.7", + "resolved": "https://registry.npmjs.org/@babel/helper-module-transforms/-/helper-module-transforms-7.29.7.tgz", + "integrity": "sha512-UPUVSyXbOh627KiCIGQSgwWzGeBKLkaJ9PJEdrngIwMSzxLR4jS4+f1f1jb7VzBbg8nFLaYotvVPFCTqdrmTAg==", + "dev": true, + "license": "MIT", + "dependencies": { + "@babel/helper-module-imports": "^7.29.7", + "@babel/helper-validator-identifier": "^7.29.7", + "@babel/traverse": "^7.29.7" + }, + "engines": { + "node": ">=6.9.0" + }, + "peerDependencies": { + "@babel/core": "^7.0.0" + } + }, "node_modules/@babel/helper-string-parser": { - "version": "7.27.1", - "resolved": "https://registry.npmjs.org/@babel/helper-string-parser/-/helper-string-parser-7.27.1.tgz", - "integrity": "sha512-qMlSxKbpRlAridDExk92nSobyDdpPijUq2DW6oDnUqd0iOGxmQjyqhMIihI9+zv4LPyZdRje2cavWPbCbWm3eA==", + "version": "7.29.7", + "resolved": "https://registry.npmjs.org/@babel/helper-string-parser/-/helper-string-parser-7.29.7.tgz", + "integrity": "sha512-Pb5ijPrZ89GDH8223L4UP8i6QApWxs04RbPQJTeWDV0/keR2E36MeKnyr6LYmUUvqRRI+Iv87SuF1W6ErINzYw==", "dev": true, "license": "MIT", "engines": { @@ -291,23 +453,47 @@ } }, "node_modules/@babel/helper-validator-identifier": { - "version": "7.28.5", - "resolved": "https://registry.npmjs.org/@babel/helper-validator-identifier/-/helper-validator-identifier-7.28.5.tgz", - "integrity": "sha512-qSs4ifwzKJSV39ucNjsvc6WVHs6b7S03sOh2OcHF9UHfVPqWWALUsNUVzhSBiItjRZoLHx7nIarVjqKVusUZ1Q==", + "version": "7.29.7", + "resolved": "https://registry.npmjs.org/@babel/helper-validator-identifier/-/helper-validator-identifier-7.29.7.tgz", + "integrity": "sha512-qehxGkRj55h/ff8EMaJ+cYhyaKlHIxqYDn682wQD7RNp9UujOQsHog2uS0r2vzr4pW+sXf90NeeayjcNaX3fFg==", "dev": true, "license": "MIT", "engines": { "node": ">=6.9.0" } }, - "node_modules/@babel/parser": { - "version": "7.29.3", - "resolved": "https://registry.npmjs.org/@babel/parser/-/parser-7.29.3.tgz", - "integrity": "sha512-b3ctpQwp+PROvU/cttc4OYl4MzfJUWy6FZg+PMXfzmt/+39iHVF0sDfqay8TQM3JA2EUOyKcFZt75jWriQijsA==", + "node_modules/@babel/helper-validator-option": { + "version": "7.29.7", + "resolved": "https://registry.npmjs.org/@babel/helper-validator-option/-/helper-validator-option-7.29.7.tgz", + "integrity": "sha512-N9ZErrD+yW5geCDtBqnOoxmR8+tNKiGuxKlDpuJxfsqpa2dFcexaziGAE/qoHLiDDreVNMupxGmSoNlyvsA3gw==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=6.9.0" + } + }, + "node_modules/@babel/helpers": { + "version": "7.29.7", + "resolved": "https://registry.npmjs.org/@babel/helpers/-/helpers-7.29.7.tgz", + "integrity": "sha512-1k2lAGRMfHTcwuNYcCNUmaUffmQv8KWMfh2iJUUeRlwlwH4FdNG7mfPI10NPfLHJFThE4Tyr4mv7kTNZOiPuBg==", "dev": true, "license": "MIT", "dependencies": { - "@babel/types": "^7.29.0" + "@babel/template": "^7.29.7", + "@babel/types": "^7.29.7" + }, + "engines": { + "node": ">=6.9.0" + } + }, + "node_modules/@babel/parser": { + "version": "7.29.7", + "resolved": "https://registry.npmjs.org/@babel/parser/-/parser-7.29.7.tgz", + "integrity": "sha512-hnORnjP/1P/zFEndoeX+n+t1RwWRJiJpM/jO7FW32Kn9r5+sJB2JWOdYo4L6k78j15eCwY3Gm/7364B1EMwtNg==", + "dev": true, + "license": "MIT", + "dependencies": { + "@babel/types": "^7.29.7" }, "bin": { "parser": "bin/babel-parser.js" @@ -325,15 +511,49 @@ "node": ">=6.9.0" } }, - "node_modules/@babel/types": { - "version": "7.29.0", - "resolved": "https://registry.npmjs.org/@babel/types/-/types-7.29.0.tgz", - "integrity": "sha512-LwdZHpScM4Qz8Xw2iKSzS+cfglZzJGvofQICy7W7v4caru4EaAmyUuO6BGrbyQ2mYV11W0U8j5mBhd14dd3B0A==", + "node_modules/@babel/template": { + "version": "7.29.7", + "resolved": "https://registry.npmjs.org/@babel/template/-/template-7.29.7.tgz", + "integrity": "sha512-puq+Gf35oI24FeN11LkoUQFqv9uwNeWpxXZi/Ji3rRIoKAzKnxRaZ+Gkj0vKS9ZCiTESfng1N9LyOyXvo+m+Gg==", "dev": true, "license": "MIT", "dependencies": { - "@babel/helper-string-parser": "^7.27.1", - "@babel/helper-validator-identifier": "^7.28.5" + "@babel/code-frame": "^7.29.7", + "@babel/parser": "^7.29.7", + "@babel/types": "^7.29.7" + }, + "engines": { + "node": ">=6.9.0" + } + }, + "node_modules/@babel/traverse": { + "version": "7.29.7", + "resolved": "https://registry.npmjs.org/@babel/traverse/-/traverse-7.29.7.tgz", + "integrity": "sha512-EhlfNQtZ+NK22w5BM61ciuiq1m58ed33Wr1Xan//ZRTy6hgjnwyCffRYwzsGXdASJSUJ1guZILsErh1eQcl+zw==", + "dev": true, + "license": "MIT", + "dependencies": { + "@babel/code-frame": "^7.29.7", + "@babel/generator": "^7.29.7", + "@babel/helper-globals": "^7.29.7", + "@babel/parser": "^7.29.7", + "@babel/template": "^7.29.7", + "@babel/types": "^7.29.7", + "debug": "^4.3.1" + }, + "engines": { + "node": ">=6.9.0" + } + }, + "node_modules/@babel/types": { + "version": "7.29.7", + "resolved": "https://registry.npmjs.org/@babel/types/-/types-7.29.7.tgz", + "integrity": "sha512-4zBIxpPzowiZpusoFkyGVwakdRJUyuH5PxQ/PrqghfdFWWasvnCdPfQXHrenDai+gyLARulZjZowCOj6fjT4pA==", + "dev": true, + "license": "MIT", + "dependencies": { + "@babel/helper-string-parser": "^7.29.7", + "@babel/helper-validator-identifier": "^7.29.7" }, "engines": { "node": ">=6.9.0" @@ -1838,6 +2058,17 @@ "@jridgewell/trace-mapping": "^0.3.24" } }, + "node_modules/@jridgewell/remapping": { + "version": "2.3.5", + "resolved": "https://registry.npmjs.org/@jridgewell/remapping/-/remapping-2.3.5.tgz", + "integrity": "sha512-LI9u/+laYG4Ds1TDKSJW2YPrIlcVYOwi2fUC6xB43lueCjgxV4lffOCZCtYFiH6TNOX+tQKXx97T4IKHbhyHEQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "@jridgewell/gen-mapping": "^0.3.5", + "@jridgewell/trace-mapping": "^0.3.24" + } + }, "node_modules/@jridgewell/resolve-uri": { "version": "3.1.2", "resolved": "https://registry.npmjs.org/@jridgewell/resolve-uri/-/resolve-uri-3.1.2.tgz", @@ -1889,9 +2120,9 @@ "license": "MIT" }, "node_modules/@next/eslint-plugin-next": { - "version": "15.5.10", - "resolved": "https://registry.npmjs.org/@next/eslint-plugin-next/-/eslint-plugin-next-15.5.10.tgz", - "integrity": "sha512-fDpxcy6G7Il4lQVVsaJD0fdC2/+SmuBGTF+edRLlsR4ZFOE3W2VyzrrGYdg/pHW8TydeAdSVM+mIzITGtZ3yWA==", + "version": "16.2.6", + "resolved": "https://registry.npmjs.org/@next/eslint-plugin-next/-/eslint-plugin-next-16.2.6.tgz", + "integrity": "sha512-Z8l6o4JWKUl755x4R+wogD86KPeU+Ckw4K+SYG4kHeOJtRenDeK+OSbGcqZpDtbwn9DsJVdir2UxmwXuinUbUw==", "dev": true, "license": "MIT", "dependencies": { @@ -2928,13 +3159,6 @@ "dev": true, "license": "MIT" }, - "node_modules/@rushstack/eslint-patch": { - "version": "1.16.1", - "resolved": "https://registry.npmjs.org/@rushstack/eslint-patch/-/eslint-patch-1.16.1.tgz", - "integrity": "sha512-TvZbIpeKqGQQ7X0zSCvPH9riMSFQFSggnfBjFZ1mEoILW+UuXCKwOoPcgjMwiUtRqFZ8jWhPJc4um14vC6I4ag==", - "dev": true, - "license": "MIT" - }, "node_modules/@swc/helpers": { "version": "0.5.21", "resolved": "https://registry.npmjs.org/@swc/helpers/-/helpers-0.5.21.tgz", @@ -3467,17 +3691,17 @@ "license": "MIT" }, "node_modules/@typescript-eslint/eslint-plugin": { - "version": "8.59.2", - "resolved": "https://registry.npmjs.org/@typescript-eslint/eslint-plugin/-/eslint-plugin-8.59.2.tgz", - "integrity": "sha512-j/bwmkBvHUtPNxzuWe5z6BEk3q54YRyGlBXkSsmfoih7zNrBvl5A9A98anlp/7JbyZcWIJ8KXo/3Tq/DjFLtuQ==", + "version": "8.60.1", + "resolved": "https://registry.npmjs.org/@typescript-eslint/eslint-plugin/-/eslint-plugin-8.60.1.tgz", + "integrity": "sha512-JQ4S5GB0tfjO8BuJ4fcX+HodkzJjYBV+7OJ+wLygaX7OGQ7FudyHL4NSCA6ob+w3Yn+5MkKIozOwQhXeM7opVg==", "dev": true, "license": "MIT", "dependencies": { "@eslint-community/regexpp": "^4.12.2", - "@typescript-eslint/scope-manager": "8.59.2", - "@typescript-eslint/type-utils": "8.59.2", - "@typescript-eslint/utils": "8.59.2", - "@typescript-eslint/visitor-keys": "8.59.2", + "@typescript-eslint/scope-manager": "8.60.1", + "@typescript-eslint/type-utils": "8.60.1", + "@typescript-eslint/utils": "8.60.1", + "@typescript-eslint/visitor-keys": "8.60.1", "ignore": "^7.0.5", "natural-compare": "^1.4.0", "ts-api-utils": "^2.5.0" @@ -3490,7 +3714,7 @@ "url": "https://opencollective.com/typescript-eslint" }, "peerDependencies": { - "@typescript-eslint/parser": "^8.59.2", + "@typescript-eslint/parser": "^8.60.1", "eslint": "^8.57.0 || ^9.0.0 || ^10.0.0", "typescript": ">=4.8.4 <6.1.0" } @@ -3506,16 +3730,16 @@ } }, "node_modules/@typescript-eslint/parser": { - "version": "8.59.2", - "resolved": "https://registry.npmjs.org/@typescript-eslint/parser/-/parser-8.59.2.tgz", - "integrity": "sha512-plR3pp6D+SSUn1HM7xvSkx12/DhoHInI2YF35KAcVFNZvlC0gtrWqx7Qq1oH2Ssgi0vlFRCTbP+DZc7B9+TtsQ==", + "version": "8.60.1", + "resolved": "https://registry.npmjs.org/@typescript-eslint/parser/-/parser-8.60.1.tgz", + "integrity": "sha512-A0M6ua6H252bVjPvvtSgl2QA4+ET9S5Mtkb2GDyTxIhH/C4qDItT7RQNO5PhMC6NXGYXOR9dIalcDDgBKT7oFA==", "dev": true, "license": "MIT", "dependencies": { - "@typescript-eslint/scope-manager": "8.59.2", - "@typescript-eslint/types": "8.59.2", - "@typescript-eslint/typescript-estree": "8.59.2", - "@typescript-eslint/visitor-keys": "8.59.2", + "@typescript-eslint/scope-manager": "8.60.1", + "@typescript-eslint/types": "8.60.1", + "@typescript-eslint/typescript-estree": "8.60.1", + "@typescript-eslint/visitor-keys": "8.60.1", "debug": "^4.4.3" }, "engines": { @@ -3531,14 +3755,14 @@ } }, "node_modules/@typescript-eslint/project-service": { - "version": "8.59.2", - "resolved": "https://registry.npmjs.org/@typescript-eslint/project-service/-/project-service-8.59.2.tgz", - "integrity": "sha512-+2hqvEkeyf/0FBor67duF0Ll7Ot8jyKzDQOSrxazF/danillRq2DwR9dLptsXpoZQqxE1UisSmoZewrlPas9Vw==", + "version": "8.60.1", + "resolved": "https://registry.npmjs.org/@typescript-eslint/project-service/-/project-service-8.60.1.tgz", + "integrity": "sha512-eXkTH2bxmXlqD1RnOPmLZ9ZM9D3VwSx04JOwBnP9RQ+yUA5a2Mu7SfW8uaV2Aon53NJzZlZYuX7tn91Izf+xaw==", "dev": true, "license": "MIT", "dependencies": { - "@typescript-eslint/tsconfig-utils": "^8.59.2", - "@typescript-eslint/types": "^8.59.2", + "@typescript-eslint/tsconfig-utils": "^8.60.1", + "@typescript-eslint/types": "^8.60.1", "debug": "^4.4.3" }, "engines": { @@ -3553,14 +3777,14 @@ } }, "node_modules/@typescript-eslint/scope-manager": { - "version": "8.59.2", - "resolved": "https://registry.npmjs.org/@typescript-eslint/scope-manager/-/scope-manager-8.59.2.tgz", - "integrity": "sha512-JzfyEpEtOU89CcFSwyNS3mu4MLvLSXqnmX05+aKBDM+TdR5jzcGOEBwxwGNxrEQ7p/z6kK2WyioCGBf2zZBnvg==", + "version": "8.60.1", + "resolved": "https://registry.npmjs.org/@typescript-eslint/scope-manager/-/scope-manager-8.60.1.tgz", + "integrity": "sha512-gvI5OQoptnxQnchOirukCuQ55svJSTuD/4k5+pC267xyBtYry748R9/c3tYUzb/iE6RZfllRz2lVulLCHkTm4w==", "dev": true, "license": "MIT", "dependencies": { - "@typescript-eslint/types": "8.59.2", - "@typescript-eslint/visitor-keys": "8.59.2" + "@typescript-eslint/types": "8.60.1", + "@typescript-eslint/visitor-keys": "8.60.1" }, "engines": { "node": "^18.18.0 || ^20.9.0 || >=21.1.0" @@ -3571,9 +3795,9 @@ } }, "node_modules/@typescript-eslint/tsconfig-utils": { - "version": "8.59.2", - "resolved": "https://registry.npmjs.org/@typescript-eslint/tsconfig-utils/-/tsconfig-utils-8.59.2.tgz", - "integrity": "sha512-BKK4alN7oi4C/zv4VqHQ+uRU+lTa6JGIZ7s1juw7b3RHo9OfKB+bKX3u0iVZetdsUCBBkSbdWbarJbmN0fTeSw==", + "version": "8.60.1", + "resolved": "https://registry.npmjs.org/@typescript-eslint/tsconfig-utils/-/tsconfig-utils-8.60.1.tgz", + "integrity": "sha512-nh8w4qAteiKuZu3pSSzG/yGKpw0OlkrKnzFmbVRenKaD4qc+7i1GrmZaLVkr8rk4uipiPGMOW4YsM6WmKZ5CvA==", "dev": true, "license": "MIT", "engines": { @@ -3588,15 +3812,15 @@ } }, "node_modules/@typescript-eslint/type-utils": { - "version": "8.59.2", - "resolved": "https://registry.npmjs.org/@typescript-eslint/type-utils/-/type-utils-8.59.2.tgz", - "integrity": "sha512-nhqaj1nmTdVVl/BP5omXNRGO38jn5iosis2vbdmupF2txCf8ylWT8lx+JlvMYYVqzGVKtjojUFoQ3JRWK+mfzQ==", + "version": "8.60.1", + "resolved": "https://registry.npmjs.org/@typescript-eslint/type-utils/-/type-utils-8.60.1.tgz", + "integrity": "sha512-sdwTrpjosW7ANQYJ39ZBF1ZyEMEGVB2UsikrserVM/30a/F1dTLnu9bGxEdosugyu5caigjLrR2qiD11asjI1A==", "dev": true, "license": "MIT", "dependencies": { - "@typescript-eslint/types": "8.59.2", - "@typescript-eslint/typescript-estree": "8.59.2", - "@typescript-eslint/utils": "8.59.2", + "@typescript-eslint/types": "8.60.1", + "@typescript-eslint/typescript-estree": "8.60.1", + "@typescript-eslint/utils": "8.60.1", "debug": "^4.4.3", "ts-api-utils": "^2.5.0" }, @@ -3613,9 +3837,9 @@ } }, "node_modules/@typescript-eslint/types": { - "version": "8.59.2", - "resolved": "https://registry.npmjs.org/@typescript-eslint/types/-/types-8.59.2.tgz", - "integrity": "sha512-e82GVOE8Ps3E++Egvb6Y3Dw0S10u8NkQ9KXmtRhCWJJ8kDhOJTvtMAWnFL16kB1583goCWXsr0NieKCZMs2/0Q==", + "version": "8.60.1", + "resolved": "https://registry.npmjs.org/@typescript-eslint/types/-/types-8.60.1.tgz", + "integrity": "sha512-4h0tY8ppCkdCzcrl2YM5M3my0xsE1Tf8om3owEu5oPWmXwkKRmk0j0LGDzYBGUcAlesEbxBhazqu/K4cu3Ug7w==", "dev": true, "license": "MIT", "engines": { @@ -3627,16 +3851,16 @@ } }, "node_modules/@typescript-eslint/typescript-estree": { - "version": "8.59.2", - "resolved": "https://registry.npmjs.org/@typescript-eslint/typescript-estree/-/typescript-estree-8.59.2.tgz", - "integrity": "sha512-o0XPGNwcWw+FIwStOWn+BwBuEmL6QXP0rsvAFg7ET1dey1Nr6Wb1ac8p5HEsK0ygO/6mUxlk+YWQD9xcb/nnXg==", + "version": "8.60.1", + "resolved": "https://registry.npmjs.org/@typescript-eslint/typescript-estree/-/typescript-estree-8.60.1.tgz", + "integrity": "sha512-alpRkfG8hlVE5kdJW2GkfgDgXxold3e8e4l6EnmhRmRLbekgAPCCGDVD++sABy9FcgPFroq+uFcCSM1vR57Cew==", "dev": true, "license": "MIT", "dependencies": { - "@typescript-eslint/project-service": "8.59.2", - "@typescript-eslint/tsconfig-utils": "8.59.2", - "@typescript-eslint/types": "8.59.2", - "@typescript-eslint/visitor-keys": "8.59.2", + "@typescript-eslint/project-service": "8.60.1", + "@typescript-eslint/tsconfig-utils": "8.60.1", + "@typescript-eslint/types": "8.60.1", + "@typescript-eslint/visitor-keys": "8.60.1", "debug": "^4.4.3", "minimatch": "^10.2.2", "semver": "^7.7.3", @@ -3655,16 +3879,16 @@ } }, "node_modules/@typescript-eslint/utils": { - "version": "8.59.2", - "resolved": "https://registry.npmjs.org/@typescript-eslint/utils/-/utils-8.59.2.tgz", - "integrity": "sha512-Juw3EinkXqjaffxz6roowvV7GZT/kET5vSKKZT6upl5TXdWkLkYmNPXwDDL2Vkt2DPn0nODIS4egC/0AGxKo/Q==", + "version": "8.60.1", + "resolved": "https://registry.npmjs.org/@typescript-eslint/utils/-/utils-8.60.1.tgz", + "integrity": "sha512-h2MPBLoNtjc3qZWfY3Tl51yPorQ2McHn8pJfcMNTcIvrrZrr90Ykffit0yjrPFWQcRcUxzH20+6OcVdW4yHtUg==", "dev": true, "license": "MIT", "dependencies": { "@eslint-community/eslint-utils": "^4.9.1", - "@typescript-eslint/scope-manager": "8.59.2", - "@typescript-eslint/types": "8.59.2", - "@typescript-eslint/typescript-estree": "8.59.2" + "@typescript-eslint/scope-manager": "8.60.1", + "@typescript-eslint/types": "8.60.1", + "@typescript-eslint/typescript-estree": "8.60.1" }, "engines": { "node": "^18.18.0 || ^20.9.0 || >=21.1.0" @@ -3679,13 +3903,13 @@ } }, "node_modules/@typescript-eslint/visitor-keys": { - "version": "8.59.2", - "resolved": "https://registry.npmjs.org/@typescript-eslint/visitor-keys/-/visitor-keys-8.59.2.tgz", - "integrity": "sha512-NwjLUnGy8/Zfx23fl50tRC8rYaYnM52xNRYFAXvmiil9yh1+K6aRVQMnzW6gQB/1DLgWt977lYQn7C+wtgXZiA==", + "version": "8.60.1", + "resolved": "https://registry.npmjs.org/@typescript-eslint/visitor-keys/-/visitor-keys-8.60.1.tgz", + "integrity": "sha512-EbGRQg4FhrmwLodl+t3JNAnXHWVr9Vp+Zl1QBZVPY4ByfkzIT8cX3K6QWODHtkIZqqJVEWvhHSx3v5PDHsaQag==", "dev": true, "license": "MIT", "dependencies": { - "@typescript-eslint/types": "8.59.2", + "@typescript-eslint/types": "8.60.1", "eslint-visitor-keys": "^5.0.0" }, "engines": { @@ -5103,6 +5327,13 @@ "integrity": "sha512-VRhuHOLoKYOy4UbilLbUzbYg93XLjv2PncJC50EuTWPA3gaja1UjBsUP/D/9/juV3vQFr6XBEzn9KCAHdUvOHw==", "license": "MIT" }, + "node_modules/convert-source-map": { + "version": "2.0.0", + "resolved": "https://registry.npmjs.org/convert-source-map/-/convert-source-map-2.0.0.tgz", + "integrity": "sha512-Kvp459HrV2FEJ1CAsi1Ku+MY3kasH19TFykTz2xWmMeq6bk2NU3XXvfJ+Q61m0xktWwt+1HSYf3JZsTms3aRJg==", + "dev": true, + "license": "MIT" + }, "node_modules/copy-to-clipboard": { "version": "3.3.3", "resolved": "https://registry.npmjs.org/copy-to-clipboard/-/copy-to-clipboard-3.3.3.tgz", @@ -5963,25 +6194,24 @@ } }, "node_modules/eslint-config-next": { - "version": "15.5.10", - "resolved": "https://registry.npmjs.org/eslint-config-next/-/eslint-config-next-15.5.10.tgz", - "integrity": "sha512-AeYOVGiSbIfH4KXFT3d0fIDm7yTslR/AWGoHLdsXQ99MH0zFWmkRIin1H7I9SFlkKgf4PKm9ncsyWHq1aAfHBA==", + "version": "16.2.6", + "resolved": "https://registry.npmjs.org/eslint-config-next/-/eslint-config-next-16.2.6.tgz", + "integrity": "sha512-z2ELYSkyrrJ6cuunTU8vhsT/RpouPkjaSah06nVW6Rg2Hpg0Vs8s497/e5s8G8qtdp4ccsiovz5P1rv+5VSW2Q==", "dev": true, "license": "MIT", "dependencies": { - "@next/eslint-plugin-next": "15.5.10", - "@rushstack/eslint-patch": "^1.10.3", - "@typescript-eslint/eslint-plugin": "^5.4.2 || ^6.0.0 || ^7.0.0 || ^8.0.0", - "@typescript-eslint/parser": "^5.4.2 || ^6.0.0 || ^7.0.0 || ^8.0.0", + "@next/eslint-plugin-next": "16.2.6", "eslint-import-resolver-node": "^0.3.6", "eslint-import-resolver-typescript": "^3.5.2", - "eslint-plugin-import": "^2.31.0", + "eslint-plugin-import": "^2.32.0", "eslint-plugin-jsx-a11y": "^6.10.0", "eslint-plugin-react": "^7.37.0", - "eslint-plugin-react-hooks": "^5.0.0" + "eslint-plugin-react-hooks": "^7.0.0", + "globals": "16.4.0", + "typescript-eslint": "^8.46.0" }, "peerDependencies": { - "eslint": "^7.23.0 || ^8.0.0 || ^9.0.0", + "eslint": ">=9.0.0", "typescript": ">=3.3.1" }, "peerDependenciesMeta": { @@ -5990,6 +6220,19 @@ } } }, + "node_modules/eslint-config-next/node_modules/globals": { + "version": "16.4.0", + "resolved": "https://registry.npmjs.org/globals/-/globals-16.4.0.tgz", + "integrity": "sha512-ob/2LcVVaVGCYN+r14cnwnoDPUufjiYgSqRhiFD0Q1iI4Odora5RE8Iv1D24hAz5oMophRGkGz+yuvQmmUMnMw==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=18" + }, + "funding": { + "url": "https://github.com/sponsors/sindresorhus" + } + }, "node_modules/eslint-config-prettier": { "version": "10.1.8", "resolved": "https://registry.npmjs.org/eslint-config-prettier/-/eslint-config-prettier-10.1.8.tgz", @@ -6219,16 +6462,23 @@ } }, "node_modules/eslint-plugin-react-hooks": { - "version": "5.2.0", - "resolved": "https://registry.npmjs.org/eslint-plugin-react-hooks/-/eslint-plugin-react-hooks-5.2.0.tgz", - "integrity": "sha512-+f15FfK64YQwZdJNELETdn5ibXEUQmW1DZL6KXhNnc2heoy/sg9VJJeT7n8TlMWouzWqSWavFkIhHyIbIAEapg==", + "version": "7.1.1", + "resolved": "https://registry.npmjs.org/eslint-plugin-react-hooks/-/eslint-plugin-react-hooks-7.1.1.tgz", + "integrity": "sha512-f2I7Gw6JbvCexzIInuSbZpfdQ44D7iqdWX01FKLvrPgqxoE7oMj8clOfto8U6vYiz4yd5oKu39rRSVOe1zRu0g==", "dev": true, "license": "MIT", + "dependencies": { + "@babel/core": "^7.24.4", + "@babel/parser": "^7.24.4", + "hermes-parser": "^0.25.1", + "zod": "^3.25.0 || ^4.0.0", + "zod-validation-error": "^3.5.0 || ^4.0.0" + }, "engines": { - "node": ">=10" + "node": ">=18" }, "peerDependencies": { - "eslint": "^3.0.0 || ^4.0.0 || ^5.0.0 || ^6.0.0 || ^7.0.0 || ^8.0.0-0 || ^9.0.0" + "eslint": "^3.0.0 || ^4.0.0 || ^5.0.0 || ^6.0.0 || ^7.0.0 || ^8.0.0-0 || ^9.0.0 || ^10.0.0" } }, "node_modules/eslint-plugin-react/node_modules/semver": { @@ -6734,6 +6984,16 @@ "node": ">= 0.4" } }, + "node_modules/gensync": { + "version": "1.0.0-beta.2", + "resolved": "https://registry.npmjs.org/gensync/-/gensync-1.0.0-beta.2.tgz", + "integrity": "sha512-3hN7NaskYvMDLQY55gnW3NQ+mesEAepTqlg+VEbj7zzqEMBVNhzcGYYeqFo/TlYz6eQiFcp1HcsCZO+nGgS8zg==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=6.9.0" + } + }, "node_modules/get-intrinsic": { "version": "1.3.0", "resolved": "https://registry.npmjs.org/get-intrinsic/-/get-intrinsic-1.3.0.tgz", @@ -7080,6 +7340,23 @@ "url": "https://github.com/sponsors/wooorm" } }, + "node_modules/hermes-estree": { + "version": "0.25.1", + "resolved": "https://registry.npmjs.org/hermes-estree/-/hermes-estree-0.25.1.tgz", + "integrity": "sha512-0wUoCcLp+5Ev5pDW2OriHC2MJCbwLwuRx+gAqMTOkGKJJiBCLjtrvy4PWUGn6MIVefecRpzoOZ/UV6iGdOr+Cw==", + "dev": true, + "license": "MIT" + }, + "node_modules/hermes-parser": { + "version": "0.25.1", + "resolved": "https://registry.npmjs.org/hermes-parser/-/hermes-parser-0.25.1.tgz", + "integrity": "sha512-6pEjquH3rqaI6cYAXYPcz9MS4rY6R4ngRgrgfDshRptUZIc3lw0MCIJIGDj9++mfySOuPTHB4nrSW99BCvOPIA==", + "dev": true, + "license": "MIT", + "dependencies": { + "hermes-estree": "0.25.1" + } + }, "node_modules/highlight.js": { "version": "10.7.3", "resolved": "https://registry.npmjs.org/highlight.js/-/highlight.js-10.7.3.tgz", @@ -7867,6 +8144,19 @@ } } }, + "node_modules/jsesc": { + "version": "3.1.0", + "resolved": "https://registry.npmjs.org/jsesc/-/jsesc-3.1.0.tgz", + "integrity": "sha512-/sM3dO2FOzXjKQhJuo0Q173wf2KOo8t4I8vHy6lF9poUp7bKT0/NHE8fPX23PwfhnykfqnC2xRxOnVw5XuGIaA==", + "dev": true, + "license": "MIT", + "bin": { + "jsesc": "bin/jsesc" + }, + "engines": { + "node": ">=6" + } + }, "node_modules/json-buffer": { "version": "3.0.1", "resolved": "https://registry.npmjs.org/json-buffer/-/json-buffer-3.0.1.tgz", @@ -12622,6 +12912,30 @@ "node": ">=14.17" } }, + "node_modules/typescript-eslint": { + "version": "8.60.1", + "resolved": "https://registry.npmjs.org/typescript-eslint/-/typescript-eslint-8.60.1.tgz", + "integrity": "sha512-6m5hkkRAp8lKvhVpcprAIn5KkehQEh+47oHH2VGnExEh7dhNxXlg6GPAOIu6TxbVQxhebrJDvjl3020ooiWCMA==", + "dev": true, + "license": "MIT", + "dependencies": { + "@typescript-eslint/eslint-plugin": "8.60.1", + "@typescript-eslint/parser": "8.60.1", + "@typescript-eslint/typescript-estree": "8.60.1", + "@typescript-eslint/utils": "8.60.1" + }, + "engines": { + "node": "^18.18.0 || ^20.9.0 || >=21.1.0" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/typescript-eslint" + }, + "peerDependencies": { + "eslint": "^8.57.0 || ^9.0.0 || ^10.0.0", + "typescript": ">=4.8.4 <6.1.0" + } + }, "node_modules/unbox-primitive": { "version": "1.1.0", "resolved": "https://registry.npmjs.org/unbox-primitive/-/unbox-primitive-1.1.0.tgz", @@ -13320,6 +13634,13 @@ "node": ">=0.4" } }, + "node_modules/yallist": { + "version": "3.1.1", + "resolved": "https://registry.npmjs.org/yallist/-/yallist-3.1.1.tgz", + "integrity": "sha512-a4UGQaWPH59mOXUYnAG2ewncQS4i4F43Tv3JoAM+s2VDAmS9NsK8GpDMLrCHPksFT7h3K6TOoUNn2pb7RoXx4g==", + "dev": true, + "license": "ISC" + }, "node_modules/yocto-queue": { "version": "0.1.0", "resolved": "https://registry.npmjs.org/yocto-queue/-/yocto-queue-0.1.0.tgz", @@ -13333,6 +13654,29 @@ "url": "https://github.com/sponsors/sindresorhus" } }, + "node_modules/zod": { + "version": "3.25.76", + "resolved": "https://registry.npmjs.org/zod/-/zod-3.25.76.tgz", + "integrity": "sha512-gzUt/qt81nXsFGKIFcC3YnfEAx5NkunCfnDlvuBSSFS02bcXu4Lmea0AFIUwbLWxWPx3d9p8S5QoaujKcNQxcQ==", + "devOptional": true, + "license": "MIT", + "funding": { + "url": "https://github.com/sponsors/colinhacks" + } + }, + "node_modules/zod-validation-error": { + "version": "4.0.2", + "resolved": "https://registry.npmjs.org/zod-validation-error/-/zod-validation-error-4.0.2.tgz", + "integrity": "sha512-Q6/nZLe6jxuU80qb/4uJ4t5v2VEZ44lzQjPDhYJNztRQ4wyWc6VF3D3Kb/fAuPetZQnhS3hnajCf9CsWesghLQ==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=18.0.0" + }, + "peerDependencies": { + "zod": "^3.25.0 || ^4.0.0" + } + }, "node_modules/zwitch": { "version": "2.0.4", "resolved": "https://registry.npmjs.org/zwitch/-/zwitch-2.0.4.tgz", diff --git a/ui/litellm-dashboard/package.json b/ui/litellm-dashboard/package.json index 72b9bc2a159..77731bbd693 100644 --- a/ui/litellm-dashboard/package.json +++ b/ui/litellm-dashboard/package.json @@ -7,7 +7,7 @@ "dev:webpack": "next dev --webpack", "build": "next build", "start": "next start", - "lint": "next lint", + "lint": "eslint .", "test": "vitest", "test:dot": "vitest --reporter=dot", "test:watch": "vitest -w", @@ -49,6 +49,7 @@ "uuid": "14.0.0" }, "devDependencies": { + "@eslint/js": "9.39.2", "@playwright/test": "1.58.1", "@tailwindcss/forms": "0.5.11", "@testing-library/dom": "10.4.1", @@ -68,7 +69,7 @@ "autoprefixer": "10.4.24", "dotenv": "17.2.3", "eslint": "9.39.2", - "eslint-config-next": "15.5.10", + "eslint-config-next": "16.2.6", "eslint-config-prettier": "10.1.8", "eslint-plugin-unused-imports": "4.3.0", "jsdom": "27.4.0", @@ -77,6 +78,7 @@ "prettier": "3.2.5", "tailwindcss": "3.4.19", "typescript": "5.9.3", + "typescript-eslint": "8.60.1", "vite": "7.3.2", "vitest": "3.2.4" }, From e9417603a38d53765894254fc6b588ff6700bf6a Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Wed, 3 Jun 2026 19:11:53 -0700 Subject: [PATCH 09/11] fix(key_generate): scope session-token team-key budget exemption to caller-supplied team_id (#29641) #29612 exempts UI/CLI session tokens from the key budget ceiling when they create a team key, keyed on data.team_id. That value is read after the default_key_generate_params loop can populate team_id, so on deployments that set default_key_generate_params.team_id a request the caller did not scope to a team is treated as a team key and skips the ceiling. Capture _requested_team_id before defaults run and key the exemption off it, mirroring how _requested_max_budget is already captured. Requests the caller did not scope to a team keep the ceiling. --- .../key_management_endpoints.py | 10 +++-- .../test_key_management_endpoints.py | 45 +++++++++++++++++++ 2 files changed, 51 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 771f8287b6e..c8c590af97c 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -691,10 +691,12 @@ async def _common_key_generation_helper( # noqa: PLR0915 prisma_client=prisma_client, ) - # Capture the caller-supplied max_budget before any defaults or upperbound - # params can fill it, so the ceiling check only fires when the caller - # explicitly requested a budget. + # Capture caller-supplied max_budget and team_id before any defaults or + # upperbound params can fill them, so the ceiling check and its team-key + # exemption key off what the caller explicitly requested, not a value that + # default_key_generate_params injected. _requested_max_budget = data.max_budget + _requested_team_id = data.team_id # check if user set default key/generate params on config.yaml if litellm.default_key_generate_params is not None: @@ -728,7 +730,7 @@ async def _common_key_generation_helper( # noqa: PLR0915 # at request time. Personal keys keep the ceiling; nothing else bounds them. is_ui_session_team_key = ( user_api_key_dict.team_id == UI_SESSION_TOKEN_TEAM_ID - and data.team_id is not None + and _requested_team_id is not None ) if ( user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index ca4a3f4fea0..3c212d86e65 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -11623,3 +11623,48 @@ async def test_ghsa_q775_ui_session_token_personal_key_still_capped(): msg = str(getattr(err, "detail", "")) + str(getattr(err, "message", "")) assert str(code) == "400" assert "cannot exceed" in msg.lower() + + +@pytest.mark.asyncio +async def test_ghsa_q775_default_team_id_does_not_grant_session_token_exemption(): + """ + Security regression for GHSA-q775: the team-key exemption must key off the + team_id the CALLER supplied, not one injected by default_key_generate_params. + With default_key_generate_params.team_id set, a UI session token's personal-key + request (no team_id) would otherwise have team_id auto-filled before the ceiling + check, flipping is_ui_session_team_key to True and bypassing the ceiling. The + request must still be rejected. Mirrors how _requested_max_budget is captured + before defaults run. + """ + from litellm.constants import UI_SESSION_TOKEN_TEAM_ID + + data = GenerateKeyRequest(max_budget=500) + assert data.team_id is None + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-ui-session", + user_id="user-1", + team_id=UI_SESSION_TOKEN_TEAM_ID, + max_budget=0.25, + ) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", AsyncMock()), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + patch("litellm.proxy.proxy_server.llm_router", None), + patch("litellm.proxy.proxy_server.premium_user", False), + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "default_user_id"), + patch("litellm.default_key_generate_params", {"team_id": "injected-team"}), + ): + with pytest.raises((HTTPException, ProxyException)) as exc_info: + await _common_key_generation_helper( + data=data, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + team_table=None, + ) + err = exc_info.value + code = getattr(err, "status_code", None) or getattr(err, "code", None) + msg = str(getattr(err, "detail", "")) + str(getattr(err, "message", "")) + assert str(code) == "400" + assert "cannot exceed" in msg.lower() From be7b9319d2017cd3590cee2805fe90574260a035 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 4 Jun 2026 04:53:14 -0700 Subject: [PATCH 10/11] fix(proxy): disable proxy buffering on streaming SSE responses (#29557) Streaming responses from the proxy (/chat/completions, /v1/messages, /v1/responses, assistants) all return through create_response() but never sent the headers that tell an intermediary reverse proxy not to buffer the SSE stream. nginx with the default proxy_buffering, k8s ingress-nginx, and Envoy/Istio sidecars therefore hold the whole stream and release it in one batch, which looks like a broken/buffered stream to the client even though litellm is yielding chunks incrementally. Add Cache-Control: no-cache and X-Accel-Buffering: no to every StreamingResponse create_response() returns, matching what the proxy already does for its own usage/policy SSE endpoints. Fixes #28384. --- litellm/proxy/common_request_processing.py | 13 +++++++--- .../proxy/test_common_request_processing.py | 25 +++++++++++++++++++ 2 files changed, 35 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 36acd9653e8..6558543370d 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -249,6 +249,13 @@ async def create_response( # noqa: PLR0915 If the first chunk is an error, return a standard JSON error response. Otherwise, return StreamingResponse and stream all content. """ + # Tell buffering reverse proxies (nginx, ingress-nginx, Envoy) to flush SSE + # immediately instead of releasing the whole stream in one batch (issue #28384). + streaming_headers = { + **headers, + "Cache-Control": "no-cache", + "X-Accel-Buffering": "no", + } first_chunk_value: Optional[str] = None final_status_code = default_status_code @@ -300,7 +307,7 @@ async def create_response( # noqa: PLR0915 return StreamingResponse( empty_gen(), media_type=media_type, - headers=headers, + headers=streaming_headers, status_code=default_status_code, ) except Exception as e: @@ -338,7 +345,7 @@ async def create_response( # noqa: PLR0915 return StreamingResponse( error_gen_message(), media_type=media_type, - headers=headers, + headers=streaming_headers, status_code=error_status, ) @@ -360,7 +367,7 @@ async def create_response( # noqa: PLR0915 return StreamingResponse( combined_generator(), media_type=media_type, - headers=headers, + headers=streaming_headers, status_code=final_status_code, ) diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 265a82d4a44..0f5a0cbe4b6 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -1258,6 +1258,31 @@ class TestCommonRequestProcessingHelpers: ) assert response.headers["x-custom-header"] == "TestValue" + async def test_create_streaming_response_disables_proxy_buffering(self): + """Regression for #28384: every StreamingResponse create_response returns + must carry the headers that stop nginx/ingress/Envoy from buffering the + SSE stream into one batch, while preserving caller-supplied headers.""" + + async def normal_stream(): + yield 'data: {"content": "part"}\n\n' + yield "data: [DONE]\n\n" + + async def empty_stream(): + if False: # never yields -> StopAsyncIteration + yield + + error_stream = AsyncMock() + error_stream.__anext__.side_effect = ValueError("boom") + + for generator in (normal_stream(), empty_stream(), error_stream): + response = await create_response( + generator, "text/event-stream", {"X-Custom-Header": "keep"} + ) + assert isinstance(response, StreamingResponse) + assert response.headers["x-accel-buffering"] == "no" + assert response.headers["cache-control"] == "no-cache" + assert response.headers["x-custom-header"] == "keep" + async def test_create_streaming_response_non_default_status_code(self): async def mock_generator(): yield 'data: {"content": "data"}\n\n' From 9196098e9e1d5cd7de7dc1407f5ee6af31754c9a Mon Sep 17 00:00:00 2001 From: michelligabriele Date: Thu, 4 Jun 2026 13:56:59 +0200 Subject: [PATCH 11/11] fix(mcp): gate /public/mcp_hub strictly on litellm.public_mcp_servers (#27764) * fix(mcp): gate /public/mcp_hub strictly on litellm.public_mcp_servers * fix(mcp): add public_mcp_hub_strict_whitelist flag (default True) for migration --- litellm/__init__.py | 1 + .../mcp_server/mcp_server_manager.py | 36 ++++- .../mcp_server/test_mcp_server_manager.py | 135 ++++++++++++++++++ .../public_endpoints/test_public_endpoints.py | 64 +++++++++ 4 files changed, 229 insertions(+), 7 deletions(-) diff --git a/litellm/__init__.py b/litellm/__init__.py index c954f5fd31e..98c9dcb5ddf 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -444,6 +444,7 @@ disable_copilot_system_to_assistant: bool = ( False # If false (default), converts all 'system' role messages to 'assistant' for GitHub Copilot compatibility. Set to true to disable this behavior. ) public_mcp_servers: Optional[List[str]] = None +public_mcp_hub_strict_whitelist: bool = True public_model_groups: Optional[List[str]] = None public_agent_groups: Optional[List[str]] = None # Supports both old format (Dict[str, str]) and new format (Dict[str, Dict[str, Any]]) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 739dc4a2f88..d9b112f6c21 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -3540,15 +3540,37 @@ class MCPServerManager: def get_public_mcp_servers(self) -> List[MCPServer]: """ - Get the public MCP servers (available_on_public_internet=True flag on server). - Also includes servers from litellm.public_mcp_servers for backwards compat. + Return the MCP servers published to the AI Hub via /v1/mcp/make_public. + + Default (litellm.public_mcp_hub_strict_whitelist=True): mirrors + /public/model_hub and /public/agent_hub — gates strictly on the + litellm.public_mcp_servers whitelist. Returns an empty list when no + servers have been published. The per-server available_on_public_internet + flag is unrelated — it governs IP-based access in + _is_server_accessible_from_ip, not hub visibility. + + Legacy (litellm.public_mcp_hub_strict_whitelist=False): preserves the + pre-fix behavior where any server with available_on_public_internet=True + is also included. Intended as a one-release migration window for + deployments that relied on the OR-with-default semantics; will be + removed in a future release. """ - servers: List[MCPServer] = [] + if litellm.public_mcp_hub_strict_whitelist: + if litellm.public_mcp_servers is None: + return [] + public_ids = set(litellm.public_mcp_servers) + return [ + server + for server in self.get_registry().values() + if server.server_id in public_ids + ] + public_ids = set(litellm.public_mcp_servers or []) - for server in self.get_registry().values(): - if server.available_on_public_internet or server.server_id in public_ids: - servers.append(server) - return servers + return [ + server + for server in self.get_registry().values() + if server.available_on_public_internet or server.server_id in public_ids + ] def expand_permission_list(self, identifiers: List[str]) -> List[str]: """ diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index ec690aef629..32e4ec19311 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -3984,5 +3984,140 @@ class TestApprovalStatusGate: assert "never-seen" not in manager.registry +class TestGetPublicMCPServers: + """ + /public/mcp_hub strict-whitelist semantics — mirrors /public/model_hub + and /public/agent_hub. Regression test for the PR #20607 OR-with-default + behavior that made `litellm.public_mcp_servers` ignored by the hub. + """ + + def _make_server(self, server_id, available_on_public_internet=True): + return MCPServer( + server_id=server_id, + name=server_id, + server_name=server_id, + transport=MCPTransport.http, + available_on_public_internet=available_on_public_internet, + ) + + def _make_manager(self, servers): + manager = MCPServerManager() + for s in servers: + manager.config_mcp_servers[s.server_id] = s + return manager + + @patch("litellm.public_mcp_servers", None) + def test_returns_empty_when_whitelist_is_none(self): + """No /make_public call yet → hub returns nothing, regardless of + per-server flags.""" + manager = self._make_manager( + [ + self._make_server("a", available_on_public_internet=True), + self._make_server("b", available_on_public_internet=True), + ] + ) + assert manager.get_public_mcp_servers() == [] + + @patch("litellm.public_mcp_servers", []) + def test_returns_empty_when_whitelist_is_empty(self): + """Explicit empty whitelist → hub returns nothing.""" + manager = self._make_manager( + [self._make_server("a", available_on_public_internet=True)] + ) + assert manager.get_public_mcp_servers() == [] + + @patch("litellm.public_mcp_servers", ["a"]) + def test_returns_only_whitelisted_when_flag_defaults_to_true(self): + """ + Regression: prior to the fix, every server with + available_on_public_internet=True (the default) leaked into the hub + regardless of the whitelist. Whitelist must be authoritative. + """ + manager = self._make_manager( + [ + self._make_server("a", available_on_public_internet=True), + self._make_server("b", available_on_public_internet=True), + ] + ) + result = manager.get_public_mcp_servers() + assert [s.server_id for s in result] == ["a"] + + @patch("litellm.public_mcp_servers", ["a"]) + def test_does_not_leak_servers_via_internal_flag(self): + """ + available_on_public_internet is an IP-gating flag, not a hub flag. + A server with the flag True that is not in the whitelist must not + appear in the hub. + """ + manager = self._make_manager( + [ + self._make_server("a", available_on_public_internet=False), + self._make_server("b", available_on_public_internet=True), + ] + ) + result = manager.get_public_mcp_servers() + assert [s.server_id for s in result] == ["a"] + + @patch("litellm.public_mcp_servers", ["does-not-exist"]) + def test_stale_whitelist_id_returns_empty(self): + """Whitelist references an unknown server_id → no spurious results.""" + manager = self._make_manager( + [self._make_server("a", available_on_public_internet=True)] + ) + assert manager.get_public_mcp_servers() == [] + + +class TestGetPublicMCPServersLegacyMode: + """ + Legacy migration knob: litellm.public_mcp_hub_strict_whitelist=False + preserves the pre-fix OR-with-default semantics for one release so + operators that relied on the old behavior have a window to call + /v1/mcp/make_public before /public/mcp_hub goes empty. + """ + + def _make_server(self, server_id, available_on_public_internet=True): + return MCPServer( + server_id=server_id, + name=server_id, + server_name=server_id, + transport=MCPTransport.http, + available_on_public_internet=available_on_public_internet, + ) + + def _make_manager(self, servers): + manager = MCPServerManager() + for s in servers: + manager.config_mcp_servers[s.server_id] = s + return manager + + @patch("litellm.public_mcp_hub_strict_whitelist", False) + @patch("litellm.public_mcp_servers", None) + def test_legacy_returns_default_flag_servers_when_whitelist_is_none(self): + """Legacy mode + no whitelist → every server with the default + available_on_public_internet=True appears (old behavior).""" + manager = self._make_manager( + [ + self._make_server("a", available_on_public_internet=True), + self._make_server("b", available_on_public_internet=False), + ] + ) + result = manager.get_public_mcp_servers() + assert [s.server_id for s in result] == ["a"] + + @patch("litellm.public_mcp_hub_strict_whitelist", False) + @patch("litellm.public_mcp_servers", ["b"]) + def test_legacy_unions_whitelist_and_default_flag(self): + """Legacy mode unions the whitelist with any + available_on_public_internet=True server.""" + manager = self._make_manager( + [ + self._make_server("a", available_on_public_internet=True), + self._make_server("b", available_on_public_internet=False), + ] + ) + result = manager.get_public_mcp_servers() + assert sorted(s.server_id for s in result) == ["a", "b"] + + if __name__ == "__main__": pytest.main([__file__]) diff --git a/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py b/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py index 6cff91d2c74..ecab59c10a1 100644 --- a/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py +++ b/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py @@ -703,3 +703,67 @@ def test_clean_display_name_strips_suffix(): def test_clean_display_name_passthrough_when_no_suffix(): assert _clean_display_name("OpenAI") == "OpenAI" assert _clean_display_name("") == "" + + +def test_public_mcp_hub_returns_only_whitelisted_servers(): + """Regression: /public/mcp_hub must gate strictly on + litellm.public_mcp_servers, mirroring /public/model_hub and + /public/agent_hub. Servers with available_on_public_internet=True that + are not on the whitelist must not leak.""" + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.proxy._types import MCPTransport + + app = FastAPI() + app.include_router(router) + app.dependency_overrides[user_api_key_auth] = lambda: MagicMock() + client = TestClient(app) + + listed = MCPServer( + server_id="listed", + name="listed", + server_name="listed", + transport=MCPTransport.http, + available_on_public_internet=True, + ) + + mock_manager = MagicMock() + mock_manager.get_public_mcp_servers.return_value = [listed] + + with ( + patch("litellm.public_mcp_servers", ["listed"]), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + mock_manager, + ), + ): + response = client.get("/public/mcp_hub") + + assert response.status_code == 200 + data = response.json() + assert [item["server_id"] for item in data] == ["listed"] + app.dependency_overrides.clear() + + +def test_public_mcp_hub_returns_empty_when_whitelist_unset(): + """When no servers have been published via /v1/mcp/make_public, the + hub returns an empty list (matches /public/agent_hub behavior).""" + app = FastAPI() + app.include_router(router) + app.dependency_overrides[user_api_key_auth] = lambda: MagicMock() + client = TestClient(app) + + mock_manager = MagicMock() + mock_manager.get_public_mcp_servers.return_value = [] + + with ( + patch("litellm.public_mcp_servers", None), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + mock_manager, + ), + ): + response = client.get("/public/mcp_hub") + + assert response.status_code == 200 + assert response.json() == [] + app.dependency_overrides.clear()