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 <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-03 09:20:28 -07:00 • committed by GitHub
parent 797353f13a
commit 4c6c84afb7
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
17 changed files with 637 additions and 20 deletions

View file

@ -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
# ============================================================================

View file

@ -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,

View file

@ -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()

View file

@ -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,

View file

@ -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

View file

@ -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.

View file

@ -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:

View file

@ -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(

View file

@ -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,

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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"}]}'

View file

@ -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(

View file

@ -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)]

View file

@ -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