From 4c6c84afb73013332d19fb6167c889534c92596a Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 3 Oct 2026 09:20:28 -0700 Subject: [PATCH] perf(proxy): stop prompt-cache eligibility from tokenizing the whole conversation (#44221) is_prompt_caching_valid_prompt ran the full Python token_counter over every message to compare against the deployment's prompt cache minimum, 500 to 1000 ms at 440k to 740k tokens on every request that reaches the prompt_caching pre-call check, Rust on or off. messages_reach_token_count does the same arithmetic as token_counter(...) >= threshold and stops at the first message that reaches the threshold. Groups with one healthy deployment skip the prefix hash and pin lookup, which cannot change the result for them Four fixed name span events make the pre-LLM phases measurable with OTel v2: litellm.request.body_received (with body_bytes) once per body read before parsing, on the JSON, binary and form branches, body_parsed, pre_call_completed, and deployment_selected emitted once per pick inside Router.async_get_available_deployment and get_available_deployment with attempt, reason and model group, so every router surface, retry and fallback is covered. Measured locally on /v1/chat/completions, /v1/messages and /v1/responses at 440k tokens with Rust on and off against a fake upstream Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/_lazy_imports.py | 13 ++ litellm/integrations/otel/logger.py | 12 ++ litellm/integrations/otel/runtime.py | 20 +- litellm/litellm_core_utils/token_counter.py | 32 +++ litellm/proxy/auth/user_api_key_auth.py | 3 +- litellm/proxy/common_request_processing.py | 2 + .../proxy/common_utils/http_parsing_utils.py | 20 +- litellm/router.py | 42 +++- .../prompt_caching_deployment_check.py | 2 + litellm/utils.py | 17 +- .../integrations/otel/test_otel_v2_logger.py | 25 +++ tests/unit/integrations/otel/test_runtime.py | 7 + .../litellm_core_utils/test_token_counter.py | 189 ++++++++++++++++++ .../common_utils/test_http_parsing_utils.py | 53 ++++- .../test_prompt_caching_deployment_check.py | 52 ++++- tests/unit/test_router/test_router.py | 128 ++++++++++++ tests/unit/test_utils.py | 40 ++++ 17 files changed, 637 insertions(+), 20 deletions(-) diff --git a/litellm/_lazy_imports.py b/litellm/_lazy_imports.py index 29fb46fa125..d4a12e7c1e5 100644 --- a/litellm/_lazy_imports.py +++ b/litellm/_lazy_imports.py @@ -123,6 +123,7 @@ def _get_modified_max_tokens() -> "Callable[..., int | None]": # Lazy loader for token_counter to avoid importing token_counter module at module import time _token_counter_new_func: "Callable[..., int] | None" = None +_messages_reach_token_count_func: "Callable[..., bool] | None" = None def _get_token_counter_new() -> "Callable[..., int]": @@ -145,6 +146,18 @@ def _get_token_counter_new() -> "Callable[..., int]": return _token_counter_new_func +def _get_messages_reach_token_count() -> "Callable[..., bool]": + """Lazily load ``messages_reach_token_count`` for the same reason as ``_get_token_counter_new``.""" + global _messages_reach_token_count_func + if _messages_reach_token_count_func is None: + from litellm.litellm_core_utils.token_counter import ( + messages_reach_token_count as _messages_reach_token_count_imported, + ) + + _messages_reach_token_count_func = _messages_reach_token_count_imported + return _messages_reach_token_count_func + + # ============================================================================ # MAIN LAZY IMPORT SYSTEM # ============================================================================ diff --git a/litellm/integrations/otel/logger.py b/litellm/integrations/otel/logger.py index 740580dfcf3..96df9728a73 100644 --- a/litellm/integrations/otel/logger.py +++ b/litellm/integrations/otel/logger.py @@ -771,6 +771,12 @@ class OpenTelemetryV2(CustomLogger): stamp_error(span, _span_error_from_exception(exc), record_event=False, set_status=False) raise + def add_phase_event(self, name: str, attributes: Mapping[str, str | int] | None = None) -> None: + """Mark a point in the request on its root span, or on the ambient span before the root is anchored.""" + span: Final = request_root_span() or get_current_span() + if is_recordable_span(span): + span.add_event(name, attributes) + async def async_pre_call_hook( self, user_api_key_dict: "UserAPIKeyAuth", @@ -1027,6 +1033,12 @@ def phase_span(name: str) -> "Iterator[Span | None]": yield span +def phase_event(name: str, attributes: Mapping[str, str | int] | None = None) -> None: + logger: Final = _registered_v2_logger() + if logger is not None: + logger.add_phase_event(name, attributes) + + def build_otel_v2_logger( config: OpenTelemetryV2Config, callback_name: str | None = None, diff --git a/litellm/integrations/otel/runtime.py b/litellm/integrations/otel/runtime.py index 13903597e1a..ff75bc4d800 100644 --- a/litellm/integrations/otel/runtime.py +++ b/litellm/integrations/otel/runtime.py @@ -7,17 +7,21 @@ V2 is not the active logger — so a call site can wrap a request phase or seed identity unconditionally. """ -from collections.abc import Callable, Iterator +from collections.abc import Callable, Iterator, Mapping from contextlib import AbstractContextManager, contextmanager from functools import cache -from typing import TYPE_CHECKING, Final +from typing import TYPE_CHECKING, Final, TypeAlias if TYPE_CHECKING: from opentelemetry.trace import Span +PhaseEventAttributes: TypeAlias = Mapping[str, str | int] + @cache -def _otel_runtime() -> "tuple[Callable[[str], AbstractContextManager[Span | None]], Callable[..., None]] | None": +def _otel_runtime() -> ( + "tuple[Callable[[str], AbstractContextManager[Span | None]], Callable[..., None], Callable[[str, PhaseEventAttributes | None], None]] | None" +): """Resolve the SDK-backed hooks once and cache the outcome, absence included. CPython never caches a failed import, so without this memoization every call @@ -28,7 +32,7 @@ def _otel_runtime() -> "tuple[Callable[[str], AbstractContextManager[Span | None from litellm.integrations.otel import logger except Exception: return None - return (logger.phase_span, logger.seed_request_identity) + return (logger.phase_span, logger.seed_request_identity, logger.phase_event) @contextmanager @@ -46,6 +50,14 @@ def phase_span(name: str) -> "Iterator[Span | None]": yield span +def phase_event(name: str, attributes: PhaseEventAttributes | None = None) -> None: + """Mark a point in the request on its span (no-op without V2).""" + runtime: Final = _otel_runtime() + if runtime is None: + return + runtime[2](name, attributes) + + def seed_request_identity(user_api_key_dict: object, model: object = None) -> None: """Seed request-identity Baggage at the auth boundary (no-op without V2).""" runtime: Final = _otel_runtime() diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index cdd2d0654be..b92eb74cbb1 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -4,6 +4,7 @@ import base64 import io import struct from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence +from itertools import accumulate from typing import Final, Literal, cast import anyio @@ -466,6 +467,37 @@ def token_counter( return num_tokens +def messages_reach_token_count( + model: str, + messages: Sequence[AllMessageValues | Message], + threshold: int, + tools: list[ChatCompletionToolParam] | None = None, + use_default_image_token_count: bool = False, +) -> bool: + """Whether ``messages`` plus ``tools`` hold at least ``threshold`` prompt tokens for ``model``. + + Same arithmetic as ``token_counter(messages=..., tools=...) >= threshold``, counted one message + at a time and stopped at the first message that crosses the threshold, so a prompt far above it + costs the tokenizer a few messages rather than the whole conversation. + """ + from litellm.utils import convert_list_message_to_dict + + if litellm.disable_token_counter is True: + return threshold <= 0 + new_messages: Final = cast( # cast-ok: convert_list_message_to_dict is untyped, same as token_counter + list[AllMessageValues], convert_list_message_to_dict(messages) + ) + params: Final = _MessageCountParams(model, None) + includes_system_message: Final = any(message.get("role", None) == "system" for message in new_messages) + per_message_counts: Final = ( + _count_messages(params, [message], use_default_image_token_count, None) for message in new_messages + ) + running_totals: Final = accumulate( + per_message_counts, initial=_count_extra(params.count_function, tools, None, includes_system_message) + ) + return any(total >= threshold for total in running_totals) + + def _count_function_call_tokens( key: str, value: object, diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index b7620c5f8bd..5395d817e33 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -36,7 +36,7 @@ from litellm.constants import ( MODEL_GROUP_ALIAS_RESOLVED_SCOPE_KEY, ) from litellm.integrations.otel.model.config import is_otel_v2_enabled -from litellm.integrations.otel.runtime import phase_span, seed_request_identity +from litellm.integrations.otel.runtime import phase_event, phase_span, seed_request_identity from litellm.litellm_core_utils.dd_tracing import tracer from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value from litellm.proxy._types import * @@ -3501,6 +3501,7 @@ async def user_api_key_auth( _ensure_parent_otel_span_on_request_state(request) request_data, body_parse_exception = await _read_request_body_deferring_parse_failure(request=request) + phase_event("litellm.request.body_parsed") route: Final[str] = get_request_route(request=request) ## CHECK IF ROUTE IS ALLOWED diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index a6651114cd0..78c3f53c44f 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -47,6 +47,7 @@ from litellm.constants import ( UNSAFE_PROXY_RESPONSE_HEADERS, ) from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.integrations.otel.runtime import phase_event from litellm.litellm_core_utils.bug_report import ( allowlisted, bug_report_notice, @@ -2574,6 +2575,7 @@ class ProxyBaseLLMRequestProcessing: route_type=route_type, llm_router=llm_router, ) + phase_event("litellm.request.pre_call_completed") # Defer async logging when post-call guardrails are configured so the # StandardLoggingPayload is built after guardrails write to metadata. diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index aa4d6a39f25..b72fa2edb1c 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -15,6 +15,7 @@ from litellm.constants import ( CLIENT_REQUESTED_MODEL_SCOPE_KEY, MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB, ) +from litellm.integrations.otel.runtime import phase_event from litellm.proxy._types import ProxyException from litellm.proxy.common_utils.callback_utils import ( get_metadata_variable_name_from_kwargs, @@ -168,6 +169,19 @@ def _parse_binary_body(body: bytes) -> dict: return {} +def _declared_content_length(headers: Mapping[str, str]) -> int | None: + declared: Final = headers.get("content-length") + return int(declared) if isinstance(declared, str) and declared.isdigit() else None + + +def _mark_body_received(byte_count: int | None) -> None: + """Marks the end of body transfer on the request's server span, once per body read.""" + phase_event( + "litellm.request.body_received", + None if byte_count is None else {"litellm.request.body_bytes": byte_count}, + ) + + def is_otlp_trace_request(request: Request) -> bool: return request.method == "POST" and get_route_path(request.scope) == "/v1/traces" @@ -198,10 +212,13 @@ async def _read_request_body(request: Request | None) -> dict: content_type: Final = _request_headers.get("content-type", "") if _normalize_media_type(content_type) in _BINARY_CONTENT_TYPES: - parsed_body = _parse_binary_body(await request.body()) + binary_body: Final = await request.body() + _mark_body_received(len(binary_body)) + parsed_body = _parse_binary_body(binary_body) elif _is_form_content_type(content_type): try: form_data: Final = await request.form() + _mark_body_received(_declared_content_length(request.headers)) except Exception as e: # ``request.form()`` raises on malformed multipart (missing # boundary, malformed chunk encoding, …). Surface as 400 so @@ -222,6 +239,7 @@ async def _read_request_body(request: Request | None) -> dict: else: # Read the request body body: Final = await request.body() + _mark_body_received(len(body)) # Return empty dict if body is empty or None if not body: diff --git a/litellm/router.py b/litellm/router.py index cbc7e655810..3662d1f43eb 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -72,7 +72,7 @@ from litellm.constants import ( ) from litellm.integrations.custom_guardrail import is_guardrail_intervention from litellm.integrations.custom_logger import CustomLogger -from litellm.integrations.otel.runtime import phase_span +from litellm.integrations.otel.runtime import phase_event, phase_span from litellm.litellm_core_utils.asyncify import run_async_function from litellm.litellm_core_utils.core_helpers import ( _get_parent_otel_span_from_kwargs, @@ -477,6 +477,27 @@ def _stream_chunks_have_generated_content(chunks: Sequence[ModelResponseStream]) _NO_SESSION_KWARGS: Final[Mapping[str, Mapping[str, object]]] = MappingProxyType({}) _SESSION_ADAPTER: Final = TypeAdapter(Mapping[str, object]) _SILENT_MODEL_ADAPTER: Final = TypeAdapter(str | list[str]) +_ROUTING_KWARGS_ADAPTER: Final[TypeAdapter[Mapping[str, object] | None]] = TypeAdapter(Mapping[str, object] | None) +_DEPLOYMENT_SELECTED_EVENT: Final = "litellm.request.deployment_selected" + + +def _deployment_pick_attributes(model: str, request_kwargs: Mapping[str, object] | None) -> Mapping[str, str | int]: + """Bounded attributes for one deployment pick; attempt is 1-based within the current model group.""" + kwargs: Final = request_kwargs or {} + metadata: Final = kwargs.get("litellm_metadata", kwargs.get("metadata")) + attempted_retries: Final = metadata.get("attempted_retries") if isinstance(metadata, Mapping) else None + retries: Final = attempted_retries if isinstance(attempted_retries, int) else 0 + fallback_depth: Final = kwargs.get("fallback_depth") + reason: Final = ( + "retry" if retries > 0 else "fallback" if isinstance(fallback_depth, int) and fallback_depth > 0 else "initial" + ) + return MappingProxyType( + { + "litellm.deployment.attempt": retries + 1, + "litellm.deployment.reason": reason, + "litellm.deployment.model_group": model, + } + ) def _as_retry_skipped_deployment_ids(value: object) -> tuple[str, ...]: @@ -13342,6 +13363,9 @@ class Router: # the hook can replace `model` and routing-group lookup must key # off the final model name. strategy, strategy_selector = self._get_routing_context(model, request_kwargs) + pick_attributes: Final = _deployment_pick_attributes( + model, _ROUTING_KWARGS_ADAPTER.validate_python(request_kwargs) + ) routing_read_batch: Final = RoutingReadBatch.for_strategy(strategy, strategy_selector) with RoutingReadBatch.scoped(routing_read_batch): @@ -13357,6 +13381,7 @@ class Router: await self._async_override_selector_pre_call_check( strategy, strategy_selector, healthy_deployments, parent_otel_span ) + phase_event(_DEPLOYMENT_SELECTED_EVENT, pick_attributes) return healthy_deployments # When encrypted content affinity pins to a specific deployment, @@ -13364,16 +13389,19 @@ class Router: await self._async_override_selector_pre_call_check( strategy, strategy_selector, healthy_deployments[0], parent_otel_span ) + phase_event(_DEPLOYMENT_SELECTED_EVENT, pick_attributes) return healthy_deployments[0] start_time: Final = time.time() if strategy == "simple-shuffle": - return simple_shuffle( + shuffled: Final = simple_shuffle( resolve_model_alias=self._get_model_from_alias, healthy_deployments=healthy_deployments, model=model, request_kwargs=request_kwargs, ) + phase_event(_DEPLOYMENT_SELECTED_EVENT, pick_attributes) + return shuffled with PrefetchedUsage.scoped( routing_read_batch.prefetched_usage if routing_read_batch is not None else None ): @@ -13416,6 +13444,7 @@ class Router: ) ) + phase_event(_DEPLOYMENT_SELECTED_EVENT, pick_attributes) return deployment except Exception as e: traceback_exception: Final = traceback.format_exc() @@ -14207,6 +14236,9 @@ class Router: request_kwargs=request_kwargs, ) strategy, strategy_selector = self._get_routing_context(model, request_kwargs) + pick_attributes: Final = _deployment_pick_attributes( + model, _ROUTING_KWARGS_ADAPTER.validate_python(request_kwargs) + ) if isinstance(healthy_deployments, dict): if (healthy_deployments.get("model_info") or {}).get("blocked") is True: @@ -14216,6 +14248,7 @@ class Router: llm_provider="", ) self._override_selector_pre_call_check(strategy, strategy_selector, healthy_deployments) + phase_event(_DEPLOYMENT_SELECTED_EVENT, pick_attributes) return healthy_deployments parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(request_kwargs) @@ -14295,12 +14328,14 @@ class Router: if strategy == "simple-shuffle": # if users pass rpm or tpm, we do a random weighted pick - based on rpm/tpm ############## Check 'weight' param set for weighted pick ################# - return simple_shuffle( + shuffled: Final = simple_shuffle( resolve_model_alias=self._get_model_from_alias, healthy_deployments=healthy_deployments, model=model, request_kwargs=request_kwargs, ) + phase_event(_DEPLOYMENT_SELECTED_EVENT, pick_attributes) + return shuffled deployment: Final = self._select_deployment_sync( strategy=strategy, selector=strategy_selector, @@ -14332,6 +14367,7 @@ class Router: self.print_deployment(deployment), model, ) + phase_event(_DEPLOYMENT_SELECTED_EVENT, pick_attributes) return deployment def get_available_deployment_for_pass_through( diff --git a/litellm/router_utils/pre_call_checks/prompt_caching_deployment_check.py b/litellm/router_utils/pre_call_checks/prompt_caching_deployment_check.py index eabd79f1847..9f1558dedae 100644 --- a/litellm/router_utils/pre_call_checks/prompt_caching_deployment_check.py +++ b/litellm/router_utils/pre_call_checks/prompt_caching_deployment_check.py @@ -66,6 +66,8 @@ class PromptCachingDeploymentCheck(CustomLogger): return healthy_deployments if request_kwargs is not None and request_kwargs.get("_target_order") is not None: return healthy_deployments + if not healthy_deployments[1:]: + return healthy_deployments if messages is not None and await offload_token_count(is_prompt_caching_valid_prompt)( messages=messages, diff --git a/litellm/utils.py b/litellm/utils.py index e398f4eaf8a..9200844a2e3 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -55,6 +55,7 @@ import litellm.litellm_core_utils.json_validation_rule from litellm._internal_context import is_internal_call from litellm._lazy_imports import ( _get_default_encoding, + _get_messages_reach_token_count, _get_modified_max_tokens, _get_token_counter_new, ) @@ -10119,19 +10120,19 @@ def is_prompt_caching_valid_prompt( OpenAI's minimum is a flat 1024 across models, which the default already covers. """ try: - if messages is None and tools is None: + if messages is None: return False if custom_llm_provider is not None and not model.startswith(custom_llm_provider): model = custom_llm_provider + "/" + model - token_count: Final = token_counter( - messages=messages, - tools=tools, - model=model, - use_default_image_token_count=True, - ) if min_token_count is None: min_token_count = get_prompt_cache_min_tokens(model=model) - return token_count >= min_token_count + return _get_messages_reach_token_count()( + model=model, + messages=messages, + threshold=min_token_count, + tools=tools, + use_default_image_token_count=True, + ) except Exception as e: verbose_logger.error("Error in is_prompt_caching_valid_prompt: %s", e) return False diff --git a/tests/unit/integrations/otel/test_otel_v2_logger.py b/tests/unit/integrations/otel/test_otel_v2_logger.py index c0fa890e8a2..6c99e8bf14d 100644 --- a/tests/unit/integrations/otel/test_otel_v2_logger.py +++ b/tests/unit/integrations/otel/test_otel_v2_logger.py @@ -1078,6 +1078,31 @@ def test_llm_span_anchors_to_root_even_inside_active_phase_span(): assert llm_span.parent.span_id != auth_span.get_span_context().span_id +def test_phase_event_lands_on_root_span_even_inside_active_phase_span(): + """Request phase marks (body parsed, pre-call done, deployment selected) are + events on the server span: they must land on the anchored root even while the + ``auth`` phase span is active, and on the ambient server span before the root + is anchored (the body is parsed before auth anchors it).""" + logger, exporter = _logger() + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) + with trace.use_span(server, end_on_exit=False): + logger.add_phase_event("litellm.request.body_parsed") + set_request_root_span(server) + with logger.start_phase_span("auth /chat/completions"): + logger.add_phase_event("litellm.request.pre_call_completed") + logger.add_phase_event("litellm.request.deployment_selected", {"litellm.deployment.attempt": 1}) + server.end() + by_name = {s.name: s for s in exporter.get_finished_spans()} + root_events = by_name[LITELLM_PROXY_REQUEST_SPAN_NAME].events + assert [e.name for e in root_events] == [ + "litellm.request.body_parsed", + "litellm.request.pre_call_completed", + "litellm.request.deployment_selected", + ] + assert dict(root_events[2].attributes or {}) == {"litellm.deployment.attempt": 1} + assert by_name["auth /chat/completions"].events == () + + def test_live_llm_span_anchors_to_root_with_no_active_span(): """Bug 2 (pass-through), live path: even with no span active at ``pre_call``, the anchor is a recordable parent, so the span opens live under the server root diff --git a/tests/unit/integrations/otel/test_runtime.py b/tests/unit/integrations/otel/test_runtime.py index d11f31b2523..285759c7553 100644 --- a/tests/unit/integrations/otel/test_runtime.py +++ b/tests/unit/integrations/otel/test_runtime.py @@ -62,3 +62,10 @@ def test_wrappers_no_op_when_runtime_absent(monkeypatch): assert span is None assert runtime.seed_request_identity({"token": "sk-x"}, model="gpt-4o") is None + + +def test_phase_event_no_ops_when_runtime_absent(monkeypatch): + monkeypatch.setattr(runtime, "_otel_runtime", lambda: None) + + assert runtime.phase_event("litellm.request.body_parsed") is None + assert runtime.phase_event("litellm.request.body_received", {"litellm.request.body_bytes": 3}) is None diff --git a/tests/unit/litellm_core_utils/test_token_counter.py b/tests/unit/litellm_core_utils/test_token_counter.py index c71b1496bdd..e0c5c22d420 100644 --- a/tests/unit/litellm_core_utils/test_token_counter.py +++ b/tests/unit/litellm_core_utils/test_token_counter.py @@ -1633,3 +1633,192 @@ def test_token_counter_uses_the_tokenizer_of_each_model_family_and_of_a_custom_t "custom": expected["Xenova/llama-3-tokenizer"], "requested": sorted(served), } + + +def _threshold_test_messages(turns: int) -> list[dict]: + messages: list[dict] = [{"role": "system", "content": "You are a terse assistant. " * 20}] + for index in range(turns): + messages.append({"role": "user", "content": f"Question {index}: what is the capital of country number {index}?"}) + messages.append( + { + "role": "assistant", + "content": [{"type": "text", "text": f"Answer {index}: the capital is city number {index}."}], + } + ) + return messages + + +_THRESHOLD_TEST_TOOLS: Final = [ + { + "type": "function", + "function": { + "name": "lookup_capital", + "description": "Look up the capital of a country", + "parameters": {"type": "object", "properties": {"country": {"type": "string"}}}, + }, + } +] + + +def test_messages_reach_token_count_agrees_with_token_counter_at_every_threshold() -> None: + """The threshold check is the same arithmetic as token_counter(...) >= threshold, including the + tools and system-message adjustments, so the boundary values must agree exactly.""" + from litellm.litellm_core_utils.token_counter import messages_reach_token_count + + messages = _threshold_test_messages(turns=12) + total = token_counter_new( + model="claude-3-5-sonnet-20240620", + messages=messages, + tools=_THRESHOLD_TEST_TOOLS, + use_default_image_token_count=True, + ) + assert total > 100 + for threshold in (0, 1, total - 1, total, total + 1, 10 * total): + assert messages_reach_token_count( + model="claude-3-5-sonnet-20240620", + messages=messages, + threshold=threshold, + tools=_THRESHOLD_TEST_TOOLS, + use_default_image_token_count=True, + ) is (total >= threshold), threshold + + +_SHAPE_IMAGE: Final = "data:image/png;base64," + "iVBORw0KGgo=" * 4 + +_OPENAI_SHAPE_MESSAGES: Final = [ + {"role": "system", "content": "You are a careful assistant. " * 20}, + { + "role": "user", + "content": [ + {"type": "text", "text": "Describe this screenshot. " * 30}, + {"type": "image_url", "image_url": {"url": _SHAPE_IMAGE}}, + ], + }, + { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": "c1", "type": "function", "function": {"name": "read_file", "arguments": '{"path": "/a/b"}'}} + ], + }, + {"role": "tool", "tool_call_id": "c1", "content": "file body line\n" * 40}, + {"role": "assistant", "content": "Here is what the file does. " * 20}, +] + +_ANTHROPIC_SHAPE_MESSAGES: Final = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Describe this screenshot. " * 30, "cache_control": {"type": "ephemeral"}}, + {"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": "iVBORw0KGgo=" * 4}}, + ], + }, + { + "role": "assistant", + "content": [ + {"type": "text", "text": "Let me look."}, + {"type": "tool_use", "id": "t1", "name": "read_file", "input": {"path": "/a/b"}}, + ], + }, + { + "role": "user", + "content": [{"type": "tool_result", "tool_use_id": "t1", "content": [{"type": "text", "text": "file body line\n" * 40}]}], + }, + {"role": "assistant", "content": [{"type": "text", "text": "Here is what the file does. " * 20}]}, +] + +_RESPONSES_SHAPE_INPUT: Final = [ + { + "type": "message", + "role": "user", + "content": [ + {"type": "input_text", "text": "Describe this screenshot. " * 30}, + {"type": "input_image", "image_url": _SHAPE_IMAGE}, + ], + }, + {"type": "function_call", "call_id": "c1", "name": "read_file", "arguments": '{"path": "/a/b"}'}, + {"type": "function_call_output", "call_id": "c1", "output": "file body line\n" * 40}, +] + + +@pytest.mark.parametrize( + "messages", + [ + pytest.param(_OPENAI_SHAPE_MESSAGES, id="openai_chat_shape"), + pytest.param(_ANTHROPIC_SHAPE_MESSAGES, id="anthropic_messages_shape"), + ], +) +def test_messages_reach_token_count_agrees_with_token_counter_per_message_shape(messages: list[dict]) -> None: + """Content lists, images, tool calls, tool results and cache_control blocks in the OpenAI chat shape + (/v1/chat/completions) and the Anthropic shape (/v1/messages) count the same bounded as in full.""" + from litellm.litellm_core_utils.token_counter import messages_reach_token_count + + total = token_counter_new( + model="claude-3-5-sonnet-20240620", + messages=messages, + tools=_THRESHOLD_TEST_TOOLS, + use_default_image_token_count=True, + ) + assert total > 100 + for threshold in (0, 1, total - 1, total, total + 1, 10 * total): + assert messages_reach_token_count( + model="claude-3-5-sonnet-20240620", + messages=messages, + threshold=threshold, + tools=_THRESHOLD_TEST_TOOLS, + use_default_image_token_count=True, + ) is (total >= threshold), threshold + + +def test_messages_reach_token_count_rejects_responses_items_exactly_like_token_counter() -> None: + """Responses API input items are not chat messages; the full counter raises on them and the + bounded counter raises the same error rather than silently returning a verdict.""" + from litellm.litellm_core_utils.token_counter import messages_reach_token_count + + with pytest.raises(ValueError, match="input_text") as full: + token_counter_new(model="gpt-4o", messages=_RESPONSES_SHAPE_INPUT, use_default_image_token_count=True) + with pytest.raises(ValueError, match="input_text") as bounded: + messages_reach_token_count( + model="gpt-4o", messages=_RESPONSES_SHAPE_INPUT, threshold=10**6, use_default_image_token_count=True + ) + assert str(bounded.value) == str(full.value) + + +def test_messages_reach_token_count_stops_at_the_first_message_past_the_threshold( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Regression: the prompt-cache eligibility check used to tokenize every message of a 700k-token + Claude Code conversation to compare against a 1024-token minimum, costing hundreds of + milliseconds per request before routing. Counting must stop once the threshold is crossed.""" + import litellm.litellm_core_utils.token_counter as token_counter_module + + messages = _threshold_test_messages(turns=500) + counted_batches: list[int] = [] # mutable-ok: recorder for the _count_messages double + real_count_messages = token_counter_module._count_messages + + def counting(params, batch, use_default_image_token_count, default_token_count): + counted_batches.append(len(batch)) + return real_count_messages(params, batch, use_default_image_token_count, default_token_count) + + monkeypatch.setattr(token_counter_module, "_count_messages", counting) + assert token_counter_module.messages_reach_token_count( + model="claude-3-5-sonnet-20240620", messages=messages, threshold=1024 + ) + assert all(size == 1 for size in counted_batches) + bounded_calls: Final = len(counted_batches) + assert bounded_calls < len(messages) // 4, bounded_calls + + assert not token_counter_module.messages_reach_token_count( + model="claude-3-5-sonnet-20240620", messages=messages, threshold=10**9 + ) + assert len(counted_batches) - bounded_calls == len(messages) + + +def test_messages_reach_token_count_honours_disable_token_counter(monkeypatch: pytest.MonkeyPatch) -> None: + """With the counter disabled token_counter reports 0, so only a non-positive threshold is reached.""" + from litellm.litellm_core_utils.token_counter import messages_reach_token_count + + monkeypatch.setattr(litellm, "disable_token_counter", True) + messages = _threshold_test_messages(turns=3) + assert messages_reach_token_count(model="gpt-4o", messages=messages, threshold=0) is True + assert messages_reach_token_count(model="gpt-4o", messages=messages, threshold=1) is False diff --git a/tests/unit/proxy/common_utils/test_http_parsing_utils.py b/tests/unit/proxy/common_utils/test_http_parsing_utils.py index d68771eb847..7e42bf70671 100644 --- a/tests/unit/proxy/common_utils/test_http_parsing_utils.py +++ b/tests/unit/proxy/common_utils/test_http_parsing_utils.py @@ -14,6 +14,7 @@ from starlette.requests import Request import litellm +import litellm.proxy.common_utils.http_parsing_utils as http_parsing_utils from litellm.proxy._types import ProxyException from litellm.proxy.common_utils.http_parsing_utils import ( _is_form_content_type, @@ -33,13 +34,21 @@ from litellm.proxy.common_utils.http_parsing_utils import ( def _starlette_request( - body: bytes, content_type: str, path: str = "/v1/messages", content_encoding: str = "" + body: bytes, + content_type: str, + path: str = "/v1/messages", + content_encoding: str = "", + content_length: str = "", ) -> Request: scope = { "type": "http", "method": "POST", "path": path, - "headers": [(b"content-type", content_type.encode()), (b"content-encoding", content_encoding.encode())], + "headers": [ + (b"content-type", content_type.encode()), + (b"content-encoding", content_encoding.encode()), + (b"content-length", content_length.encode()), + ], "query_string": b"", } chunks = iter((body,)) @@ -50,6 +59,46 @@ def _starlette_request( return Request(scope, receive) +@pytest.mark.asyncio +async def test_read_request_body_marks_body_received_once_with_its_size(monkeypatch: pytest.MonkeyPatch): + events: list[tuple[str, dict[str, str | int]]] = [] # mutable-ok: recorder for the injected phase_event double + + def record(name: str, attributes: dict[str, str | int]) -> None: + events.append((name, dict(attributes))) + + monkeypatch.setattr(http_parsing_utils, "phase_event", record) + body: Final = orjson.dumps({"model": "claude-sonnet-4-5", "messages": [{"role": "user", "content": "x" * 4096}]}) + request: Final = _starlette_request(body, "application/json") + + assert await _read_request_body(request) == orjson.loads(body) + assert await _read_request_body(request) == orjson.loads(body) + + assert events == [("litellm.request.body_received", {"litellm.request.body_bytes": len(body)})] + + +@pytest.mark.asyncio +async def test_read_request_body_marks_body_received_for_binary_and_form_bodies(monkeypatch: pytest.MonkeyPatch): + events: list[tuple[str, dict[str, str | int] | None]] = [] # mutable-ok: recorder for the phase_event double + + def record(name: str, attributes: dict[str, str | int] | None) -> None: + events.append((name, None if attributes is None else dict(attributes))) + + monkeypatch.setattr(http_parsing_utils, "phase_event", record) + protobuf: Final = b"\x08\x96\x01" * 50 + form: Final = b"model=whisper-1&language=en" + form_type: Final = "application/x-www-form-urlencoded" + + await _read_request_body(_starlette_request(protobuf, "application/x-protobuf")) + await _read_request_body(_starlette_request(form, form_type, content_length=str(len(form)))) + await _read_request_body(_starlette_request(form, form_type)) + + assert events == [ + ("litellm.request.body_received", {"litellm.request.body_bytes": len(protobuf)}), + ("litellm.request.body_received", {"litellm.request.body_bytes": len(form)}), + ("litellm.request.body_received", None), + ] + + @pytest.mark.asyncio async def test_read_raw_json_body_returns_the_bytes_the_parsed_body_came_from(): body = b'{"model": "claude-sonnet-4-5", "messages": [{"role": "user", "content": "hi"}]}' diff --git a/tests/unit/router_utils/pre_call_checks/test_prompt_caching_deployment_check.py b/tests/unit/router_utils/pre_call_checks/test_prompt_caching_deployment_check.py index 00462b65bc2..c412546153a 100644 --- a/tests/unit/router_utils/pre_call_checks/test_prompt_caching_deployment_check.py +++ b/tests/unit/router_utils/pre_call_checks/test_prompt_caching_deployment_check.py @@ -181,6 +181,56 @@ async def test_async_filter_deployments_narrows_prompt_above_model_minimum(): assert filtered == [deployments[1]] +class _PinLookupCounter(DualCache): + def __init__(self) -> None: + super().__init__() + self.pin_lookups = 0 + + async def async_batch_get_cache( + self, + keys: list[str], + parent_otel_span: object = None, + local_only: bool = False, + throttle_redis: bool = True, + **kwargs: object, + ): + self.pin_lookups += 1 + return await super().async_batch_get_cache( + keys, + parent_otel_span=parent_otel_span, + local_only=local_only, + throttle_redis=throttle_redis, + **kwargs, + ) + + +@pytest.mark.asyncio +async def test_async_filter_deployments_skips_prefix_hash_for_a_single_deployment(): + """ + With one healthy deployment there is nothing to pin to, so the check must hand the group + back without hashing the prefix or probing the pin cache: on a 400k-token Claude Code + prompt that hash alone is ~30 ms of GIL-holding work per request. + """ + cache = _PinLookupCounter() + check = PromptCachingDeploymentCheck(cache=cache) + deployments = _deployments("anthropic/claude-opus-4-6") + messages = _messages(word_count=5000) + await PromptCachingCache(cache=cache).async_add_model_id(model_id="dep-1", messages=messages, tools=None) + + filtered = await check.async_filter_deployments( + model=MODEL_GROUP_ALIAS, healthy_deployments=deployments, messages=messages + ) + + assert filtered == deployments + assert cache.pin_lookups == 0 + + two = _deployments("anthropic/claude-opus-4-6", "anthropic/claude-opus-4-6") + assert await check.async_filter_deployments( + model=MODEL_GROUP_ALIAS, healthy_deployments=two, messages=messages + ) == [two[0]] + assert cache.pin_lookups == 1 + + @pytest.mark.asyncio async def test_async_filter_deployments_does_not_pin_when_target_order_is_set(): cache = DualCache() @@ -573,7 +623,7 @@ async def test_async_filter_deployments_counts_the_prompt_off_the_event_loop(): warm_tokenizer("anthropic/claude-fable-5") check = PromptCachingDeploymentCheck(cache=DualCache()) - deployments = _deployments("anthropic/claude-fable-5") + deployments = _deployments("anthropic/claude-fable-5", "anthropic/claude-fable-5") messages = cast(list[AllMessageValues], [{"role": "user", "content": text * 100}]) result, took, lags = await timed_with_loop_lags( diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index aeab488d270..924375a18b7 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -18991,3 +18991,131 @@ async def test_deployment_selection_runs_inside_a_route_phase_named_after_the_mo assert route_span.parent is not None and route_span.parent.span_id == server_span.get_span_context().span_id assert route_span.end_time is not None assert recording_cache.active_span_names and set(recording_cache.active_span_names) == {"route gpt-group"} + + +def _record_phase_events(monkeypatch: pytest.MonkeyPatch) -> list[tuple[str, dict[str, str | int]]]: + events: list[tuple[str, dict[str, str | int]]] = [] # mutable-ok: recorder for the injected phase_event double + + def record(name: str, attributes: dict[str, str | int]) -> None: + events.append((name, dict(attributes))) + + monkeypatch.setattr(litellm.router, "phase_event", record) + return events + + +def _pick(model_group: str, reason: str, attempt: int) -> tuple[str, dict[str, str | int]]: + return ( + "litellm.request.deployment_selected", + { + "litellm.deployment.attempt": attempt, + "litellm.deployment.reason": reason, + "litellm.deployment.model_group": model_group, + }, + ) + + +@pytest.mark.parametrize( + "request_kwargs, expected_reason, expected_attempt", + [ + (None, "initial", 1), + ({"metadata": {"attempted_retries": 0}, "fallback_depth": 0}, "initial", 1), + ({"metadata": {"attempted_retries": 2}}, "retry", 3), + ({"litellm_metadata": {"attempted_retries": 1}, "metadata": {"attempted_retries": 4}}, "retry", 2), + ({"metadata": {}, "fallback_depth": 1}, "fallback", 1), + ({"metadata": {"attempted_retries": 1}, "fallback_depth": 1}, "retry", 2), + ], +) +def test_deployment_pick_attributes_derive_attempt_and_reason( + request_kwargs: dict[str, object] | None, expected_reason: str, expected_attempt: int +): + attributes: Final = litellm.router._deployment_pick_attributes("gpt-4o", request_kwargs) + + assert dict(attributes) == { + "litellm.deployment.attempt": expected_attempt, + "litellm.deployment.reason": expected_reason, + "litellm.deployment.model_group": "gpt-4o", + } + + +@pytest.mark.asyncio +async def test_acompletion_marks_deployment_selected_once(monkeypatch: pytest.MonkeyPatch): + events: Final = _record_phase_events(monkeypatch) + router: Final = Router( + model_list=[ + { + "model_name": "gpt-4o", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake", "mock_response": "hi"}, + } + ] + ) + + await router.acompletion(model="gpt-4o", messages=[{"role": "user", "content": "hi"}]) + + assert events == [_pick("gpt-4o", "initial", 1)] + + +@pytest.mark.asyncio +async def test_acompletion_marks_every_retry_pick(monkeypatch: pytest.MonkeyPatch): + events: Final = _record_phase_events(monkeypatch) + router: Final = Router( + model_list=[ + { + "model_name": "flaky", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake", "mock_response": Exception("boom")}, + } + ], + num_retries=2, + retry_after=0, + ) + + with pytest.raises(Exception, match="boom"): + await router.acompletion(model="flaky", messages=[{"role": "user", "content": "hi"}]) + + assert events == [_pick("flaky", "initial", 1), _pick("flaky", "retry", 2), _pick("flaky", "retry", 3)] + + +@pytest.mark.asyncio +async def test_acompletion_marks_fallback_pick_with_its_model_group(monkeypatch: pytest.MonkeyPatch): + events: Final = _record_phase_events(monkeypatch) + router: Final = Router( + model_list=[ + { + "model_name": "primary", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake", "mock_response": Exception("boom")}, + }, + { + "model_name": "backup", + "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "fake", "mock_response": "hi"}, + }, + ], + fallbacks=[{"primary": ["backup"]}], + num_retries=0, + ) + + response: Final = await router.acompletion(model="primary", messages=[{"role": "user", "content": "hi"}]) + + assert response.choices[0].message.content == "hi" + assert events == [_pick("primary", "initial", 1), _pick("backup", "fallback", 1)] + + +@pytest.mark.asyncio +async def test_non_chat_surfaces_mark_their_deployment_pick(monkeypatch: pytest.MonkeyPatch): + """The event is emitted where the router picks, so embeddings and the sync path report it too.""" + events: Final = _record_phase_events(monkeypatch) + router: Final = Router( + model_list=[ + { + "model_name": "embed", + "litellm_params": {"model": "openai/text-embedding-3-small", "api_key": "fake", "mock_response": [0.1]}, + }, + { + "model_name": "gpt-4o", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake", "mock_response": "hi"}, + }, + ] + ) + + await router.aembedding(model="embed", input="hi") + router.completion(model="gpt-4o", messages=[{"role": "user", "content": "hi"}]) + + assert events == [_pick("embed", "initial", 1), _pick("gpt-4o", "initial", 1)] diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index 404fe4fa6f6..a72aec07766 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -4118,6 +4118,46 @@ def test_is_prompt_caching_valid_prompt_explicit_min_token_count_overrides_model ) +def test_is_prompt_caching_valid_prompt_stops_counting_once_the_minimum_is_reached( + local_model_cost_map: None, monkeypatch: pytest.MonkeyPatch +) -> None: + """Regression: the router's prompt-cache deployment check tokenized the whole 400k to 700k token + Claude Code conversation on every request only to compare it with a 1024-token minimum, which + sat on the request's wall clock between auth and the LLM call. The check must decide after the + first few messages and still agree with the full count on both sides of the minimum.""" + import litellm.litellm_core_utils.token_counter as token_counter_module + + long_prompt = PROMPT_CACHE_MESSAGES * 50 + counted_messages: list[int] = [] # mutable-ok: recorder for the _count_messages double + real_count_messages = token_counter_module._count_messages + + def counting(params, batch, use_default_image_token_count, default_token_count): + counted_messages.append(len(batch)) + return real_count_messages(params, batch, use_default_image_token_count, default_token_count) + + monkeypatch.setattr(token_counter_module, "_count_messages", counting) + + assert is_prompt_caching_valid_prompt(model="claude-opus-4-8", messages=long_prompt, min_token_count=1024) is True + assert sum(counted_messages) < len(long_prompt), sum(counted_messages) + + full_count = litellm.token_counter(model="claude-opus-4-8", messages=long_prompt, use_default_image_token_count=True) + assert ( + is_prompt_caching_valid_prompt(model="claude-opus-4-8", messages=long_prompt, min_token_count=full_count) + is True + ) + assert ( + is_prompt_caching_valid_prompt(model="claude-opus-4-8", messages=long_prompt, min_token_count=full_count + 1) + is False + ) + + +def test_is_prompt_caching_valid_prompt_without_messages_is_not_cacheable(local_model_cost_map: None) -> None: + """A tools-only call has no cacheable prefix, matching the pre-existing result for messages=None.""" + tools = [{"type": "function", "function": {"name": "f", "parameters": {"type": "object", "properties": {}}}}] + assert is_prompt_caching_valid_prompt(model="claude-opus-4-8", messages=None, tools=tools) is False + assert is_prompt_caching_valid_prompt(model="claude-opus-4-8", messages=None) is False + + def test_custom_logger_guards_ignore_subclass_instances(monkeypatch: pytest.MonkeyPatch) -> None: """Regression LIT-4392: the success/failure existence guards used isinstance, so a user subclass of a built-in logger already promoted into the callback lists made the guard