diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index fa047cd4477..13381c7a6c9 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -120,6 +120,22 @@ def _log_budget_lookup_failure(entity: str, error: Exception) -> None: ) +def _get_router_zero_cost_cache(llm_router: Router) -> Optional[Dict[str, bool]]: + """ + Return the router's per-instance zero-cost cache, or ``None`` for objects + that don't expose one (e.g. ``MagicMock`` stand-ins in unit tests). + + The cache lives on the ``Router`` instance so it: + * is invalidated by ``Router._invalidate_model_group_info_cache`` on + any model add/remove/upsert (including in-place pricing changes via + ``/model/update``, which go through ``upsert_deployment``); + * dies with the router itself — no risk of CPython reusing the + previous router's ``id()`` and serving its cached entries. + """ + cache = getattr(llm_router, "_zero_cost_cache", None) + return cache if isinstance(cache, dict) else None + + def _is_model_cost_zero( model: Optional[Union[str, List[str]]], llm_router: Optional[Router] ) -> bool: @@ -141,7 +157,15 @@ def _is_model_cost_zero( # Handle list of models model_list = [model] if isinstance(model, str) else model + zero_cost_cache = _get_router_zero_cost_cache(llm_router) + for model_name in model_list: + if zero_cost_cache is not None: + cached = zero_cost_cache.get(model_name) + if cached is not None: + if cached is False: + return False + continue try: # Use router's get_model_group_info method directly for better reliability model_group_info = llm_router.get_model_group_info(model_group=model_name) @@ -152,6 +176,8 @@ def _is_model_cost_zero( verbose_proxy_logger.debug( f"No model group info found for {model_name}, assuming it has cost" ) + if zero_cost_cache is not None: + zero_cost_cache[model_name] = False return False # Check costs for this model @@ -164,6 +190,8 @@ def _is_model_cost_zero( verbose_proxy_logger.debug( f"Model {model_name} has undefined cost (input: {input_cost}, output: {output_cost}), assuming it has cost" ) + if zero_cost_cache is not None: + zero_cost_cache[model_name] = False return False # If either cost is non-zero, return False @@ -171,6 +199,8 @@ def _is_model_cost_zero( verbose_proxy_logger.debug( f"Model {model_name} has non-zero cost (input: {input_cost}, output: {output_cost})" ) + if zero_cost_cache is not None: + zero_cost_cache[model_name] = False return False # Costs are 0 — verify this is from explicit configuration, @@ -184,6 +214,8 @@ def _is_model_cost_zero( "cost (enforce budget)", safe_name, ) + if zero_cost_cache is not None: + zero_cost_cache[model_name] = False return False verbose_proxy_logger.debug( @@ -192,6 +224,8 @@ def _is_model_cost_zero( input_cost, output_cost, ) + if zero_cost_cache is not None: + zero_cost_cache[model_name] = True except Exception as e: # If we can't determine the cost, assume it has cost (conservative approach) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 0114774cc2d..46ebff515b7 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -6691,6 +6691,9 @@ def _restamp_streaming_chunk_model( downstream_model = ( chunk.get("model") if isinstance(chunk, dict) else getattr(chunk, "model", None) ) + if downstream_model == requested_model_from_client: + return chunk, model_mismatch_logged + if not model_mismatch_logged and downstream_model != requested_model_from_client: verbose_proxy_logger.debug( "litellm_call_id=%s: streaming chunk model mismatch - requested=%r downstream=%r. Overriding model to requested.", @@ -6719,7 +6722,125 @@ def _restamp_streaming_chunk_model( return chunk, model_mismatch_logged -async def async_data_generator( +def _fast_serialize_simple_model_response_stream( + chunk: ModelResponseStream, +) -> Optional[bytes]: + """ + Serialize the common OpenAI text streaming chunk without the full Pydantic + serializer. Fall back for richer chunks so tool calls, logprobs, usage, and + provider-specific fields keep the canonical model_dump_json behavior. + """ + if ( + getattr(chunk, "provider_specific_fields", None) is not None + or getattr(chunk, "system_fingerprint", None) is not None + or getattr(chunk, "usage", None) is not None + ): + return None + + choices = getattr(chunk, "choices", None) + if not isinstance(choices, list) or len(choices) != 1: + return None + + choice = choices[0] + if ( + getattr(choice, "logprobs", None) is not None + or getattr(choice, "enhancements", None) is not None + ): + return None + + delta = getattr(choice, "delta", None) + if delta is None: + return None + + unsupported_delta_fields = ( + "function_call", + "tool_calls", + "audio", + "images", + "annotations", + "reasoning_content", + "thinking_blocks", + "provider_specific_fields", + "refusal", + ) + if any( + getattr(delta, field, None) is not None for field in unsupported_delta_fields + ): + return None + + delta_dict: dict = {} + role = getattr(delta, "role", None) + content = getattr(delta, "content", None) + if role is not None: + delta_dict["role"] = role + if content is not None: + delta_dict["content"] = content + + choice_dict = {"index": getattr(choice, "index", 0), "delta": delta_dict} + finish_reason = getattr(choice, "finish_reason", None) + if finish_reason is not None: + choice_dict["finish_reason"] = finish_reason + + # Match the canonical ``model_dump_json(exclude_none=True)`` shape — if a + # field is None, omit it entirely rather than emitting ``"key": null``. + # Strict OpenAI-compatible clients reject ``null`` for optional fields like + # ``model``, so diverging here would surface as a client-side regression + # only on the fast path. Fall back to the slow path if a required-looking + # top-level identifier is missing. + model = getattr(chunk, "model", None) + if model is None: + return None + + payload: dict = { + "id": getattr(chunk, "id", None), + "object": getattr(chunk, "object", None), + "created": getattr(chunk, "created", None), + "model": model, + "choices": [choice_dict], + } + for top_level_key in ("id", "object", "created"): + if payload[top_level_key] is None: + payload.pop(top_level_key) + return orjson.dumps(payload) + + +def _serialize_streaming_chunk(chunk: BaseModel) -> Union[str, bytes]: + if isinstance(chunk, ModelResponseStream): + serialized_chunk = _fast_serialize_simple_model_response_stream(chunk) + if serialized_chunk is not None: + return serialized_chunk + + return chunk.model_dump_json(exclude_none=True, exclude_unset=True) + + +async def _apply_streaming_chunk_hooks( + *, + chunk: Any, + user_api_key_dict: UserAPIKeyAuth, + request_data: dict, + str_so_far: str, +) -> Tuple[Any, str]: + chunk = await proxy_logging_obj.async_post_call_streaming_hook( + user_api_key_dict=user_api_key_dict, + response=chunk, + data=request_data, + str_so_far=str_so_far if str_so_far else None, + ) + + if isinstance(chunk, (ModelResponse, ModelResponseStream)): + response_str = litellm.get_response_string(response_obj=chunk) + str_so_far += response_str + + return chunk, str_so_far + + +def _format_streaming_sse_chunk(chunk: Union[str, bytes]) -> Union[str, bytes]: + if isinstance(chunk, bytes): + return b"data: " + chunk + b"\n\n" + return f"data: {chunk}\n\n" + + +async def async_data_generator( # noqa: PLR0915 response, user_api_key_dict: UserAPIKeyAuth, request_data: dict ): verbose_proxy_logger.debug("inside generator") @@ -6733,22 +6854,36 @@ async def async_data_generator( # Previously "".join(str_so_far_parts) was called every chunk, re-joining # the entire accumulated response. String += is O(n) amortized total. _str_so_far: str = "" - async for chunk in proxy_logging_obj.async_post_call_streaming_iterator_hook( - user_api_key_dict=user_api_key_dict, - response=response, - request_data=request_data, - ): - ### CALL HOOKS ### - modify outgoing data - chunk = await proxy_logging_obj.async_post_call_streaming_hook( - user_api_key_dict=user_api_key_dict, - response=chunk, - data=request_data, - str_so_far=_str_so_far if _str_so_far else None, - ) + # Separate iterator-level vs per-chunk hook decisions. The iterator + # wrap is needed when any callback overrides + # ``async_post_call_streaming_iterator_hook`` or has + # ``apply_guardrail``; the per-chunk hook (which builds ``str_so_far`` + # and calls ``async_post_call_streaming_hook``) is only needed when + # there is an active CustomGuardrail or a class that overrides the + # per-chunk hook. Coalescing them into a single flag forced wasted + # ``get_response_string`` work per chunk on every deployment that + # happened to ship a streaming-iterator override (the default). + needs_iterator_wrap = proxy_logging_obj.needs_iterator_wrap() + needs_per_chunk_hook = proxy_logging_obj.needs_per_chunk_streaming_hook() - if isinstance(chunk, (ModelResponse, ModelResponseStream)): - response_str = litellm.get_response_string(response_obj=chunk) - _str_so_far += response_str + if needs_iterator_wrap: + stream_iterator = proxy_logging_obj.async_post_call_streaming_iterator_hook( + user_api_key_dict=user_api_key_dict, + response=response, + request_data=request_data, + ) + else: + stream_iterator = response + + async for chunk in stream_iterator: + if needs_per_chunk_hook: + ### CALL HOOKS ### - modify outgoing data + chunk, _str_so_far = await _apply_streaming_chunk_hooks( + chunk=chunk, + user_api_key_dict=user_api_key_dict, + request_data=request_data, + str_so_far=_str_so_far, + ) chunk, model_mismatch_logged = _restamp_streaming_chunk_model( chunk=chunk, @@ -6758,16 +6893,22 @@ async def async_data_generator( ) if isinstance(chunk, BaseModel): - chunk = chunk.model_dump_json(exclude_none=True, exclude_unset=True) + chunk = _serialize_streaming_chunk(chunk) elif isinstance(chunk, str) and chunk.startswith("data: "): error_message = chunk break try: - yield f"data: {chunk}\n\n" + yield _format_streaming_sse_chunk(chunk=chunk) except Exception as e: yield f"data: {str(e)}\n\n" + if not needs_iterator_wrap: + # The iterator-wrap path fires deferred logging itself; fire it + # here for the no-wrap fast path so non-callback deployments + # still flush their post-stream logging. + ProxyLogging._fire_deferred_stream_logging(request_data) + # Streaming is done, yield the [DONE] chunk if error_message is not None: yield error_message diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 559d5c99b9d..871f084c58d 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -12,15 +12,18 @@ import traceback from datetime import date, datetime, timedelta, timezone from email.mime.multipart import MIMEMultipart from email.mime.text import MIMEText +from dataclasses import dataclass, field from typing import ( TYPE_CHECKING, Any, AsyncGenerator, Awaitable, + ClassVar, Dict, List, Literal, Optional, + Tuple, Union, cast, overload, @@ -335,6 +338,30 @@ def _enrich_http_exception_with_guardrail_context( detail.setdefault("guardrail_mode", event_hook) +@dataclass(frozen=True) +class _CallbackCapabilities: + """Cached per-hook capability flags derived from ``litellm.callbacks``. + + Recomputing this per request walked the callback list and resolved every + string entry via ``get_custom_logger_compatible_class`` — a measurable + chunk of overhead on streaming and non-streaming chat completions. + """ + + has_post_call_response_headers: bool = False + has_iterator_override: bool = False + has_streaming_chunk_override: bool = False + has_guardrail: bool = False + has_pre_call_override: bool = False + # Tuple[(resolved_callback, "override" | "apply_guardrail"), ...] + # Ordered the same as ``litellm.callbacks``; used to build the streaming + # iterator chain without re-scanning per request. + iterator_overrides: Tuple[Tuple[Any, str], ...] = field(default_factory=tuple) + # Resolved CustomLogger callbacks in original order. Pre-resolving once + # avoids the per-request ``get_custom_logger_compatible_class`` walk for + # every string entry in ``litellm.callbacks``. + resolved_callbacks: Tuple[Any, ...] = field(default_factory=tuple) + + class ProxyLogging: """ Logging/Custom Handlers for proxy. @@ -1397,20 +1424,20 @@ class ProxyLogging: metadata = data.get("metadata", data.get("litellm_metadata", {})) or {} pipeline_managed: set = metadata.get("_pipeline_managed_guardrails", set()) - for callback in litellm.callbacks: + caps = ProxyLogging._callback_capabilities() + # Skip the per-request callback walk entirely when nothing in + # ``litellm.callbacks`` overrides ``async_pre_call_hook`` and no + # CustomGuardrail is configured. Saves the loop overhead + + # ``time.time()`` x2 per registered callback for the common + # "callbacks=[]" case on small / dev deployments. + if not caps.has_guardrail and not caps.has_pre_call_override: + if data is not None: + self._process_guardrail_metadata(data) + return data + + for _callback in caps.resolved_callbacks: start_time = time.time() - _callback = None - if isinstance(callback, str): - _callback = litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class( - cast(_custom_logger_compatible_callbacks_literal, callback) - ) - else: - _callback = callback # type: ignore - if ( - _callback is not None - and isinstance(_callback, CustomGuardrail) - and data is not None - ): + if isinstance(_callback, CustomGuardrail) and data is not None: # Skip guardrails managed by a pipeline if ( _callback.guardrail_name @@ -1505,6 +1532,131 @@ class ProxyLogging: _enrich_http_exception_with_guardrail_context(e, callback) raise + # Cache for callback-capability detection. Keyed on a signature of + # litellm.callbacks (length + each item's id) so we recompute when the + # callback list mutates (add/remove) without iterating every request. + _callback_capabilities_cache: ClassVar[ + Dict[Tuple[int, Tuple[int, ...]], "_CallbackCapabilities"] + ] = {} + + @staticmethod + def _callback_capabilities() -> "_CallbackCapabilities": + """ + Inspect ``litellm.callbacks`` once and answer the per-hook capability + questions used to short-circuit no-op work on the chat-completions hot + path. Per-request callers iterated ``litellm.callbacks`` and called + ``get_custom_logger_compatible_class`` for every string entry — that + scanning cost dominated the proxy overhead on low-config deployments. + + Cache invalidates whenever the list length or member identities change. + """ + callbacks = litellm.callbacks + sig = (len(callbacks), tuple(id(c) for c in callbacks)) + cache = ProxyLogging._callback_capabilities_cache + cached = cache.get(sig) + if cached is not None: + return cached + + has_post_call_response_headers = False + has_iterator_override = False + has_streaming_chunk_override = False + has_guardrail = False + has_pre_call_override = False + iterator_overrides: List[Tuple[Any, str]] = [] # (callback, kind) + resolved_callbacks: List[Any] = [] + + for callback in callbacks: + if isinstance(callback, str): + resolved: Any = ( + litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class( + cast(_custom_logger_compatible_callbacks_literal, callback) + ) + ) + else: + resolved = callback + if resolved is None or not isinstance(resolved, CustomLogger): + continue + resolved_callbacks.append(resolved) + cls = type(resolved) + if cls is CustomLogger: + continue + if isinstance(resolved, CustomGuardrail): + has_guardrail = True + # Use the same leaf-class ``__dict__`` check as the other hook + # capabilities: only callbacks that actually override the hook + # contribute to the flag. Setting this for every ``CustomLogger`` + # instance (the prior behaviour) forced the full + # ``post_call_response_headers_hook`` body to run on every request + # even when no registered callback customized response headers. + cls_attrs = cls.__dict__ + if "async_post_call_response_headers_hook" in cls_attrs: + has_post_call_response_headers = True + if "async_post_call_streaming_iterator_hook" in cls_attrs: + has_iterator_override = True + iterator_overrides.append((resolved, "override")) + elif "apply_guardrail" in cls_attrs: + iterator_overrides.append((resolved, "apply_guardrail")) + if "async_post_call_streaming_hook" in cls_attrs: + has_streaming_chunk_override = True + if "async_pre_call_hook" in cls_attrs: + has_pre_call_override = True + + caps = _CallbackCapabilities( + has_post_call_response_headers=has_post_call_response_headers, + has_iterator_override=has_iterator_override + or any(kind == "apply_guardrail" for _, kind in iterator_overrides), + has_streaming_chunk_override=has_streaming_chunk_override, + has_guardrail=has_guardrail, + has_pre_call_override=has_pre_call_override, + iterator_overrides=tuple(iterator_overrides), + resolved_callbacks=tuple(resolved_callbacks), + ) + # Limit cache to handle test churn without leaking; production + # callback lists are stable so this rarely grows past 1 entry. + if len(cache) >= 32: + cache.clear() + cache[sig] = caps + return caps + + @staticmethod + def has_post_call_response_headers_callbacks() -> bool: + return ProxyLogging._callback_capabilities().has_post_call_response_headers + + @staticmethod + def has_streaming_callbacks() -> bool: + caps = ProxyLogging._callback_capabilities() + return ( + caps.has_iterator_override + or caps.has_streaming_chunk_override + or caps.has_guardrail + ) + + @staticmethod + def has_streaming_chunk_hook_overrides() -> bool: + """True iff any callback overrides ``async_post_call_streaming_hook`` + (the per-chunk hook, distinct from the iterator wrapper).""" + caps = ProxyLogging._callback_capabilities() + return caps.has_streaming_chunk_override or caps.has_guardrail + + def needs_iterator_wrap(self) -> bool: + """Whether ``async_data_generator`` needs to wrap the upstream stream + through ``async_post_call_streaming_iterator_hook``. Instance method + so tests can override the gate via ``MagicMock(spec=ProxyLogging)``. + """ + return ProxyLogging._callback_capabilities().has_iterator_override + + def needs_per_chunk_streaming_hook(self) -> bool: + """Whether ``async_data_generator`` needs to call the per-chunk + ``_apply_streaming_chunk_hooks`` for every emitted chunk. Instance + method for the same reason as :py:meth:`needs_iterator_wrap`. + """ + caps = ProxyLogging._callback_capabilities() + return caps.has_streaming_chunk_override or caps.has_guardrail + + @staticmethod + def has_during_call_guardrails() -> bool: + return ProxyLogging._callback_capabilities().has_guardrail + async def during_call_hook( self, data: dict, @@ -1514,6 +1666,12 @@ class ProxyLogging: """ Runs the CustomGuardrail's async_moderation_hook() in parallel """ + # Fast path: skip the entire guardrail scan when no CustomGuardrail + # callbacks are registered. Saves per-request iteration over + # ``litellm.callbacks`` plus an ``asyncio.gather([])`` round trip on + # deployments with no guardrails configured. + if not ProxyLogging._callback_capabilities().has_guardrail: + return data # Step 1: Collect all guardrail tasks to run in parallel guardrail_tasks = [] @@ -2122,6 +2280,14 @@ class ProxyLogging: Dict[str, str]: Merged headers from all callbacks. """ merged_headers: Dict[str, str] = {} + # Outer call sites in common_request_processing.py already gate this + # call with ``has_post_call_response_headers_callbacks()``. The + # cached detection makes the redundant interior guard cheap, but the + # guard would still iterate every code path through this function so + # keep it cheap and rely on the cached capability lookup. + if not ProxyLogging._callback_capabilities().has_post_call_response_headers: + return merged_headers + try: # Build litellm_call_info — normalized routing metadata for callbacks litellm_call_info = self._build_litellm_call_info( @@ -2203,6 +2369,16 @@ class ProxyLogging: Covers: 1. /chat/completions """ + # Per-chunk fast path: skip the response-string materialization and + # callback scan when no configured callback overrides + # ``async_post_call_streaming_hook`` AND no CustomGuardrail is + # active. ``get_response_string`` walks every choice/delta on the + # chunk so paying it per chunk for no-op callbacks dominated stream + # CPU time even after the iterator-chain fix. + caps = ProxyLogging._callback_capabilities() + if not caps.has_streaming_chunk_override and not caps.has_guardrail: + return response + from litellm.proxy.proxy_server import llm_router response_str: Optional[str] = None @@ -2278,6 +2454,18 @@ class ProxyLogging: Covers: 1. /chat/completions """ + caps = ProxyLogging._callback_capabilities() + # Fast path: no real overrides. Internal proxy CustomLogger callbacks + # (e.g. _PROXY_MaxBudgetLimiter, ManagedFiles) inherit the default + # ``async for chunk: yield chunk`` body, so wrapping the iterator + # through each of them adds N pass-through trampolines per chunk for + # zero behavior change. Skip the chain entirely and stream through. + if not caps.iterator_overrides: + async for chunk in response: + yield chunk + ProxyLogging._fire_deferred_stream_logging(request_data) + return + from litellm.proxy.proxy_server import llm_router # Merge model-level guardrails before checking which guardrails to run @@ -2287,55 +2475,35 @@ class ProxyLogging: current_response = response - for callback in litellm.callbacks: - _callback: Optional[CustomLogger] = None - if isinstance(callback, str): - _callback = litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class( - cast(_custom_logger_compatible_callbacks_literal, callback) + for resolved_callback, kind in caps.iterator_overrides: + if isinstance(resolved_callback, CustomGuardrail): + if ( + resolved_callback.should_run_guardrail( + data=request_data, event_type=GuardrailEventHooks.post_call + ) + is not True + ): + continue + if kind == "override": + current_response = self._wrap_streaming_iterator_with_enrichment( + resolved_callback, + resolved_callback.async_post_call_streaming_iterator_hook( + user_api_key_dict=user_api_key_dict, + response=current_response, + request_data=request_data, + ), ) else: - _callback = callback # type: ignore - if _callback is not None and isinstance(_callback, CustomLogger): - if not isinstance( - _callback, CustomGuardrail - ) or _callback.should_run_guardrail( - data=request_data, event_type=GuardrailEventHooks.post_call - ): - if ( - "async_post_call_streaming_iterator_hook" - in type(callback).__dict__ - ): - current_response = ( - self._wrap_streaming_iterator_with_enrichment( - _callback, - _callback.async_post_call_streaming_iterator_hook( - user_api_key_dict=user_api_key_dict, - response=current_response, - request_data=request_data, - ), - ) - ) - elif "apply_guardrail" in type(callback).__dict__: - request_data["guardrail_to_apply"] = callback - current_response = self._wrap_streaming_iterator_with_enrichment( - _callback, - unified_guardrail.async_post_call_streaming_iterator_hook( - user_api_key_dict=user_api_key_dict, - request_data=request_data, - response=current_response, - ), - ) - else: - current_response = ( - self._wrap_streaming_iterator_with_enrichment( - _callback, - _callback.async_post_call_streaming_iterator_hook( - user_api_key_dict=user_api_key_dict, - response=current_response, - request_data=request_data, - ), - ) - ) + # kind == "apply_guardrail": route through unified_guardrail + request_data["guardrail_to_apply"] = resolved_callback + current_response = self._wrap_streaming_iterator_with_enrichment( + resolved_callback, + unified_guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=user_api_key_dict, + request_data=request_data, + response=current_response, + ), + ) # Actually iterate through the chained async generator and yield chunks async for chunk in current_response: diff --git a/litellm/router.py b/litellm/router.py index 1d070b3af8f..3340e06627e 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -491,6 +491,17 @@ class Router: # Maps (team_id, team_public_model_name) -> list of indices in model_list self.team_model_to_deployment_indices: Dict[Tuple[str, str], List[int]] = {} + # Initialize cache attributes that ``_invalidate_model_group_info_cache`` + # touches *before* the first ``set_model_list`` below (which calls + # that invalidation as part of building the model index). + self._access_groups_cache: Optional[Dict[str, List[str]]] = None + # Per-router cache for the proxy auth-layer "is this model explicitly + # zero-cost?" check. Lives on the router so it is invalidated alongside + # ``_cached_get_model_group_info`` and dies with the router (no + # ``id()``-reuse risk after GC). See + # ``litellm.proxy.auth.auth_checks._is_model_cost_zero``. + self._zero_cost_cache: Dict[str, bool] = {} + if model_list is not None: # set_model_list will build indices automatically self.set_model_list(model_list) @@ -503,8 +514,6 @@ class Router: [] ) # initialize an empty list - to allow _add_deployment and delete_deployment to work - self._access_groups_cache: Optional[Dict[str, List[str]]] = None - if allowed_fails is not None: self.allowed_fails = allowed_fails else: @@ -9228,8 +9237,13 @@ class Router: """Invalidate the cached model group info. Call this whenever self.model_list is modified to ensure the cache is rebuilt. + Also clears the auth-layer zero-cost cache, which depends on the same + ``ModelGroupInfo`` data — without this, an in-place pricing update on + an existing deployment (same model count) would keep a stale ``True`` + result and bypass budget enforcement. """ self._cached_get_model_group_info.cache_clear() + self._zero_cost_cache.clear() def _invalidate_access_groups_cache(self) -> None: """Invalidate the cached access groups. diff --git a/scripts/benchmark_chat_completions_perf.py b/scripts/benchmark_chat_completions_perf.py new file mode 100644 index 00000000000..2c211f674fe --- /dev/null +++ b/scripts/benchmark_chat_completions_perf.py @@ -0,0 +1,842 @@ +#!/usr/bin/env python3 +"""Benchmark LiteLLM proxy /v1/chat/completions overhead and streaming TTFT. + +The script can run a local OpenAI-compatible mock provider plus a LiteLLM proxy +from any checkout. That makes it useful for comparing tags/commits without +depending on real provider latency. + +Example: + uv run python scripts/benchmark_chat_completions_perf.py \ + --label current --requests 500 --concurrency 100 + +Compare another checkout: + uv run python scripts/benchmark_chat_completions_perf.py \ + --label v1.83.14-stable --litellm-dir /tmp/litellm-v1.83.14-stable +""" + +from __future__ import annotations + +import argparse +import asyncio +import json +import os +import shlex +import signal +import statistics +import subprocess +import sys +import tempfile +import time +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Optional + +import aiohttp +from aiohttp import web + + +DEFAULT_MODEL = "perf-test-model" +DEFAULT_API_KEY = "sk-1234" + + +@dataclass +class RequestSample: + success: bool + latency_ms: float + status_code: int + overhead_header_ms: Optional[float] = None + error: str = "" + + +@dataclass +class SummaryStats: + requests: int + failures: int + rps: float + mean_ms: float + p50_ms: float + p95_ms: float + p99_ms: float + overhead_header_mean_ms: Optional[float] = None + overhead_header_p50_ms: Optional[float] = None + overhead_header_p95_ms: Optional[float] = None + + +class MockOpenAIProvider: + def __init__( + self, + host: str, + port: int, + first_token_delay_ms: float, + stream_content_chunks: int, + ) -> None: + self.host = host + self.port = port + self.first_token_delay_ms = first_token_delay_ms + self.stream_content_chunks = stream_content_chunks + self.runner: Optional[web.AppRunner] = None + + @property + def base_url(self) -> str: + return f"http://{self.host}:{self.port}" + + async def start(self) -> None: + app = web.Application() + app.router.add_post("/v1/chat/completions", self.handle_chat_completions) + self.runner = web.AppRunner(app, access_log=None) + await self.runner.setup() + site = web.TCPSite(self.runner, self.host, self.port) + await site.start() + + async def stop(self) -> None: + if self.runner is not None: + await self.runner.cleanup() + + async def handle_chat_completions(self, request: web.Request) -> web.StreamResponse: + body = await request.json() + if body.get("stream"): + return await self._streaming_response(request=request, body=body) + return self._json_response(body) + + def _json_response(self, body: dict[str, Any]) -> web.Response: + now = int(time.time()) + payload = { + "id": "chatcmpl-perf", + "object": "chat.completion", + "created": now, + "model": body.get("model", DEFAULT_MODEL), + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "hello"}, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 1, + "completion_tokens": 1, + "total_tokens": 2, + }, + } + return web.json_response(payload) + + async def _streaming_response( + self, request: web.Request, body: dict[str, Any] + ) -> web.StreamResponse: + response = web.StreamResponse( + status=200, + headers={ + "Content-Type": "text/event-stream", + "Cache-Control": "no-cache", + }, + ) + await response.prepare(request) + if self.first_token_delay_ms > 0: + await asyncio.sleep(self.first_token_delay_ms / 1000) + + created = int(time.time()) + chunks = [{"role": "assistant"}] + chunks.extend({"content": "hello"} for _ in range(self.stream_content_chunks)) + for delta in chunks: + event = { + "id": "chatcmpl-perf", + "object": "chat.completion.chunk", + "created": created, + "model": body.get("model", DEFAULT_MODEL), + "choices": [{"index": 0, "delta": delta, "finish_reason": None}], + } + await response.write(f"data: {json.dumps(event)}\n\n".encode()) + + done_event = { + "id": "chatcmpl-perf", + "object": "chat.completion.chunk", + "created": created, + "model": body.get("model", DEFAULT_MODEL), + "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}], + } + await response.write(f"data: {json.dumps(done_event)}\n\n".encode()) + await response.write(b"data: [DONE]\n\n") + await response.write_eof() + return response + + +def percentile(values: list[float], pct: float) -> float: + if not values: + return 0.0 + sorted_values = sorted(values) + index = min(int(len(sorted_values) * pct / 100), len(sorted_values) - 1) + return sorted_values[index] + + +def summarize(samples: list[RequestSample], wall_time_s: float) -> SummaryStats: + latencies = [sample.latency_ms for sample in samples if sample.success] + overhead_headers = [ + sample.overhead_header_ms + for sample in samples + if sample.success and sample.overhead_header_ms is not None + ] + failures = len(samples) - len(latencies) + return SummaryStats( + requests=len(samples), + failures=failures, + rps=(len(latencies) / wall_time_s) if wall_time_s > 0 else 0.0, + mean_ms=statistics.mean(latencies) if latencies else 0.0, + p50_ms=percentile(latencies, 50), + p95_ms=percentile(latencies, 95), + p99_ms=percentile(latencies, 99), + overhead_header_mean_ms=( + statistics.mean(overhead_headers) if overhead_headers else None + ), + overhead_header_p50_ms=( + percentile(overhead_headers, 50) if overhead_headers else None + ), + overhead_header_p95_ms=( + percentile(overhead_headers, 95) if overhead_headers else None + ), + ) + + +def format_optional_ms(value: Optional[float]) -> str: + return "n/a" if value is None else f"{value:.2f}" + + +def get_git_revision(litellm_dir: Path) -> str: + try: + result = subprocess.run( + ["git", "rev-parse", "--short", "HEAD"], + cwd=litellm_dir, + check=True, + capture_output=True, + text=True, + ) + return result.stdout.strip() + except Exception: + return "unknown" + + +def write_proxy_config(config_path: Path, provider_base_url: str, api_key: str) -> None: + config_path.write_text( + f"""model_list: + - model_name: {DEFAULT_MODEL} + litellm_params: + model: openai/{DEFAULT_MODEL} + api_key: fake-provider-key + api_base: {provider_base_url}/v1 + +general_settings: + master_key: {api_key} + +litellm_settings: + drop_params: true + telemetry: false +""", + encoding="utf-8", + ) + + +async def wait_for_proxy(base_url: str, timeout_s: float) -> None: + deadline = time.perf_counter() + timeout_s + last_error = "" + async with aiohttp.ClientSession() as session: + while time.perf_counter() < deadline: + try: + async with session.get(f"{base_url}/health") as response: + if response.status < 500: + return + last_error = f"HTTP {response.status}: {await response.text()}" + except Exception as exc: + last_error = str(exc) + await asyncio.sleep(0.5) + raise TimeoutError(f"Timed out waiting for proxy at {base_url}: {last_error}") + + +def start_proxy_process( + litellm_dir: Path, + proxy_command: str, + config_path: Path, + port: int, + log_path: Path, +) -> subprocess.Popen: + command = shlex.split(proxy_command) + [ + "--config", + str(config_path), + "--port", + str(port), + ] + env = { + **os.environ, + "LITELLM_TELEMETRY": "False", + "PYTHONUNBUFFERED": "1", + } + log_file = log_path.open("w", encoding="utf-8") + return subprocess.Popen( + command, + cwd=litellm_dir, + env=env, + stdout=log_file, + stderr=subprocess.STDOUT, + start_new_session=True, + ) + + +def stop_proxy_process(process: subprocess.Popen) -> None: + if process.poll() is not None: + return + try: + os.killpg(process.pid, signal.SIGTERM) + process.wait(timeout=10) + except Exception: + try: + os.killpg(process.pid, signal.SIGKILL) + except Exception: + pass + + +def extract_overhead_header(headers: aiohttp.typedefs.LooseHeaders) -> Optional[float]: + raw_value = headers.get("x-litellm-overhead-duration-ms") # type: ignore[union-attr] + if raw_value is None: + return None + try: + return float(raw_value) + except ValueError: + return None + + +async def post_non_streaming( + session: aiohttp.ClientSession, + url: str, + headers: dict[str, str], + payload: dict[str, Any], + semaphore: asyncio.Semaphore, +) -> RequestSample: + async with semaphore: + start = time.perf_counter() + try: + async with session.post(url, headers=headers, json=payload) as response: + body = await response.read() + latency_ms = (time.perf_counter() - start) * 1000 + if response.status != 200: + return RequestSample( + success=False, + latency_ms=latency_ms, + status_code=response.status, + error=body.decode("utf-8", errors="ignore")[:200], + ) + return RequestSample( + success=True, + latency_ms=latency_ms, + status_code=response.status, + overhead_header_ms=extract_overhead_header(response.headers), + ) + except Exception as exc: + return RequestSample( + success=False, + latency_ms=(time.perf_counter() - start) * 1000, + status_code=0, + error=str(exc)[:200], + ) + + +async def run_non_streaming_benchmark( + url: str, + headers: dict[str, str], + payload: dict[str, Any], + requests: int, + concurrency: int, + warmup: int, + timeout_s: float, +) -> SummaryStats: + timeout = aiohttp.ClientTimeout(total=timeout_s) + connector = aiohttp.TCPConnector( + limit=max(concurrency * 2, 10), + limit_per_host=max(concurrency, 10), + force_close=False, + ) + semaphore = asyncio.Semaphore(concurrency) + async with aiohttp.ClientSession(connector=connector, timeout=timeout) as session: + if warmup > 0: + await asyncio.gather( + *[ + post_non_streaming(session, url, headers, payload, semaphore) + for _ in range(warmup) + ] + ) + wall_start = time.perf_counter() + samples = await asyncio.gather( + *[ + post_non_streaming(session, url, headers, payload, semaphore) + for _ in range(requests) + ] + ) + wall_time_s = time.perf_counter() - wall_start + return summarize(samples, wall_time_s) + + +async def measure_stream_ttft( + session: aiohttp.ClientSession, + url: str, + headers: dict[str, str], + payload: dict[str, Any], + semaphore: asyncio.Semaphore, +) -> RequestSample: + async with semaphore: + start = time.perf_counter() + try: + async with session.post(url, headers=headers, json=payload) as response: + if response.status != 200: + body = await response.read() + return RequestSample( + success=False, + latency_ms=(time.perf_counter() - start) * 1000, + status_code=response.status, + error=body.decode("utf-8", errors="ignore")[:200], + ) + + while raw_line := await response.content.readline(): + line = raw_line.strip() + if not line or not line.startswith(b"data:"): + continue + event_payload = line[5:].strip() + if event_payload == b"[DONE]": + break + event = json.loads(event_payload) + choice = (event.get("choices") or [{}])[0] + delta = choice.get("delta") or {} + content = delta.get("content") or choice.get("text") + if content: + return RequestSample( + success=True, + latency_ms=(time.perf_counter() - start) * 1000, + status_code=response.status, + overhead_header_ms=extract_overhead_header( + response.headers + ), + ) + return RequestSample( + success=False, + latency_ms=(time.perf_counter() - start) * 1000, + status_code=response.status, + error="stream ended before a content token", + ) + except Exception as exc: + return RequestSample( + success=False, + latency_ms=(time.perf_counter() - start) * 1000, + status_code=0, + error=str(exc)[:200], + ) + + +async def run_streaming_ttft_benchmark( + url: str, + headers: dict[str, str], + payload: dict[str, Any], + requests: int, + concurrency: int, + warmup: int, + timeout_s: float, +) -> SummaryStats: + timeout = aiohttp.ClientTimeout(total=timeout_s) + connector = aiohttp.TCPConnector( + limit=max(concurrency * 2, 10), + limit_per_host=max(concurrency, 10), + force_close=False, + ) + semaphore = asyncio.Semaphore(concurrency) + async with aiohttp.ClientSession(connector=connector, timeout=timeout) as session: + if warmup > 0: + await asyncio.gather( + *[ + measure_stream_ttft(session, url, headers, payload, semaphore) + for _ in range(warmup) + ] + ) + wall_start = time.perf_counter() + samples = await asyncio.gather( + *[ + measure_stream_ttft(session, url, headers, payload, semaphore) + for _ in range(requests) + ] + ) + wall_time_s = time.perf_counter() - wall_start + return summarize(samples, wall_time_s) + + +async def measure_stream_full_response( + session: aiohttp.ClientSession, + url: str, + headers: dict[str, str], + payload: dict[str, Any], + semaphore: asyncio.Semaphore, +) -> RequestSample: + async with semaphore: + start = time.perf_counter() + try: + async with session.post(url, headers=headers, json=payload) as response: + if response.status != 200: + body = await response.read() + return RequestSample( + success=False, + latency_ms=(time.perf_counter() - start) * 1000, + status_code=response.status, + error=body.decode("utf-8", errors="ignore")[:200], + ) + + saw_content = False + while raw_line := await response.content.readline(): + line = raw_line.strip() + if not line or not line.startswith(b"data:"): + continue + event_payload = line[5:].strip() + if event_payload == b"[DONE]": + return RequestSample( + success=saw_content, + latency_ms=(time.perf_counter() - start) * 1000, + status_code=response.status, + overhead_header_ms=extract_overhead_header( + response.headers + ), + error="" if saw_content else "stream ended without content", + ) + if b'"content"' in event_payload or b'"text"' in event_payload: + saw_content = True + + return RequestSample( + success=False, + latency_ms=(time.perf_counter() - start) * 1000, + status_code=response.status, + error="stream ended before [DONE]", + ) + except Exception as exc: + return RequestSample( + success=False, + latency_ms=(time.perf_counter() - start) * 1000, + status_code=0, + error=str(exc)[:200], + ) + + +async def run_streaming_full_benchmark( + url: str, + headers: dict[str, str], + payload: dict[str, Any], + requests: int, + concurrency: int, + warmup: int, + timeout_s: float, +) -> SummaryStats: + timeout = aiohttp.ClientTimeout(total=timeout_s) + connector = aiohttp.TCPConnector( + limit=max(concurrency * 2, 10), + limit_per_host=max(concurrency, 10), + force_close=False, + ) + semaphore = asyncio.Semaphore(concurrency) + async with aiohttp.ClientSession(connector=connector, timeout=timeout) as session: + if warmup > 0: + await asyncio.gather( + *[ + measure_stream_full_response( + session, url, headers, payload, semaphore + ) + for _ in range(warmup) + ] + ) + wall_start = time.perf_counter() + samples = await asyncio.gather( + *[ + measure_stream_full_response(session, url, headers, payload, semaphore) + for _ in range(requests) + ] + ) + wall_time_s = time.perf_counter() - wall_start + return summarize(samples, wall_time_s) + + +def stats_to_dict(stats: SummaryStats) -> dict[str, Any]: + return { + "requests": stats.requests, + "failures": stats.failures, + "rps": stats.rps, + "mean_ms": stats.mean_ms, + "p50_ms": stats.p50_ms, + "p95_ms": stats.p95_ms, + "p99_ms": stats.p99_ms, + "overhead_header_mean_ms": stats.overhead_header_mean_ms, + "overhead_header_p50_ms": stats.overhead_header_p50_ms, + "overhead_header_p95_ms": stats.overhead_header_p95_ms, + } + + +def _median_run( + runs: list[tuple[SummaryStats, SummaryStats, SummaryStats, Optional[SummaryStats]]], +) -> tuple[SummaryStats, SummaryStats, SummaryStats, Optional[SummaryStats]]: + # Pick the run whose proxy non-stream p50 is the median across repeats. + # Choosing a single representative run (rather than aggregating each metric + # separately) keeps related metrics from the same execution context so + # client-overhead deltas stay internally consistent. + sorted_runs = sorted(runs, key=lambda r: r[1].p50_ms) + return sorted_runs[len(sorted_runs) // 2] + + +def print_summary( + label: str, + revision: str, + direct: SummaryStats, + proxy: SummaryStats, + stream: SummaryStats, + stream_full: Optional[SummaryStats], +) -> None: + client_overhead_p50 = proxy.p50_ms - direct.p50_ms + client_overhead_p95 = proxy.p95_ms - direct.p95_ms + print("\n=== Benchmark summary ===") + print(f"Label: {label}") + print(f"Revision: {revision}") + print(f"Direct provider non-stream p50: {direct.p50_ms:.2f} ms") + print(f"Proxy non-stream p50: {proxy.p50_ms:.2f} ms") + print(f"Proxy non-stream p95: {proxy.p95_ms:.2f} ms") + print(f"Proxy non-stream RPS: {proxy.rps:.2f}") + print(f"Client-observed overhead p50: {client_overhead_p50:.2f} ms") + print(f"Client-observed overhead p95: {client_overhead_p95:.2f} ms") + print( + "x-litellm-overhead-duration-ms p50: " + f"{format_optional_ms(proxy.overhead_header_p50_ms)} ms" + ) + print(f"Streaming TTFT p50: {stream.p50_ms:.2f} ms") + print(f"Streaming TTFT p95: {stream.p95_ms:.2f} ms") + print(f"Streaming TTFT RPS: {stream.rps:.2f}") + if stream_full is not None: + print(f"Streaming full response p50: {stream_full.p50_ms:.2f} ms") + print(f"Streaming full response p95: {stream_full.p95_ms:.2f} ms") + print(f"Streaming full response RPS: {stream_full.rps:.2f}") + print("\nMarkdown row:") + print( + "| " + + " | ".join( + [ + label, + revision, + f"{stream.p50_ms:.2f}", + f"{stream.p95_ms:.2f}", + f"{proxy.rps:.2f}", + f"{client_overhead_p50:.2f}", + f"{client_overhead_p95:.2f}", + format_optional_ms(proxy.overhead_header_p50_ms), + f"{stream_full.p50_ms:.2f}" if stream_full is not None else "n/a", + f"{stream_full.rps:.2f}" if stream_full is not None else "n/a", + ] + ) + + " |" + ) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--label", default="current", help="Label for this run") + parser.add_argument( + "--litellm-dir", + default=str(Path.cwd()), + help="Checkout directory used to start the LiteLLM proxy", + ) + parser.add_argument( + "--proxy-command", + default="uv run litellm", + help="Command used to start the proxy inside --litellm-dir", + ) + parser.add_argument("--proxy-host", default="127.0.0.1") + parser.add_argument("--proxy-port", type=int, default=4000) + parser.add_argument("--provider-host", default="127.0.0.1") + parser.add_argument("--provider-port", type=int, default=8099) + parser.add_argument("--api-key", default=DEFAULT_API_KEY) + parser.add_argument("--requests", type=int, default=500) + parser.add_argument("--concurrency", type=int, default=100) + parser.add_argument("--stream-requests", type=int, default=200) + parser.add_argument("--stream-concurrency", type=int, default=20) + parser.add_argument("--warmup", type=int, default=100) + parser.add_argument("--stream-warmup", type=int, default=20) + parser.add_argument("--timeout", type=float, default=30) + parser.add_argument("--proxy-start-timeout", type=float, default=90) + parser.add_argument("--provider-first-token-delay-ms", type=float, default=0) + parser.add_argument( + "--provider-stream-content-chunks", + type=int, + default=20, + help="Streaming chunks the mock provider emits. Default 20 (realistic).", + ) + parser.add_argument( + "--measure-full-stream", + action="store_true", + default=True, + help="Measure time to consume the complete streaming response (on by default).", + ) + parser.add_argument( + "--no-measure-full-stream", + dest="measure_full_stream", + action="store_false", + help="Skip the full-stream RPS measurement.", + ) + parser.add_argument( + "--repeats", + type=int, + default=1, + help="Run the entire suite N times against the same proxy and report the median run.", + ) + parser.add_argument( + "--no-start-proxy", + action="store_true", + help="Benchmark an already-running proxy at --proxy-host/--proxy-port", + ) + parser.add_argument( + "--provider-url", + help="Use an already-running provider instead of starting the mock provider", + ) + parser.add_argument("--output-json", help="Write machine-readable results") + return parser.parse_args() + + +async def async_main() -> None: + args = parse_args() + litellm_dir = Path(args.litellm_dir).resolve() + revision = get_git_revision(litellm_dir) + proxy_base_url = f"http://{args.proxy_host}:{args.proxy_port}" + proxy_url = f"{proxy_base_url}/v1/chat/completions" + headers = { + "Authorization": f"Bearer {args.api_key}", + "Content-Type": "application/json", + } + provider_headers = { + "Authorization": "Bearer fake-provider-key", + "Content-Type": "application/json", + } + non_stream_payload = { + "model": DEFAULT_MODEL, + "messages": [{"role": "user", "content": "hi"}], + "max_tokens": 1, + } + stream_payload = {**non_stream_payload, "stream": True} + + provider: Optional[MockOpenAIProvider] = None + proxy_process: Optional[subprocess.Popen] = None + with tempfile.TemporaryDirectory(prefix="litellm-perf-") as tmp_dir_name: + tmp_dir = Path(tmp_dir_name) + proxy_log_path = tmp_dir / "proxy.log" + if args.provider_url: + provider_base_url = args.provider_url.rstrip("/") + else: + provider = MockOpenAIProvider( + host=args.provider_host, + port=args.provider_port, + first_token_delay_ms=args.provider_first_token_delay_ms, + stream_content_chunks=args.provider_stream_content_chunks, + ) + await provider.start() + provider_base_url = provider.base_url + + config_path = tmp_dir / "config.yaml" + write_proxy_config(config_path, provider_base_url, args.api_key) + + try: + if not args.no_start_proxy: + proxy_process = start_proxy_process( + litellm_dir=litellm_dir, + proxy_command=args.proxy_command, + config_path=config_path, + port=args.proxy_port, + log_path=proxy_log_path, + ) + await wait_for_proxy(proxy_base_url, args.proxy_start_timeout) + + runs: list[ + tuple[ + SummaryStats, + SummaryStats, + SummaryStats, + Optional[SummaryStats], + ] + ] = [] + for run_idx in range(max(1, args.repeats)): + if args.repeats > 1: + print(f"\n--- Run {run_idx + 1}/{args.repeats} ---") + _direct = await run_non_streaming_benchmark( + url=f"{provider_base_url}/v1/chat/completions", + headers=provider_headers, + payload=non_stream_payload, + requests=args.requests, + concurrency=args.concurrency, + warmup=args.warmup, + timeout_s=args.timeout, + ) + _proxy = await run_non_streaming_benchmark( + url=proxy_url, + headers=headers, + payload=non_stream_payload, + requests=args.requests, + concurrency=args.concurrency, + warmup=args.warmup, + timeout_s=args.timeout, + ) + _stream = await run_streaming_ttft_benchmark( + url=proxy_url, + headers=headers, + payload=stream_payload, + requests=args.stream_requests, + concurrency=args.stream_concurrency, + warmup=args.stream_warmup, + timeout_s=args.timeout, + ) + _stream_full = ( + await run_streaming_full_benchmark( + url=proxy_url, + headers=headers, + payload=stream_payload, + requests=args.stream_requests, + concurrency=args.stream_concurrency, + warmup=args.stream_warmup, + timeout_s=args.timeout, + ) + if args.measure_full_stream + else None + ) + runs.append((_direct, _proxy, _stream, _stream_full)) + if args.repeats > 1: + print( + f" run {run_idx + 1}: non-stream p50={_proxy.p50_ms:.2f}ms " + f"rps={_proxy.rps:.2f} | TTFT p50={_stream.p50_ms:.2f}ms " + f"full RPS=" + + (f"{_stream_full.rps:.2f}" if _stream_full else "n/a") + ) + + direct, proxy, stream, stream_full = _median_run(runs) + finally: + if proxy_process is not None: + stop_proxy_process(proxy_process) + if provider is not None: + await provider.stop() + + print_summary(args.label, revision, direct, proxy, stream, stream_full) + + if args.output_json: + output = { + "label": args.label, + "revision": revision, + "direct_non_streaming": stats_to_dict(direct), + "proxy_non_streaming": stats_to_dict(proxy), + "proxy_streaming_ttft": stats_to_dict(stream), + "proxy_streaming_full": ( + stats_to_dict(stream_full) if stream_full is not None else None + ), + "client_observed_overhead_p50_ms": proxy.p50_ms - direct.p50_ms, + "client_observed_overhead_p95_ms": proxy.p95_ms - direct.p95_ms, + "proxy_log_path": str(proxy_log_path), + } + Path(args.output_json).write_text( + json.dumps(output, indent=2, sort_keys=True), encoding="utf-8" + ) + + +def main() -> None: + asyncio.run(async_main()) + + +if __name__ == "__main__": + main() diff --git a/tests/test_litellm/proxy/auth/test_unmapped_model_budget_enforcement.py b/tests/test_litellm/proxy/auth/test_unmapped_model_budget_enforcement.py index be4f534040d..d7e32cf1c16 100644 --- a/tests/test_litellm/proxy/auth/test_unmapped_model_budget_enforcement.py +++ b/tests/test_litellm/proxy/auth/test_unmapped_model_budget_enforcement.py @@ -104,3 +104,82 @@ class TestUnmappedModelBudgetEnforcement: assert ( result is True ), "Model with explicit cost=0 in litellm_params should bypass budget" + + def test_cache_invalidates_on_in_place_pricing_update(self): + """ + Regression test for the stale-cache bug surfaced in PR review: + upgrading an explicitly free deployment to paid via ``upsert_deployment`` + (same deployment count, same router instance) must invalidate the + cached ``_is_model_cost_zero=True`` answer so budget checks resume + immediately — not after the next proxy restart. + """ + from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo + + router = Router( + model_list=[ + { + "model_name": "ramping-model", + "litellm_params": { + "model": "openai/ramping-deploy", + "api_key": "sk-fake", + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + }, + "model_info": { + "id": "ramping-deploy-id", + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + }, + }, + ] + ) + # Warm the cache as zero-cost. + assert _is_model_cost_zero(model="ramping-model", llm_router=router) is True + assert router._zero_cost_cache.get("ramping-model") is True + + # In-place pricing update: same deployment count, same router id, + # same model name. The pre-fix cache key was + # ``(id(router), len(model_list), model_name)`` and would not change. + router.upsert_deployment( + deployment=Deployment( + model_name="ramping-model", + litellm_params=LiteLLM_Params( + model="openai/ramping-deploy", + api_key="sk-fake", + input_cost_per_token=0.000002, + output_cost_per_token=0.000008, + ), + model_info=ModelInfo( + id="ramping-deploy-id", + input_cost_per_token=0.000002, + output_cost_per_token=0.000008, + ), + ) + ) + + # Cache must have been cleared by ``_invalidate_model_group_info_cache``. + assert router._zero_cost_cache == {} + # Subsequent call sees the new pricing and enforces budget. + assert _is_model_cost_zero(model="ramping-model", llm_router=router) is False + + def test_handles_router_without_zero_cost_cache_attribute(self): + """Tolerate router-like objects (e.g. ``MagicMock`` stand-ins) that + do not expose ``_zero_cost_cache`` — the auth check must still + compute a correct answer, just without caching.""" + from unittest.mock import MagicMock + + from litellm.types.router import ModelGroupInfo + + mock_router = MagicMock(spec=Router) + mock_router.model_list = [] + mock_router.get_model_group_info.return_value = ModelGroupInfo( + model_group="paid-model", + providers=["openai"], + input_cost_per_token=0.001, + output_cost_per_token=0.002, + ) + # Strip the attribute so the helper falls back to the no-cache path. + del mock_router._zero_cost_cache + + result = _is_model_cost_zero(model="paid-model", llm_router=mock_router) + assert result is False diff --git a/tests/test_litellm/proxy/test_proxy_logging_hook_detection.py b/tests/test_litellm/proxy/test_proxy_logging_hook_detection.py new file mode 100644 index 00000000000..4aebcf40aa5 --- /dev/null +++ b/tests/test_litellm/proxy/test_proxy_logging_hook_detection.py @@ -0,0 +1,128 @@ +import pytest + +import litellm +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.integrations.custom_logger import CustomLogger +from litellm.proxy.utils import ProxyLogging + + +def test_has_post_call_response_headers_callbacks_ignores_empty_callbacks( + monkeypatch, +): + monkeypatch.setattr(litellm, "callbacks", []) + + assert ProxyLogging.has_post_call_response_headers_callbacks() is False + + +def test_has_post_call_response_headers_callbacks_requires_override( + monkeypatch, +): + """A vanilla ``CustomLogger`` inherits the no-op response-headers hook; + the capability flag must stay False so the proxy can skip the headers + loop entirely. Only callbacks that *override* the hook should flip it.""" + monkeypatch.setattr(litellm, "callbacks", [CustomLogger()]) + assert ProxyLogging.has_post_call_response_headers_callbacks() is False + + class _AddsHeaders(CustomLogger): + async def async_post_call_response_headers_hook(self, **kwargs): + return {"x-custom": "1"} + + monkeypatch.setattr(litellm, "callbacks", [_AddsHeaders()]) + assert ProxyLogging.has_post_call_response_headers_callbacks() is True + + +def test_has_streaming_callbacks_uses_custom_logger_detection(monkeypatch): + monkeypatch.setattr(litellm, "callbacks", []) + assert ProxyLogging.has_streaming_callbacks() is False + + monkeypatch.setattr(litellm, "callbacks", [CustomLogger()]) + assert ProxyLogging.has_streaming_callbacks() is False + + class StreamingLogger(CustomLogger): + async def async_post_call_streaming_hook(self, **kwargs): + return kwargs.get("response") + + monkeypatch.setattr(litellm, "callbacks", [StreamingLogger()]) + assert ProxyLogging.has_streaming_callbacks() is True + + +def test_has_streaming_callbacks_detects_guardrails(monkeypatch): + monkeypatch.setattr(litellm, "callbacks", [CustomGuardrail()]) + assert ProxyLogging.has_streaming_callbacks() is True + + +@pytest.mark.asyncio +async def test_post_call_response_headers_hook_returns_early_without_callbacks( + monkeypatch, +): + monkeypatch.setattr(litellm, "callbacks", []) + proxy_logging_obj = ProxyLogging(user_api_key_cache={}) # type: ignore[arg-type] + + result = await proxy_logging_obj.post_call_response_headers_hook( + data={}, + user_api_key_dict=None, # type: ignore[arg-type] + response=None, + request_headers={}, + ) + + assert result == {} + + +def test_callback_capabilities_skips_default_custom_logger(monkeypatch): + """ + Internal proxy hooks (e.g. _PROXY_MaxBudgetLimiter, ManagedFiles) inherit + the default ``async_post_call_streaming_iterator_hook`` body. The + capability scanner must NOT report them as iterator overrides — wrapping + the chunk stream through every no-op layer was responsible for ~10x + streaming overhead on default deployments. + """ + + class _InternalNoopHook(CustomLogger): + pass + + monkeypatch.setattr(litellm, "callbacks", [_InternalNoopHook()]) + + caps = ProxyLogging._callback_capabilities() + # Subclass inherits the base no-op for every hook — every capability flag + # must stay False so the proxy short-circuits the corresponding loops. + assert caps.has_post_call_response_headers is False + assert caps.iterator_overrides == () + assert caps.has_iterator_override is False + assert caps.has_streaming_chunk_override is False + assert caps.has_guardrail is False + + +def test_callback_capabilities_captures_iterator_override(monkeypatch): + class _OverridesIterator(CustomLogger): + async def async_post_call_streaming_iterator_hook( # type: ignore[override] + self, user_api_key_dict, response, request_data + ): + async for item in response: + yield item + + override = _OverridesIterator() + monkeypatch.setattr(litellm, "callbacks", [override]) + + caps = ProxyLogging._callback_capabilities() + assert caps.has_iterator_override is True + assert len(caps.iterator_overrides) == 1 + resolved, kind = caps.iterator_overrides[0] + assert resolved is override + assert kind == "override" + + +def test_callback_capabilities_cache_invalidates_on_list_change(monkeypatch): + """The cache key includes (length, id-of-each-callback). Mutating the + callback list must produce a fresh capability snapshot.""" + monkeypatch.setattr(litellm, "callbacks", []) + assert ProxyLogging._callback_capabilities().resolved_callbacks == () + + class _OverridesPreCall(CustomLogger): + async def async_pre_call_hook(self, *args, **kwargs): + return kwargs.get("data") + + pre = _OverridesPreCall() + monkeypatch.setattr(litellm, "callbacks", [pre]) + caps = ProxyLogging._callback_capabilities() + assert caps.has_pre_call_override is True + assert pre in caps.resolved_callbacks diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index e66dbcc3495..6eade957014 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -4515,6 +4515,69 @@ async def test_async_data_generator_cleanup_on_early_exit(): mock_response.aclose.assert_awaited_once() +@pytest.mark.asyncio +async def test_async_data_generator_uses_direct_stream_fast_path_without_callbacks(): + """ + When there are no streaming callbacks, async_data_generator should avoid + per-chunk hook machinery and iterate the provider stream directly. + """ + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.proxy_server import async_data_generator + from litellm.proxy.utils import ProxyLogging + + mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) + mock_request_data = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "test"}], + } + mock_chunks = [ + {"choices": [{"delta": {"content": "Hello"}}]}, + {"choices": [{"delta": {"content": " world"}}]}, + ] + + class MockStream: + def __aiter__(self): + return self._stream() + + async def _stream(self): + for chunk in mock_chunks: + yield chunk + + async def aclose(self): + pass + + mock_response = MockStream() + mock_response.aclose = AsyncMock() + mock_proxy_logging_obj = MagicMock(spec=ProxyLogging) + mock_proxy_logging_obj.has_streaming_callbacks.return_value = False + mock_proxy_logging_obj.needs_iterator_wrap.return_value = False + mock_proxy_logging_obj.needs_per_chunk_streaming_hook.return_value = False + mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = MagicMock() + mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock() + mock_proxy_logging_obj.post_call_failure_hook = AsyncMock() + + with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj): + with patch.object( + ProxyLogging, "_fire_deferred_stream_logging" + ) as mock_deferred_logging: + yielded_data = [] + async for data in async_data_generator( + mock_response, mock_user_api_key_dict, mock_request_data + ): + yielded_data.append(data) + + yielded_text = [ + chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk + for chunk in yielded_data + ] + assert len([chunk for chunk in yielded_text if chunk.startswith("data: {")]) == 2 + assert yielded_text[-1] == "data: [DONE]\n\n" + mock_proxy_logging_obj.async_post_call_streaming_iterator_hook.assert_not_called() + mock_proxy_logging_obj.async_post_call_streaming_hook.assert_not_awaited() + mock_deferred_logging.assert_called_once_with(mock_request_data) + mock_response.aclose.assert_awaited_once() + + @pytest.mark.asyncio async def test_async_data_generator_cleanup_on_normal_completion(): """ diff --git a/tests/test_litellm/proxy/test_response_model_sanitization.py b/tests/test_litellm/proxy/test_response_model_sanitization.py index 91792f62d6c..621291b8331 100644 --- a/tests/test_litellm/proxy/test_response_model_sanitization.py +++ b/tests/test_litellm/proxy/test_response_model_sanitization.py @@ -66,6 +66,69 @@ def _make_model_response_stream_chunk(model: str) -> litellm.ModelResponseStream return litellm.ModelResponseStream(**chunk_dict) +def _decode_sse_chunk(chunk) -> str: + return chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk + + +def test_restamp_streaming_chunk_skips_matching_model(): + from litellm.proxy.proxy_server import _restamp_streaming_chunk_model + + chunk = _make_model_response_stream_chunk("client-model") + + result, model_mismatch_logged = _restamp_streaming_chunk_model( + chunk=chunk, + requested_model_from_client="client-model", + request_data={"litellm_call_id": "test-call-id"}, + model_mismatch_logged=False, + ) + + assert result is chunk + assert result.model == "client-model" + assert model_mismatch_logged is False + + +def test_fast_serialize_simple_streaming_chunk_matches_model_dump_json(): + from litellm.proxy.proxy_server import _serialize_streaming_chunk + + chunk = _make_model_response_stream_chunk("client-model") + + assert json.loads(_serialize_streaming_chunk(chunk)) == json.loads( + chunk.model_dump_json(exclude_none=True, exclude_unset=True) + ) + + +def test_fast_serialize_returns_none_when_model_field_is_missing(): + """ + The fast path must mirror ``model_dump_json(exclude_none=True)``: when + ``chunk.model`` is ``None`` the slow path omits the field entirely. + Emitting ``"model": null`` would diverge and trip strict OpenAI- + compatible clients that reject ``null`` for optional string fields. + Falling back to ``None`` lets the canonical serializer handle the edge. + """ + from litellm.proxy.proxy_server import ( + _fast_serialize_simple_model_response_stream, + _serialize_streaming_chunk, + ) + + chunk = _make_model_response_stream_chunk("client-model") + chunk.model = None # type: ignore[assignment] + + assert _fast_serialize_simple_model_response_stream(chunk) is None + + # Going through the public ``_serialize_streaming_chunk`` should still + # produce a serialized result via the slow-path fallback, and it must + # not contain ``"model": null``. + serialized = _serialize_streaming_chunk(chunk) + payload_str = ( + serialized.decode("utf-8") if isinstance(serialized, bytes) else serialized + ) + assert '"model": null' not in payload_str + assert '"model":null' not in payload_str + assert json.loads(payload_str) == json.loads( + chunk.model_dump_json(exclude_none=True, exclude_unset=True) + ) + + def test_proxy_chat_completion_does_not_return_provider_prefixed_model( tmp_path, monkeypatch ): @@ -164,6 +227,21 @@ async def test_proxy_streaming_chunks_do_not_return_provider_prefixed_model( "async_post_call_streaming_hook", AsyncMock(side_effect=lambda **kwargs: kwargs["response"]), ) + monkeypatch.setattr( + proxy_server.proxy_logging_obj, + "has_streaming_callbacks", + MagicMock(return_value=True), + ) + monkeypatch.setattr( + proxy_server.proxy_logging_obj, + "needs_iterator_wrap", + MagicMock(return_value=True), + ) + monkeypatch.setattr( + proxy_server.proxy_logging_obj, + "needs_per_chunk_streaming_hook", + MagicMock(return_value=True), + ) user_api_key_dict = UserAPIKeyAuth(api_key="sk-1234") @@ -179,7 +257,7 @@ async def test_proxy_streaming_chunks_do_not_return_provider_prefixed_model( # First chunk is expected to be JSON, last chunk is [DONE] assert len(chunks) >= 2 - first = chunks[0] + first = _decode_sse_chunk(chunks[0]) assert first.startswith("data: ") payload = json.loads(first[len("data: ") :].strip()) @@ -222,6 +300,21 @@ async def test_proxy_streaming_chunks_use_client_requested_model_before_alias_ma "async_post_call_streaming_hook", AsyncMock(side_effect=lambda **kwargs: kwargs["response"]), ) + monkeypatch.setattr( + proxy_server.proxy_logging_obj, + "has_streaming_callbacks", + MagicMock(return_value=True), + ) + monkeypatch.setattr( + proxy_server.proxy_logging_obj, + "needs_iterator_wrap", + MagicMock(return_value=True), + ) + monkeypatch.setattr( + proxy_server.proxy_logging_obj, + "needs_per_chunk_streaming_hook", + MagicMock(return_value=True), + ) user_api_key_dict = UserAPIKeyAuth(api_key="sk-1234") @@ -239,7 +332,7 @@ async def test_proxy_streaming_chunks_use_client_requested_model_before_alias_ma chunks.append(item) assert len(chunks) >= 2 - first = chunks[0] + first = _decode_sse_chunk(chunks[0]) assert first.startswith("data: ") payload = json.loads(first[len("data: ") :].strip()) @@ -279,6 +372,21 @@ async def test_proxy_streaming_azure_model_router_preserves_actual_model(monkeyp "async_post_call_streaming_hook", AsyncMock(side_effect=lambda **kwargs: kwargs["response"]), ) + monkeypatch.setattr( + proxy_server.proxy_logging_obj, + "has_streaming_callbacks", + MagicMock(return_value=True), + ) + monkeypatch.setattr( + proxy_server.proxy_logging_obj, + "needs_iterator_wrap", + MagicMock(return_value=True), + ) + monkeypatch.setattr( + proxy_server.proxy_logging_obj, + "needs_per_chunk_streaming_hook", + MagicMock(return_value=True), + ) user_api_key_dict = UserAPIKeyAuth(api_key="sk-1234") @@ -296,7 +404,7 @@ async def test_proxy_streaming_azure_model_router_preserves_actual_model(monkeyp chunks.append(item) assert len(chunks) >= 2 - first = chunks[0] + first = _decode_sse_chunk(chunks[0]) assert first.startswith("data: ") payload = json.loads(first[len("data: ") :].strip()) @@ -337,6 +445,21 @@ async def test_proxy_streaming_fastest_response_preserves_winning_model(monkeypa "async_post_call_streaming_hook", AsyncMock(side_effect=lambda **kwargs: kwargs["response"]), ) + monkeypatch.setattr( + proxy_server.proxy_logging_obj, + "has_streaming_callbacks", + MagicMock(return_value=True), + ) + monkeypatch.setattr( + proxy_server.proxy_logging_obj, + "needs_iterator_wrap", + MagicMock(return_value=True), + ) + monkeypatch.setattr( + proxy_server.proxy_logging_obj, + "needs_per_chunk_streaming_hook", + MagicMock(return_value=True), + ) user_api_key_dict = UserAPIKeyAuth(api_key="sk-1234") @@ -355,7 +478,7 @@ async def test_proxy_streaming_fastest_response_preserves_winning_model(monkeypa chunks.append(item) assert len(chunks) >= 2 - first = chunks[0] + first = _decode_sse_chunk(chunks[0]) assert first.startswith("data: ") payload = json.loads(first[len("data: ") :].strip())