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