mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_/patch-endpoint-1d2646
This commit is contained in:
commit
472b64cc62
43 changed files with 2959 additions and 214 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)))
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 ###
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -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))",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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))",
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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)"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
});
|
||||
});
|
||||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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, <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 }}>
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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" && (
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue