Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_/patch-endpoint-1d2646

This commit is contained in:
Yuneng Jiang 2026-07-10 23:14:13 -07:00
commit 472b64cc62
No known key found for this signature in database
43 changed files with 2959 additions and 214 deletions

View file

@ -19,7 +19,7 @@ import os
import time
import traceback
from datetime import datetime as datetimeObj
from typing import Any, Dict, List, Optional, Union
from typing import Any, Dict, List, Optional, Sequence, Union
import httpx
from httpx import Response
@ -50,6 +50,7 @@ from litellm.types.integrations.base_health_check import IntegrationHealthCheckS
from litellm.types.integrations.datadog import (
DD_ERRORS,
DD_MAX_BATCH_SIZE,
DD_MAX_PAYLOAD_SIZE_BYTES,
DataDogStatus,
DatadogInitParams,
DatadogPayload,
@ -384,8 +385,10 @@ class DataDogLogger(
async def _send_with_413_split(self, batch: List) -> List:
"""
Send a batch, halving any sub-batch that 413s (payload too large) and retrying the
halves, since Datadog enforces a 5MB uncompressed limit per request.
Send a batch, halving any sub-batch that exceeds Datadog's intake limits before
sending, and halving again on a 413 (payload too large) response, since Datadog
enforces a 5MB uncompressed limit per request. The proactive split avoids paying
a serialize + gzip + round trip for a payload the intake is guaranteed to reject.
A 413 surfaces as a raised MaskedHTTPStatusError (httpx raise_for_status), not a
returned response, so both paths are handled. A lone event that still 413s is
@ -398,6 +401,11 @@ class DataDogLogger(
chunk = pending.pop()
if not chunk:
continue
if len(chunk) > 1 and self._exceeds_intake_limits(chunk):
mid = len(chunk) // 2
pending.append(chunk[mid:])
pending.append(chunk[:mid])
continue
try:
response = await self.async_send_compressed_data(chunk)
except Exception as e:
@ -436,6 +444,21 @@ class DataDogLogger(
def _undelivered(chunk: List, pending: List[List]) -> List:
return chunk + [event for remaining in reversed(pending) for event in remaining]
@staticmethod
def _exceeds_intake_limits(chunk: Sequence[DatadogPayload]) -> bool:
"""
True when a chunk would breach Datadog's log intake limits: more than
DD_MAX_BATCH_SIZE events per payload, or a serialized size above
DD_MAX_PAYLOAD_SIZE_BYTES (held under Datadog's 5MB uncompressed cap so
the batch is split before the intake rejects it with a 413).
"""
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
if len(chunk) > DD_MAX_BATCH_SIZE:
return True
payload_size_bytes = len(safe_dumps(chunk).encode("utf-8"))
return payload_size_bytes > DD_MAX_PAYLOAD_SIZE_BYTES
async def flush_queue(self):
if self.flush_lock is None:
return

View file

@ -1618,6 +1618,14 @@ class PrometheusLogger(CustomLogger):
user_id: Optional[str] = None,
user_api_key_org_id: Optional[str] = None,
):
if (
isinstance(self.litellm_remaining_team_budget_metric, NoOpMetric)
and isinstance(self.litellm_remaining_api_key_budget_metric, NoOpMetric)
and isinstance(self.litellm_remaining_user_budget_metric, NoOpMetric)
and isinstance(self.litellm_remaining_org_budget_metric, NoOpMetric)
):
return
_metadata = litellm_params.get("metadata") or {}
_team_spend = _metadata.get("user_api_key_team_spend", None)
_team_max_budget = _metadata.get("user_api_key_team_max_budget", None)
@ -3332,6 +3340,9 @@ class PrometheusLogger(CustomLogger):
- looks up team info from db if not available in metadata
- Set team budget metrics
"""
if isinstance(self.litellm_remaining_team_budget_metric, NoOpMetric):
return
if user_api_team:
team_object = await self._assemble_team_object(
team_id=user_api_team,
@ -3453,6 +3464,9 @@ class PrometheusLogger(CustomLogger):
- Fetches org info via cache (get_org_object)
- Sets org budget metrics
"""
if isinstance(self.litellm_remaining_org_budget_metric, NoOpMetric):
return
if not org_id:
return
@ -3582,6 +3596,9 @@ class PrometheusLogger(CustomLogger):
key_max_budget: Optional[float],
key_spend: Optional[float],
):
if isinstance(self.litellm_remaining_api_key_budget_metric, NoOpMetric):
return
if user_api_key:
user_api_key_dict = await self._assemble_key_object(
user_api_key=user_api_key,
@ -3642,6 +3659,9 @@ class PrometheusLogger(CustomLogger):
- looks up user info from db if not available in metadata
- Set user budget metrics
"""
if isinstance(self.litellm_remaining_user_budget_metric, NoOpMetric):
return
if user_id:
user_object = await self._assemble_user_object(
user_id=user_id,

View file

@ -8,9 +8,12 @@ The metadata is a partial cost-map entry: ``litellm_provider`` drives provider
routing, and the remaining fields (``mode``, ``supports_*``, context window,
pricing, ...) drive ``get_model_info`` / ``supports_*``.
Precedence: rules are evaluated in file order and the first match wins. They are
consulted only after exact and case-insensitive lookups miss, so an exact entry
always takes precedence over a rule.
Precedence: rules are evaluated in file order and the first match wins. Callers
with extra constraints (model-info resolution checks the provider) use
``match_all_fallback_generalizations`` to skip inapplicable earlier rules instead
of discarding the model name. Rules are consulted only after exact and
case-insensitive lookups miss, so an exact entry always takes precedence over a
rule.
Patterns are matched case-insensitively with ``re.search`` and are not implicitly
anchored: a rule must include ``^`` and ``$`` (as the shipped rules do) to bind to
@ -105,15 +108,15 @@ class _FallbackGeneralizations:
)
return compiled
def match(self, model: str) -> Optional[dict]:
def matches(self, model: str) -> list[dict]:
if not model:
return None
return []
if self._compiled is None:
self._compiled = self._compile()
for pattern, model_info in self._compiled:
if pattern.search(model) is not None:
return dict(model_info)
return None
return [dict(model_info) for pattern, model_info in self._compiled if pattern.search(model) is not None]
def match(self, model: str) -> Optional[dict]:
return next(iter(self.matches(model)), None)
_registry = _FallbackGeneralizations()
@ -139,3 +142,12 @@ def match_fallback_generalization(model: str) -> Optional[dict]:
O(number of rules). Only call this once exact lookups have missed.
"""
return _registry.match(model)
def match_all_fallback_generalizations(model: str) -> list[dict]:
"""Return the ``model_info`` of every rule whose regex matches ``model``, in rule order.
Lets a caller with extra constraints (e.g. a provider match) skip an
inapplicable earlier rule instead of discarding the whole candidate.
"""
return _registry.matches(model)

View file

@ -289,6 +289,13 @@ class AnthropicModelInfo(BaseLLMModelInfo):
status_code=400,
)
@staticmethod
def _strip_version_suffix(model: str) -> str:
at = model.rfind("@")
if at > 0:
return model[:at]
return model
@staticmethod
def _model_map_lookup_candidates(model: str) -> List[str]:
"""Model-map keys to try for ``model``: the id itself, the same id with a
@ -324,6 +331,7 @@ class AnthropicModelInfo(BaseLLMModelInfo):
_DATED_RELEASE_SUFFIX_RE.sub("", cand),
_DOTTED_VERSION_RE.sub(r"\1-\2", cand),
_strip_bedrock_id_suffixes(cand),
AnthropicModelInfo._strip_version_suffix(cand),
)
)
return list(dict.fromkeys((*primary, *normalized)))

View file

@ -3,6 +3,7 @@ from typing import Any, AsyncIterator, Dict, List, Optional, Tuple
import httpx
from litellm.constants import (
ANTHROPIC_MIN_THINKING_BUDGET_TOKENS,
DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET,
@ -32,6 +33,12 @@ from ...common_utils import (
DEFAULT_ANTHROPIC_API_VERSION = "2023-06-01"
DROP_UNSUPPORTED_ADAPTIVE_EFFORT_WARNING = (
"Dropping adaptive `thinking`/`output_config.effort` for model=%s: the model "
"does not support extended thinking, or max_tokens is too small to fit the "
"minimum thinking budget."
)
class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
def get_supported_anthropic_messages_params(self, model: str) -> list:
@ -253,6 +260,111 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
existing_output_config.setdefault("effort", effort)
optional_params["output_config"] = existing_output_config
@staticmethod
def _translate_adaptive_effort_for_non_adaptive_model(
model: str, optional_params: Dict, max_tokens: Optional[int]
) -> None:
"""Translate the 4.6+ adaptive-thinking interface (``thinking.type=adaptive``
and/or ``output_config.effort``) down to what an older Anthropic model
supports. Clients like Claude Code send this interface unconditionally, so
without translation it reaches a pre-4.6 model and Anthropic rejects it with
"This model does not support the effort parameter".
The reshape is silent, matching how the messages path already strips
unsupported ``output_config`` for older models (bedrock invoke, issue
#22797): the goal is to keep the request working, not to fail it.
``thinking.type=adaptive`` and ``output_config.effort`` are independent
capabilities. Adaptive thinking needs ``supports_adaptive_thinking`` (4.6+);
``output_config.effort`` needs ``supports_output_config``, which some
non-adaptive models (e.g. Claude Opus 4.5) advertise on its own. So the two
are handled separately:
- Adaptive-thinking models (4.6+): both are native, left untouched.
- ``supports_output_config`` but non-adaptive (Opus 4.5): keep
``output_config.effort`` (native), only drop the unsupported adaptive
``thinking`` block. When adaptive thinking is being dropped and the
effort level itself isn't supported by the model (e.g. ``xhigh``/``max``
on Opus 4.5, which only accepts low/medium/high, while ``xhigh`` is
Claude Code's default), fall through to the legacy translation below
instead of forwarding a level Anthropic would reject. Effort-only
requests are always left untouched: provider subclasses own their level
normalization (bedrock clamps ``xhigh`` to the model's ceiling after
this base transform runs).
- Thinking-capable but neither (``supports_reasoning``, e.g. Haiku/Sonnet
4.5): map effort to legacy ``thinking={type: enabled, budget_tokens}`` via
``AnthropicConfig._map_reasoning_effort``, capped below ``max_tokens``
(Anthropic requires ``max_tokens > budget_tokens``) and dropped when
``max_tokens`` can't fit even the minimum budget.
- No reasoning support: ``thinking`` is dropped.
For the last two, only the consumed ``effort`` key is removed from
``output_config``; any residual (e.g. ``format``) is left for provider
subclasses to handle.
"""
from litellm.exceptions import BadRequestError as _BadRequestError
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
if AnthropicConfig._is_adaptive_thinking_model(model):
return
output_config = optional_params.get("output_config")
thinking = optional_params.get("thinking")
effort = output_config.get("effort") if isinstance(output_config, dict) else None
adaptive_thinking = isinstance(thinking, dict) and thinking.get("type") == "adaptive"
if effort is None and not adaptive_thinking:
return
if AnthropicConfig._model_supports_effort_param(model) and (
not adaptive_thinking or AnthropicConfig._validate_effort_for_model(model, effort) is None
):
if adaptive_thinking:
optional_params.pop("thinking", None)
return
supports_thinking = AnthropicModelInfo._supports_model_capability(model, "supports_reasoning")
try:
legacy_thinking = (
AnthropicConfig._map_reasoning_effort(reasoning_effort=effort or "medium", model=model)
if supports_thinking
else None
)
except _BadRequestError as e:
raise AnthropicError(message=str(e.message), status_code=400)
capped_thinking = (
AnthropicMessagesConfig._cap_thinking_budget_to_max_tokens(legacy_thinking, max_tokens)
if legacy_thinking is not None
else None
)
if capped_thinking is not None:
optional_params["thinking"] = capped_thinking
else:
verbose_logger.warning(DROP_UNSUPPORTED_ADAPTIVE_EFFORT_WARNING, model)
optional_params.pop("thinking", None)
if isinstance(output_config, dict) and "effort" in output_config:
residual = {k: v for k, v in output_config.items() if k != "effort"}
if residual:
optional_params["output_config"] = residual
else:
optional_params.pop("output_config", None)
@staticmethod
def _cap_thinking_budget_to_max_tokens(thinking: Dict, max_tokens: Optional[int]) -> Optional[Dict]:
"""Cap a legacy ``thinking.budget_tokens`` below ``max_tokens`` (Anthropic
requires ``max_tokens > budget_tokens``). Returns the (possibly capped)
thinking dict, or ``None`` when ``max_tokens`` is too small to fit even the
minimum thinking budget and thinking should be dropped."""
budget = thinking.get("budget_tokens")
if max_tokens is None or not isinstance(budget, int):
return thinking
if max_tokens <= ANTHROPIC_MIN_THINKING_BUDGET_TOKENS:
return None
if budget < max_tokens:
return thinking
return {**thinking, "budget_tokens": max_tokens - 1}
def transform_anthropic_messages_request(
self,
model: str,
@ -284,6 +396,12 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
optional_params=anthropic_messages_optional_request_params,
)
self._translate_adaptive_effort_for_non_adaptive_model(
model=model,
optional_params=anthropic_messages_optional_request_params,
max_tokens=max_tokens,
)
system_param = anthropic_messages_optional_request_params.get("system")
if self.should_strip_billing_metadata() and system_param is not None:
filtered_system = self._filter_billing_headers_from_system(system_param)

View file

@ -93,31 +93,48 @@ class AmazonAnthropicClaudeMessagesConfig(
return [{"type": "text", "text": value}]
return [value]
def _normalize_system_role_messages_for_bedrock(self, anthropic_messages_request: dict) -> None:
"""Bedrock Invoke rejects a conversation that opens with ``role: "system"``
entries inside ``messages`` ("messages.0: use the top-level 'system'
parameter for the initial system prompt"); Anthropic Messages carries that
content in the top-level ``system`` field, so hoist the leading run of
system entries there. Mid-conversation system entries (e.g. Claude Code's
``mid-conversation-system-2026-04-07`` reminders) are accepted by Invoke in
place and MUST stay in place: hoisting one mutates the ``system`` prefix
and invalidates the prompt cache for the entire message history.
@staticmethod
def _is_system_role_message(message: Any) -> bool:
return isinstance(message, dict) and message.get("role") == "system"
def _normalize_system_role_messages_for_bedrock(self, anthropic_messages_request: dict, model: str) -> None:
"""Bedrock Invoke validates ``role: "system"`` entries inside ``messages``
per model. Models carrying ``supports_mid_conversation_system`` in the
cost map (the Opus 4.8 family) only reject a leading run ("messages.0:
use the top-level 'system' parameter for the initial system prompt") and
accept mid-conversation entries (e.g. Claude Code's
``mid-conversation-system-2026-04-07`` reminders) in place, where they
MUST stay: hoisting one mutates the ``system`` prefix and invalidates the
prompt cache for the entire message history. Older Claude models (Opus
4.7, Sonnet 4.6, Haiku 4.5, ...) reject the role in every position
("role 'system' is not supported on this model"), so without the flag
every system entry is hoisted into the top-level ``system`` field.
Billing-header system blocks are stripped from the top-level ``system``
field regardless of whether anything was hoisted."""
messages = anthropic_messages_request.get("messages")
if not isinstance(messages, list):
return
leading_count = next(
(i for i, m in enumerate(messages) if not (isinstance(m, dict) and m.get("role") == "system")),
len(messages),
)
if leading_count:
anthropic_messages_request["messages"] = messages[leading_count:]
if _supports_factory(
model=model,
custom_llm_provider="bedrock",
key="supports_mid_conversation_system",
):
leading_count = next(
(i for i, m in enumerate(messages) if not self._is_system_role_message(m)),
len(messages),
)
hoisted = messages[:leading_count]
remaining = messages[leading_count:]
else:
hoisted = [m for m in messages if self._is_system_role_message(m)]
remaining = [m for m in messages if not self._is_system_role_message(m)]
if hoisted:
anthropic_messages_request["messages"] = remaining
system_content = [
block
for source in (
anthropic_messages_request.get("system"),
*(m.get("content") for m in messages[:leading_count]),
*(m.get("content") for m in hoisted),
)
for block in self._as_system_content_blocks(source)
]
@ -674,7 +691,7 @@ class AmazonAnthropicClaudeMessagesConfig(
litellm_params=litellm_params,
headers=headers,
)
self._normalize_system_role_messages_for_bedrock(anthropic_messages_request)
self._normalize_system_role_messages_for_bedrock(anthropic_messages_request, model=model)
#########################################################
############## BEDROCK Invoke SPECIFIC TRANSFORMATION ###
#########################################################

View file

@ -1359,6 +1359,7 @@
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1393,6 +1394,7 @@
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1427,6 +1429,7 @@
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1461,6 +1464,7 @@
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1481,6 +1485,7 @@
"anthropic.claude-opus-4-8": {
"bedrock_converse_supports_strict_tools": false,
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -1516,6 +1521,7 @@
"global.anthropic.claude-opus-4-8": {
"bedrock_converse_supports_strict_tools": false,
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -1551,6 +1557,7 @@
"us.anthropic.claude-opus-4-8": {
"bedrock_converse_supports_strict_tools": false,
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1586,6 +1593,7 @@
"eu.anthropic.claude-opus-4-8": {
"bedrock_converse_supports_strict_tools": false,
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1621,6 +1629,43 @@
"au.anthropic.claude-opus-4-8": {
"bedrock_converse_supports_strict_tools": false,
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
"input_cost_per_token": 5.5e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.75e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_sampling_params": false,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_native_structured_output": true,
"supports_max_reasoning_effort": true,
"supports_output_config": true,
"bedrock_output_config_effort_ceiling": "xhigh",
"supports_parallel_tool_use_config": true
},
"jp.anthropic.claude-opus-4-8": {
"bedrock_converse_supports_strict_tools": false,
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1703,6 +1748,7 @@
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1737,6 +1783,7 @@
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1771,6 +1818,7 @@
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1805,6 +1853,7 @@
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1839,6 +1888,7 @@
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1873,6 +1923,7 @@
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -44998,6 +45049,17 @@
},
"fallback_generalizations": {
"rules": [
{
"name": "bedrock-anthropic-claude-mid-conversation-system",
"pattern": "anthropic\\.claude-[a-z]+-(?:4[-._](?:[89]|[1-9]\\d)(?!\\d)|(?:[5-9]|[1-9]\\d)(?!\\d)(?:[-._]\\d{1,2}(?!\\d))?)",
"description": "Bedrock Invoke ids for Claude 4.8 or higher: anthropic.claude-<family> with minor 4.8 through 4.99, any 5.x or later major-minor, or a bare 5+ major, which also admits new families such as fable. These models accept mid-conversation role system messages in place (verified live on Opus 4.8, Sonnet 5 and Fable 5), so unmapped future Bedrock Claudes keep the cache-preserving in-place handling instead of the hoist-all default. Listed first so bare-id provider inference, which takes the first pattern hit, resolves these Bedrock ids to bedrock; model-info resolution skips provider-mismatched rules either way.",
"extends": "anthropic-claude",
"model_info": {
"litellm_provider": "bedrock",
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true
}
},
{
"name": "anthropic-claude-adaptive-thinking",
"pattern": "(?:opus|sonnet|haiku)[-._](?:4[-._](?:[6-9]|[1-9]\\d)(?!\\d)|(?:[5-9]|[1-9]\\d{1,})[-._]\\d{1,2}(?!\\d))",

View file

@ -76,6 +76,10 @@ def cli(ctx: click.Context, base_url: str, api_key: Optional[str]) -> None:
"""LiteLLM Proxy CLI - Manage your LiteLLM proxy server"""
ctx.ensure_object(dict)
# Normalize once here so every downstream command (login, agents, http, ...) can safely
# do f"{base_url}/some/path" without producing a double slash.
base_url = base_url.rstrip("/")
# If no API key provided via flag or environment variable, try to load from saved token.
# Pass base_url so we only use the stored key when it was issued for this server.
if api_key is None:

View file

@ -2020,8 +2020,6 @@ class LiteLLMCompletionResponsesConfig:
output_details_dict: dict[str, int] = {}
if hasattr(completion_details, "reasoning_tokens") and completion_details.reasoning_tokens is not None:
output_details_dict["reasoning_tokens"] = completion_details.reasoning_tokens
else:
output_details_dict["reasoning_tokens"] = 0
if hasattr(completion_details, "text_tokens") and completion_details.text_tokens is not None:
output_details_dict["text_tokens"] = completion_details.text_tokens

View file

@ -127,7 +127,7 @@ def mock_responses_api_response(
"input_tokens": 36,
"input_tokens_details": {"cached_tokens": 0},
"output_tokens": 87,
"output_tokens_details": {"reasoning_tokens": 0},
"output_tokens_details": {},
"total_tokens": 123,
},
"user": None,

View file

@ -17,6 +17,7 @@ from litellm.constants import (
LITELLM_MAX_STREAMING_DURATION_SECONDS,
STREAM_SSE_DONE_STRING,
)
from litellm.exceptions import MidStreamFallbackError
from litellm.litellm_core_utils.asyncify import run_async_function
from litellm.litellm_core_utils.core_helpers import process_response_headers
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -26,7 +27,7 @@ from litellm.litellm_core_utils.llm_response_utils.response_metadata import (
)
from litellm.litellm_core_utils.thread_pool_executor import executor
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.responses.utils import ResponseAPILoggingUtils, ResponsesAPIRequestUtils
from litellm.types.llms.openai import ResponsesAPIStreamEvents
from litellm.types.utils import CallTypes
from litellm.utils import async_post_call_success_deployment_hook
@ -47,6 +48,44 @@ def _log_background_task_failure(task: "asyncio.Task[Any]", *, task_name: str) -
verbose_logger.error("%s failed: %s", task_name, exception)
_CLIENT_ERROR_CODES: frozenset[str] = frozenset(
(
"invalid_request_error",
"context_length_exceeded",
"content_policy_violation",
"model_not_found",
)
)
def _error_event_fields(error_obj: object) -> tuple[str, Optional[str], Optional[str]]:
if isinstance(error_obj, dict):
raw_message = error_obj.get("message")
raw_type = error_obj.get("type")
raw_code = error_obj.get("code")
elif error_obj is not None:
raw_message = getattr(error_obj, "message", None)
raw_type = getattr(error_obj, "type", None)
raw_code = getattr(error_obj, "code", None)
else:
raw_message = None
raw_type = None
raw_code = None
message = str(raw_message) if raw_message is not None else "Response API in-stream error"
error_type = raw_type if isinstance(raw_type, str) else None
code = raw_code if isinstance(raw_code, str) else None
return message, error_type, code
def _status_code_for_error_fields(error_type: Optional[str], error_code: Optional[str]) -> int:
fields = tuple(field for field in (error_type, error_code) if field is not None)
if any(field.startswith("rate_limit") or field == "insufficient_quota" for field in fields):
return 429
if any(field in _CLIENT_ERROR_CODES for field in fields):
return 400
return 500
class BaseResponsesAPIStreamingIterator:
"""
Base class for streaming iterators that process responses from the Responses API.
@ -73,6 +112,8 @@ class BaseResponsesAPIStreamingIterator:
self.completed_response: Optional[Any] = None
self.start_time = getattr(logging_obj, "start_time", datetime.now())
self._failure_handled = False # Track if failure handler has been called
self._yielded_first_chunk = False
self._generated_content = ""
self._completed_response_cached = False
self._completed_response_logged = False
self._completed_response_cache_hit: Optional[bool] = None
@ -160,6 +201,10 @@ class BaseResponsesAPIStreamingIterator:
# Encode container_id on streaming events so proxy/UI follow-ups route correctly
_event_type = getattr(openai_responses_api_chunk, "type", None)
if _event_type == ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA:
_delta = getattr(openai_responses_api_chunk, "delta", None)
if isinstance(_delta, str):
self._generated_content += _delta
_stream_model_id = (
self.litellm_metadata.get("model_info", {}).get("id") if self.litellm_metadata else None
)
@ -327,17 +372,66 @@ class BaseResponsesAPIStreamingIterator:
"""
response_obj = getattr(self.completed_response, "response", None) if self.completed_response else None
error_info = getattr(response_obj, "error", None) if response_obj else None
error_message = "Response failed"
if isinstance(error_info, dict):
error_message = error_info.get("message", str(error_info))
error_message, error_type, error_code = _error_event_fields(error_info)
self._record_failed_response_usage(response_obj)
exception = litellm.APIError(
status_code=500,
status_code=_status_code_for_error_fields(error_type, error_code),
message=error_message,
llm_provider=self.custom_llm_provider or "",
model=self.model or "",
)
self._handle_failure(exception)
def _record_failed_response_usage(self, response_obj: Optional[Any]) -> None:
if response_obj is None or self.logging_obj is None:
return
usage_obj = getattr(response_obj, "usage", None)
if usage_obj is None:
return
try:
self.logging_obj.model_call_details["combined_usage_object"] = (
ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage_obj)
)
except (TypeError, ValueError) as usage_error:
verbose_logger.debug(
"could not record usage for failed responses stream: %s",
usage_error,
)
return
self.logging_obj.model_call_details["response_cost"] = (
self.logging_obj._response_cost_calculator(result=response_obj) or 0.0
)
def _maybe_raise_for_error_event(self, result: object) -> None:
chunk_type = getattr(result, "type", None)
if chunk_type not in ("error", "response.failed"):
return
error_obj: object = (
getattr(getattr(result, "response", None), "error", None)
if chunk_type == "response.failed"
else getattr(result, "error", None)
)
error_message, error_type, error_code = _error_event_fields(error_obj)
status_code = _status_code_for_error_fields(error_type, error_code)
mapped_exception = litellm.APIError(
status_code=status_code,
message=error_message,
llm_provider=self.custom_llm_provider or "",
model=self.model or "",
)
if 400 <= status_code < 500 and status_code != 429:
raise mapped_exception
raise MidStreamFallbackError(
message=str(mapped_exception),
model=self.model or "",
llm_provider=self.custom_llm_provider or "",
original_exception=mapped_exception,
generated_content=self._generated_content,
is_pre_first_chunk=not self._yielded_first_chunk,
)
def _get_completed_response_object(self) -> Optional[Any]:
openai_types = _get_openai_response_types()
completed_response = self.completed_response
@ -611,11 +705,13 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
if self.finished:
raise StopAsyncIteration
elif result is not None:
self._maybe_raise_for_error_event(result)
# Await hook directly instead of run_async_function
# (which spawns a thread + event loop per call)
result = await self._call_post_streaming_deployment_hook(
chunk=result,
)
self._yielded_first_chunk = True
return result
# If result is None, continue the loop to get the next chunk
@ -685,11 +781,13 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
if self.finished:
raise StopIteration
elif result is not None:
self._maybe_raise_for_error_event(result)
# Sync path: use run_async_function for the hook
result = run_async_function(
async_function=self._call_post_streaming_deployment_hook,
chunk=result,
)
self._yielded_first_chunk = True
return result
# If result is None, continue the loop to get the next chunk

View file

@ -2344,6 +2344,8 @@ class Router:
self.completed_response = None
self.start_time = getattr(source_iterator, "start_time", datetime.now())
self._failure_handled = False
self._yielded_first_chunk = False
self._generated_content = ""
self._completed_response_cached = False
self._completed_response_logged = False
self._completed_response_cache_hit = None

View file

@ -4,16 +4,21 @@ Complexity-based Auto Router
A rule-based routing strategy that uses weighted scoring across multiple dimensions
to classify requests by complexity and route them to appropriate models.
No external API calls - all scoring is local and <1ms.
By default, scoring is local (regex/keyword-based) with no external API calls and <1ms
latency. Optionally, classifier_type="llm" routes classification through a configured
model instead, trading that latency/cost guarantee for potentially better accuracy.
Inspired by ClawRouter: https://github.com/BlockRunAI/ClawRouter
"""
import re
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union
from pydantic import BaseModel
from litellm._logging import verbose_router_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.utils import ModelResponse
from .config import (
DEFAULT_CODE_KEYWORDS,
@ -32,6 +37,24 @@ else:
PreRoutingHookResponse = Any
class TierClassification(BaseModel):
"""Structured response schema for the LLM-based complexity classifier."""
tier: Literal["SIMPLE", "MEDIUM", "COMPLEX", "REASONING"]
_CLASSIFICATION_PROMPT_TEMPLATE = """Classify the complexity of the following user request into exactly one tier.
Tiers:
- SIMPLE: factual lookups, greetings, short direct questions with no reasoning or code involved.
- MEDIUM: everyday requests needing some explanation or minor code/technical content.
- COMPLEX: requests involving non-trivial code, architecture, or multi-step technical work.
- REASONING: requests explicitly requiring step-by-step reasoning, analysis, or weighing tradeoffs.
{system_context}Request:
{prompt}"""
def _append_custom_keywords(base_keywords: list[str], custom_keywords: Optional[list[str]]) -> list[str]:
if not custom_keywords:
return base_keywords
@ -40,6 +63,20 @@ def _append_custom_keywords(base_keywords: list[str], custom_keywords: Optional[
return [*base_keywords, *deduped_custom.values()]
# Metadata keys that carry the parent request's budget reservation. These must not
# reach the classifier's internal acompletion call: the reservation belongs to the
# routed completion that the classifier is deciding on, not to the classifier call
# itself, and forwarding it would let the classifier's cost-tracking reconcile
# against a reservation it isn't responsible for.
_BUDGET_RESERVATION_METADATA_KEYS = frozenset({"user_api_key_budget_reservation", "user_api_key_auth"})
def _classifier_call_metadata(metadata: Optional[dict[str, Any]]) -> Optional[dict[str, Any]]:
if not metadata:
return metadata
return {k: v for k, v in metadata.items() if k not in _BUDGET_RESERVATION_METADATA_KEYS}
class DimensionScore:
"""Represents a score for a single dimension with optional signal."""
@ -53,10 +90,10 @@ class DimensionScore:
class ComplexityRouter(CustomLogger):
"""
Rule-based complexity router that classifies requests and routes to appropriate models.
Complexity router that classifies requests and routes to appropriate models.
Handles requests in <1ms with zero external API calls by using weighted scoring
across multiple dimensions:
By default, handles requests in <1ms with zero external API calls, using weighted
scoring across multiple dimensions:
- Token count (short=simple, long=complex)
- Code presence (code keywords → complex)
- Reasoning markers ("step by step", "think through" → reasoning tier)
@ -297,6 +334,63 @@ class ComplexityRouter(CustomLogger):
return tier, weighted_score, signals
async def aclassify(
self,
prompt: str,
system_prompt: Optional[str] = None,
request_kwargs: Optional[dict[str, Any]] = None,
) -> tuple[ComplexityTier, float, list[str]]:
"""
Classify a prompt by complexity, using the LLM classifier when configured.
Falls back to the local heuristic scorer if classifier_type is "heuristic",
or if the LLM call fails, times out, or returns an unparseable response.
"""
if self.config.classifier_type != "llm" or self.config.classifier_llm_config is None:
return self.classify(prompt, system_prompt)
try:
tier = await self._classify_with_llm(prompt, system_prompt, request_kwargs)
return tier, 1.0, [f"llm-classifier:{tier.value}"]
except Exception as e: # noqa: BLE001 -- external LLM call can fail in many distinct ways (timeout, provider error, validation, parse error); any failure must fall back to the heuristic scorer
verbose_router_logger.warning(
f"ComplexityRouter: LLM classifier failed ({e}), falling back to heuristic scoring"
)
return self.classify(prompt, system_prompt)
async def _classify_with_llm(
self,
prompt: str,
system_prompt: Optional[str] = None,
request_kwargs: Optional[dict[str, Any]] = None,
) -> ComplexityTier:
"""Call the configured classifier model and parse its structured tier response."""
llm_config = self.config.classifier_llm_config
if llm_config is None:
raise ValueError("classifier_llm_config is not set")
system_context = f"Context: {system_prompt}\n\n" if system_prompt else ""
classification_prompt = _CLASSIFICATION_PROMPT_TEMPLATE.format(system_context=system_context, prompt=prompt)
# Forward the original request's metadata so the classifier call's spend is
# attributed to the calling key/team instead of being dropped. Excludes the
# parent request's budget reservation, which the routed completion (not this
# internal classifier call) is responsible for reconciling.
metadata = _classifier_call_metadata((request_kwargs or {}).get("litellm_metadata"))
response: ModelResponse = await self.litellm_router_instance.acompletion(
model=llm_config.model,
messages=[{"role": "user", "content": classification_prompt}],
response_format=TierClassification,
timeout=llm_config.timeout_ms / 1000,
metadata=metadata,
)
content = response.choices[0].message.content
if not content:
raise ValueError("LLM classifier returned empty content")
result = TierClassification.model_validate_json(content)
return ComplexityTier[result.tier]
def get_model_for_tier(self, tier: ComplexityTier) -> str:
"""
Get the model name for a given complexity tier.
@ -445,7 +539,7 @@ class ComplexityRouter(CustomLogger):
messages=messages if has_original_messages else None,
)
tier, score, signals = self.classify(user_message, system_prompt)
tier, score, signals = await self.aclassify(user_message, system_prompt, request_kwargs)
routed_model = self.get_model_for_tier(tier)
verbose_router_logger.info(

View file

@ -6,9 +6,9 @@ All values are configurable via proxy config.yaml.
"""
from enum import Enum
from typing import Dict, List, Optional
from typing import Dict, List, Literal, Optional
from pydantic import BaseModel, ConfigDict, Field
from pydantic import BaseModel, ConfigDict, Field, model_validator
class ComplexityTier(str, Enum):
@ -197,6 +197,18 @@ DEFAULT_TIER_MODELS: Dict[str, str] = {
}
class ClassifierLLMConfig(BaseModel):
"""Configuration for the LLM-based complexity classifier."""
model: str = Field(
description="Model name (from the router's model_list) to call for classification",
)
timeout_ms: int = Field(
default=3000,
description="Timeout budget for the classification call, in milliseconds",
)
class ComplexityRouterConfig(BaseModel):
"""Configuration for the ComplexityRouter."""
@ -257,8 +269,24 @@ class ComplexityRouterConfig(BaseModel):
description="Default model to use if tier cannot be determined",
)
# Classifier strategy
classifier_type: Literal["heuristic", "llm"] = Field(
default="heuristic",
description="Classification strategy: local regex/keyword scoring, or an LLM call",
)
classifier_llm_config: Optional[ClassifierLLMConfig] = Field(
default=None,
description="Configuration for the LLM classifier; required when classifier_type is 'llm'",
)
model_config = ConfigDict(extra="allow") # Allow additional fields
@model_validator(mode="after")
def _validate_llm_classifier_config(self) -> "ComplexityRouterConfig":
if self.classifier_type == "llm" and self.classifier_llm_config is None:
raise ValueError("classifier_llm_config is required when classifier_type is 'llm'")
return self
# Combined default config
DEFAULT_COMPLEXITY_CONFIG = ComplexityRouterConfig()

View file

@ -6,6 +6,7 @@ from typing_extensions import NotRequired, TypedDict
from litellm.types.integrations.custom_logger import StandardCustomLoggerInitParams
DD_MAX_BATCH_SIZE = 1000
DD_MAX_PAYLOAD_SIZE_BYTES = 4_000_000
class DataDogStatus(str, Enum):

View file

@ -1185,7 +1185,7 @@ class ResponsesAPIRequestParams(ResponsesAPIOptionalRequestParams, total=False):
class OutputTokensDetails(BaseLiteLLMOpenAIResponseObject):
reasoning_tokens: int = 0
reasoning_tokens: Optional[int] = None
text_tokens: Optional[int] = None
@ -1720,7 +1720,7 @@ class ErrorEventError(BaseLiteLLMOpenAIResponseObject):
type: str # e.g., 'invalid_request_error'
code: str # e.g., 'context_length_exceeded'
message: str
param: Optional[str] = None
param: Optional[Union[str, Dict[str, Any]]] = None
class ErrorEvent(BaseLiteLLMOpenAIResponseObject):

View file

@ -143,6 +143,7 @@ class ProviderSpecificModelInfo(TypedDict, total=False):
supports_web_search: Optional[bool]
supports_reasoning: Optional[bool]
supports_adaptive_thinking: Optional[bool]
supports_mid_conversation_system: Optional[bool]
supports_url_context: Optional[bool]
supports_none_reasoning_effort: Optional[bool]
supports_minimal_reasoning_effort: Optional[bool]

View file

@ -61,7 +61,7 @@ from litellm._lazy_imports import (
)
from litellm._uuid import uuid
from litellm.litellm_core_utils.fallback_generalizations import (
match_fallback_generalization,
match_all_fallback_generalizations,
)
from litellm.constants import (
DEFAULT_CHAT_COMPLETION_PARAM_VALUES,
@ -5046,9 +5046,10 @@ def _get_model_info_from_generalization(
"""Resolve an unmapped model via a declarative fallback-generalization rule.
Tries the same name candidates as the exact lookups, in the same order, and
returns ``(matched_name, model_info)`` for the first candidate whose rule also
satisfies the provider constraint. O(number of rules); only call after the
exact lookups have missed.
returns ``(matched_name, model_info)`` for the first matching rule that also
satisfies the provider constraint; a rule scoped to another provider is
skipped in favor of later rules rather than discarding the candidate.
O(number of rules); only call after the exact lookups have missed.
"""
candidates = [
potential_model_names["combined_model_name"],
@ -5058,11 +5059,9 @@ def _get_model_info_from_generalization(
potential_model_names["stripped_model_name"],
]
for candidate in candidates:
generalized_info = match_fallback_generalization(candidate)
if generalized_info is not None and _check_provider_match(
model_info=generalized_info, custom_llm_provider=custom_llm_provider
):
return candidate, generalized_info
for generalized_info in match_all_fallback_generalizations(candidate):
if _check_provider_match(model_info=generalized_info, custom_llm_provider=custom_llm_provider):
return candidate, generalized_info
return None
@ -5472,6 +5471,7 @@ def _get_model_info_helper(
supports_url_context=_model_info.get("supports_url_context", None),
supports_reasoning=_model_info.get("supports_reasoning", None),
supports_adaptive_thinking=_model_info.get("supports_adaptive_thinking", None),
supports_mid_conversation_system=_model_info.get("supports_mid_conversation_system", None),
supports_none_reasoning_effort=_model_info.get("supports_none_reasoning_effort", None),
supports_minimal_reasoning_effort=_model_info.get("supports_minimal_reasoning_effort", None),
supports_low_reasoning_effort=_model_info.get("supports_low_reasoning_effort", None),

View file

@ -1359,6 +1359,7 @@
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1393,6 +1394,7 @@
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1427,6 +1429,7 @@
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1461,6 +1464,7 @@
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1481,6 +1485,7 @@
"anthropic.claude-opus-4-8": {
"bedrock_converse_supports_strict_tools": false,
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -1516,6 +1521,7 @@
"global.anthropic.claude-opus-4-8": {
"bedrock_converse_supports_strict_tools": false,
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -1551,6 +1557,7 @@
"us.anthropic.claude-opus-4-8": {
"bedrock_converse_supports_strict_tools": false,
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1586,6 +1593,7 @@
"eu.anthropic.claude-opus-4-8": {
"bedrock_converse_supports_strict_tools": false,
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1621,6 +1629,43 @@
"au.anthropic.claude-opus-4-8": {
"bedrock_converse_supports_strict_tools": false,
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
"input_cost_per_token": 5.5e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.75e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_sampling_params": false,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_native_structured_output": true,
"supports_max_reasoning_effort": true,
"supports_output_config": true,
"bedrock_output_config_effort_ceiling": "xhigh",
"supports_parallel_tool_use_config": true
},
"jp.anthropic.claude-opus-4-8": {
"bedrock_converse_supports_strict_tools": false,
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1703,6 +1748,7 @@
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1737,6 +1783,7 @@
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1771,6 +1818,7 @@
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1805,6 +1853,7 @@
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1839,6 +1888,7 @@
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -1873,6 +1923,7 @@
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
@ -45231,6 +45282,17 @@
},
"fallback_generalizations": {
"rules": [
{
"name": "bedrock-anthropic-claude-mid-conversation-system",
"pattern": "anthropic\\.claude-[a-z]+-(?:4[-._](?:[89]|[1-9]\\d)(?!\\d)|(?:[5-9]|[1-9]\\d)(?!\\d)(?:[-._]\\d{1,2}(?!\\d))?)",
"description": "Bedrock Invoke ids for Claude 4.8 or higher: anthropic.claude-<family> with minor 4.8 through 4.99, any 5.x or later major-minor, or a bare 5+ major, which also admits new families such as fable. These models accept mid-conversation role system messages in place (verified live on Opus 4.8, Sonnet 5 and Fable 5), so unmapped future Bedrock Claudes keep the cache-preserving in-place handling instead of the hoist-all default. Listed first so bare-id provider inference, which takes the first pattern hit, resolves these Bedrock ids to bedrock; model-info resolution skips provider-mismatched rules either way.",
"extends": "anthropic-claude",
"model_info": {
"litellm_provider": "bedrock",
"supports_adaptive_thinking": true,
"supports_mid_conversation_system": true
}
},
{
"name": "anthropic-claude-adaptive-thinking",
"pattern": "(?:opus|sonnet|haiku)[-._](?:4[-._](?:[6-9]|[1-9]\\d)(?!\\d)|(?:[5-9]|[1-9]\\d{1,})[-._]\\d{1,2}(?!\\d))",

View file

@ -1628,23 +1628,27 @@ async def test_openai_responses_api_token_limit_error():
"""
Relevant issue: https://github.com/BerriAI/litellm/issues/15785
When this fails you'll see:
"pydantic_core._pydantic_core.ValidationError: 3 validation errors for ErrorEvent"
in the console.
Parsing the in-stream ErrorEvent must not raise
"pydantic_core._pydantic_core.ValidationError: 3 validation errors for ErrorEvent".
The iterator now surfaces the event as litellm.APIError with status 400
(invalid_request_error is a non-retriable client error, so no
MidStreamFallbackError wrapping) carrying the provider's message.
"""
litellm._turn_on_debug()
# Generate text with >400k tokens to trigger token limit error
oversized_text = "This is a test sentence. " * 50000 # ~400k tokens
# This will raise ValidationError instead of showing the real error
response = await litellm.aresponses(
model="gpt-5-mini", input=oversized_text, stream=True
)
async for event in response:
print(event) # Never reaches here - ValidationError is raised
with pytest.raises(litellm.APIError) as exc_info:
async for event in response:
print(event)
assert exc_info.value.status_code == 400
assert "exceeds the context window" in str(exc_info.value)
async def test_openai_streaming_logging():

View file

@ -266,3 +266,220 @@ async def test_aresponses_with_streaming_fallbacks_wraps_streaming_iterator():
)
assert out is wrapped
mock_wrap.assert_awaited_once()
@pytest.mark.asyncio
async def test_aresponses_fallback_on_in_stream_error_event():
"""A retriable in-stream error event (429) must trigger the router's mid-stream
fallback path: the wrapper catches MidStreamFallbackError raised by the source
iterator and yields the fallback stream instead of surfacing the error."""
import json
from unittest.mock import Mock
import litellm
from litellm.exceptions import MidStreamFallbackError
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator
from litellm.types.llms.openai import ErrorEvent, ErrorEventError
router = _make_router()
error_payload = {
"type": "error",
"error": {"type": "tokens", "code": "rate_limit_exceeded", "message": "rate limited"},
}
sse_bytes = f"data: {json.dumps(error_payload)}\n\n".encode()
async def mock_aiter_bytes():
yield sse_bytes
mock_response = Mock()
mock_response.headers = {}
mock_response.aiter_bytes = mock_aiter_bytes
mock_logging_obj = MagicMock(spec=LiteLLMLoggingObj)
mock_logging_obj.model_call_details = {"litellm_params": {}}
mock_logging_obj.completion_start_time = None
mock_config = Mock(spec=BaseResponsesAPIConfig)
mock_config.transform_streaming_response.return_value = ErrorEvent(
type=ResponsesAPIStreamEvents.ERROR,
sequence_number=0,
error=ErrorEventError(type="tokens", code="rate_limit_exceeded", message="rate limited"),
)
source = ResponsesAPIStreamingIterator(
response=mock_response,
model="gpt-5",
responses_api_provider_config=mock_config,
logging_obj=mock_logging_obj,
custom_llm_provider="openai",
)
fallback_event = _make_completed_event(1, 1, 2)
class _FallbackStream:
def __init__(self) -> None:
self._done = False
def __aiter__(self):
return self
async def __anext__(self):
if self._done:
raise StopAsyncIteration
self._done = True
return fallback_event
with patch.object(
router,
"async_function_with_fallbacks_common_utils",
new=AsyncMock(return_value=_FallbackStream()),
) as mock_fallback:
wrapped = await router._aresponses_streaming_iterator(
response=source,
initial_kwargs={"model": "primary", "input": "original question"},
)
collected = [ev async for ev in wrapped]
assert collected == [fallback_event]
mock_fallback.assert_awaited_once()
raised = mock_fallback.await_args.kwargs["e"]
assert isinstance(raised, MidStreamFallbackError)
assert raised.status_code == 429
assert isinstance(raised.original_exception, litellm.APIError)
assert raised.original_exception.status_code == 429
assert mock_fallback.await_args.kwargs["kwargs"]["input"] == "original question"
@pytest.mark.asyncio
async def test_aresponses_fallback_uses_continuation_input_after_partial_content():
"""When output text was already streamed before the error, the fallback re-entry
must carry a continuation input with the partial assistant text instead of
retrying the original input from scratch (which would duplicate streamed content)."""
import json
from unittest.mock import Mock
from litellm.exceptions import MidStreamFallbackError
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator
from litellm.types.llms.openai import ErrorEvent, ErrorEventError
router = _make_router()
events = [
{"type": "response.output_text.delta", "delta": "partial answer"},
{"type": "error", "error": {"type": "server_error", "code": "internal_error", "message": "boom"}},
]
sse_payload = b"".join(f"data: {json.dumps(event)}\n\n".encode() for event in events)
async def mock_aiter_bytes():
yield sse_payload
mock_response = Mock()
mock_response.headers = {}
mock_response.aiter_bytes = mock_aiter_bytes
mock_logging_obj = MagicMock(spec=LiteLLMLoggingObj)
mock_logging_obj.model_call_details = {"litellm_params": {}}
mock_logging_obj.completion_start_time = None
mock_config = Mock(spec=BaseResponsesAPIConfig)
def transform(model, parsed_chunk, logging_obj):
if parsed_chunk.get("type") == "error":
return ErrorEvent(
type=ResponsesAPIStreamEvents.ERROR,
sequence_number=0,
error=ErrorEventError(**parsed_chunk["error"]),
)
delta_event = Mock()
delta_event.type = ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA
delta_event.delta = parsed_chunk["delta"]
return delta_event
mock_config.transform_streaming_response.side_effect = transform
source = ResponsesAPIStreamingIterator(
response=mock_response,
model="gpt-5",
responses_api_provider_config=mock_config,
logging_obj=mock_logging_obj,
custom_llm_provider="openai",
)
fallback_event = _make_completed_event(1, 1, 2)
class _FallbackStream:
def __init__(self) -> None:
self._done = False
def __aiter__(self):
return self
async def __anext__(self):
if self._done:
raise StopAsyncIteration
self._done = True
return fallback_event
with patch.object(
router,
"async_function_with_fallbacks_common_utils",
new=AsyncMock(return_value=_FallbackStream()),
) as mock_fallback:
wrapped = await router._aresponses_streaming_iterator(
response=source,
initial_kwargs={"model": "primary", "input": "original question"},
)
collected = [ev async for ev in wrapped]
assert collected[-1] == fallback_event
raised = mock_fallback.await_args.kwargs["e"]
assert isinstance(raised, MidStreamFallbackError)
assert raised.is_pre_first_chunk is False
assert raised.generated_content == "partial answer"
continuation = mock_fallback.await_args.kwargs["kwargs"]["input"]
assert isinstance(continuation, list)
assert continuation[0]["content"][0]["text"] == "original question"
assert continuation[-2]["role"] == "developer"
assert continuation[-1]["role"] == "assistant"
assert continuation[-1]["content"][0]["text"] == "partial answer"
@pytest.mark.asyncio
async def test_aresponses_client_error_event_skips_fallback():
"""A 400-mapped in-stream error (raised as APIError, not MidStreamFallbackError)
must surface to the caller without invoking the router's fallback path."""
import litellm
router = _make_router()
class _ClientErrorSource:
completed_response = None
def __aiter__(self):
return self
async def __anext__(self):
raise litellm.APIError(
status_code=400,
message="bad request",
llm_provider="openai",
model="gpt-5",
)
wrapped = await router._aresponses_streaming_iterator(
response=_ClientErrorSource(),
initial_kwargs={"model": "primary"},
)
with patch.object(
router,
"async_function_with_fallbacks_common_utils",
new=AsyncMock(),
) as mock_fallback:
with pytest.raises(litellm.APIError) as exc_info:
async for _ in wrapped:
pass
assert exc_info.value.status_code == 400
mock_fallback.assert_not_awaited()

View file

@ -1,3 +1,4 @@
import asyncio
from unittest.mock import AsyncMock, Mock, patch
import httpx
@ -6,16 +7,20 @@ from httpx import Request, Response
from litellm.integrations.datadog.datadog import DataDogLogger
from litellm.llms.custom_httpx.http_handler import MaskedHTTPStatusError
from litellm.types.integrations.datadog import DD_MAX_BATCH_SIZE, DatadogPayload
from litellm.types.integrations.datadog import (
DD_MAX_BATCH_SIZE,
DD_MAX_PAYLOAD_SIZE_BYTES,
DatadogPayload,
)
def _payloads(n):
def _payloads(n, message=None):
return [
DatadogPayload(
ddsource="litellm",
ddtags="env:test",
hostname="host",
message=f'{{"event": {i}}}',
message=f"{message}{i}" if message else f'{{"event": {i}}}',
service="svc",
status="info",
)
@ -177,6 +182,87 @@ async def test_413_returned_response_also_splits(datadog_env):
assert logger.log_queue == []
def _make_recording_send(sent_batches, delivered):
async def _send(data):
sent_batches.append(list(data))
delivered.extend(data)
return Response(
202, request=Request("POST", "https://example.com"), text="Accepted"
)
return _send
@pytest.mark.asyncio
async def test_oversized_payload_splits_before_any_send(datadog_env):
"""Regression for LIT-4325: a batch above Datadog's uncompressed payload limit is
split proactively, so the intake never has to reject it with a 413."""
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
with patch("asyncio.create_task"):
logger = DataDogLogger()
events = _payloads(3, message="x" * 3_000_000)
logger.log_queue = list(events)
sent_batches: list = []
delivered: list = []
logger.async_send_compressed_data = AsyncMock(
side_effect=_make_recording_send(sent_batches, delivered)
)
await logger.async_send_batch()
assert delivered == events
assert len(sent_batches) == 3
assert all(
len(safe_dumps(batch).encode("utf-8")) <= DD_MAX_PAYLOAD_SIZE_BYTES
for batch in sent_batches
)
assert logger.log_queue == []
@pytest.mark.asyncio
async def test_batch_over_max_event_count_splits_before_any_send(datadog_env):
"""Datadog caps a payload at 1000 events; a queue that grew past that (e.g. after
re-queues) must be sent in count-compliant chunks."""
with patch("asyncio.create_task"):
logger = DataDogLogger()
events = _payloads(DD_MAX_BATCH_SIZE + 1)
logger.log_queue = list(events)
sent_batches: list = []
delivered: list = []
logger.async_send_compressed_data = AsyncMock(
side_effect=_make_recording_send(sent_batches, delivered)
)
await logger.async_send_batch()
assert delivered == events
assert all(len(batch) <= DD_MAX_BATCH_SIZE for batch in sent_batches)
assert logger.log_queue == []
@pytest.mark.asyncio
async def test_single_event_above_payload_cap_is_still_sent(datadog_env):
"""A lone event over the byte cap cannot be split further; it must be sent once
(Datadog decides), never looped on."""
with patch("asyncio.create_task"):
logger = DataDogLogger()
logger.log_queue = _payloads(1, message="x" * (DD_MAX_PAYLOAD_SIZE_BYTES + 1))
sent_batches: list = []
delivered: list = []
send = AsyncMock(side_effect=_make_recording_send(sent_batches, delivered))
logger.async_send_compressed_data = send
await asyncio.wait_for(logger.async_send_batch(), timeout=10)
assert send.await_count == 1
assert len(delivered) == 1
assert logger.log_queue == []
@pytest.mark.asyncio
async def test_partial_delivery_then_transient_error_requeues_only_undelivered(
datadog_env,

View file

@ -0,0 +1,359 @@
"""
Unit tests for the NoOpMetric guard in _increment_remaining_budget_metrics
and the per-entity guards in _set_*_budget_metrics_after_api_request.
Regression tests that the specific bug can never happen again:
when budget gauges are excluded from prometheus_metrics_config (and therefore
created as NoOpMetric instances), the DB/cache lookup helpers must not be called.
"""
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from prometheus_client import REGISTRY
import litellm
from litellm.integrations.prometheus import PrometheusLogger
from litellm.types.integrations.prometheus import NoOpMetric
_BUDGET_EXCLUDED_CONFIG = [
{
"group": "core-only",
"metrics": [
"litellm_requests_metric",
"litellm_total_tokens_metric",
],
}
]
@pytest.fixture(autouse=True)
def cleanup_prometheus_registry():
old_config = litellm.prometheus_metrics_config
for collector in list(REGISTRY._collector_to_names.keys()):
try:
REGISTRY.unregister(collector)
except Exception:
pass
yield
litellm.prometheus_metrics_config = old_config
for collector in list(REGISTRY._collector_to_names.keys()):
try:
REGISTRY.unregister(collector)
except Exception:
pass
def make_logger_with_budget_metrics_disabled() -> PrometheusLogger:
litellm.prometheus_metrics_config = _BUDGET_EXCLUDED_CONFIG
return PrometheusLogger()
def make_logger_with_all_metrics_enabled() -> PrometheusLogger:
litellm.prometheus_metrics_config = None
return PrometheusLogger()
COMMON_KWARGS = dict(
user_api_team="team-123",
user_api_team_alias="my-team",
user_api_key="hashed-key",
user_api_key_alias="my-key",
litellm_params={"metadata": {}},
response_cost=0.001,
user_id="user-1",
user_api_key_org_id="org-1",
)
class TestBudgetGaugesAreNoopWhenExcluded:
def test_team_gauge_is_noop(self):
logger = make_logger_with_budget_metrics_disabled()
assert isinstance(logger.litellm_remaining_team_budget_metric, NoOpMetric)
def test_api_key_gauge_is_noop(self):
logger = make_logger_with_budget_metrics_disabled()
assert isinstance(logger.litellm_remaining_api_key_budget_metric, NoOpMetric)
def test_user_gauge_is_noop(self):
logger = make_logger_with_budget_metrics_disabled()
assert isinstance(logger.litellm_remaining_user_budget_metric, NoOpMetric)
def test_org_gauge_is_noop(self):
logger = make_logger_with_budget_metrics_disabled()
assert isinstance(logger.litellm_remaining_org_budget_metric, NoOpMetric)
def test_gauges_are_real_when_all_metrics_enabled(self):
logger = make_logger_with_all_metrics_enabled()
assert not isinstance(logger.litellm_remaining_team_budget_metric, NoOpMetric)
assert not isinstance(logger.litellm_remaining_api_key_budget_metric, NoOpMetric)
assert not isinstance(logger.litellm_remaining_user_budget_metric, NoOpMetric)
assert not isinstance(logger.litellm_remaining_org_budget_metric, NoOpMetric)
class TestTopLevelGuard:
@pytest.mark.asyncio
async def test_no_db_lookups_when_all_budget_gauges_are_noop(self):
"""Regression: _increment_remaining_budget_metrics must return early
without any I/O when all four budget gauges are NoOpMetric."""
logger = make_logger_with_budget_metrics_disabled()
assemble_team = AsyncMock(return_value=MagicMock())
assemble_key = AsyncMock(return_value=MagicMock())
assemble_user = AsyncMock(return_value=MagicMock())
with (
patch.object(logger, "_assemble_team_object", assemble_team),
patch.object(logger, "_assemble_key_object", assemble_key),
patch.object(logger, "_assemble_user_object", assemble_user),
):
await logger._increment_remaining_budget_metrics(**COMMON_KWARGS)
assemble_team.assert_not_called()
assemble_key.assert_not_called()
assemble_user.assert_not_called()
@pytest.mark.asyncio
async def test_db_lookups_run_when_budget_gauges_are_real(self):
"""When budget gauges are real Prometheus metrics, the assemble helpers
must be called so I/O proceeds normally."""
logger = make_logger_with_all_metrics_enabled()
assemble_team = AsyncMock(
return_value=MagicMock(
team_id="team-123",
team_alias="my-team",
spend=0.001,
max_budget=None,
budget_reset_at=None,
)
)
assemble_key = AsyncMock(
return_value=MagicMock(
token="hashed-key",
key_alias="my-key",
spend=0.001,
max_budget=None,
budget_reset_at=None,
)
)
assemble_user = AsyncMock(
return_value=MagicMock(
user_id="user-1",
spend=0.001,
max_budget=None,
budget_reset_at=None,
user_email=None,
user_alias=None,
)
)
with (
patch.object(logger, "_assemble_team_object", assemble_team),
patch.object(logger, "_assemble_key_object", assemble_key),
patch.object(logger, "_assemble_user_object", assemble_user),
patch.object(logger, "_set_team_budget_metrics", MagicMock()),
patch.object(logger, "_set_key_budget_metrics", MagicMock()),
patch.object(logger, "_set_user_budget_metrics", MagicMock()),
patch.object(logger, "_set_org_budget_metrics_after_api_request", AsyncMock()),
):
await logger._increment_remaining_budget_metrics(**COMMON_KWARGS)
assemble_team.assert_called_once()
assemble_key.assert_called_once()
assemble_user.assert_called_once()
class TestPerEntityGuards:
@pytest.mark.asyncio
async def test_team_guard_skips_lookup_when_team_gauge_is_noop(self):
"""Per-entity guard: team assemble helper is not called when team gauge is NoOp,
even when key and user gauges are real."""
logger = make_logger_with_all_metrics_enabled()
logger.litellm_remaining_team_budget_metric = NoOpMetric()
assemble_team = AsyncMock(return_value=MagicMock())
assemble_key = AsyncMock(
return_value=MagicMock(
token="hashed-key",
key_alias="my-key",
spend=0.001,
max_budget=None,
budget_reset_at=None,
)
)
assemble_user = AsyncMock(
return_value=MagicMock(
user_id="user-1",
spend=0.001,
max_budget=None,
budget_reset_at=None,
user_email=None,
user_alias=None,
)
)
with (
patch.object(logger, "_assemble_team_object", assemble_team),
patch.object(logger, "_assemble_key_object", assemble_key),
patch.object(logger, "_assemble_user_object", assemble_user),
patch.object(logger, "_set_team_budget_metrics", MagicMock()),
patch.object(logger, "_set_key_budget_metrics", MagicMock()),
patch.object(logger, "_set_user_budget_metrics", MagicMock()),
patch.object(logger, "_set_org_budget_metrics_after_api_request", AsyncMock()),
):
await logger._increment_remaining_budget_metrics(**COMMON_KWARGS)
assemble_team.assert_not_called()
assemble_key.assert_called_once()
assemble_user.assert_called_once()
@pytest.mark.asyncio
async def test_key_guard_skips_lookup_when_key_gauge_is_noop(self):
"""Per-entity guard: key assemble helper is not called when key gauge is NoOp,
even when team and user gauges are real."""
logger = make_logger_with_all_metrics_enabled()
logger.litellm_remaining_api_key_budget_metric = NoOpMetric()
assemble_team = AsyncMock(
return_value=MagicMock(
team_id="team-123",
team_alias="my-team",
spend=0.001,
max_budget=None,
budget_reset_at=None,
)
)
assemble_key = AsyncMock(return_value=MagicMock())
assemble_user = AsyncMock(
return_value=MagicMock(
user_id="user-1",
spend=0.001,
max_budget=None,
budget_reset_at=None,
user_email=None,
user_alias=None,
)
)
with (
patch.object(logger, "_assemble_team_object", assemble_team),
patch.object(logger, "_assemble_key_object", assemble_key),
patch.object(logger, "_assemble_user_object", assemble_user),
patch.object(logger, "_set_team_budget_metrics", MagicMock()),
patch.object(logger, "_set_key_budget_metrics", MagicMock()),
patch.object(logger, "_set_user_budget_metrics", MagicMock()),
patch.object(logger, "_set_org_budget_metrics_after_api_request", AsyncMock()),
):
await logger._increment_remaining_budget_metrics(**COMMON_KWARGS)
assemble_key.assert_not_called()
assemble_team.assert_called_once()
assemble_user.assert_called_once()
@pytest.mark.asyncio
async def test_user_guard_skips_lookup_when_user_gauge_is_noop(self):
"""Per-entity guard: user assemble helper is not called when user gauge is NoOp,
even when team and key gauges are real."""
logger = make_logger_with_all_metrics_enabled()
logger.litellm_remaining_user_budget_metric = NoOpMetric()
assemble_team = AsyncMock(
return_value=MagicMock(
team_id="team-123",
team_alias="my-team",
spend=0.001,
max_budget=None,
budget_reset_at=None,
)
)
assemble_key = AsyncMock(
return_value=MagicMock(
token="hashed-key",
key_alias="my-key",
spend=0.001,
max_budget=None,
budget_reset_at=None,
)
)
assemble_user = AsyncMock(return_value=MagicMock())
with (
patch.object(logger, "_assemble_team_object", assemble_team),
patch.object(logger, "_assemble_key_object", assemble_key),
patch.object(logger, "_assemble_user_object", assemble_user),
patch.object(logger, "_set_team_budget_metrics", MagicMock()),
patch.object(logger, "_set_key_budget_metrics", MagicMock()),
patch.object(logger, "_set_user_budget_metrics", MagicMock()),
patch.object(logger, "_set_org_budget_metrics_after_api_request", AsyncMock()),
):
await logger._increment_remaining_budget_metrics(**COMMON_KWARGS)
assemble_user.assert_not_called()
assemble_team.assert_called_once()
assemble_key.assert_called_once()
@pytest.mark.asyncio
async def test_set_team_budget_metrics_directly_skips_when_gauge_is_noop(self):
"""_set_team_budget_metrics_after_api_request returns early when team gauge is NoOp."""
logger = make_logger_with_budget_metrics_disabled()
assemble_team = AsyncMock(return_value=MagicMock())
with patch.object(logger, "_assemble_team_object", assemble_team):
await logger._set_team_budget_metrics_after_api_request(
user_api_team="team-123",
user_api_team_alias="my-team",
team_spend=0.5,
team_max_budget=10.0,
response_cost=0.001,
)
assemble_team.assert_not_called()
@pytest.mark.asyncio
async def test_set_api_key_budget_metrics_directly_skips_when_gauge_is_noop(self):
"""_set_api_key_budget_metrics_after_api_request returns early when key gauge is NoOp."""
logger = make_logger_with_budget_metrics_disabled()
assemble_key = AsyncMock(return_value=MagicMock())
with patch.object(logger, "_assemble_key_object", assemble_key):
await logger._set_api_key_budget_metrics_after_api_request(
user_api_key="hashed-key",
user_api_key_alias="my-key",
response_cost=0.001,
key_max_budget=10.0,
key_spend=0.5,
)
assemble_key.assert_not_called()
@pytest.mark.asyncio
async def test_set_user_budget_metrics_directly_skips_when_gauge_is_noop(self):
"""_set_user_budget_metrics_after_api_request returns early when user gauge is NoOp."""
logger = make_logger_with_budget_metrics_disabled()
assemble_user = AsyncMock(return_value=MagicMock())
with patch.object(logger, "_assemble_user_object", assemble_user):
await logger._set_user_budget_metrics_after_api_request(
user_id="user-1",
user_spend=0.5,
user_max_budget=10.0,
response_cost=0.001,
)
assemble_user.assert_not_called()
@pytest.mark.asyncio
async def test_set_org_budget_metrics_directly_skips_when_gauge_is_noop(self):
"""_set_org_budget_metrics_after_api_request returns early when org gauge is NoOp.
The guard fires before any import of auth_checks, so prisma_client is never touched."""
logger = make_logger_with_budget_metrics_disabled()
set_org_metrics = MagicMock()
with patch.object(logger, "_set_org_budget_metrics", set_org_metrics):
await logger._set_org_budget_metrics_after_api_request(
org_id="org-1",
response_cost=0.001,
)
set_org_metrics.assert_not_called()

View file

@ -16,6 +16,7 @@ sys.path.insert(0, os.path.abspath("../../.."))
import litellm
from litellm.litellm_core_utils.fallback_generalizations import (
get_fallback_generalization_rules,
match_all_fallback_generalizations,
match_fallback_generalization,
set_fallback_generalizations,
)
@ -57,6 +58,48 @@ def test_match_returns_model_info_of_first_matching_rule(restore_generalizations
assert matched["tag"] == "first"
def test_match_all_returns_every_matching_rule_in_order(restore_generalizations):
restore_generalizations(
[
{
"name": "first",
"pattern": r"^acme-",
"model_info": {"litellm_provider": "openai", "tag": "first"},
},
{
"name": "second",
"pattern": r"^acme-pro-",
"model_info": {"litellm_provider": "anthropic", "tag": "second"},
},
]
)
assert [m["tag"] for m in match_all_fallback_generalizations("acme-pro-1")] == ["first", "second"]
assert match_all_fallback_generalizations("gpt-4o") == []
def test_provider_scoped_rule_is_skipped_for_other_providers(restore_generalizations):
"""Model-info resolution must fall through a provider-mismatched earlier rule to a
later applicable one, instead of discarding the model name at the first pattern hit."""
restore_generalizations(
[
{
"name": "bedrock-scoped",
"pattern": r"^acme-",
"model_info": {"litellm_provider": "bedrock", "supports_vision": False},
},
{
"name": "openai-scoped",
"pattern": r"^acme-",
"model_info": {"litellm_provider": "openai", "mode": "chat", "supports_vision": True},
},
]
)
litellm.get_model_info.cache_clear()
info = litellm.get_model_info("acme-fast-1", custom_llm_provider="openai")
assert info["litellm_provider"] == "openai"
assert info["supports_vision"] is True
def test_match_is_case_insensitive(restore_generalizations):
restore_generalizations(
[{"name": "r", "pattern": r"^claude-opus", "model_info": {"ok": True}}]
@ -298,3 +341,47 @@ def test_shipped_adaptive_rule_gates_on_version_not_pricing(shipped_cost_map):
assert non_adaptive not in litellm.model_cost
assert AnthropicModelInfo._is_adaptive_thinking_model(adaptive) is True
assert AnthropicModelInfo._is_adaptive_thinking_model(non_adaptive) is False
def test_shipped_bedrock_rule_resolves_unmapped_future_claude_for_bedrock(shipped_cost_map):
"""An unmapped Bedrock Claude >= 4.8 resolves via the bedrock-scoped
``bedrock-anthropic-claude-mid-conversation-system`` rule even when the lookup
carries ``custom_llm_provider="bedrock"``, which the provider check uses to drop
the anthropic-scoped rules. It inherits base capabilities, gains both
version-gated flags, and stays unpriced."""
model = "us.anthropic.claude-opus-4-9"
assert model not in litellm.model_cost
info = litellm.get_model_info(model, custom_llm_provider="bedrock")
assert info["litellm_provider"] == "bedrock"
assert info["supports_mid_conversation_system"] is True
assert info["supports_adaptive_thinking"] is True
assert info["supports_function_calling"] is True
assert not info.get("input_cost_per_token")
def test_shipped_bedrock_mid_conversation_rule_gates_on_version_and_naming(shipped_cost_map):
"""The bedrock rule only claims Bedrock-style ids at 4.8+, bare 5+ majors and
new families included; pre-4.8 Bedrock ids and native ids never gain the flag,
and the rule outranks the anthropic-scoped ones for Bedrock ids because it is
listed first."""
for flagged in (
"us.anthropic.claude-opus-4-8",
"jp.anthropic.claude-opus-4-8",
"anthropic.claude-sonnet-5",
"us.anthropic.claude-fable-5",
"anthropic.claude-sonnet-5-20260101-v1:0",
):
matched = match_fallback_generalization(flagged)
assert matched is not None, flagged
assert matched["litellm_provider"] == "bedrock", flagged
assert matched["supports_mid_conversation_system"] is True, flagged
for unflagged in (
"us.anthropic.claude-opus-4-7",
"us.anthropic.claude-sonnet-4-6",
"us.anthropic.claude-haiku-4-5-20251001-v1:0",
"anthropic.claude-3-5-sonnet-20240620-v1:0",
"claude-opus-4-9",
"claude-sonnet-5",
):
matched = match_fallback_generalization(unflagged)
assert matched is None or not matched.get("supports_mid_conversation_system"), unflagged

View file

@ -0,0 +1,189 @@
import pytest
from litellm.constants import (
DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET,
)
from litellm.llms.anthropic.common_utils import AnthropicError
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
AnthropicMessagesConfig,
)
def _claude_code_payload(effort="medium", max_tokens=8192, **output_config_extra):
"""The exact adaptive-thinking shape Claude Code (claude-cli) sends."""
output_config = {"effort": effort, **output_config_extra}
return {
"max_tokens": max_tokens,
"thinking": {"type": "adaptive"},
"output_config": output_config,
}
def _transform(model, params, litellm_params=None):
return AnthropicMessagesConfig().transform_anthropic_messages_request(
model=model,
messages=[{"role": "user", "content": "Hello"}],
anthropic_messages_optional_request_params=dict(params),
litellm_params=litellm_params or {},
headers={},
)
def test_effort_translated_to_legacy_thinking_for_haiku_4_5():
"""Core regression: Claude Code sends adaptive thinking + effort to Haiku 4.5
(thinking-capable, pre-4.6). Effort must be translated to legacy extended
thinking rather than forwarded raw (which Anthropic rejects with "This model
does not support the effort parameter")."""
result = _transform("claude-haiku-4-5", _claude_code_payload(effort="medium"))
assert result["thinking"] == {
"type": "enabled",
"budget_tokens": DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET,
}
assert "output_config" not in result
def test_effort_high_maps_to_high_budget_for_sonnet_4_5():
result = _transform("claude-sonnet-4-5", _claude_code_payload(effort="high"))
assert result["thinking"] == {
"type": "enabled",
"budget_tokens": DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET,
}
assert "output_config" not in result
def test_adaptive_effort_passes_through_untouched_for_4_6():
"""4.6+ natively supports the adaptive interface, so it must not be rewritten."""
result = _transform("claude-sonnet-4-6", _claude_code_payload(effort="high"))
assert result["thinking"] == {"type": "adaptive"}
assert result["output_config"] == {"effort": "high"}
def test_thinking_and_effort_dropped_for_non_reasoning_model():
"""A model with no reasoning support cannot take thinking or effort, so both are
silently dropped (no drop_params required) so the request still succeeds."""
result = _transform("claude-3-5-haiku-latest", _claude_code_payload(effort="medium"))
assert "thinking" not in result
assert "output_config" not in result
def test_residual_output_config_preserved_after_effort_translation():
"""output_config may carry `format` (structured outputs) alongside effort. Only
the consumed effort key is removed; the residual is left for provider subclasses
(bedrock/vertex) to handle, and effort is translated to legacy thinking."""
result = _transform(
"claude-haiku-4-5",
_claude_code_payload(effort="medium", format={"type": "json_schema"}),
)
assert result["thinking"] == {
"type": "enabled",
"budget_tokens": DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET,
}
assert result["output_config"] == {"format": {"type": "json_schema"}}
def test_opus_4_5_keeps_effort_but_drops_adaptive_thinking():
"""Regression: Opus 4.5 advertises supports_output_config (accepts
output_config.effort) but is NOT adaptive, so thinking:{type:adaptive} is
rejected by Anthropic. The effort must be kept and only the adaptive thinking
block dropped, rather than early-returning and forwarding adaptive thinking raw."""
result = _transform("claude-opus-4-5", _claude_code_payload(effort="medium"))
assert result["output_config"] == {"effort": "medium"}
assert "thinking" not in result
def test_opus_4_5_preserves_native_effort_without_adaptive_thinking():
"""A caller sending output_config.effort alone (no adaptive thinking) to Opus 4.5
must pass through untouched, since the model supports it natively."""
result = AnthropicMessagesConfig().transform_anthropic_messages_request(
model="claude-opus-4-5",
messages=[{"role": "user", "content": "Hello"}],
anthropic_messages_optional_request_params={
"max_tokens": 8192,
"output_config": {"effort": "high"},
},
litellm_params={},
headers={},
)
assert result["output_config"] == {"effort": "high"}
assert "thinking" not in result
def test_opus_4_5_unsupported_effort_level_translated_to_legacy_thinking():
"""Opus 4.5 accepts output_config.effort but only levels low/medium/high;
Claude Code defaults to xhigh on newer models, and forwarding that level raw
would be rejected with "effort='xhigh' is not supported by this model". An
unsupported level must fall through to the legacy translation (budget-based
thinking, effort stripped) instead of being preserved."""
result = _transform("claude-opus-4-5", _claude_code_payload(effort="xhigh", max_tokens=64000))
assert result["thinking"] == {
"type": "enabled",
"budget_tokens": DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET,
}
assert "output_config" not in result
def test_opus_4_5_effort_only_unsupported_level_left_for_provider_normalization():
"""An effort-only request (no adaptive thinking) must pass through untouched even
when the level exceeds what the model supports: provider subclasses own their
level normalization (bedrock clamps xhigh to the model's ceiling after this base
transform runs), so consuming the effort here breaks that contract."""
result = _transform(
"claude-opus-4-5",
{"max_tokens": 4096, "output_config": {"effort": "xhigh"}},
)
assert result["output_config"] == {"effort": "xhigh"}
assert "thinking" not in result
def test_budget_capped_below_max_tokens():
"""Adaptive thinking carries no budget, so the translated legacy budget must be
capped below max_tokens (Anthropic requires max_tokens > budget_tokens). A
high-effort budget (4096) with max_tokens=3000 must be capped to 2999."""
result = _transform("claude-haiku-4-5", _claude_code_payload(effort="high", max_tokens=3000))
assert result["thinking"] == {"type": "enabled", "budget_tokens": 2999}
def test_thinking_dropped_when_max_tokens_too_small_for_min_budget():
"""When max_tokens can't fit even the minimum thinking budget, thinking is
silently dropped so the request still succeeds rather than being rejected."""
result = _transform("claude-haiku-4-5", _claude_code_payload(effort="medium", max_tokens=512))
assert "thinking" not in result
assert "output_config" not in result
def test_unrecognized_effort_raises_clean_400():
"""An unrecognized effort value (e.g. a future Anthropic tier) must surface as a
clean AnthropicError 400, matching _translate_reasoning_effort_to_anthropic,
rather than leaking litellm's internal BadRequestError."""
with pytest.raises(AnthropicError) as exc_info:
_transform("claude-haiku-4-5", _claude_code_payload(effort="turbo"))
assert exc_info.value.status_code == 400
def test_non_adaptive_request_without_effort_is_untouched():
"""A non-adaptive model receiving a request with no adaptive interface (no
effort, no adaptive thinking) must pass through untouched."""
result = AnthropicMessagesConfig().transform_anthropic_messages_request(
model="claude-haiku-4-5",
messages=[{"role": "user", "content": "Hello"}],
anthropic_messages_optional_request_params={"max_tokens": 1024},
litellm_params={},
headers={},
)
assert "thinking" not in result
assert "output_config" not in result

View file

@ -772,7 +772,6 @@ class TestProxyOAuthHeaderForwarding:
self,
):
"""OAuth Authorization header IS forwarded when x-litellm-api-key was used for proxy auth."""
from unittest.mock import patch
from starlette.datastructures import Headers
@ -1724,3 +1723,48 @@ class TestClaudeOpus48AdaptiveThinking:
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
assert AnthropicModelInfo._is_adaptive_thinking_model(model) is False
class TestDefaultSuffixAdaptiveThinking:
"""@default-suffixed Vertex AI model names (e.g. vertex_ai/claude-opus-4-8@default)
must resolve as adaptive thinking. Before the fix, _model_map_lookup_candidates
never stripped the @default suffix, so the lookup fell through to the bare
model name without @default, which may or may not have the flag, and for
provider-prefixed forms the lookup always missed (issue #31760)."""
@pytest.mark.parametrize(
"model",
[
"vertex_ai/claude-opus-4-8@default",
"vertex_ai/claude-sonnet-4-6@default",
"vertex_ai/claude-opus-4-7@default",
"vertex_ai/claude-opus-4-6@default",
"vertex_ai/claude-fable-5@default",
],
)
def test_default_suffix_models_are_adaptive_thinking(
self, local_model_cost_map, model: str
) -> None:
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
assert AnthropicModelInfo._is_adaptive_thinking_model(model) is True, (
f"{model} not classified as adaptive thinking. "
"Check _model_map_lookup_candidates strips @default suffix."
)
@pytest.mark.parametrize(
"model,expected_bare",
[
("vertex_ai/claude-opus-4-8@default", "claude-opus-4-8"),
("vertex_ai/claude-sonnet-4-6@default", "claude-sonnet-4-6"),
],
)
def test_lookup_candidates_include_bare_name(
self, model: str, expected_bare: str
) -> None:
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
candidates = AnthropicModelInfo._model_map_lookup_candidates(model)
assert expected_bare in candidates, (
f"Expected '{expected_bare}' in candidates for '{model}', got: {candidates}"
)

View file

@ -1974,14 +1974,24 @@ def test_bedrock_invoke_transform_merges_list_content_system_role_into_system():
]
def test_bedrock_invoke_transform_keeps_mid_conversation_system_role_in_place():
@pytest.mark.parametrize(
"model",
[
"anthropic.claude-opus-4-8",
"jp.anthropic.claude-opus-4-8",
"us.anthropic.claude-sonnet-5",
"us.anthropic.claude-fable-5",
],
)
def test_bedrock_invoke_transform_keeps_mid_conversation_system_role_in_place(local_model_cost_map, model):
"""Regression test for the Bedrock prompt-cache collapse: hoisting a
mid-conversation ``role: "system"`` message (e.g. Claude Code's
``mid-conversation-system-2026-04-07`` reminders) into the top-level
``system`` field mutates the cache prefix and invalidates the cached message
history, so such entries must be forwarded in place. Invoke only rejects a
system entry at ``messages.0``. Billing-header blocks must still be stripped
from the top-level ``system`` field even when nothing is hoisted."""
history, so on models flagged ``supports_mid_conversation_system`` (Claude
4.8+, which Invoke accepts the role on) such entries must be forwarded
in place. Billing-header blocks must still be stripped from the top-level
``system`` field even when nothing is hoisted."""
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
@ -1993,7 +2003,7 @@ def test_bedrock_invoke_transform_keeps_mid_conversation_system_role_in_place():
]
result = cfg.transform_anthropic_messages_request(
model="anthropic.claude-opus-4-8",
model=model,
messages=copy.deepcopy(messages),
anthropic_messages_optional_request_params={
"max_tokens": 256,
@ -2013,10 +2023,11 @@ def test_bedrock_invoke_transform_keeps_mid_conversation_system_role_in_place():
]
def test_bedrock_invoke_transform_hoists_only_leading_system_run():
"""Only the leading run of ``role: "system"`` messages is hoisted into the
top-level ``system`` field; a later system entry keeps its position in
``messages`` so the serialized prefix stays stable across turns."""
def test_bedrock_invoke_transform_hoists_only_leading_system_run(local_model_cost_map):
"""On models flagged ``supports_mid_conversation_system``, only the leading
run of ``role: "system"`` messages is hoisted into the top-level ``system``
field; a later system entry keeps its position in ``messages`` so the
serialized prefix stays stable across turns."""
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
@ -2047,6 +2058,133 @@ def test_bedrock_invoke_transform_hoists_only_leading_system_run():
]
def test_bedrock_invoke_transform_hoists_mid_conversation_system_for_older_claude(local_model_cost_map):
"""Regression test for Claude Code 400s on pre-Opus-4.8 Bedrock models:
Invoke rejects ``role: "system"`` in every position on Opus 4.7, Sonnet 4.6,
Haiku 4.5, etc. ("role 'system' is not supported on this model"), so on
models without ``supports_mid_conversation_system`` every system entry must
be hoisted into the top-level ``system`` field, mid-conversation ones
included."""
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
messages = [
{"role": "user", "content": "read the file"},
{"role": "system", "content": "[Truncated: PARTIAL view of big1.txt]"},
{"role": "assistant", "content": "reading"},
{"role": "user", "content": "continue"},
]
result = cfg.transform_anthropic_messages_request(
model="us.anthropic.claude-opus-4-7",
messages=copy.deepcopy(messages),
anthropic_messages_optional_request_params={
"max_tokens": 256,
"stream": False,
"system": [{"type": "text", "text": "Base."}],
},
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert result["messages"] == [
{"role": "user", "content": "read the file"},
{"role": "assistant", "content": "reading"},
{"role": "user", "content": "continue"},
]
assert result["system"] == [
{"type": "text", "text": "Base."},
{"type": "text", "text": "[Truncated: PARTIAL view of big1.txt]"},
]
def test_bedrock_invoke_transform_hoists_all_system_for_unmapped_model(local_model_cost_map):
"""A model with no cost-map entry and no fallback-generalization rule gets
the hoist-everything behavior: the safe default is a mutated cache prefix,
never a provider 400 from forwarding a role the model may not accept."""
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
messages = [
{"role": "user", "content": "hi"},
{"role": "system", "content": "mid-conversation reminder"},
{"role": "assistant", "content": "hello"},
{"role": "user", "content": "continue"},
]
result = cfg.transform_anthropic_messages_request(
model="us.anthropic.claude-opus-3-9",
messages=copy.deepcopy(messages),
anthropic_messages_optional_request_params={"max_tokens": 256, "stream": False},
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert result["messages"] == [
{"role": "user", "content": "hi"},
{"role": "assistant", "content": "hello"},
{"role": "user", "content": "continue"},
]
assert result["system"] == [{"type": "text", "text": "mid-conversation reminder"}]
def test_bedrock_invoke_transform_keeps_system_in_place_for_unmapped_future_claude(local_model_cost_map):
"""An unmapped Bedrock Claude at 4.8 or higher resolves through the
``bedrock-anthropic-claude-mid-conversation-system`` fallback rule, so a
future model that has not landed in the cost map yet keeps the
cache-preserving in-place behavior instead of falling back to hoist-all."""
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
messages = [
{"role": "user", "content": "hi"},
{"role": "system", "content": "mid-conversation reminder"},
{"role": "assistant", "content": "hello"},
{"role": "user", "content": "continue"},
]
result = cfg.transform_anthropic_messages_request(
model="us.anthropic.claude-opus-4-9",
messages=copy.deepcopy(messages),
anthropic_messages_optional_request_params={"max_tokens": 256, "stream": False},
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert result["messages"] == messages
assert "system" not in result
def test_bedrock_claude_4_8_plus_cost_map_entries_carry_mid_conversation_system_flag():
"""Exact cost-map hits resolve before fallback-generalization rules, so a
mapped Bedrock Claude 4.8+ entry without ``supports_mid_conversation_system``
silently loses the cache-preserving in-place handling that the
``bedrock-anthropic-claude-mid-conversation-system`` rule grants unmapped
ids. Every mapped entry the rule's own pattern matches must carry the flag
explicitly."""
import re
import litellm
cost_map_path = os.path.join(os.path.dirname(litellm.__file__), "model_prices_and_context_window_backup.json")
with open(cost_map_path) as f:
cost_map = json.load(f)
rules = cost_map["fallback_generalizations"]["rules"]
pattern = re.compile(
next(r["pattern"] for r in rules if r["name"] == "bedrock-anthropic-claude-mid-conversation-system"),
re.IGNORECASE,
)
missing = [
key
for key, info in cost_map.items()
if isinstance(info, dict)
and str(info.get("litellm_provider", "")).startswith("bedrock")
and pattern.search(key)
and info.get("supports_mid_conversation_system") is not True
]
assert missing == []
def test_as_system_content_blocks_handles_each_shape():
"""``_as_system_content_blocks`` normalizes every system shape: ``None`` -> empty,
a string -> a single text block, a list -> a shallow copy, and any other value

View file

@ -1,7 +1,7 @@
# stdlib imports
import os
import sys
from unittest.mock import patch
from unittest.mock import Mock, patch
import pytest
from click.testing import CliRunner
@ -36,6 +36,35 @@ def test_cli_version_flag(cli_runner):
assert "LiteLLM Proxy Server Version: 1.2.3" in result.output
def test_base_url_trailing_slash_normalized(cli_runner):
"""A trailing slash on --base-url must not produce a double slash (e.g. '//sso/cli/start')."""
with (
patch("webbrowser.open"),
patch(
"requests.post",
return_value=Mock(
status_code=200,
json=Mock(
return_value={
"login_id": "cli-test-uuid",
"poll_secret": "poll-secret",
"user_code": "ABCD-EFGH",
}
),
raise_for_status=Mock(),
),
) as mock_post,
patch("requests.get", side_effect=ValueError("stop after start request")),
):
cli_runner.invoke(
cli, ["--base-url", "https://gateway.litellm-sandbox.ai/", "login"]
)
mock_post.assert_called_once_with(
"https://gateway.litellm-sandbox.ai/sso/cli/start", timeout=10
)
def test_cli_version_command(cli_runner):
"""Test that 'version' command prints the correct version, server URL, and server version, and exits successfully"""
with (

View file

@ -1959,6 +1959,110 @@ class TestUsageTransformation:
assert response_usage.output_tokens_details.text_tokens == 50
assert response_usage.output_tokens_details.image_tokens == 100
def test_reasoning_tokens_not_forced_to_zero_when_absent(self):
# Regression: previously the else branch wrote reasoning_tokens=0 even when
# completion_tokens_details had no reasoning (reasoning_tokens=None). That caused
# the proxy to always report reasoning_tokens=0 for non-thinking responses.
usage = Usage(
prompt_tokens=10,
completion_tokens=50,
total_tokens=60,
completion_tokens_details=CompletionTokensDetailsWrapper(
text_tokens=50,
# reasoning_tokens intentionally absent -> None
),
)
chat_completion_response = ModelResponse(
id="test-response-id",
created=1234567890,
model="claude-haiku-4-5",
object="chat.completion",
usage=usage,
choices=[
Choices(
finish_reason="stop",
index=0,
message=Message(content="Hello!", role="assistant"),
)
],
)
response_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage(
chat_completion_response=chat_completion_response
)
assert response_usage.output_tokens_details is not None
assert response_usage.output_tokens_details.reasoning_tokens is None
def test_reasoning_tokens_preserved_when_thinking_occurred(self):
# Regression: reasoning_tokens must survive the chat->responses translation
# when the provider actually did thinking.
usage = Usage(
prompt_tokens=100,
completion_tokens=612,
total_tokens=712,
completion_tokens_details=CompletionTokensDetailsWrapper(
reasoning_tokens=512,
text_tokens=100,
),
)
chat_completion_response = ModelResponse(
id="test-response-id",
created=1234567890,
model="claude-haiku-4-5",
object="chat.completion",
usage=usage,
choices=[
Choices(
finish_reason="stop",
index=0,
message=Message(content="Hello!", role="assistant"),
)
],
)
response_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage(
chat_completion_response=chat_completion_response
)
assert response_usage.output_tokens_details is not None
assert response_usage.output_tokens_details.reasoning_tokens == 512
def test_reasoning_tokens_explicit_zero_preserved(self):
usage = Usage(
prompt_tokens=10,
completion_tokens=50,
total_tokens=60,
completion_tokens_details=CompletionTokensDetailsWrapper(
reasoning_tokens=0,
text_tokens=50,
),
)
chat_completion_response = ModelResponse(
id="test-response-id",
created=1234567890,
model="gpt-5.6",
object="chat.completion",
usage=usage,
choices=[
Choices(
finish_reason="stop",
index=0,
message=Message(content="Hello!", role="assistant"),
)
],
)
response_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage(
chat_completion_response=chat_completion_response
)
assert response_usage.output_tokens_details is not None
assert response_usage.output_tokens_details.reasoning_tokens == 0
class TestStreamingIDConsistency:
"""Test cases for consistent IDs across streaming events (issue #14962)"""

View file

@ -0,0 +1,357 @@
"""
Regression: in-stream error events (type="error", type="response.failed") must
raise instead of being returned as benign chunks, mirroring chat streaming
semantics (_handle_stream_fallback_error): non-retriable 4xx (except 429)
raise litellm.APIError directly; 429 and 5xx are wrapped in
MidStreamFallbackError so the Router's mid-stream fallback machinery fires.
Status mapping must consider both the OpenAI error `type` (e.g.
"invalid_request_error") and `code` (e.g. "invalid_prompt",
"rate_limit_exceeded") fields — previously only `code` was read, so
type-classified client errors fell through to 500.
Also covers: ErrorEventError.param must accept dict payloads without raising a
Pydantic ValidationError (previously typed as Optional[str]).
"""
import json
import os
import sys
from unittest.mock import Mock, patch
import pytest
sys.path.insert(0, os.path.abspath("../.."))
import litellm
from litellm.exceptions import MidStreamFallbackError
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.responses.streaming_iterator import (
BaseResponsesAPIStreamingIterator,
ResponsesAPIStreamingIterator,
SyncResponsesAPIStreamingIterator,
)
from litellm.types.llms.openai import (
ErrorEvent,
ErrorEventError,
ResponseAPIUsage,
ResponsesAPIStreamEvents,
)
def _make_iterator() -> BaseResponsesAPIStreamingIterator:
mock_logging_obj = Mock(spec=LiteLLMLoggingObj)
mock_logging_obj.model_call_details = {"litellm_params": {}}
mock_config = Mock(spec=BaseResponsesAPIConfig)
mock_response = Mock()
mock_response.headers = {}
return BaseResponsesAPIStreamingIterator(
response=mock_response,
model="gpt-5",
responses_api_provider_config=mock_config,
logging_obj=mock_logging_obj,
custom_llm_provider="openai",
)
def _make_error_chunk(error_type: str, code: str, message: str = "err") -> ErrorEvent:
error_obj = ErrorEventError(type=error_type, code=code, message=message)
return ErrorEvent(type=ResponsesAPIStreamEvents.ERROR, sequence_number=0, error=error_obj)
def test_maybe_raise_for_error_event_wraps_unknown_error_in_mid_stream_fallback():
iterator = _make_iterator()
chunk = _make_error_chunk("server_error", "internal_error", "something went wrong")
with pytest.raises(MidStreamFallbackError) as exc_info:
iterator._maybe_raise_for_error_event(chunk)
assert exc_info.value.status_code == 500
assert isinstance(exc_info.value.original_exception, litellm.APIError)
assert exc_info.value.original_exception.status_code == 500
def test_maybe_raise_for_error_event_maps_rate_limit_code_to_429_mid_stream_fallback():
"""429 is retriable: it must be wrapped so the Router can fall back, carrying the mapped APIError."""
iterator = _make_iterator()
chunk = _make_error_chunk("tokens", "rate_limit_exceeded", "Too many requests")
with pytest.raises(MidStreamFallbackError) as exc_info:
iterator._maybe_raise_for_error_event(chunk)
assert exc_info.value.status_code == 429
assert exc_info.value.generated_content == ""
assert exc_info.value.is_pre_first_chunk is True
assert isinstance(exc_info.value.original_exception, litellm.APIError)
assert exc_info.value.original_exception.status_code == 429
def test_maybe_raise_for_error_event_maps_invalid_request_type_to_400():
"""Client errors classified via the `type` field must raise APIError directly (no fallback)."""
iterator = _make_iterator()
chunk = _make_error_chunk("invalid_request_error", "invalid_prompt", "bad request")
with pytest.raises(litellm.APIError) as exc_info:
iterator._maybe_raise_for_error_event(chunk)
assert exc_info.value.status_code == 400
assert not isinstance(exc_info.value, MidStreamFallbackError)
def test_maybe_raise_for_error_event_maps_context_length_code_to_400():
"""Client errors classified via the `code` field alone must still map to 400."""
iterator = _make_iterator()
chunk = Mock()
chunk.type = "error"
chunk.error = {"code": "context_length_exceeded", "message": "too long"}
with pytest.raises(litellm.APIError) as exc_info:
iterator._maybe_raise_for_error_event(chunk)
assert exc_info.value.status_code == 400
assert not isinstance(exc_info.value, MidStreamFallbackError)
def test_maybe_raise_for_error_event_maps_insufficient_quota_to_429():
"""OpenAI returns HTTP 429 for insufficient_quota; it must not map to 400 even though its type
is invalid_request_error-adjacent, and it must be wrapped for fallback."""
iterator = _make_iterator()
chunk = _make_error_chunk("invalid_request_error", "insufficient_quota", "You exceeded your current quota")
with pytest.raises(MidStreamFallbackError) as exc_info:
iterator._maybe_raise_for_error_event(chunk)
assert exc_info.value.status_code == 429
def test_maybe_raise_for_error_event_passes_through_normal_chunk():
iterator = _make_iterator()
chunk = Mock()
chunk.type = ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA
iterator._maybe_raise_for_error_event(chunk) # must not raise
def test_error_event_error_param_accepts_dict():
error_obj = ErrorEventError(
type="invalid_request_error",
code="context_length_exceeded",
message="too long",
param={"field": "messages", "index": 0},
)
assert isinstance(error_obj.param, dict)
def _make_async_iterator_with_events(events: list) -> ResponsesAPIStreamingIterator:
sse_payload = b"".join(f"data: {json.dumps(event)}\n\n".encode() for event in events)
async def mock_aiter_bytes():
yield sse_payload
mock_response = Mock()
mock_response.headers = {}
mock_response.aiter_bytes = mock_aiter_bytes
mock_logging_obj = Mock(spec=LiteLLMLoggingObj)
mock_logging_obj.model_call_details = {"litellm_params": {}}
mock_logging_obj.completion_start_time = None
mock_config = Mock(spec=BaseResponsesAPIConfig)
def transform(model, parsed_chunk, logging_obj):
if parsed_chunk.get("type") == "error":
return ErrorEvent(
type=ResponsesAPIStreamEvents.ERROR,
sequence_number=0,
error=ErrorEventError(**parsed_chunk["error"]),
)
delta_event = Mock()
delta_event.type = ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA
delta_event.delta = parsed_chunk.get("delta", "")
return delta_event
mock_config.transform_streaming_response.side_effect = transform
return ResponsesAPIStreamingIterator(
response=mock_response,
model="gpt-5",
responses_api_provider_config=mock_config,
logging_obj=mock_logging_obj,
custom_llm_provider="openai",
)
@pytest.mark.asyncio
async def test_async_iterator_raises_mid_stream_fallback_on_rate_limit_error_event():
iterator = _make_async_iterator_with_events(
[
{
"type": "error",
"error": {"type": "tokens", "code": "rate_limit_exceeded", "message": "rate limited"},
}
]
)
with pytest.raises(MidStreamFallbackError) as exc_info:
async for _ in iterator:
pass
assert exc_info.value.status_code == 429
assert exc_info.value.is_pre_first_chunk is True
assert exc_info.value.generated_content == ""
assert isinstance(exc_info.value.original_exception, litellm.APIError)
assert exc_info.value.original_exception.status_code == 429
@pytest.mark.asyncio
async def test_async_iterator_error_after_first_chunk_carries_generated_content():
"""An error after streamed output must expose the accumulated text so the router's
fallback can build a continuation input instead of restarting from scratch."""
iterator = _make_async_iterator_with_events(
[
{"type": "response.output_text.delta", "delta": "hello "},
{"type": "response.output_text.delta", "delta": "world"},
{
"type": "error",
"error": {"type": "server_error", "code": "internal_error", "message": "boom"},
},
]
)
chunks = []
with pytest.raises(MidStreamFallbackError) as exc_info:
async for chunk in iterator:
chunks.append(chunk)
assert len(chunks) == 2
assert exc_info.value.status_code == 500
assert exc_info.value.is_pre_first_chunk is False
assert exc_info.value.generated_content == "hello world"
def test_maybe_raise_for_response_failed_event_with_dict_error():
"""response.failed chunks carry a dict error on .response.error; covers dict branch."""
iterator = _make_iterator()
mock_response_obj = Mock()
mock_response_obj.error = {"type": "tokens", "code": "rate_limit_exceeded", "message": "throttled"}
chunk = Mock()
chunk.type = "response.failed"
chunk.response = mock_response_obj
with pytest.raises(MidStreamFallbackError) as exc_info:
iterator._maybe_raise_for_error_event(chunk)
assert exc_info.value.status_code == 429
def test_maybe_raise_for_error_event_null_error_obj():
"""error chunk with no error field: message and code default; wrapped as 500."""
iterator = _make_iterator()
chunk = Mock()
chunk.type = "error"
chunk.error = None
with pytest.raises(MidStreamFallbackError) as exc_info:
iterator._maybe_raise_for_error_event(chunk)
assert exc_info.value.status_code == 500
assert "Response API in-stream error" in str(exc_info.value)
def _make_failed_chunk(error: dict, usage: ResponseAPIUsage | None = None) -> Mock:
mock_response_obj = Mock()
mock_response_obj.error = error
mock_response_obj.usage = usage
chunk = Mock()
chunk.type = "response.failed"
chunk.response = mock_response_obj
return chunk
def test_handle_logging_failed_response_maps_rate_limit_to_429():
"""The exception logged to failure handlers must carry the mapped status, not a hardcoded 500."""
iterator = _make_iterator()
iterator.completed_response = _make_failed_chunk(
{"type": "tokens", "code": "rate_limit_exceeded", "message": "throttled"}
)
with (
patch("litellm.responses.streaming_iterator.run_async_function") as mock_run_async,
patch("litellm.responses.streaming_iterator.executor"),
):
iterator._handle_logging_failed_response()
logged_exception = mock_run_async.call_args.kwargs["exception"]
assert isinstance(logged_exception, litellm.APIError)
assert logged_exception.status_code == 429
assert "throttled" in str(logged_exception)
def test_handle_logging_failed_response_maps_type_field_to_400():
"""Status derivation for failed-response logging must also read the error `type` field."""
iterator = _make_iterator()
iterator.completed_response = _make_failed_chunk(
{"type": "invalid_request_error", "code": "invalid_prompt", "message": "bad prompt"}
)
with (
patch("litellm.responses.streaming_iterator.run_async_function") as mock_run_async,
patch("litellm.responses.streaming_iterator.executor"),
):
iterator._handle_logging_failed_response()
logged_exception = mock_run_async.call_args.kwargs["exception"]
assert isinstance(logged_exception, litellm.APIError)
assert logged_exception.status_code == 400
def test_handle_logging_failed_response_records_usage_and_cost():
"""Usage on a response.failed event must reach failure spend accounting via combined_usage_object."""
iterator = _make_iterator()
usage = ResponseAPIUsage(input_tokens=10, output_tokens=5, total_tokens=15)
chunk = _make_failed_chunk(
{"type": "server_error", "code": "server_error", "message": "boom"},
usage=usage,
)
iterator.completed_response = chunk
iterator.logging_obj._response_cost_calculator.return_value = 0.0042
with (
patch("litellm.responses.streaming_iterator.run_async_function"),
patch("litellm.responses.streaming_iterator.executor"),
):
iterator._handle_logging_failed_response()
combined_usage = iterator.logging_obj.model_call_details["combined_usage_object"]
assert isinstance(combined_usage, litellm.Usage)
assert combined_usage.prompt_tokens == 10
assert combined_usage.completion_tokens == 5
assert combined_usage.total_tokens == 15
assert iterator.logging_obj.model_call_details["response_cost"] == 0.0042
iterator.logging_obj._response_cost_calculator.assert_called_once_with(result=chunk.response)
def test_handle_logging_failed_response_without_usage_skips_recording():
iterator = _make_iterator()
iterator.completed_response = _make_failed_chunk(
{"type": "server_error", "code": "server_error", "message": "boom"}
)
with (
patch("litellm.responses.streaming_iterator.run_async_function"),
patch("litellm.responses.streaming_iterator.executor"),
):
iterator._handle_logging_failed_response()
assert "combined_usage_object" not in iterator.logging_obj.model_call_details
iterator.logging_obj._response_cost_calculator.assert_not_called()
def test_sync_iterator_raises_mid_stream_fallback_on_rate_limit_error_event():
"""SyncResponsesAPIStreamingIterator must wrap retriable error events for fallback."""
error_payload = {
"type": "error",
"error": {"type": "tokens", "code": "rate_limit_exceeded", "message": "throttled"},
}
sse_bytes = f"data: {json.dumps(error_payload)}\n\n".encode()
mock_response = Mock()
mock_response.headers = {}
mock_response.iter_bytes.return_value = iter([sse_bytes])
mock_logging_obj = Mock(spec=LiteLLMLoggingObj)
mock_logging_obj.model_call_details = {"litellm_params": {}}
mock_logging_obj.completion_start_time = None
mock_config = Mock(spec=BaseResponsesAPIConfig)
error_obj = ErrorEventError(type="tokens", code="rate_limit_exceeded", message="throttled")
mock_config.transform_streaming_response.return_value = ErrorEvent(
type=ResponsesAPIStreamEvents.ERROR, sequence_number=0, error=error_obj
)
iterator = SyncResponsesAPIStreamingIterator(
response=mock_response,
model="gpt-5",
responses_api_provider_config=mock_config,
logging_obj=mock_logging_obj,
custom_llm_provider="openai",
)
with pytest.raises(MidStreamFallbackError) as exc_info:
for _ in iterator:
pass
assert exc_info.value.status_code == 429
assert isinstance(exc_info.value.original_exception, litellm.APIError)

View file

@ -6,10 +6,11 @@ Tests the rule-based complexity scoring and tier assignment logic.
import os
import sys
from typing import Dict, List
from unittest.mock import MagicMock, patch
from typing import Dict
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from pydantic import ValidationError
sys.path.insert(
0, os.path.abspath("../../..")
@ -743,7 +744,7 @@ class TestKeywordFalsePositives:
# Should NOT detect code presence from 'api' in 'capital'
assert not any(
"code" in s.lower() for s in signals
), f"False positive: got code signal from 'capital'"
), "False positive: got code signal from 'capital'"
# Should be SIMPLE (definition question)
assert tier == ComplexityTier.SIMPLE
@ -754,7 +755,7 @@ class TestKeywordFalsePositives:
# Should NOT detect code presence from 'git' in 'digital'
assert not any(
"code" in s.lower() for s in signals
), f"False positive: got code signal from 'digital'"
), "False positive: got code signal from 'digital'"
def test_try_not_in_entry(self, complexity_router):
"""'try' should not match in 'entry'."""
@ -770,7 +771,7 @@ class TestKeywordFalsePositives:
tier, score, signals = complexity_router.classify(prompt)
assert not any(
"code" in s.lower() for s in signals
), f"False positive: got code signal from 'terrorism'"
), "False positive: got code signal from 'terrorism'"
def test_class_not_in_classical(self, complexity_router):
"""'class' should not match in 'classical'."""
@ -778,7 +779,7 @@ class TestKeywordFalsePositives:
tier, score, signals = complexity_router.classify(prompt)
assert not any(
"code" in s.lower() for s in signals
), f"False positive: got code signal from 'classical'"
), "False positive: got code signal from 'classical'"
def test_merge_not_in_emerged(self, complexity_router):
"""'merge' should not match in 'emerged'."""
@ -786,7 +787,7 @@ class TestKeywordFalsePositives:
tier, score, signals = complexity_router.classify(prompt)
assert not any(
"code" in s.lower() for s in signals
), f"False positive: got code signal from 'emerged'"
), "False positive: got code signal from 'emerged'"
def test_actual_api_keyword_detected(self, complexity_router):
"""Actual 'api' usage should be detected."""
@ -1131,3 +1132,180 @@ class TestExtractUserMessageAndSystemPrompt:
)
assert user_msg is None
assert sys_prompt is None
def _llm_response(content: str):
"""Build a fake acompletion response with the given message content."""
response = MagicMock()
response.choices = [MagicMock()]
response.choices[0].message.content = content
return response
@pytest.fixture
def llm_classifier_config() -> Dict:
"""Config with an LLM-based classifier wired to a 'haiku-classifier' model."""
return {
"tiers": {
"SIMPLE": "gpt-4o-mini",
"MEDIUM": "gpt-4o",
"COMPLEX": "claude-sonnet-4-20250514",
"REASONING": "o1-preview",
},
"classifier_type": "llm",
"classifier_llm_config": {"model": "haiku-classifier", "timeout_ms": 400},
}
@pytest.fixture
def llm_complexity_router(mock_router_instance, llm_classifier_config):
"""ComplexityRouter configured to classify via an LLM call."""
return ComplexityRouter(
model_name="test-complexity-router",
litellm_router_instance=mock_router_instance,
complexity_router_config=llm_classifier_config,
)
class TestLLMClassifierConfig:
"""Test config validation for the LLM classifier option."""
def test_llm_classifier_type_requires_config(self):
"""classifier_type='llm' without classifier_llm_config must raise."""
with pytest.raises(ValidationError):
ComplexityRouterConfig(classifier_type="llm")
def test_heuristic_classifier_type_needs_no_llm_config(self):
"""classifier_type='heuristic' (the default) needs no classifier_llm_config."""
config = ComplexityRouterConfig()
assert config.classifier_type == "heuristic"
assert config.classifier_llm_config is None
class TestLLMClassifier:
"""Test the LLM-based classifier path (aclassify) and its fallback behavior."""
@pytest.mark.asyncio
async def test_aclassify_heuristic_skips_llm_call(self, complexity_router, mock_router_instance):
"""When classifier_type is 'heuristic' (default), aclassify must not call the LLM."""
mock_router_instance.acompletion = AsyncMock()
tier, score, signals = await complexity_router.aclassify("Hello!")
mock_router_instance.acompletion.assert_not_called()
assert tier == ComplexityTier.SIMPLE
@pytest.mark.asyncio
async def test_aclassify_llm_success_routes_by_llm_verdict(
self, llm_complexity_router, mock_router_instance
):
"""A well-formed structured LLM response should decide the tier directly.
Uses a prompt that heuristic scoring alone would classify as SIMPLE, to prove
the LLM verdict -- not the heuristic scorer -- is what decided the tier.
"""
mock_router_instance.acompletion = AsyncMock(
return_value=_llm_response('{"tier": "COMPLEX"}')
)
tier, score, signals = await llm_complexity_router.aclassify("hi")
assert tier == ComplexityTier.COMPLEX
assert "llm-classifier:COMPLEX" in signals
mock_router_instance.acompletion.assert_awaited_once()
call_kwargs = mock_router_instance.acompletion.call_args.kwargs
assert call_kwargs["model"] == "haiku-classifier"
assert call_kwargs["timeout"] == 0.4
@pytest.mark.asyncio
async def test_aclassify_forwards_request_metadata_for_spend_tracking(
self, llm_complexity_router, mock_router_instance
):
"""The classifier call must carry the original request's metadata.
Without this, the proxy's cost-tracking gate (_should_track_cost_callback)
sees no user_api_key/team_id/user_id and silently drops all spend logging
and budget accounting for the classifier call.
"""
mock_router_instance.acompletion = AsyncMock(
return_value=_llm_response('{"tier": "SIMPLE"}')
)
request_metadata = {"user_api_key": "sk-abc", "user_api_key_team_id": "team-1"}
await llm_complexity_router.aclassify(
"hi", request_kwargs={"litellm_metadata": request_metadata}
)
call_kwargs = mock_router_instance.acompletion.call_args.kwargs
assert call_kwargs["metadata"] == request_metadata
@pytest.mark.asyncio
async def test_aclassify_strips_budget_reservation_from_classifier_metadata(
self, llm_complexity_router, mock_router_instance
):
"""The classifier call must not receive the parent request's budget reservation.
The reservation belongs to the routed completion the classifier is deciding
on, not to this internal classifier call. Forwarding it would let the
classifier's own cost-tracking reconcile against a reservation it has no
business touching, so it must be stripped while the rest of the attribution
metadata (key/team) is preserved.
"""
mock_router_instance.acompletion = AsyncMock(
return_value=_llm_response('{"tier": "SIMPLE"}')
)
request_metadata = {
"user_api_key": "sk-abc",
"user_api_key_team_id": "team-1",
"user_api_key_budget_reservation": {"reserved_cost": 1.0},
"user_api_key_auth": {"budget_reservation": {"reserved_cost": 1.0}},
}
await llm_complexity_router.aclassify(
"hi", request_kwargs={"litellm_metadata": request_metadata}
)
call_kwargs = mock_router_instance.acompletion.call_args.kwargs
assert call_kwargs["metadata"] == {
"user_api_key": "sk-abc",
"user_api_key_team_id": "team-1",
}
@pytest.mark.asyncio
async def test_aclassify_falls_back_to_heuristic_on_llm_exception(
self, llm_complexity_router, mock_router_instance
):
"""A timeout/error from the classifier model must fall back to heuristic scoring."""
mock_router_instance.acompletion = AsyncMock(side_effect=TimeoutError("classifier timed out"))
tier, score, signals = await llm_complexity_router.aclassify("Hello!")
assert tier == llm_complexity_router.classify("Hello!")[0]
assert tier == ComplexityTier.SIMPLE
@pytest.mark.asyncio
async def test_aclassify_falls_back_to_heuristic_on_unparseable_response(
self, llm_complexity_router, mock_router_instance
):
"""Non-JSON or schema-violating output must fall back to heuristic scoring, not raise."""
mock_router_instance.acompletion = AsyncMock(return_value=_llm_response("not json"))
tier, score, signals = await llm_complexity_router.aclassify("Hello!")
assert tier == ComplexityTier.SIMPLE
@pytest.mark.asyncio
async def test_aclassify_falls_back_to_heuristic_on_empty_content(
self, llm_complexity_router, mock_router_instance
):
"""Empty/None message content (e.g. provider quirk) must fall back, not raise."""
mock_router_instance.acompletion = AsyncMock(return_value=_llm_response(None))
tier, score, signals = await llm_complexity_router.aclassify("Hello!")
assert tier == ComplexityTier.SIMPLE
@pytest.mark.asyncio
async def test_pre_routing_hook_uses_llm_classifier_end_to_end(
self, llm_complexity_router, mock_router_instance
):
"""The full pre-routing hook should route using the LLM classifier's verdict."""
mock_router_instance.acompletion = AsyncMock(
return_value=_llm_response('{"tier": "REASONING"}')
)
request_metadata = {"user_api_key": "sk-abc", "user_api_key_team_id": "team-1"}
result = await llm_complexity_router.async_pre_routing_hook(
model="test-model",
request_kwargs={"litellm_metadata": request_metadata},
messages=[{"role": "user", "content": "hi"}],
)
assert result is not None
assert result.model == "o1-preview" # REASONING tier model
call_kwargs = mock_router_instance.acompletion.call_args.kwargs
assert call_kwargs["metadata"] == request_metadata

View file

@ -122,6 +122,8 @@ def test_azure_gpt_5_6_regional_model_info(model):
assert info is not None, f"{model} not found in model_prices_and_context_window.json"
assert info["litellm_provider"] == "azure"
assert info["mode"] == "chat"
input_cost, output_cost, cache_read_cost, _ = STANDARD_PRICING[_tier_key(model)]
assert info["input_cost_per_token"] == pytest.approx(input_cost * 1.1)
@ -132,6 +134,13 @@ def test_azure_gpt_5_6_regional_model_info(model):
assert info["input_cost_per_token_priority"] == pytest.approx(input_cost * 2.75)
assert info["output_cost_per_token_priority"] == pytest.approx(output_cost * 2.75)
assert info["max_input_tokens"] == 1050000
assert info["max_output_tokens"] == 128000
assert info["supports_reasoning"] is True
_, provider, _, _ = get_llm_provider(model=model)
assert provider == "azure"
def test_gpt_5_6_backup_matches_main():
"""Ensure the bundled model cost map stays in sync with the canonical file."""

View file

@ -269,29 +269,30 @@ def test_transform_usage_with_zero_values():
"""
Test transformation when token details are explicitly set to 0.
This ensures 0 values are preserved and not treated as None.
cached_tokens=0 is preserved (cache was available; nothing was cached).
reasoning_tokens=0 is preserved the same way: an explicit provider-reported
zero passes through, while an absent value (None) is omitted.
"""
completion_response = create_mock_completion_response(
model="gpt-4",
prompt_tokens=100,
completion_tokens=50,
total_tokens=150,
cached_tokens=0, # Explicitly 0
reasoning_tokens=0, # Explicitly 0
cached_tokens=0, # Explicitly 0 — preserved
reasoning_tokens=0, # Explicitly 0 — preserved
)
responses_usage = LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage(
completion_response
)
# Should preserve 0 values
assert responses_usage.input_tokens_details is not None
assert responses_usage.input_tokens_details.cached_tokens == 0
assert responses_usage.output_tokens_details is not None
assert responses_usage.output_tokens_details.reasoning_tokens == 0
print("✓ Transformation preserves explicit 0 values")
print("✓ Transformation preserves explicit reasoning_tokens=0 and omits absent values")
def test_input_tokens_details_requires_cached_tokens():
@ -315,25 +316,23 @@ def test_input_tokens_details_requires_cached_tokens():
print("✓ InputTokensDetails correctly defaults cached_tokens to 0")
def test_output_tokens_details_requires_reasoning_tokens():
def test_output_tokens_details_reasoning_tokens():
"""
Test that OutputTokensDetails has reasoning_tokens as an int with default value 0.
Test OutputTokensDetails.reasoning_tokens field semantics.
This ensures backward compatibility while making the field non-optional.
reasoning_tokens is Optional[int] = None: present only when reasoning actually occurred.
"""
# Should work with reasoning_tokens=0
details1 = OutputTokensDetails(reasoning_tokens=0)
assert details1.reasoning_tokens == 0
details_explicit_zero = OutputTokensDetails(reasoning_tokens=0)
assert details_explicit_zero.reasoning_tokens == 0
# Should work with reasoning_tokens=100
details2 = OutputTokensDetails(reasoning_tokens=100)
assert details2.reasoning_tokens == 100
details_positive = OutputTokensDetails(reasoning_tokens=100)
assert details_positive.reasoning_tokens == 100
# Should work without reasoning_tokens (defaults to 0)
details3 = OutputTokensDetails()
assert details3.reasoning_tokens == 0
# Default is None — absence means reasoning did not occur (or was not tracked)
details_default = OutputTokensDetails()
assert details_default.reasoning_tokens is None
print("✓ OutputTokensDetails correctly defaults reasoning_tokens to 0")
print("✓ OutputTokensDetails.reasoning_tokens defaults to None")
def test_all_providers_transformation_scenarios():
@ -419,7 +418,7 @@ if __name__ == "__main__":
test_transform_usage_with_both_token_details()
test_transform_usage_with_zero_values()
test_input_tokens_details_requires_cached_tokens()
test_output_tokens_details_requires_reasoning_tokens()
test_output_tokens_details_reasoning_tokens()
test_all_providers_transformation_scenarios()
print("\n" + "=" * 60)

View file

@ -855,6 +855,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"supports_xhigh_reasoning_effort": {"type": "boolean"},
"supports_max_reasoning_effort": {"type": "boolean"},
"supports_adaptive_thinking": {"type": "boolean"},
"supports_mid_conversation_system": {"type": "boolean"},
"supports_sampling_params": {"type": "boolean"},
"supports_output_config": {"type": "boolean"},
"supports_speed": {"type": "boolean"},
@ -1091,8 +1092,8 @@ def test_get_model_info_bedrock_regional_inference_profile_pricing(local_model_c
def test_get_model_info_bedrock_regional_profile_without_entry_falls_back_to_base(local_model_cost_map):
"""A regional profile with no dedicated cost-map entry must still resolve to its
region-stripped base entry."""
assert "jp.anthropic.claude-opus-4-8" not in litellm.model_cost
info = litellm.get_model_info(model="bedrock/jp.anthropic.claude-opus-4-8")
assert "apac.anthropic.claude-opus-4-8" not in litellm.model_cost
info = litellm.get_model_info(model="bedrock/apac.anthropic.claude-opus-4-8")
assert info["key"] == "anthropic.claude-opus-4-8"

View file

@ -1,8 +1,8 @@
{
"@typescript-eslint/no-explicit-any": 1977,
"complexity": 129,
"@typescript-eslint/no-explicit-any": 1978,
"complexity": 130,
"local/no-large-inline-object-arg": 509,
"local/no-long-condition-chain": 233,
"local/no-long-condition-chain": 234,
"max-depth": 59,
"no-console": 16
}

View file

@ -0,0 +1,135 @@
import React from "react";
import { describe, it, expect, vi, beforeEach } from "vitest";
import { screen, waitFor, within } from "@testing-library/react";
import { renderWithProviders } from "../../../../../tests/test-utils";
import CacheDashboard from "./cache_dashboard";
const { adminGlobalCacheActivity, cachingHealthCheckCall } = vi.hoisted(() => ({
adminGlobalCacheActivity: vi.fn(),
cachingHealthCheckCall: vi.fn(),
}));
vi.mock("@/components/networking", () => ({
adminGlobalCacheActivity,
cachingHealthCheckCall,
}));
const cacheActivity = [
{
api_key: "sk-1",
model: "gpt-5.1",
call_type: "acompletion",
total_rows: 1500,
cache_hit_true_rows: 300,
cached_completion_tokens: 12000,
generated_completion_tokens: 48000,
},
{
api_key: "sk-2",
model: "text-embedding-3-large",
call_type: "aembedding",
total_rows: 700,
cache_hit_true_rows: 100,
cached_completion_tokens: 2000,
generated_completion_tokens: 9000,
},
];
const renderDashboard = () =>
renderWithProviders(
<CacheDashboard accessToken="sk-test" token="tok" userRole="Admin" userID="u1" premiumUser={false} />,
);
const findChartCards = async () => {
await screen.findByText("Cache Hits vs API Requests");
await waitFor(() => {
expect(document.querySelectorAll("path.recharts-rectangle").length).toBeGreaterThan(0);
});
const cards = Array.from(document.querySelectorAll('[data-slot="card"]'));
expect(cards).toHaveLength(2);
return { requestsCard: cards[0] as HTMLElement, tokensCard: cards[1] as HTMLElement };
};
const barFills = (card: HTMLElement) =>
Array.from(card.querySelectorAll(".recharts-bar")).map((bar) =>
bar.querySelector("path.recharts-rectangle")?.getAttribute("fill"),
);
const legendFillByCategory = (card: HTMLElement) =>
Object.fromEntries(
Array.from(card.querySelectorAll('.recharts-legend-wrapper [style*="background-color"]')).map((swatch) => [
swatch.parentElement?.textContent,
swatch.getAttribute("style")?.match(/background-color:\s*([^;]+);?/)?.[1],
]),
);
describe("CacheDashboard cache analytics charts", () => {
beforeEach(() => {
vi.clearAllMocks();
adminGlobalCacheActivity.mockResolvedValue(cacheActivity);
});
it("renders both chart card titles", async () => {
renderDashboard();
expect(await screen.findByText("Cache Hits vs API Requests")).toBeInTheDocument();
expect(screen.getByText("Cached Completion Tokens vs Generated Completion Tokens")).toBeInTheDocument();
});
it("renders the requests chart with each category legend-bound to its fill and stacked in order", async () => {
renderDashboard();
const { requestsCard } = await findChartCards();
expect(legendFillByCategory(requestsCard)).toEqual({
"LLM API requests": "var(--color-sky-500, #0ea5e9)",
"Cache hit": "var(--color-teal-500, #14b8a6)",
});
expect(barFills(requestsCard)).toEqual(["var(--color-sky-500, #0ea5e9)", "var(--color-teal-500, #14b8a6)"]);
});
it("renders the tokens chart with each category legend-bound to its fill and stacked in order", async () => {
renderDashboard();
const { tokensCard } = await findChartCards();
expect(legendFillByCategory(tokensCard)).toEqual({
"Generated Completion Tokens": "var(--color-sky-500, #0ea5e9)",
"Cached Completion Tokens": "var(--color-teal-500, #14b8a6)",
});
expect(barFills(tokensCard)).toEqual(["var(--color-sky-500, #0ea5e9)", "var(--color-teal-500, #14b8a6)"]);
});
it("indexes bars by call_type name on the x axis", async () => {
renderDashboard();
const { requestsCard, tokensCard } = await findChartCards();
for (const card of [requestsCard, tokensCard]) {
expect(within(card).getAllByText("acompletion").length).toBeGreaterThan(0);
expect(within(card).getAllByText("aembedding").length).toBeGreaterThan(0);
}
});
it("stacks the two categories into one column per call_type", async () => {
renderDashboard();
const { requestsCard, tokensCard } = await findChartCards();
for (const card of [requestsCard, tokensCard]) {
const rects = Array.from(card.querySelectorAll("path.recharts-rectangle"));
expect(rects).toHaveLength(4);
const xPositions = rects.map((rect) => rect.getAttribute("d")?.split(",")[0]);
expect(new Set(xPositions).size).toBe(2);
}
});
it("formats y-axis ticks with compact notation", async () => {
renderDashboard();
const { requestsCard, tokensCard } = await findChartCards();
const compactTicks = (card: HTMLElement) =>
within(card)
.getAllByText(/^\d+(\.\d+)?K$/)
.map((tick) => tick.textContent);
expect(compactTicks(requestsCard).length).toBeGreaterThan(0);
expect(compactTicks(tokensCard)).toContain("60K");
});
});

View file

@ -1,5 +1,4 @@
import {
BarChart,
Card,
Col,
DateRangePickerValue,
@ -7,7 +6,6 @@ import {
Icon,
MultiSelect,
MultiSelectItem,
Subtitle,
Tab,
TabGroup,
TabList,
@ -18,6 +16,8 @@ import {
import React, { useEffect, useState } from "react";
import NotificationsManager from "@/components/molecules/notifications_manager";
import UsageDatePicker from "@/components/shared/usage_date_picker";
import { BarChart } from "@/components/shared/charts";
import { Card as ChartCard, CardContent, CardHeader, CardTitle } from "@/components/ui/card";
import { RefreshIcon } from "@heroicons/react/outline";
import { adminGlobalCacheActivity, cachingHealthCheckCall } from "@/components/networking";
@ -62,13 +62,13 @@ interface cacheDataItem {
// Add other properties as needed
}
interface uiData {
type uiData = {
name: string;
"LLM API requests": number;
"Cache hit": number;
"Cached Completion Tokens": number;
"Generated Completion Tokens": number;
}
};
interface CacheHealthResponse {
status?: string;
@ -350,29 +350,41 @@ const CacheDashboard: React.FC<CachePageProps> = ({ accessToken, token, userRole
</Card>
</div>
<Subtitle className="mt-4">Cache Hits vs API Requests</Subtitle>
<BarChart
title="Cache Hits vs API Requests"
data={filteredData}
stack={true}
index="name"
valueFormatter={valueFormatterNumbers}
categories={["LLM API requests", "Cache hit"]}
colors={["sky", "teal"]}
yAxisWidth={48}
/>
<ChartCard className="mt-4">
<CardHeader>
<CardTitle className="text-base font-semibold">Cache Hits vs API Requests</CardTitle>
</CardHeader>
<CardContent>
<BarChart
data={filteredData}
stack={true}
index="name"
valueFormatter={valueFormatterNumbers}
categories={["LLM API requests", "Cache hit"]}
colors={["sky", "teal"]}
yAxisWidth={48}
/>
</CardContent>
</ChartCard>
<Subtitle className="mt-4">Cached Completion Tokens vs Generated Completion Tokens</Subtitle>
<BarChart
className="mt-6"
data={filteredData}
stack={true}
index="name"
valueFormatter={valueFormatterNumbers}
categories={["Generated Completion Tokens", "Cached Completion Tokens"]}
colors={["sky", "teal"]}
yAxisWidth={48}
/>
<ChartCard className="mt-6">
<CardHeader>
<CardTitle className="text-base font-semibold">
Cached Completion Tokens vs Generated Completion Tokens
</CardTitle>
</CardHeader>
<CardContent>
<BarChart
data={filteredData}
stack={true}
index="name"
valueFormatter={valueFormatterNumbers}
categories={["Generated Completion Tokens", "Cached Completion Tokens"]}
colors={["sky", "teal"]}
yAxisWidth={48}
/>
</CardContent>
</ChartCard>
</Card>
</TabPanel>
<TabPanel>

View file

@ -1,7 +1,7 @@
import { renderWithProviders, screen, within } from "../../../tests/test-utils";
import { fireEvent, renderWithProviders, screen, within } from "../../../tests/test-utils";
import userEvent from "@testing-library/user-event";
import { vi } from "vitest";
import ComplexityRouterConfig from "./ComplexityRouterConfig";
import ComplexityRouterConfig, { ComplexityRouterConfigValue } from "./ComplexityRouterConfig";
const mockModelInfo = [
{ model_group: "gpt-4" },
@ -9,21 +9,24 @@ const mockModelInfo = [
{ model_group: "claude-3-opus" },
] as any[];
const defaultTiers = {
SIMPLE: "gpt-3.5-turbo",
MEDIUM: "gpt-3.5-turbo",
COMPLEX: "gpt-4",
REASONING: "claude-3-opus",
const defaultValue: ComplexityRouterConfigValue = {
tiers: {
SIMPLE: "gpt-3.5-turbo",
MEDIUM: "gpt-3.5-turbo",
COMPLEX: "gpt-4",
REASONING: "claude-3-opus",
},
classifier_type: "heuristic",
};
describe("ComplexityRouterConfig", () => {
it("should render", () => {
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultTiers} onChange={vi.fn()} />);
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultValue} onChange={vi.fn()} />);
expect(screen.getByText("Complexity Tier Configuration")).toBeInTheDocument();
});
it("should display all four tier labels", () => {
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultTiers} onChange={vi.fn()} />);
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultValue} onChange={vi.fn()} />);
expect(screen.getByText("Simple Tier")).toBeInTheDocument();
expect(screen.getByText("Medium Tier")).toBeInTheDocument();
expect(screen.getByText("Complex Tier")).toBeInTheDocument();
@ -31,7 +34,7 @@ describe("ComplexityRouterConfig", () => {
});
it("should show example queries for each tier", () => {
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultTiers} onChange={vi.fn()} />);
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultValue} onChange={vi.fn()} />);
expect(screen.getByText(/Hello!/)).toBeInTheDocument();
expect(screen.getByText(/Explain how REST APIs work/)).toBeInTheDocument();
expect(screen.getByText(/Design a microservices architecture/)).toBeInTheDocument();
@ -39,20 +42,56 @@ describe("ComplexityRouterConfig", () => {
});
it("should display the how classification works section", () => {
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultTiers} onChange={vi.fn()} />);
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultValue} onChange={vi.fn()} />);
expect(screen.getByText("How Classification Works")).toBeInTheDocument();
});
it("should show score thresholds in the classification section", () => {
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultTiers} onChange={vi.fn()} />);
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultValue} onChange={vi.fn()} />);
expect(screen.getByText(/Score < 0.15/)).toBeInTheDocument();
expect(screen.getByText(/Score 0.15 - 0.35/)).toBeInTheDocument();
expect(screen.getByText(/Score 0.35 - 0.60/)).toBeInTheDocument();
expect(screen.getByText(/Score > 0.60/)).toBeInTheDocument();
});
it("should default to heuristic and hide classifier model/timeout fields", () => {
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultValue} onChange={vi.fn()} />);
expect(screen.getByText("Advanced: Classification Method")).toBeInTheDocument();
expect(screen.queryByText("Classifier Model")).not.toBeInTheDocument();
});
it("should reveal classifier model and timeout fields when llm is selected", () => {
const onChange = vi.fn();
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultValue} onChange={onChange} />);
// Collapse panel content isn't rendered until first expanded.
fireEvent.click(screen.getByText("Advanced: Classification Method"));
fireEvent.click(screen.getByText("LLM Classifier"));
expect(onChange).toHaveBeenCalledWith({
...defaultValue,
classifier_type: "llm",
classifier_llm_config: { model: "", timeout_ms: 3000 },
});
});
it("should show classifier fields and use the configured values when classifier_type is llm", () => {
const llmValue: ComplexityRouterConfigValue = {
...defaultValue,
classifier_type: "llm",
classifier_llm_config: { model: "gpt-3.5-turbo", timeout_ms: 750 },
};
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={llmValue} onChange={vi.fn()} />);
fireEvent.click(screen.getByText("Advanced: Classification Method"));
expect(screen.getByText("Classifier Model")).toBeInTheDocument();
expect(screen.getByText("Timeout (ms)")).toBeInTheDocument();
expect(screen.getByDisplayValue("750")).toBeInTheDocument();
});
it("should render the custom technical keywords field", () => {
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultTiers} onChange={vi.fn()} />);
renderWithProviders(<ComplexityRouterConfig modelInfo={mockModelInfo} value={defaultValue} onChange={vi.fn()} />);
expect(screen.getByText("Custom Technical Keywords")).toBeInTheDocument();
});
@ -60,7 +99,7 @@ describe("ComplexityRouterConfig", () => {
renderWithProviders(
<ComplexityRouterConfig
modelInfo={mockModelInfo}
value={defaultTiers}
value={defaultValue}
onChange={vi.fn()}
customTechnicalKeywords={["udp", "kafka"]}
onCustomTechnicalKeywordsChange={vi.fn()}
@ -76,7 +115,7 @@ describe("ComplexityRouterConfig", () => {
renderWithProviders(
<ComplexityRouterConfig
modelInfo={mockModelInfo}
value={defaultTiers}
value={defaultValue}
onChange={vi.fn()}
customTechnicalKeywords={[]}
onCustomTechnicalKeywordsChange={onCustomTechnicalKeywordsChange}

View file

@ -1,21 +1,36 @@
import { InfoCircleOutlined } from "@ant-design/icons";
import { Select as AntdSelect, Card, Divider, Space, Tooltip, Typography } from "antd";
import { Select as AntdSelect, Card, Collapse, Divider, InputNumber, Radio, Space, Tooltip, Typography } from "antd";
import React from "react";
import { ModelGroup } from "@/components/llm_calls/fetch_models";
const { Text } = Typography;
interface ComplexityTiers {
export const DEFAULT_CLASSIFIER_TIMEOUT_MS = 3000;
export interface ComplexityTiers {
SIMPLE: string;
MEDIUM: string;
COMPLEX: string;
REASONING: string;
}
export interface ClassifierLLMConfig {
model: string;
timeout_ms: number;
}
export type ClassifierType = "heuristic" | "llm";
export interface ComplexityRouterConfigValue {
tiers: ComplexityTiers;
classifier_type: ClassifierType;
classifier_llm_config?: ClassifierLLMConfig;
}
interface ComplexityRouterConfigProps {
modelInfo: ModelGroup[];
value: ComplexityTiers;
onChange: (tiers: ComplexityTiers) => void;
value: ComplexityRouterConfigValue;
onChange: (value: ComplexityRouterConfigValue) => void;
customTechnicalKeywords?: string[];
onCustomTechnicalKeywordsChange?: (keywords: string[]) => void;
}
@ -59,7 +74,38 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
const handleTierChange = (tier: keyof ComplexityTiers, model: string) => {
onChange({
...value,
[tier]: model,
tiers: { ...value.tiers, [tier]: model },
});
};
const handleClassifierTypeChange = (classifierType: ClassifierType) => {
onChange({
...value,
classifier_type: classifierType,
classifier_llm_config:
classifierType === "llm"
? value.classifier_llm_config ?? { model: "", timeout_ms: DEFAULT_CLASSIFIER_TIMEOUT_MS }
: undefined,
});
};
const handleClassifierModelChange = (model: string) => {
onChange({
...value,
classifier_llm_config: {
model,
timeout_ms: value.classifier_llm_config?.timeout_ms ?? DEFAULT_CLASSIFIER_TIMEOUT_MS,
},
});
};
const handleClassifierTimeoutChange = (timeoutMs: number | null) => {
onChange({
...value,
classifier_llm_config: {
model: value.classifier_llm_config?.model ?? "",
timeout_ms: timeoutMs ?? DEFAULT_CLASSIFIER_TIMEOUT_MS,
},
});
};
@ -98,7 +144,7 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
Examples: {tierInfo.examples}
</Text>
<AntdSelect
value={value[tier]}
value={value.tiers[tier]}
onChange={(model) => handleTierChange(tier, model)}
placeholder={`Select model for ${tierInfo.label.toLowerCase()} queries`}
showSearch
@ -113,6 +159,76 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
<Divider />
<Collapse
ghost
style={{ background: "#f9fafb", borderRadius: 8, border: "1px solid #e5e7eb" }}
items={[
{
key: "classifier",
label: (
<Text strong style={{ color: "#374151" }}>
Advanced: Classification Method
</Text>
),
children: (
<>
<Radio.Group
value={value.classifier_type}
onChange={(e) => handleClassifierTypeChange(e.target.value)}
className="w-full"
>
<Space direction="vertical" className="w-full">
<Radio value="heuristic">
<Text strong>Heuristic</Text>{" "}
<Text type="secondary">(default) — rule-based scoring, no API calls, &lt;1ms latency</Text>
</Radio>
<Radio value="llm">
<Text strong>LLM Classifier</Text>{" "}
<Text type="secondary">— use a model to decide the tier (e.g. a small/fast model)</Text>
</Radio>
</Space>
</Radio.Group>
{value.classifier_type === "llm" && (
<div className="mt-4 space-y-3">
<div>
<Text strong style={{ display: "block", marginBottom: 4 }}>
Classifier Model
</Text>
<AntdSelect
value={value.classifier_llm_config?.model || undefined}
onChange={handleClassifierModelChange}
placeholder="Select the model that will classify request complexity"
showSearch
style={{ width: "100%" }}
options={modelOptions}
/>
</div>
<div>
<Text strong style={{ display: "block", marginBottom: 4 }}>
Timeout (ms)
</Text>
<InputNumber
value={value.classifier_llm_config?.timeout_ms ?? DEFAULT_CLASSIFIER_TIMEOUT_MS}
onChange={handleClassifierTimeoutChange}
min={1}
style={{ width: "100%" }}
/>
<Text type="secondary" style={{ fontSize: 12 }}>
Falls back to the heuristic scorer if the classifier call errors, times out, or returns an
unparseable response.
</Text>
</div>
</div>
)}
</>
),
},
]}
/>
<Divider />
<Card>
<div className="flex items-center gap-2 mb-2">
<Text strong style={{ fontSize: 16 }}>

View file

@ -8,7 +8,7 @@ import { all_admin_roles } from "@/utils/roles";
import { handleAddAutoRouterSubmit } from "./handle_add_auto_router_submit";
import { fetchAvailableModels, ModelGroup } from "@/components/llm_calls/fetch_models";
import RouterConfigBuilder from "./RouterConfigBuilder";
import ComplexityRouterConfig from "./ComplexityRouterConfig";
import ComplexityRouterConfig, { ComplexityRouterConfigValue } from "./ComplexityRouterConfig";
import NotificationManager from "../molecules/notifications_manager";
import { ThunderboltOutlined, BranchesOutlined } from "@ant-design/icons";
@ -21,13 +21,6 @@ interface AddAutoRouterTabProps {
type RouterType = "complexity" | "semantic";
interface ComplexityTiers {
SIMPLE: string;
MEDIUM: string;
COMPLEX: string;
REASONING: string;
}
const { Title, Link } = Typography;
const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, accessToken, userRole }) => {
@ -48,11 +41,9 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
const [routerConfig, setRouterConfig] = useState<any>(null);
// Complexity router config (new)
const [complexityTiers, setComplexityTiers] = useState<ComplexityTiers>({
SIMPLE: "",
MEDIUM: "",
COMPLEX: "",
REASONING: "",
const [complexityRouterConfig, setComplexityRouterConfig] = useState<ComplexityRouterConfigValue>({
tiers: { SIMPLE: "", MEDIUM: "", COMPLEX: "", REASONING: "" },
classifier_type: "heuristic",
});
const [customTechnicalKeywords, setCustomTechnicalKeywords] = useState<string[]>([]);
@ -99,15 +90,20 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
// Validation differs based on router type
if (routerType === "complexity") {
// Complexity Router validation
const filledTiers = Object.values(complexityTiers).filter(Boolean);
const { tiers, classifier_type, classifier_llm_config } = complexityRouterConfig;
const filledTiers = Object.values(tiers).filter(Boolean);
if (filledTiers.length === 0) {
NotificationManager.fromBackend("Please select at least one model for a complexity tier");
return;
}
if (classifier_type === "llm" && !classifier_llm_config?.model) {
NotificationManager.fromBackend("Please select a classifier model, or switch back to Heuristic");
return;
}
// For complexity router, use the first non-empty tier as default
const defaultModel =
complexityTiers.MEDIUM || complexityTiers.SIMPLE || complexityTiers.COMPLEX || complexityTiers.REASONING;
const defaultModel = tiers.MEDIUM || tiers.SIMPLE || tiers.COMPLEX || tiers.REASONING;
// Set form values for complexity router
form.setFieldsValue({
@ -128,7 +124,9 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
// Use special model prefix for complexity router
model_type: "complexity_router",
complexity_router_config: {
tiers: complexityTiers,
tiers,
classifier_type,
...(classifier_type === "llm" ? { classifier_llm_config } : {}),
...(customTechnicalKeywords.length > 0 && { custom_technical_keywords: customTechnicalKeywords }),
},
model_access_group: currentFormValues.model_access_group,
@ -279,9 +277,9 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
<div className="w-full mb-4">
<ComplexityRouterConfig
modelInfo={modelInfo}
value={complexityTiers}
onChange={(tiers) => {
setComplexityTiers(tiers);
value={complexityRouterConfig}
onChange={(config) => {
setComplexityRouterConfig(config);
}}
customTechnicalKeywords={customTechnicalKeywords}
onCustomTechnicalKeywordsChange={setCustomTechnicalKeywords}

View file

@ -4,8 +4,13 @@ import { Text, TextInput } from "@tremor/react";
import { modelAvailableCall, modelPatchUpdateCall } from "../networking";
import { fetchAvailableModels, ModelGroup } from "@/components/llm_calls/fetch_models";
import RouterConfigBuilder from "../add_model/RouterConfigBuilder";
import ComplexityRouterConfig, { ComplexityRouterConfigValue } from "../add_model/ComplexityRouterConfig";
import NotificationsManager from "../molecules/notifications_manager";
const isComplexityRouterModel = (modelData: any): boolean =>
modelData?.litellm_params?.model?.startsWith("auto_router/complexity_router") ||
modelData?.litellm_params?.complexity_router_config != null;
interface EditAutoRouterModalProps {
isVisible: boolean;
onCancel: () => void;
@ -30,6 +35,11 @@ const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
const [showCustomDefaultModel, setShowCustomDefaultModel] = useState<boolean>(false);
const [showCustomEmbeddingModel, setShowCustomEmbeddingModel] = useState<boolean>(false);
const [routerConfig, setRouterConfig] = useState<any>(null);
const [complexityRouterConfig, setComplexityRouterConfig] = useState<ComplexityRouterConfigValue>({
tiers: { SIMPLE: "", MEDIUM: "", COMPLEX: "", REASONING: "" },
classifier_type: "heuristic",
});
const isComplexityRouter = isComplexityRouterModel(modelData);
useEffect(() => {
if (isVisible && modelData) {
@ -66,6 +76,31 @@ const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
const initializeForm = () => {
try {
if (isComplexityRouterModel(modelData)) {
// Parse the complexity_router_config if it exists and is a string
let parsedConfig = modelData.litellm_params?.complexity_router_config || {};
if (typeof parsedConfig === "string") {
parsedConfig = JSON.parse(parsedConfig);
}
setComplexityRouterConfig({
tiers: {
SIMPLE: parsedConfig.tiers?.SIMPLE || "",
MEDIUM: parsedConfig.tiers?.MEDIUM || "",
COMPLEX: parsedConfig.tiers?.COMPLEX || "",
REASONING: parsedConfig.tiers?.REASONING || "",
},
classifier_type: parsedConfig.classifier_type || "heuristic",
classifier_llm_config: parsedConfig.classifier_llm_config,
});
form.setFieldsValue({
auto_router_name: modelData.model_name,
model_access_group: modelData.model_info?.access_groups || [],
});
return;
}
// Parse the auto_router_config if it exists and is a string
let parsedConfig = null;
if (modelData.litellm_params?.auto_router_config) {
@ -101,6 +136,49 @@ const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
setLoading(true);
const values = await form.validateFields();
if (isComplexityRouter) {
const { tiers, classifier_type, classifier_llm_config } = complexityRouterConfig;
if (Object.values(tiers).filter(Boolean).length === 0) {
NotificationsManager.fromBackend("Please select at least one model for a complexity tier");
return;
}
if (classifier_type === "llm" && !classifier_llm_config?.model) {
NotificationsManager.fromBackend("Please select a classifier model, or switch back to Heuristic");
return;
}
const defaultModel = tiers.MEDIUM || tiers.SIMPLE || tiers.COMPLEX || tiers.REASONING;
const updatedLitellmParams = {
...modelData.litellm_params,
complexity_router_config: {
tiers,
classifier_type,
...(classifier_type === "llm" ? { classifier_llm_config } : {}),
},
complexity_router_default_model: defaultModel,
};
const updatedModelInfo = {
...modelData.model_info,
access_groups: values.model_access_group || [],
};
await modelPatchUpdateCall(
accessToken,
{ model_name: values.auto_router_name, litellm_params: updatedLitellmParams, model_info: updatedModelInfo },
modelData.model_info.id,
);
NotificationsManager.success("Auto router configuration updated successfully");
onSuccess({
...modelData,
model_name: values.auto_router_name,
litellm_params: updatedLitellmParams,
model_info: updatedModelInfo,
});
onCancel();
return;
}
// Prepare the updated litellm_params
const updatedLitellmParams = {
...modelData.litellm_params,
@ -177,45 +255,60 @@ const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
<TextInput placeholder="e.g., auto_router_1, smart_routing" />
</Form.Item>
{/* Router Configuration Builder */}
<div className="w-full">
<RouterConfigBuilder
modelInfo={modelInfo}
value={routerConfig}
onChange={(config) => {
setRouterConfig(config);
}}
/>
</div>
{isComplexityRouter ? (
/* Complexity Router Configuration */
<div className="w-full">
<ComplexityRouterConfig
modelInfo={modelInfo}
value={complexityRouterConfig}
onChange={(config) => {
setComplexityRouterConfig(config);
}}
/>
</div>
) : (
<>
{/* Router Configuration Builder */}
<div className="w-full">
<RouterConfigBuilder
modelInfo={modelInfo}
value={routerConfig}
onChange={(config) => {
setRouterConfig(config);
}}
/>
</div>
{/* Default Model */}
<Form.Item
label="Default Model"
name="auto_router_default_model"
rules={[{ required: true, message: "Default model is required" }]}
>
<AntdSelect
placeholder="Select a default model"
onChange={(value) => {
setShowCustomDefaultModel(value === "custom");
}}
options={[...modelOptions, { value: "custom", label: "Enter custom model name" }]}
showSearch={true}
/>
</Form.Item>
{/* Default Model */}
<Form.Item
label="Default Model"
name="auto_router_default_model"
rules={[{ required: true, message: "Default model is required" }]}
>
<AntdSelect
placeholder="Select a default model"
onChange={(value) => {
setShowCustomDefaultModel(value === "custom");
}}
options={[...modelOptions, { value: "custom", label: "Enter custom model name" }]}
showSearch={true}
/>
</Form.Item>
{/* Embedding Model */}
<Form.Item label="Embedding Model" name="auto_router_embedding_model">
<AntdSelect
placeholder="Select an embedding model (optional)"
onChange={(value) => {
setShowCustomEmbeddingModel(value === "custom");
}}
options={[...modelOptions, { value: "custom", label: "Enter custom model name" }]}
showSearch={true}
allowClear
/>
</Form.Item>
{/* Embedding Model */}
<Form.Item label="Embedding Model" name="auto_router_embedding_model">
<AntdSelect
placeholder="Select an embedding model (optional)"
onChange={(value) => {
setShowCustomEmbeddingModel(value === "custom");
}}
options={[...modelOptions, { value: "custom", label: "Enter custom model name" }]}
showSearch={true}
allowClear
/>
</Form.Item>
</>
)}
{/* Model Access Groups - Admin only */}
{userRole === "Admin" && (

View file

@ -124,7 +124,10 @@ export default function ModelInfoView({
const canEditModel =
(userRole === "Admin" || modelData?.model_info?.created_by === userID) && modelData?.model_info?.db_model;
const isAdmin = userRole === "Admin";
const isAutoRouter = modelData?.litellm_params?.auto_router_config != null;
const isAutoRouter =
modelData?.litellm_params?.auto_router_config != null ||
modelData?.litellm_params?.complexity_router_config != null ||
modelData?.litellm_params?.model?.startsWith("auto_router/complexity_router");
const usingExistingCredential =
modelData?.litellm_params?.litellm_credential_name != null &&

File diff suppressed because one or more lines are too long