mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
797353f13a
commit
4c6c84afb7
17 changed files with 637 additions and 20 deletions
|
|
@ -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
|
||||
# ============================================================================
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"}]}'
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue