mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge branch 'litellm_internal_staging' into litellm_e2e_coverage_generated_denominator
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
commit
2fe9d30f3b
38 changed files with 1900 additions and 152 deletions
|
|
@ -315,6 +315,11 @@ disable_token_counter: bool = False
|
|||
disable_add_transform_inline_image_block: bool = False
|
||||
disable_add_user_agent_to_request_tags: bool = False
|
||||
disable_anthropic_gemini_context_caching_transform: bool = False
|
||||
enable_anthropic_prompt_caching: bool = os.getenv("LITELLM_ENABLE_ANTHROPIC_PROMPT_CACHING", "false").lower() == "true"
|
||||
_anthropic_prompt_caching_ttl_env: Optional[str] = os.getenv("LITELLM_ANTHROPIC_PROMPT_CACHING_TTL")
|
||||
anthropic_prompt_caching_ttl: Optional[Literal["5m", "1h"]] = (
|
||||
"1h" if _anthropic_prompt_caching_ttl_env == "1h" else "5m" if _anthropic_prompt_caching_ttl_env == "5m" else None
|
||||
)
|
||||
disable_vertex_batch_output_transformation: bool = False
|
||||
extra_spend_tag_headers: Optional[List[str]] = None
|
||||
in_memory_llm_clients_cache: "LLMClientCache"
|
||||
|
|
|
|||
|
|
@ -296,18 +296,148 @@ class AnthropicCacheControlHook(CustomPromptManagement):
|
|||
|
||||
return processed_messages, processed_system, remaining_points
|
||||
|
||||
@staticmethod
|
||||
def _default_control() -> ChatCompletionCachedContent:
|
||||
"""Build the cache_control block for auto-injected breakpoints.
|
||||
|
||||
Defaults to Anthropic's 5-minute ephemeral cache; honors the optional
|
||||
``litellm.anthropic_prompt_caching_ttl`` override ("5m" or "1h").
|
||||
"""
|
||||
import litellm
|
||||
|
||||
ttl = litellm.anthropic_prompt_caching_ttl
|
||||
if ttl == "5m" or ttl == "1h":
|
||||
return ChatCompletionCachedContent(type="ephemeral", ttl=ttl)
|
||||
return ChatCompletionCachedContent(type="ephemeral")
|
||||
|
||||
@staticmethod
|
||||
def _request_has_cache_control(
|
||||
messages: list[AllMessageValues],
|
||||
system: str | list | None,
|
||||
tools: list | None = None,
|
||||
) -> bool:
|
||||
"""Return True if the request already carries any client-supplied cache_control.
|
||||
|
||||
When the client (e.g. Claude Code) already marks its own breakpoints we
|
||||
stand down entirely rather than add more, per the auto-caching contract.
|
||||
Tools count: they are a breakpoint the client can mark, they count toward
|
||||
the provider's four-block limit, and caching only the tool definitions is
|
||||
a common pattern, so injecting alongside them can exceed the cap.
|
||||
"""
|
||||
if any(AnthropicCacheControlHook._count_cache_control_blocks(msg) for msg in messages):
|
||||
return True
|
||||
if isinstance(system, list):
|
||||
if any(isinstance(block, dict) and block.get("cache_control") is not None for block in system):
|
||||
return True
|
||||
if tools is not None:
|
||||
return any(isinstance(tool, dict) and tool.get("cache_control") is not None for tool in tools)
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def get_default_injection_points(
|
||||
messages: list[AllMessageValues],
|
||||
system: str | list | None,
|
||||
model: str,
|
||||
custom_llm_provider: str | None,
|
||||
tools: list | None = None,
|
||||
) -> list[CacheControlInjectionPoint]:
|
||||
"""Default breakpoints when ``litellm.enable_anthropic_prompt_caching`` is on.
|
||||
|
||||
Caches the system prompt and the trailing turn, so the stable prefix
|
||||
(system + tools + history) is reused while the breakpoint advances with
|
||||
the conversation. Returns [] (stand down) when the flag is off, the
|
||||
provider does not consume cache_control breakpoints (only anthropic /
|
||||
bedrock do), the model lacks prompt-caching support, or the request
|
||||
already carries client-supplied cache_control.
|
||||
"""
|
||||
import litellm
|
||||
|
||||
if litellm.enable_anthropic_prompt_caching is not True:
|
||||
return []
|
||||
|
||||
provider = custom_llm_provider
|
||||
if provider is None:
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import (
|
||||
get_llm_provider,
|
||||
)
|
||||
|
||||
try:
|
||||
_, provider, _, _ = get_llm_provider(model=model)
|
||||
except Exception: # noqa: BLE001 # unroutable model must never block the call, just skip auto-caching
|
||||
return []
|
||||
|
||||
if provider not in ("anthropic", "bedrock"):
|
||||
return []
|
||||
|
||||
from litellm.utils import supports_prompt_caching
|
||||
|
||||
if not supports_prompt_caching(model=model, custom_llm_provider=provider):
|
||||
return []
|
||||
|
||||
if AnthropicCacheControlHook._request_has_cache_control(messages, system, tools):
|
||||
return []
|
||||
|
||||
control = AnthropicCacheControlHook._default_control()
|
||||
points: list[CacheControlInjectionPoint] = [
|
||||
CacheControlMessageInjectionPoint(location="message", role="system", index=None, control=control),
|
||||
CacheControlMessageInjectionPoint(location="message", role=None, index=-1, control=control),
|
||||
]
|
||||
return points
|
||||
|
||||
@staticmethod
|
||||
def maybe_seed_default_injection_points(
|
||||
non_default_params: dict[str, Any],
|
||||
messages: list[AllMessageValues],
|
||||
model: str,
|
||||
custom_llm_provider: str | None,
|
||||
tools: list | None = None,
|
||||
) -> None:
|
||||
"""For /chat/completions: add default injection points to the request params.
|
||||
|
||||
No-op when injection points are already configured (explicit config wins).
|
||||
Seeding the param lets the existing prompt-management gate and the
|
||||
AnthropicCacheControlHook run unchanged.
|
||||
"""
|
||||
if non_default_params.get("cache_control_injection_points"):
|
||||
return
|
||||
points = AnthropicCacheControlHook.get_default_injection_points(
|
||||
messages=messages,
|
||||
system=None,
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
tools=tools,
|
||||
)
|
||||
if points:
|
||||
non_default_params["cache_control_injection_points"] = points
|
||||
|
||||
@staticmethod
|
||||
def maybe_inject_cache_control(
|
||||
messages: List[Dict],
|
||||
system: str | list | None,
|
||||
kwargs: Dict[str, Any],
|
||||
model: str | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
tools: list[dict] | None = None,
|
||||
) -> Tuple[List[Dict], str | list | None]:
|
||||
"""Extract cache_control_injection_points from kwargs and apply if present.
|
||||
|
||||
When none are configured but ``litellm.enable_anthropic_prompt_caching``
|
||||
is on, synthesize default breakpoints for the native /v1/messages path.
|
||||
Pops the key from kwargs; if remaining (non-message) points exist they
|
||||
are written back so downstream transforms can handle them.
|
||||
"""
|
||||
injection_points = kwargs.pop("cache_control_injection_points", None)
|
||||
configured = cast( # cast-ok: kwargs is untyped; this key only holds the documented injection-point list
|
||||
list[CacheControlInjectionPoint] | None, kwargs.pop("cache_control_injection_points", None)
|
||||
)
|
||||
injection_points: list[CacheControlInjectionPoint] = configured or []
|
||||
if not injection_points and model is not None:
|
||||
injection_points = AnthropicCacheControlHook.get_default_injection_points(
|
||||
messages=cast(list[AllMessageValues], messages), # cast-ok: Anthropic-shaped dicts from v1/messages
|
||||
system=system,
|
||||
tools=tools,
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
if not injection_points:
|
||||
return messages, system
|
||||
|
||||
|
|
|
|||
|
|
@ -1453,6 +1453,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
response_cost = litellm.response_cost_calculator(**response_cost_calculator_kwargs)
|
||||
|
||||
verbose_logger.debug(f"response_cost: {response_cost}")
|
||||
additional_response_cost: object = self.model_call_details.get("additional_response_cost")
|
||||
if isinstance(additional_response_cost, (int, float)) and additional_response_cost > 0:
|
||||
return (response_cost or 0.0) + additional_response_cost
|
||||
return response_cost
|
||||
except Exception as e: # error calculating cost
|
||||
debug_info = StandardLoggingModelCostFailureDebugInformation(
|
||||
|
|
|
|||
|
|
@ -906,16 +906,17 @@ def strip_advisor_blocks_from_messages(messages: List[Any], replace_with_text: b
|
|||
|
||||
def is_anthropic_invalid_thinking_signature_error(error_text: str) -> bool:
|
||||
"""
|
||||
Detect Anthropic 400 when encrypted thinking signatures in history do not match
|
||||
the current deployment (e.g. user rotated API key or switched model endpoint).
|
||||
Detect Anthropic 400 errors caused by missing or invalid thinking signatures.
|
||||
|
||||
Example API message:
|
||||
Known error formats:
|
||||
{"message":"messages.2.content.0.thinking.signature.str: Input should be a valid string"}
|
||||
messages.N.content.M.thinking.signature.str: Input should be a valid string
|
||||
messages.N.content.M: Invalid `signature` in `thinking` block
|
||||
"""
|
||||
if not error_text:
|
||||
return False
|
||||
lower = error_text.lower()
|
||||
return "invalid" in lower and "signature" in lower and "thinking" in lower and "block" in lower
|
||||
return "thinking" in lower and "signature" in lower and ("invalid" in lower or "valid string" in lower)
|
||||
|
||||
|
||||
def strip_thinking_blocks_from_anthropic_messages(messages: List[Any]) -> List[Any]:
|
||||
|
|
|
|||
|
|
@ -237,7 +237,9 @@ async def anthropic_messages(
|
|||
AnthropicCacheControlHook,
|
||||
)
|
||||
|
||||
messages, system = AnthropicCacheControlHook.maybe_inject_cache_control(messages, system, kwargs)
|
||||
messages, system = AnthropicCacheControlHook.maybe_inject_cache_control(
|
||||
messages, system, kwargs, model=model, custom_llm_provider=custom_llm_provider, tools=tools
|
||||
)
|
||||
|
||||
original_stream = stream or kwargs.get("_websearch_interception_converted_stream", False)
|
||||
|
||||
|
|
@ -426,7 +428,9 @@ def anthropic_messages_handler(
|
|||
AnthropicCacheControlHook,
|
||||
)
|
||||
|
||||
messages, system = AnthropicCacheControlHook.maybe_inject_cache_control(messages, system, kwargs)
|
||||
messages, system = AnthropicCacheControlHook.maybe_inject_cache_control(
|
||||
messages, system, kwargs, model=model, custom_llm_provider=custom_llm_provider, tools=tools
|
||||
)
|
||||
|
||||
metadata = validate_anthropic_api_metadata(metadata)
|
||||
|
||||
|
|
|
|||
|
|
@ -75,10 +75,23 @@ def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]:
|
|||
model_info = get_model_info(model=base_model, custom_llm_provider="fireworks_ai")
|
||||
|
||||
## CALCULATE INPUT COST
|
||||
prompt_tokens_details = usage.prompt_tokens_details
|
||||
cached_tokens: int = (
|
||||
prompt_tokens_details.cached_tokens
|
||||
if prompt_tokens_details is not None and prompt_tokens_details.cached_tokens is not None
|
||||
else 0
|
||||
)
|
||||
input_cost_per_token: float = model_info["input_cost_per_token"] or 0.0
|
||||
cache_read_input_token_cost = model_info.get("cache_read_input_token_cost")
|
||||
cache_read_cost_per_token: float = (
|
||||
cache_read_input_token_cost if cache_read_input_token_cost is not None else input_cost_per_token
|
||||
)
|
||||
non_cached_prompt_tokens: int = max(usage.prompt_tokens - cached_tokens, 0)
|
||||
|
||||
prompt_cost: float = usage["prompt_tokens"] * model_info["input_cost_per_token"]
|
||||
prompt_cost: float = non_cached_prompt_tokens * input_cost_per_token + cached_tokens * cache_read_cost_per_token
|
||||
|
||||
## CALCULATE OUTPUT COST
|
||||
completion_cost = usage["completion_tokens"] * model_info["output_cost_per_token"]
|
||||
output_cost_per_token: float = model_info["output_cost_per_token"] or 0.0
|
||||
completion_cost: float = usage.completion_tokens * output_cost_per_token
|
||||
|
||||
return prompt_cost, completion_cost
|
||||
|
|
|
|||
|
|
@ -510,6 +510,20 @@ async def acompletion(
|
|||
#########################################################
|
||||
#########################################################
|
||||
litellm_logging_obj = kwargs.get("litellm_logging_obj", None)
|
||||
|
||||
from litellm.integrations.anthropic_cache_control_hook import (
|
||||
AnthropicCacheControlHook,
|
||||
)
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
AnthropicCacheControlHook.maybe_seed_default_injection_points(
|
||||
non_default_params=kwargs,
|
||||
messages=cast(list[AllMessageValues], messages), # cast-ok: acompletion types messages as a bare List
|
||||
model=model,
|
||||
custom_llm_provider=cast(Optional[str], custom_llm_provider), # cast-ok: read from untyped kwargs
|
||||
tools=tools,
|
||||
)
|
||||
|
||||
if isinstance(litellm_logging_obj, LiteLLMLoggingObj) and (
|
||||
litellm_logging_obj.should_run_prompt_management_hooks(
|
||||
prompt_id=kwargs.get("prompt_id", None),
|
||||
|
|
@ -5055,6 +5069,19 @@ def completion( # type: ignore
|
|||
litellm_params = {} # used to prevent unbound var errors
|
||||
## PROMPT MANAGEMENT HOOKS ##
|
||||
|
||||
from litellm.integrations.anthropic_cache_control_hook import (
|
||||
AnthropicCacheControlHook,
|
||||
)
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
AnthropicCacheControlHook.maybe_seed_default_injection_points(
|
||||
non_default_params=non_default_params,
|
||||
messages=cast(list[AllMessageValues], messages), # cast-ok: completion types messages as a bare List
|
||||
model=model,
|
||||
custom_llm_provider=cast(Optional[str], kwargs.get("custom_llm_provider")), # cast-ok: untyped kwargs
|
||||
tools=tools,
|
||||
)
|
||||
|
||||
if isinstance(litellm_logging_obj, LiteLLMLoggingObj) and (
|
||||
litellm_logging_obj.should_run_prompt_management_hooks(
|
||||
prompt_id=prompt_id, non_default_params=non_default_params
|
||||
|
|
|
|||
|
|
@ -3451,7 +3451,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2.2e-05,
|
||||
"output_cost_per_token": 2.64e-06,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -3470,7 +3470,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 0.00022,
|
||||
"output_cost_per_token": 2.2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -3489,7 +3489,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2.2e-05,
|
||||
"supported_modalities": [
|
||||
|
|
@ -4687,7 +4687,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -4707,7 +4707,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -4739,7 +4739,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -4771,7 +4771,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -4832,7 +4832,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 0.0002,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -4850,7 +4850,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"supported_modalities": [
|
||||
|
|
@ -7922,7 +7922,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2.2e-05,
|
||||
"output_cost_per_token": 2.64e-06,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -7941,7 +7941,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 0.00022,
|
||||
"output_cost_per_token": 2.2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -7960,7 +7960,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2.2e-05,
|
||||
"supported_modalities": [
|
||||
|
|
@ -22094,7 +22094,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -22113,7 +22113,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -22207,7 +22207,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -22225,7 +22225,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -22243,7 +22243,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -24438,7 +24438,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -24470,7 +24470,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -24502,7 +24502,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -24535,7 +24535,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 32000,
|
||||
"max_tokens": 32000,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 2.4e-05,
|
||||
"regional_processing_uplift_multiplier_eu": 1.1,
|
||||
|
|
@ -24570,7 +24570,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"regional_processing_uplift_multiplier_eu": 1.1,
|
||||
|
|
@ -24603,7 +24603,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -24635,7 +24635,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -43573,7 +43573,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -43606,7 +43606,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supported_endpoints": [
|
||||
|
|
|
|||
|
|
@ -3708,22 +3708,22 @@ def _attach_redis_usage_cache(redis_cache: RedisCache, enable_redis_auth_cache:
|
|||
litellm_config_cache.redis_cache = redis_cache
|
||||
|
||||
|
||||
def resolve_complexity_router_plugins(
|
||||
model_name: str,
|
||||
complexity_router_config: dict,
|
||||
def resolve_routing_plugins(
|
||||
plugin_paths: list,
|
||||
config_file_path: str | None,
|
||||
) -> None:
|
||||
source_label: str,
|
||||
) -> list:
|
||||
"""
|
||||
Resolves `complexity_router_config["plugins"]` dotted-path strings to live
|
||||
instances via `get_instance_fn` (the same convention `litellm_settings.callbacks`
|
||||
uses), in place. Raises at config-load time if a path resolves to something that
|
||||
doesn't implement `RoutingPlugin`, rather than deferring to a confusing
|
||||
`AttributeError` on the first request that reaches the plugin pipeline.
|
||||
Resolves a list of routing-plugin entries to live `RoutingPlugin` instances.
|
||||
Each string entry is resolved through `get_instance_fn` (the same dotted-path
|
||||
convention `litellm_settings.callbacks` uses, which resolves both local module
|
||||
files next to the config and modules installed as Python packages); non-string
|
||||
entries are assumed to already be instances and passed through. Raises at
|
||||
config-load time if any entry resolves to something that doesn't implement
|
||||
`RoutingPlugin`, rather than deferring to a confusing `AttributeError` on the
|
||||
first request that reaches the plugin pipeline. `source_label` names the config
|
||||
key being resolved so the error points the operator at the right place.
|
||||
"""
|
||||
plugin_paths = complexity_router_config.get("plugins")
|
||||
if not isinstance(plugin_paths, list):
|
||||
return
|
||||
|
||||
resolved_plugins = [
|
||||
get_instance_fn(value=plugin_path, config_file_path=config_file_path)
|
||||
if isinstance(plugin_path, str)
|
||||
|
|
@ -3739,12 +3739,31 @@ def resolve_complexity_router_plugins(
|
|||
getattr(resolved_plugin, "run", None)
|
||||
):
|
||||
raise ValueError(
|
||||
f"complexity_router_config.plugins entry {plugin_path!r} on model {model_name!r} "
|
||||
f"resolved to {resolved_plugin!r}, which does not implement the RoutingPlugin "
|
||||
"interface (an async `run(context)` method). Fix the referenced module before "
|
||||
"starting the proxy."
|
||||
f"{source_label} entry {plugin_path!r} resolved to {resolved_plugin!r}, which does "
|
||||
"not implement the RoutingPlugin interface (an async `run(context)` method). Fix the "
|
||||
"referenced module before starting the proxy."
|
||||
)
|
||||
complexity_router_config["plugins"] = resolved_plugins
|
||||
return resolved_plugins
|
||||
|
||||
|
||||
def resolve_complexity_router_plugins(
|
||||
model_name: str,
|
||||
complexity_router_config: dict,
|
||||
config_file_path: str | None,
|
||||
) -> None:
|
||||
"""
|
||||
Resolves `complexity_router_config["plugins"]` dotted-path strings to live
|
||||
instances in place, via `resolve_routing_plugins`.
|
||||
"""
|
||||
plugin_paths = complexity_router_config.get("plugins")
|
||||
if not isinstance(plugin_paths, list):
|
||||
return
|
||||
|
||||
complexity_router_config["plugins"] = resolve_routing_plugins(
|
||||
plugin_paths=plugin_paths,
|
||||
config_file_path=config_file_path,
|
||||
source_label=f"complexity_router_config.plugins on model {model_name!r}",
|
||||
)
|
||||
|
||||
|
||||
class ProxyConfig:
|
||||
|
|
@ -4874,6 +4893,12 @@ class ProxyConfig:
|
|||
|
||||
for k, v in router_settings.items():
|
||||
if k in available_args:
|
||||
if k == "plugins" and isinstance(v, list):
|
||||
v = resolve_routing_plugins(
|
||||
plugin_paths=v,
|
||||
config_file_path=config_file_path,
|
||||
source_label="router_settings.plugins",
|
||||
)
|
||||
router_params[k] = v
|
||||
elif k in {"health_check_interval", "health_check_concurrency"}:
|
||||
raise ValueError(
|
||||
|
|
|
|||
|
|
@ -11,14 +11,16 @@ from typing import Any, Dict, Optional, Tuple
|
|||
|
||||
import orjson
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
|
||||
from fastapi.responses import ORJSONResponse
|
||||
from fastapi.responses import ORJSONResponse, StreamingResponse
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.auth_utils import is_request_body_safe
|
||||
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
_read_request_body,
|
||||
_safe_get_request_headers,
|
||||
|
|
@ -604,6 +606,7 @@ async def rag_query(
|
|||
general_settings,
|
||||
llm_router,
|
||||
proxy_config,
|
||||
select_data_generator,
|
||||
version,
|
||||
)
|
||||
|
||||
|
|
@ -673,6 +676,31 @@ async def rag_query(
|
|||
**request_data,
|
||||
)
|
||||
|
||||
hidden_params = getattr(response, "_hidden_params", {}) or {}
|
||||
custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_id=hidden_params.get("litellm_call_id", None) or "",
|
||||
model_id=hidden_params.get("model_id", None) or "",
|
||||
cache_key=hidden_params.get("cache_key", None) or "",
|
||||
api_base=hidden_params.get("api_base", None) or "",
|
||||
version=version,
|
||||
response_cost=hidden_params.get("response_cost", None),
|
||||
request_data=request_data,
|
||||
)
|
||||
|
||||
if isinstance(response, CustomStreamWrapper):
|
||||
return StreamingResponse(
|
||||
select_data_generator(
|
||||
response=response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=request_data,
|
||||
request=request,
|
||||
),
|
||||
media_type="text/event-stream",
|
||||
headers=custom_headers,
|
||||
)
|
||||
|
||||
fastapi_response.headers.update(custom_headers)
|
||||
return response
|
||||
|
||||
except HTTPException:
|
||||
|
|
|
|||
|
|
@ -34,16 +34,12 @@ def get_instance_fn(value: str, config_file_path: Optional[str] = None) -> Any:
|
|||
module_name = ".".join(parts[:-1])
|
||||
instance_name = parts[-1]
|
||||
|
||||
# If config_file_path is provided, use it to determine the module spec and load the module
|
||||
module_file_path = None
|
||||
if config_file_path is not None:
|
||||
directory = os.path.dirname(config_file_path)
|
||||
module_file_path = os.path.join(directory, *module_name.split("."))
|
||||
module_file_path += ".py"
|
||||
|
||||
# Check if the file exists before trying to load it
|
||||
if not os.path.exists(module_file_path):
|
||||
raise ImportError(f"Could not find module file {module_file_path}")
|
||||
module_file_path = os.path.join(directory, *module_name.split(".")) + ".py"
|
||||
|
||||
if module_file_path is not None and os.path.exists(module_file_path):
|
||||
spec = importlib.util.spec_from_file_location(module_name, module_file_path) # type: ignore
|
||||
if spec is None:
|
||||
raise ImportError(f"Could not find a module specification for {module_file_path}")
|
||||
|
|
@ -52,7 +48,6 @@ def get_instance_fn(value: str, config_file_path: Optional[str] = None) -> Any:
|
|||
raise ImportError(f"Could not find a module loader for {module_file_path}")
|
||||
spec.loader.exec_module(module) # type: ignore
|
||||
else:
|
||||
# Dynamically import the module
|
||||
module = importlib.import_module(module_name)
|
||||
|
||||
# Get the instance from the module
|
||||
|
|
|
|||
|
|
@ -11,12 +11,14 @@ __all__ = ["ingest", "aingest", "query", "aquery"]
|
|||
|
||||
import asyncio
|
||||
import contextvars
|
||||
from contextlib import contextmanager
|
||||
from functools import partial
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Coroutine,
|
||||
Dict,
|
||||
Iterator,
|
||||
List,
|
||||
Optional,
|
||||
Tuple,
|
||||
|
|
@ -27,6 +29,9 @@ from typing import (
|
|||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import is_internal_call
|
||||
from litellm.cost_calculator import vector_store_search_cost
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.rag.ingestion.base_ingestion import BaseRAGIngestion
|
||||
from litellm.rag.ingestion.bedrock_ingestion import BedrockRAGIngestion
|
||||
from litellm.rag.ingestion.gemini_ingestion import GeminiRAGIngestion
|
||||
|
|
@ -188,6 +193,25 @@ async def aingest(
|
|||
)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _suppressed_sub_call_billing() -> Iterator[None]:
|
||||
"""
|
||||
Suppress a sub-call's own billing event so the parent aquery event bills it.
|
||||
|
||||
Every suppressed sub-call's cost must be folded into the parent event:
|
||||
into the response's hidden response_cost on the non-streaming path, or via
|
||||
the logging object's additional_response_cost on the streaming path (the
|
||||
streamed cost is computed from assembled chunks after this pipeline
|
||||
returns, so there is no response object to fold into here).
|
||||
"""
|
||||
previous = is_internal_call.get()
|
||||
is_internal_call.set(True)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
is_internal_call.set(previous)
|
||||
|
||||
|
||||
async def _execute_query_pipeline(
|
||||
model: str,
|
||||
messages: List[Any],
|
||||
|
|
@ -209,27 +233,46 @@ async def _execute_query_pipeline(
|
|||
raise ValueError("No query found in messages for RAG query")
|
||||
|
||||
# 2. Search vector store
|
||||
search_response = await litellm.vector_stores.asearch(
|
||||
vector_store_id=retrieval_config["vector_store_id"],
|
||||
query=query_text,
|
||||
max_num_results=retrieval_config.get("top_k", 10),
|
||||
custom_llm_provider=retrieval_config.get("custom_llm_provider", "openai"),
|
||||
**kwargs,
|
||||
)
|
||||
with _suppressed_sub_call_billing():
|
||||
search_response = await litellm.vector_stores.asearch(
|
||||
vector_store_id=retrieval_config["vector_store_id"],
|
||||
query=query_text,
|
||||
max_num_results=retrieval_config.get("top_k", 10),
|
||||
custom_llm_provider=retrieval_config.get("custom_llm_provider", "openai"),
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
search_provider = retrieval_config.get("custom_llm_provider", "openai")
|
||||
try:
|
||||
search_cost = sum(
|
||||
vector_store_search_cost(
|
||||
model=search_provider if "/" in search_provider else None,
|
||||
custom_llm_provider=search_provider,
|
||||
response=search_response,
|
||||
)
|
||||
)
|
||||
except Exception: # noqa: BLE001 - cost accounting must never break the query path
|
||||
search_cost = 0.0
|
||||
|
||||
rerank_response = None
|
||||
rerank_cost = 0.0
|
||||
context_chunks = search_response.get("data", [])
|
||||
|
||||
# 3. Optional rerank
|
||||
if rerank and rerank.get("enabled"):
|
||||
documents = RAGQuery.extract_documents_from_search(search_response)
|
||||
if documents:
|
||||
rerank_response = await litellm.arerank(
|
||||
model=rerank["model"],
|
||||
query=query_text,
|
||||
documents=documents,
|
||||
top_n=rerank.get("top_n", 5),
|
||||
)
|
||||
with _suppressed_sub_call_billing():
|
||||
rerank_response = await litellm.arerank(
|
||||
model=rerank["model"],
|
||||
query=query_text,
|
||||
documents=documents,
|
||||
top_n=rerank.get("top_n", 5),
|
||||
)
|
||||
rerank_hidden_params = getattr(rerank_response, "_hidden_params", None)
|
||||
if isinstance(rerank_hidden_params, dict):
|
||||
rerank_response_cost: float | None = rerank_hidden_params.get("response_cost")
|
||||
rerank_cost = rerank_response_cost or 0.0
|
||||
context_chunks = RAGQuery.get_top_chunks_from_rerank(search_response, rerank_response)
|
||||
|
||||
# 4. Build context message and call completion
|
||||
|
|
@ -237,28 +280,40 @@ async def _execute_query_pipeline(
|
|||
modified_messages = messages[:-1] + [context_message] + [messages[-1]]
|
||||
|
||||
# Use router if available to properly resolve virtual model names
|
||||
if router is not None:
|
||||
response = await router.acompletion(
|
||||
model=model,
|
||||
messages=modified_messages,
|
||||
stream=stream,
|
||||
**kwargs,
|
||||
)
|
||||
else:
|
||||
response = await litellm.acompletion(
|
||||
model=model,
|
||||
messages=modified_messages,
|
||||
stream=stream,
|
||||
**kwargs,
|
||||
)
|
||||
with _suppressed_sub_call_billing():
|
||||
if router is not None:
|
||||
response = await router.acompletion(
|
||||
model=model,
|
||||
messages=modified_messages,
|
||||
stream=stream,
|
||||
**kwargs,
|
||||
)
|
||||
else:
|
||||
response = await litellm.acompletion(
|
||||
model=model,
|
||||
messages=modified_messages,
|
||||
stream=stream,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
# 5. Attach search results to response
|
||||
sub_call_cost = search_cost + rerank_cost
|
||||
if not stream and isinstance(response, ModelResponse):
|
||||
response = RAGQuery.add_search_results_to_response(
|
||||
response=response,
|
||||
search_results=search_response,
|
||||
rerank_results=rerank_response,
|
||||
)
|
||||
if sub_call_cost > 0:
|
||||
hidden_params = getattr(response, "_hidden_params", None)
|
||||
if isinstance(hidden_params, dict):
|
||||
completion_response_cost: float | None = hidden_params.get("response_cost")
|
||||
if completion_response_cost is not None:
|
||||
hidden_params["response_cost"] = completion_response_cost + sub_call_cost
|
||||
elif sub_call_cost > 0:
|
||||
logging_obj: object = kwargs.get("litellm_logging_obj")
|
||||
if isinstance(logging_obj, LiteLLMLoggingObj):
|
||||
logging_obj.model_call_details["additional_response_cost"] = sub_call_cost
|
||||
|
||||
return response # type: ignore[return-value]
|
||||
|
||||
|
|
|
|||
|
|
@ -28,6 +28,7 @@ from litellm.types.utils import ModelResponse
|
|||
|
||||
from .config import (
|
||||
DEFAULT_CODE_KEYWORDS,
|
||||
DEFAULT_ESCALATION_KEYWORDS,
|
||||
DEFAULT_REASONING_KEYWORDS,
|
||||
DEFAULT_SIMPLE_KEYWORDS,
|
||||
DEFAULT_TECHNICAL_KEYWORDS,
|
||||
|
|
@ -173,6 +174,11 @@ class ComplexityRouter(CustomLogger):
|
|||
self.config.custom_technical_keywords,
|
||||
)
|
||||
self.simple_keywords = self.config.simple_keywords or DEFAULT_SIMPLE_KEYWORDS
|
||||
self.escalation_keywords = (
|
||||
self.config.escalation_keywords
|
||||
if self.config.escalation_keywords is not None
|
||||
else DEFAULT_ESCALATION_KEYWORDS
|
||||
)
|
||||
|
||||
# Lazily built on first semantic request and cached for reuse (route
|
||||
# embeddings are static, only the prompt is embedded per request). The lock
|
||||
|
|
@ -668,6 +674,53 @@ class ComplexityRouter(CustomLogger):
|
|||
}
|
||||
return best_model
|
||||
|
||||
def _escalation_triggered(self, user_message: str) -> bool:
|
||||
"""Whether the prompt asks to escalate to a stronger model.
|
||||
|
||||
Matching is a case-sensitive substring test so the default "LITELLM ESCALATE"
|
||||
only fires on the deliberate, shouted form and not on incidental lowercase
|
||||
mentions of the word (e.g. "how do I escalate this ticket").
|
||||
"""
|
||||
if not self.escalation_keywords:
|
||||
return False
|
||||
return any(keyword in user_message for keyword in self.escalation_keywords)
|
||||
|
||||
def _tier_for_model(self, model: str) -> ComplexityTier | None:
|
||||
"""Return the most-severe configured tier whose pool contains this model."""
|
||||
pools = self._tier_pools()
|
||||
matched = tuple(ComplexityTier(tier_name) for tier_name, models in pools.items() if model in models)
|
||||
if not matched:
|
||||
return None
|
||||
return max(matched, key=TIER_SEVERITY_ORDER.index)
|
||||
|
||||
def _escalate_tier(self, tier: ComplexityTier) -> ComplexityTier:
|
||||
"""Bump a tier one step up to the next-higher configured tier.
|
||||
|
||||
Returns the input tier unchanged when it is already the highest configured
|
||||
tier, so escalation can never route below the model the user would otherwise
|
||||
have received.
|
||||
"""
|
||||
configured = frozenset(self.config.tiers)
|
||||
current_index = TIER_SEVERITY_ORDER.index(tier)
|
||||
higher_tiers = tuple(
|
||||
candidate for candidate in TIER_SEVERITY_ORDER[current_index + 1 :] if candidate.value in configured
|
||||
)
|
||||
return higher_tiers[0] if higher_tiers else tier
|
||||
|
||||
def _escalated_pin(self, pinned_model: str) -> str | None:
|
||||
"""Bump a session's pinned model to the next-higher configured tier.
|
||||
|
||||
Returns None when the pin no longer maps to any configured tier, signalling
|
||||
a full reclassification instead.
|
||||
"""
|
||||
pinned_tier = self._tier_for_model(pinned_model)
|
||||
if pinned_tier is None:
|
||||
return None
|
||||
escalated_tier = self._escalate_tier(pinned_tier)
|
||||
if escalated_tier == pinned_tier:
|
||||
return pinned_model
|
||||
return self.get_model_for_tier(escalated_tier)
|
||||
|
||||
def _lexical_tier_override(self, user_message: str) -> ComplexityTier | None:
|
||||
"""When keyword_tier_rules match literally, the most-severe matched tier wins.
|
||||
|
||||
|
|
@ -910,29 +963,41 @@ class ComplexityRouter(CustomLogger):
|
|||
if cache_key is not None:
|
||||
pinned_model = await self.litellm_router_instance.cache.async_get_cache(key=cache_key)
|
||||
if isinstance(pinned_model, str):
|
||||
# Refresh the TTL on every hit so an active session doesn't lose its
|
||||
# pin mid-conversation just because it outlives the original write.
|
||||
await self.litellm_router_instance.cache.async_set_cache(
|
||||
key=cache_key,
|
||||
value=pinned_model,
|
||||
ttl=self.config.session_affinity_ttl_seconds,
|
||||
)
|
||||
if self.config.adaptive:
|
||||
from litellm.router_strategy.adaptive_router.config import (
|
||||
ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY,
|
||||
routed_model: str | None = pinned_model
|
||||
if self.escalation_keywords:
|
||||
resolved_messages = self._resolve_messages(messages, request_kwargs)
|
||||
user_message = (
|
||||
self._extract_user_message_and_system_prompt(resolved_messages)[0]
|
||||
if resolved_messages
|
||||
else None
|
||||
)
|
||||
if user_message is not None and self._escalation_triggered(user_message):
|
||||
routed_model = self._escalated_pin(pinned_model)
|
||||
if routed_model is not None:
|
||||
# Refresh the TTL on every hit so an active session doesn't lose its
|
||||
# pin mid-conversation just because it outlives the original write.
|
||||
await self.litellm_router_instance.cache.async_set_cache(
|
||||
key=cache_key,
|
||||
value=routed_model,
|
||||
ttl=self.config.session_affinity_ttl_seconds,
|
||||
)
|
||||
if self.config.adaptive:
|
||||
from litellm.router_strategy.adaptive_router.config import (
|
||||
ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY,
|
||||
)
|
||||
|
||||
kwargs_metadata = request_kwargs.setdefault("metadata", {})
|
||||
if isinstance(kwargs_metadata, dict):
|
||||
kwargs_metadata[ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY] = pinned_model
|
||||
verbose_router_logger.info(
|
||||
f"ComplexityRouter: routing decision cause=session_affinity_pin, routed_model={pinned_model}"
|
||||
)
|
||||
has_original_messages = messages is not None and len(messages) > 0
|
||||
return PreRoutingHookResponse(
|
||||
model=pinned_model,
|
||||
messages=messages if has_original_messages else None,
|
||||
)
|
||||
kwargs_metadata = request_kwargs.setdefault("metadata", {})
|
||||
if isinstance(kwargs_metadata, dict):
|
||||
kwargs_metadata[ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY] = routed_model
|
||||
cause = "session_affinity_escalation" if routed_model != pinned_model else "session_affinity_pin"
|
||||
verbose_router_logger.info(
|
||||
f"ComplexityRouter: routing decision cause={cause}, routed_model={routed_model}"
|
||||
)
|
||||
has_original_messages = messages is not None and len(messages) > 0
|
||||
return PreRoutingHookResponse(
|
||||
model=routed_model,
|
||||
messages=messages if has_original_messages else None,
|
||||
)
|
||||
|
||||
response = await self._classify_and_route(
|
||||
model=model,
|
||||
|
|
@ -1004,13 +1069,17 @@ class ComplexityRouter(CustomLogger):
|
|||
messages=messages if has_original_messages else None,
|
||||
)
|
||||
|
||||
escalate = self._escalation_triggered(user_message)
|
||||
|
||||
override_tier = await self._resolve_keyword_tier_override(user_message, request_kwargs)
|
||||
if override_tier is not None:
|
||||
routed_model = await self._pick_model_for_tier(override_tier, messages, resolved_messages, request_kwargs)
|
||||
cause = "semantic_keyword_match" if self.config.semantic_keyword_matching else "literal_keyword_match"
|
||||
routed_tier = self._escalate_tier(override_tier) if escalate else override_tier
|
||||
routed_model = await self._pick_model_for_tier(routed_tier, messages, resolved_messages, request_kwargs)
|
||||
base_cause = "semantic_keyword_match" if self.config.semantic_keyword_matching else "literal_keyword_match"
|
||||
cause = f"{base_cause}+escalation" if escalate else base_cause
|
||||
verbose_router_logger.info(
|
||||
f"ComplexityRouter: routing decision cause={cause}, "
|
||||
f"tier={override_tier.value}, routed_model={routed_model}"
|
||||
f"tier={routed_tier.value}, routed_model={routed_model}"
|
||||
)
|
||||
return PreRoutingHookResponse(
|
||||
model=routed_model,
|
||||
|
|
@ -1018,6 +1087,9 @@ class ComplexityRouter(CustomLogger):
|
|||
)
|
||||
|
||||
tier, score, signals = await self.aclassify(user_message, system_prompt, request_kwargs)
|
||||
if escalate:
|
||||
tier = self._escalate_tier(tier)
|
||||
signals = [*signals, "escalation"]
|
||||
if self.config.adaptive:
|
||||
routed_model = self._soft_floor_pick(tier, user_message, request_kwargs)
|
||||
adaptive = self._ensure_adaptive_router()
|
||||
|
|
|
|||
|
|
@ -162,6 +162,9 @@ DEFAULT_TECHNICAL_KEYWORDS: list[str] = [
|
|||
# Note: "async", "kubernetes", "docker" are in DEFAULT_CODE_KEYWORDS
|
||||
]
|
||||
|
||||
DEFAULT_ESCALATION_KEYWORDS: list[str] = ["LITELLM ESCALATE"]
|
||||
|
||||
|
||||
DEFAULT_SIMPLE_KEYWORDS: list[str] = [
|
||||
"what is",
|
||||
"what's",
|
||||
|
|
@ -339,6 +342,16 @@ class ComplexityRouterConfig(BaseModel):
|
|||
),
|
||||
)
|
||||
|
||||
escalation_keywords: list[str] | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Case-sensitive phrases a user can include to force a bump to the next-higher "
|
||||
"complexity tier when they aren't satisfied with results (they can force a stronger "
|
||||
"model, but not choose which one). Defaults to ['LITELLM ESCALATE'] when unset; "
|
||||
"set to an empty list to disable."
|
||||
),
|
||||
)
|
||||
|
||||
# Deterministic keyword -> tier overrides, evaluated before weighted scoring
|
||||
keyword_tier_rules: list[KeywordTierRule] | None = Field(
|
||||
default=None,
|
||||
|
|
@ -400,6 +413,13 @@ class ComplexityRouterConfig(BaseModel):
|
|||
coerced[key] = item
|
||||
return coerced
|
||||
|
||||
@field_validator("escalation_keywords")
|
||||
@classmethod
|
||||
def _normalize_escalation_keywords(cls, value: list[str] | None) -> list[str] | None:
|
||||
if value is None:
|
||||
return None
|
||||
return [stripped for keyword in value if (stripped := keyword.strip())]
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _validate_llm_classifier_config(self) -> "ComplexityRouterConfig":
|
||||
if self.classifier_type == "llm" and self.classifier_llm_config is None:
|
||||
|
|
|
|||
|
|
@ -529,6 +529,7 @@ class ChatCompletionDeltaToolCallChunk(TypedDict, total=False):
|
|||
|
||||
class ChatCompletionCachedContent(TypedDict):
|
||||
type: Literal["ephemeral"]
|
||||
ttl: NotRequired[Literal["5m", "1h"]]
|
||||
|
||||
|
||||
class ChatCompletionThinkingBlock(TypedDict, total=False):
|
||||
|
|
|
|||
|
|
@ -266,6 +266,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
|
|||
"audio_transcription",
|
||||
"responses",
|
||||
"ocr",
|
||||
"realtime",
|
||||
]
|
||||
]
|
||||
tpm: Optional[int]
|
||||
|
|
@ -402,6 +403,11 @@ class CallTypes(str, Enum):
|
|||
vector_store_search = "vector_store_search"
|
||||
avector_store_search = "avector_store_search"
|
||||
|
||||
ingest = "ingest"
|
||||
aingest = "aingest"
|
||||
query = "query"
|
||||
aquery = "aquery"
|
||||
|
||||
#########################################################
|
||||
# Container Call Types
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -3451,7 +3451,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2.2e-05,
|
||||
"output_cost_per_token": 2.64e-06,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -3470,7 +3470,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 0.00022,
|
||||
"output_cost_per_token": 2.2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -3489,7 +3489,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2.2e-05,
|
||||
"supported_modalities": [
|
||||
|
|
@ -4687,7 +4687,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -4707,7 +4707,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -4739,7 +4739,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -4771,7 +4771,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -4832,7 +4832,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 0.0002,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -4850,7 +4850,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"supported_modalities": [
|
||||
|
|
@ -7922,7 +7922,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2.2e-05,
|
||||
"output_cost_per_token": 2.64e-06,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -7941,7 +7941,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 0.00022,
|
||||
"output_cost_per_token": 2.2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -7960,7 +7960,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2.2e-05,
|
||||
"supported_modalities": [
|
||||
|
|
@ -22169,7 +22169,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -22188,7 +22188,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -22282,7 +22282,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -22300,7 +22300,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -22318,7 +22318,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 8e-05,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"supports_audio_input": true,
|
||||
|
|
@ -24513,7 +24513,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -24545,7 +24545,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -24577,7 +24577,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -24610,7 +24610,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 32000,
|
||||
"max_tokens": 32000,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 2.4e-05,
|
||||
"regional_processing_uplift_multiplier_eu": 1.1,
|
||||
|
|
@ -24645,7 +24645,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"regional_processing_uplift_multiplier_eu": 1.1,
|
||||
|
|
@ -24678,7 +24678,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -24710,7 +24710,7 @@
|
|||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 6.4e-05,
|
||||
"output_cost_per_token": 1.6e-05,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -43694,7 +43694,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -43727,7 +43727,7 @@
|
|||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supported_endpoints": [
|
||||
|
|
|
|||
|
|
@ -131,7 +131,7 @@ Quota Management - behavior features (entity- or config-driven caps and their ac
|
|||
quota_management.<behavior>.<variant>.<assertion>
|
||||
behavior : ratelimit | budget | spend_tracking
|
||||
variant : <ratelimit> rpm | tpm | priority_generous | priority_strict
|
||||
<budget> key | internal_user | end_user | organization | team_member | tag
|
||||
<budget> key | internal_user | end_user | organization | team | team_member | tag
|
||||
| model_max | soft | key_multi_window | team_multi_window
|
||||
| fallback | spend_counter
|
||||
<spend_tracking> chat_completions | stream | embeddings | cache_hit | key_rollup
|
||||
|
|
|
|||
|
|
@ -352,6 +352,7 @@ quota_management.ratelimit.rpm.headers_report_remaining: {tier: P1, source: para
|
|||
quota_management.ratelimit.priority_generous.picks_under_tpm: {tier: P1, source: 'dynamic_rate_limiter_v3.py:36-52', rationale: Generous mode (<80% sat) allows priority borrowing}
|
||||
quota_management.ratelimit.priority_strict.picks_under_tpm: {tier: P1, source: 'dynamic_rate_limiter_v3.py:53-71', rationale: Strict mode (>=80% sat) enforces priority fairness}
|
||||
quota_management.budget.key.blocks_over_limit: {tier: P0, source: proxy/auth/auth_checks.py, rationale: A key's max_budget blocks further paid calls once spend crosses it}
|
||||
quota_management.budget.team.blocks_over_limit: {tier: P0, source: proxy/auth/auth_checks.py, rationale: "A team's max_budget blocks every key on the team once combined spend crosses it, including keys that spent nothing themselves"}
|
||||
quota_management.budget.internal_user.blocks_over_limit: {tier: P1, source: proxy/auth/auth_checks.py, rationale: An internal user's max_budget governs personal keys}
|
||||
quota_management.budget.end_user.blocks_over_limit: {tier: P1, source: proxy/auth/auth_checks.py, rationale: A customer (end-user) max_budget blocks calls attributed via user=}
|
||||
quota_management.budget.organization.blocks_over_limit: {tier: P1, source: proxy/auth/auth_checks.py, rationale: An organization's max_budget blocks keys under its teams}
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ Each entity is an E2ECase (lifecycle.E2ECase) driven by run_case: init() creates
|
|||
the budgeted entity + a key, run() drives spend until a `budget_exceeded` block,
|
||||
teardown() deletes everything init() created (always runs, even on failure/skip).
|
||||
Covers the entities with no prior live coverage - internal user, end-user,
|
||||
organization, team member. See BUDGET_TEST_COVERAGE_MATRIX.md.
|
||||
organization, team member - plus key and team. See BUDGET_TEST_COVERAGE_MATRIX.md.
|
||||
|
||||
A non-budget error fails hard (never a skip); if calls never get blocked, budget
|
||||
enforcement is broken -> fail.
|
||||
|
|
@ -18,16 +18,17 @@ import pytest
|
|||
|
||||
from budget_client import BudgetClient, is_budget_block
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import require_successful_call
|
||||
from e2e_http import StreamingResponse, require_successful_call
|
||||
from lifecycle import run_case
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") -> None:
|
||||
"""Send paid calls until the entity's budget blocks one. Key/user/org/member
|
||||
block within a couple calls off real-time reservation counters; the end-user
|
||||
budget enforces off table spend that lands on the batch write, so it takes a
|
||||
few more. A non-budget error fails hard (never a skip)."""
|
||||
def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") -> StreamingResponse:
|
||||
"""Send paid calls until the entity's budget blocks one; return the blocked
|
||||
response so callers can assert on its shape. Key/user/org/member block within
|
||||
a couple calls off real-time reservation counters; the end-user budget
|
||||
enforces off table spend that lands on the batch write, so it takes a few
|
||||
more. A non-budget error fails hard (never a skip)."""
|
||||
for _ in range(40):
|
||||
result = client.chat(
|
||||
key,
|
||||
|
|
@ -37,7 +38,7 @@ def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") ->
|
|||
user=user or None,
|
||||
)
|
||||
if is_budget_block(result):
|
||||
return
|
||||
return result
|
||||
require_successful_call(result)
|
||||
time.sleep(2)
|
||||
pytest.fail("budget never enforced within the call budget")
|
||||
|
|
@ -69,10 +70,53 @@ class _BudgetCase:
|
|||
|
||||
|
||||
class KeyBudgetCase(_BudgetCase):
|
||||
"""A bare key (no team_id / user_id) carrying its own max_budget, so only the
|
||||
key-level budget can be the thing that blocks. The refusal must be a 429
|
||||
budget_exceeded; any other error already fails via _assert_budget_blocks."""
|
||||
|
||||
def init(self) -> None:
|
||||
self.key = self.client.generate_key(max_budget=3e-6)
|
||||
self._undo.append(lambda: self.client.delete_key(self.key))
|
||||
|
||||
def run(self) -> None:
|
||||
blocked = _assert_budget_blocks(self.client, self.key)
|
||||
assert blocked.status_code == 429, (
|
||||
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
|
||||
)
|
||||
|
||||
|
||||
class TeamBudgetCase(_BudgetCase):
|
||||
"""An admin caps a whole team: two keys under a tiny-budget team, neither with
|
||||
a key-level budget. Key A is driven until the team cap blocks it; key B's very
|
||||
first call must then be refused too, proving the cap sits on the team, not the
|
||||
key that spent. Both refusals must be 429 budget_exceeded."""
|
||||
|
||||
def init(self) -> None:
|
||||
team_id = self.client.create_team(
|
||||
alias=f"e2e-budget-team-{unique_marker()}", max_budget=3e-6
|
||||
)
|
||||
self._undo.append(lambda: self.client.delete_team(team_id))
|
||||
self.key = self.client.generate_key(team_id=team_id)
|
||||
self._undo.append(lambda: self.client.delete_key(self.key))
|
||||
self._sibling_key = self.client.generate_key(team_id=team_id)
|
||||
self._undo.append(lambda: self.client.delete_key(self._sibling_key))
|
||||
|
||||
def run(self) -> None:
|
||||
blocked = _assert_budget_blocks(self.client, self.key)
|
||||
assert blocked.status_code == 429, (
|
||||
f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}"
|
||||
)
|
||||
sibling = self.client.chat(
|
||||
self._sibling_key,
|
||||
"claude-haiku-4-5",
|
||||
f"spend {unique_marker()}",
|
||||
max_tokens=16,
|
||||
)
|
||||
assert is_budget_block(sibling) and sibling.status_code == 429, (
|
||||
f"a sibling key on the capped team must get the same 429 budget_exceeded, "
|
||||
f"got {sibling.status_code}: {sibling.body[:200]}"
|
||||
)
|
||||
|
||||
|
||||
class InternalUserBudgetCase(_BudgetCase):
|
||||
def init(self) -> None:
|
||||
|
|
@ -138,6 +182,10 @@ def _case_id(case_cls: Type[_BudgetCase]) -> str:
|
|||
KeyBudgetCase,
|
||||
marks=pytest.mark.covers("quota_management.budget.key.blocks_over_limit"),
|
||||
),
|
||||
pytest.param(
|
||||
TeamBudgetCase,
|
||||
marks=pytest.mark.covers("quota_management.budget.team.blocks_over_limit"),
|
||||
),
|
||||
pytest.param(
|
||||
InternalUserBudgetCase,
|
||||
marks=pytest.mark.covers("quota_management.budget.internal_user.blocks_over_limit"),
|
||||
|
|
|
|||
|
|
@ -172,7 +172,7 @@ def anthropic_messages():
|
|||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Here is the full text of a complex legal agreement" * 400,
|
||||
"text": "Here is the full text of a complex legal agreement" * 500,
|
||||
"cache_control": {"type": "ephemeral"},
|
||||
}
|
||||
],
|
||||
|
|
|
|||
|
|
@ -2,7 +2,9 @@ import copy
|
|||
import datetime
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import textwrap
|
||||
import unittest
|
||||
from typing import List, Optional, Tuple
|
||||
from unittest.mock import ANY, MagicMock, Mock, patch
|
||||
|
|
@ -1533,3 +1535,242 @@ class TestApplyToAnthropicMessagesRequest:
|
|||
sys_blocks = sum(1 for b in (result_sys or []) if isinstance(b, dict) and b.get("cache_control") is not None)
|
||||
total_blocks = sys_blocks + sum(AnthropicCacheControlHook._count_cache_control_blocks(m) for m in result_msgs)
|
||||
assert total_blocks <= 4
|
||||
|
||||
|
||||
class TestEnableAnthropicPromptCaching:
|
||||
"""Auto-injected default breakpoints via litellm.enable_anthropic_prompt_caching."""
|
||||
|
||||
MESSAGES: List[AllMessageValues] = [
|
||||
{"role": "system", "content": "a long system prompt"},
|
||||
{"role": "user", "content": "first turn"},
|
||||
{"role": "assistant", "content": "a reply"},
|
||||
{"role": "user", "content": "latest turn"},
|
||||
]
|
||||
|
||||
def _points(self, model="claude-sonnet-4-5", provider="anthropic", messages=None, system=None, tools=None):
|
||||
return AnthropicCacheControlHook.get_default_injection_points(
|
||||
messages=copy.deepcopy(self.MESSAGES) if messages is None else messages,
|
||||
system=system,
|
||||
model=model,
|
||||
custom_llm_provider=provider,
|
||||
tools=tools,
|
||||
)
|
||||
|
||||
def test_disabled_by_default(self):
|
||||
assert litellm.enable_anthropic_prompt_caching is False
|
||||
assert self._points() == []
|
||||
|
||||
def test_injects_system_and_trailing_turn(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
assert self._points() == [
|
||||
{"location": "message", "role": "system", "index": None, "control": {"type": "ephemeral"}},
|
||||
{"location": "message", "role": None, "index": -1, "control": {"type": "ephemeral"}},
|
||||
]
|
||||
|
||||
def test_bedrock_claude_is_injected(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
points = self._points(model="us.anthropic.claude-sonnet-4-5-20250929-v1:0", provider="bedrock")
|
||||
assert [p["index"] for p in points] == [None, -1]
|
||||
|
||||
@pytest.mark.parametrize("model, provider", [("gpt-4o", "openai"), ("gemini-2.0-flash", "gemini")])
|
||||
def test_non_anthropic_providers_never_injected(self, monkeypatch, model, provider):
|
||||
"""These report supports_prompt_caching=True but never consume cache_control markers."""
|
||||
from litellm.utils import supports_prompt_caching
|
||||
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
assert supports_prompt_caching(model=model, custom_llm_provider=provider) is True
|
||||
assert self._points(model=model, provider=provider) == []
|
||||
|
||||
def test_model_without_caching_support_not_injected(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
assert self._points(model="anthropic.claude-3-5-sonnet-20240620-v1:0", provider="bedrock") == []
|
||||
|
||||
def test_stands_down_when_client_sent_cache_control(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
messages = [
|
||||
{"role": "system", "content": [{"type": "text", "text": "s", "cache_control": {"type": "ephemeral"}}]},
|
||||
{"role": "user", "content": "latest turn"},
|
||||
]
|
||||
assert self._points(messages=messages) == []
|
||||
|
||||
def test_stands_down_when_system_block_has_cache_control(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
system = [{"type": "text", "text": "s", "cache_control": {"type": "ephemeral"}}]
|
||||
assert self._points(messages=[{"role": "user", "content": "hi"}], system=system) == []
|
||||
|
||||
@staticmethod
|
||||
def _tools(count: int, cached: bool) -> List[dict]:
|
||||
tool: dict = {"type": "function", "function": {"name": "t", "description": "d", "parameters": {}}}
|
||||
if cached:
|
||||
tool["cache_control"] = {"type": "ephemeral"}
|
||||
return [{**tool, "function": {**tool["function"], "name": f"t{i}"}} for i in range(count)]
|
||||
|
||||
def test_stands_down_when_only_tools_carry_cache_control(self, monkeypatch):
|
||||
"""Caching just the tool definitions is a normal client pattern, and those
|
||||
breakpoints count toward the provider's four-block limit. Three of them plus
|
||||
our two would be five, which Anthropic rejects outright."""
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
assert self._points(tools=self._tools(3, cached=True)) == []
|
||||
|
||||
def test_injects_when_tools_carry_no_cache_control(self, monkeypatch):
|
||||
"""Tools alone must not suppress injection; only client-marked ones do."""
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
assert [p["index"] for p in self._points(tools=self._tools(3, cached=False))] == [None, -1]
|
||||
|
||||
@pytest.mark.parametrize("tools", [None, []])
|
||||
def test_absent_tools_do_not_suppress_injection(self, monkeypatch, tools):
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
assert [p["index"] for p in self._points(tools=tools)] == [None, -1]
|
||||
|
||||
def test_seed_stands_down_when_only_tools_carry_cache_control(self, monkeypatch):
|
||||
"""Same guard on the /chat/completions seeding path."""
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
params: dict = {}
|
||||
AnthropicCacheControlHook.maybe_seed_default_injection_points(
|
||||
non_default_params=params,
|
||||
messages=copy.deepcopy(self.MESSAGES),
|
||||
model="claude-sonnet-4-5",
|
||||
custom_llm_provider="anthropic",
|
||||
tools=self._tools(3, cached=True),
|
||||
)
|
||||
assert "cache_control_injection_points" not in params
|
||||
|
||||
def test_v1_messages_stands_down_when_only_tools_carry_cache_control(self, monkeypatch):
|
||||
"""Same guard on the /v1/messages path, where tools reach the hook directly."""
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
messages = [{"role": "user", "content": [{"type": "text", "text": "hi"}]}]
|
||||
result_msgs, result_sys = AnthropicCacheControlHook.maybe_inject_cache_control(
|
||||
copy.deepcopy(messages),
|
||||
"sys",
|
||||
{},
|
||||
model="claude-sonnet-4-5",
|
||||
custom_llm_provider="anthropic",
|
||||
tools=self._tools(3, cached=True),
|
||||
)
|
||||
assert result_sys == "sys"
|
||||
assert result_msgs == messages
|
||||
|
||||
def test_default_ttl_is_anthropics_five_minute_cache(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
assert all(p["control"] == {"type": "ephemeral"} for p in self._points())
|
||||
|
||||
@pytest.mark.parametrize("ttl", ["5m", "1h"])
|
||||
def test_ttl_override_applied(self, monkeypatch, ttl):
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
monkeypatch.setattr(litellm, "anthropic_prompt_caching_ttl", ttl)
|
||||
assert all(p["control"] == {"type": "ephemeral", "ttl": ttl} for p in self._points())
|
||||
|
||||
def test_seed_does_not_override_configured_points(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
configured = [{"location": "message", "role": "user", "index": 0}]
|
||||
params = {"cache_control_injection_points": configured}
|
||||
AnthropicCacheControlHook.maybe_seed_default_injection_points(
|
||||
non_default_params=params,
|
||||
messages=copy.deepcopy(self.MESSAGES),
|
||||
model="claude-sonnet-4-5",
|
||||
custom_llm_provider="anthropic",
|
||||
)
|
||||
assert params["cache_control_injection_points"] is configured
|
||||
|
||||
def test_seed_adds_defaults_when_enabled(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
params: dict = {}
|
||||
AnthropicCacheControlHook.maybe_seed_default_injection_points(
|
||||
non_default_params=params,
|
||||
messages=copy.deepcopy(self.MESSAGES),
|
||||
model="claude-sonnet-4-5",
|
||||
custom_llm_provider="anthropic",
|
||||
)
|
||||
assert [p["index"] for p in params["cache_control_injection_points"]] == [None, -1]
|
||||
|
||||
def test_seed_is_noop_when_disabled(self):
|
||||
params: dict = {}
|
||||
AnthropicCacheControlHook.maybe_seed_default_injection_points(
|
||||
non_default_params=params,
|
||||
messages=copy.deepcopy(self.MESSAGES),
|
||||
model="claude-sonnet-4-5",
|
||||
custom_llm_provider="anthropic",
|
||||
)
|
||||
assert params == {}
|
||||
|
||||
def test_v1_messages_applies_defaults_end_to_end(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||||
messages = [
|
||||
{"role": "user", "content": [{"type": "text", "text": "first"}]},
|
||||
{"role": "assistant", "content": [{"type": "text", "text": "reply"}]},
|
||||
{"role": "user", "content": [{"type": "text", "text": "latest"}]},
|
||||
]
|
||||
result_msgs, result_sys = AnthropicCacheControlHook.maybe_inject_cache_control(
|
||||
messages,
|
||||
"a system prompt",
|
||||
{},
|
||||
model="claude-sonnet-4-5",
|
||||
custom_llm_provider="anthropic",
|
||||
)
|
||||
|
||||
assert result_sys == [{"type": "text", "text": "a system prompt", "cache_control": {"type": "ephemeral"}}]
|
||||
assert result_msgs[-1]["content"][-1]["cache_control"] == {"type": "ephemeral"}
|
||||
assert "cache_control" not in result_msgs[0]["content"][-1]
|
||||
|
||||
def test_v1_messages_is_noop_when_disabled(self):
|
||||
messages = [{"role": "user", "content": [{"type": "text", "text": "hi"}]}]
|
||||
result_msgs, result_sys = AnthropicCacheControlHook.maybe_inject_cache_control(
|
||||
messages,
|
||||
"sys",
|
||||
{},
|
||||
model="claude-sonnet-4-5",
|
||||
custom_llm_provider="anthropic",
|
||||
)
|
||||
|
||||
assert result_sys == "sys"
|
||||
assert result_msgs == messages
|
||||
|
||||
|
||||
class TestAnthropicPromptCachingEnvVars:
|
||||
"""Both settings are read from the environment at import, so an admin can enable
|
||||
auto-caching without a config file. Each case re-imports litellm in a subprocess
|
||||
so the env is read fresh without contaminating this process's module graph.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _import_litellm_with_env(env_override: dict) -> Tuple[bool, Optional[str]]:
|
||||
env = os.environ.copy()
|
||||
env.pop("LITELLM_ENABLE_ANTHROPIC_PROMPT_CACHING", None)
|
||||
env.pop("LITELLM_ANTHROPIC_PROMPT_CACHING_TTL", None)
|
||||
env.update(env_override)
|
||||
script = textwrap.dedent(
|
||||
"""
|
||||
import json, litellm
|
||||
print(json.dumps([litellm.enable_anthropic_prompt_caching, litellm.anthropic_prompt_caching_ttl]))
|
||||
"""
|
||||
)
|
||||
result = subprocess.run(
|
||||
[sys.executable, "-c", script], capture_output=True, text=True, env=env, timeout=300
|
||||
)
|
||||
assert result.returncode == 0, result.stderr
|
||||
enabled, ttl = json.loads(result.stdout.strip().splitlines()[-1])
|
||||
return enabled, ttl
|
||||
|
||||
def test_unset_env_leaves_auto_caching_off(self):
|
||||
assert self._import_litellm_with_env({}) == (False, None)
|
||||
|
||||
@pytest.mark.parametrize("value", ["true", "True", "TRUE"])
|
||||
def test_env_enables_auto_caching_case_insensitively(self, value):
|
||||
enabled, _ = self._import_litellm_with_env({"LITELLM_ENABLE_ANTHROPIC_PROMPT_CACHING": value})
|
||||
assert enabled is True
|
||||
|
||||
@pytest.mark.parametrize("value", ["false", "0", "yes", ""])
|
||||
def test_env_only_enables_on_true(self, value):
|
||||
enabled, _ = self._import_litellm_with_env({"LITELLM_ENABLE_ANTHROPIC_PROMPT_CACHING": value})
|
||||
assert enabled is False
|
||||
|
||||
@pytest.mark.parametrize("value", ["5m", "1h"])
|
||||
def test_ttl_env_is_applied(self, value):
|
||||
_, ttl = self._import_litellm_with_env({"LITELLM_ANTHROPIC_PROMPT_CACHING_TTL": value})
|
||||
assert ttl == value
|
||||
|
||||
@pytest.mark.parametrize("value", ["10m", "1H", "3600", "ephemeral"])
|
||||
def test_unsupported_ttl_env_falls_back_to_provider_default(self, value):
|
||||
"""An unparseable TTL must fall back to Anthropic's 5m default, never reach the provider verbatim."""
|
||||
_, ttl = self._import_litellm_with_env({"LITELLM_ANTHROPIC_PROMPT_CACHING_TTL": value})
|
||||
assert ttl is None
|
||||
|
|
|
|||
|
|
@ -1261,6 +1261,23 @@ class TestAnthropicThinkingSignatureSelfHeal:
|
|||
)
|
||||
assert is_anthropic_invalid_thinking_signature_error(raw) is True
|
||||
|
||||
def test_is_anthropic_invalid_thinking_signature_error_positive_bedrock(self):
|
||||
from litellm.llms.anthropic.common_utils import (
|
||||
is_anthropic_invalid_thinking_signature_error,
|
||||
)
|
||||
|
||||
# Real user-reported Bedrock scenario
|
||||
raw = '{"message":"messages.2.content.0.thinking.signature.str: Input should be a valid string"}'
|
||||
assert is_anthropic_invalid_thinking_signature_error(raw) is True
|
||||
|
||||
def test_is_anthropic_invalid_thinking_signature_error_positive_vertex(self):
|
||||
from litellm.llms.anthropic.common_utils import (
|
||||
is_anthropic_invalid_thinking_signature_error,
|
||||
)
|
||||
|
||||
raw = "messages.4.content.1.thinking.signature.str: Input should be a valid string"
|
||||
assert is_anthropic_invalid_thinking_signature_error(raw) is True
|
||||
|
||||
def test_is_anthropic_invalid_thinking_signature_error_negative(self):
|
||||
from litellm.llms.anthropic.common_utils import (
|
||||
is_anthropic_invalid_thinking_signature_error,
|
||||
|
|
@ -1271,6 +1288,11 @@ class TestAnthropicThinkingSignatureSelfHeal:
|
|||
is_anthropic_invalid_thinking_signature_error("rate limit exceeded")
|
||||
is False
|
||||
)
|
||||
assert (
|
||||
is_anthropic_invalid_thinking_signature_error("invalid_request_error: model not found")
|
||||
is False
|
||||
)
|
||||
assert is_anthropic_invalid_thinking_signature_error("thinking signature is malformed") is False
|
||||
|
||||
def test_strip_thinking_blocks_from_anthropic_messages(self):
|
||||
from litellm.llms.anthropic.common_utils import (
|
||||
|
|
|
|||
|
|
@ -0,0 +1,66 @@
|
|||
import os
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../../../.."))
|
||||
|
||||
from litellm.llms.fireworks_ai.cost_calculator import cost_per_token
|
||||
from litellm.types.utils import PromptTokensDetailsWrapper, Usage
|
||||
|
||||
MODEL = "accounts/fireworks/models/glm-5p2"
|
||||
INPUT_COST = 1.4e-06
|
||||
CACHE_READ_COST = 2.6e-07
|
||||
OUTPUT_COST = 4.4e-06
|
||||
|
||||
|
||||
def _usage(prompt_tokens: int, cached_tokens: int, completion_tokens: int) -> Usage:
|
||||
return Usage(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=prompt_tokens + completion_tokens,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=cached_tokens),
|
||||
)
|
||||
|
||||
|
||||
def test_cached_prompt_tokens_billed_at_cache_read_rate():
|
||||
prompt_tokens = 7036
|
||||
cached_tokens = 7020
|
||||
completion_tokens = 8
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(
|
||||
model=MODEL, usage=_usage(prompt_tokens, cached_tokens, completion_tokens)
|
||||
)
|
||||
|
||||
expected_prompt_cost = (prompt_tokens - cached_tokens) * INPUT_COST + cached_tokens * CACHE_READ_COST
|
||||
assert prompt_cost == pytest.approx(expected_prompt_cost)
|
||||
assert completion_cost == pytest.approx(completion_tokens * OUTPUT_COST)
|
||||
|
||||
full_rate_cost = prompt_tokens * INPUT_COST
|
||||
assert prompt_cost < full_rate_cost
|
||||
|
||||
|
||||
def test_warm_call_cheaper_than_cold_call():
|
||||
prompt_tokens = 7036
|
||||
completion_tokens = 8
|
||||
|
||||
cold_prompt_cost, _ = cost_per_token(
|
||||
model=MODEL, usage=_usage(prompt_tokens, 16, completion_tokens)
|
||||
)
|
||||
warm_prompt_cost, _ = cost_per_token(
|
||||
model=MODEL, usage=_usage(prompt_tokens, 7020, completion_tokens)
|
||||
)
|
||||
|
||||
assert warm_prompt_cost < cold_prompt_cost
|
||||
|
||||
|
||||
def test_no_cached_tokens_matches_full_input_rate():
|
||||
prompt_tokens = 100
|
||||
completion_tokens = 10
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(
|
||||
model=MODEL, usage=_usage(prompt_tokens, 0, completion_tokens)
|
||||
)
|
||||
|
||||
assert prompt_cost == pytest.approx(prompt_tokens * INPUT_COST)
|
||||
assert completion_cost == pytest.approx(completion_tokens * OUTPUT_COST)
|
||||
|
|
@ -22,6 +22,7 @@ from litellm.proxy.proxy_server import (
|
|||
_scrub_db_overlay_remote_module_loads,
|
||||
_scrub_guardrail_inner,
|
||||
resolve_complexity_router_plugins,
|
||||
resolve_routing_plugins,
|
||||
)
|
||||
|
||||
from .conftest import normalize
|
||||
|
|
@ -185,6 +186,75 @@ def test_resolve_complexity_router_plugins_rejects_synchronous_run_method(tmp_pa
|
|||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# resolve_routing_plugins
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_resolve_routing_plugins_resolves_dotted_paths(tmp_path):
|
||||
plugin_file = tmp_path / "rs_plugin.py"
|
||||
plugin_file.write_text(
|
||||
"class _Plugin:\n"
|
||||
" async def run(self, context):\n"
|
||||
" return context\n"
|
||||
"\n"
|
||||
"rs_plugin_instance = _Plugin()\n"
|
||||
)
|
||||
|
||||
resolved = resolve_routing_plugins(
|
||||
plugin_paths=["rs_plugin.rs_plugin_instance"],
|
||||
config_file_path=str(tmp_path / "config.yaml"),
|
||||
source_label="router_settings.plugins",
|
||||
)
|
||||
|
||||
assert len(resolved) == 1
|
||||
assert type(resolved[0]).__name__ == "_Plugin"
|
||||
|
||||
|
||||
def test_resolve_routing_plugins_passes_through_instances(tmp_path):
|
||||
class _Plugin:
|
||||
async def run(self, context):
|
||||
return context
|
||||
|
||||
instance = _Plugin()
|
||||
resolved = resolve_routing_plugins(
|
||||
plugin_paths=[instance],
|
||||
config_file_path=None,
|
||||
source_label="router_settings.plugins",
|
||||
)
|
||||
assert resolved == [instance]
|
||||
|
||||
|
||||
def test_resolve_routing_plugins_rejects_non_routing_plugin(tmp_path):
|
||||
plugin_file = tmp_path / "bad_rs_plugin.py"
|
||||
plugin_file.write_text("not_a_plugin = object()\n")
|
||||
|
||||
with pytest.raises(ValueError, match="router_settings.plugins"):
|
||||
resolve_routing_plugins(
|
||||
plugin_paths=["bad_rs_plugin.not_a_plugin"],
|
||||
config_file_path=str(tmp_path / "config.yaml"),
|
||||
source_label="router_settings.plugins",
|
||||
)
|
||||
|
||||
|
||||
def test_resolve_routing_plugins_rejects_synchronous_run(tmp_path):
|
||||
plugin_file = tmp_path / "sync_rs_plugin.py"
|
||||
plugin_file.write_text(
|
||||
"class _SyncPlugin:\n"
|
||||
" def run(self, context):\n"
|
||||
" return context\n"
|
||||
"\n"
|
||||
"sync_plugin_instance = _SyncPlugin()\n"
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="does not implement the RoutingPlugin interface"):
|
||||
resolve_routing_plugins(
|
||||
plugin_paths=["sync_rs_plugin.sync_plugin_instance"],
|
||||
config_file_path=str(tmp_path / "config.yaml"),
|
||||
source_label="router_settings.plugins",
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ProxyConfig.__init__
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -793,6 +863,62 @@ async def test_ProxyConfig_load_config_minimal_yaml(tmp_path, monkeypatch):
|
|||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ProxyConfig_load_config_resolves_router_settings_plugins(tmp_path, monkeypatch):
|
||||
"""Regression: router_settings.plugins dotted-path strings must be resolved to
|
||||
live RoutingPlugin instances on the created Router. Previously they were passed
|
||||
through as raw strings and only blew up at request time when the pipeline tried
|
||||
to `await "some.string".run(context)`."""
|
||||
plugin_file = tmp_path / "rs_plugin.py"
|
||||
plugin_file.write_text(
|
||||
"class _Plugin:\n"
|
||||
" async def run(self, context):\n"
|
||||
" return context\n"
|
||||
"\n"
|
||||
"rs_plugin_instance = _Plugin()\n"
|
||||
)
|
||||
f = tmp_path / "c.yaml"
|
||||
f.write_text(
|
||||
"model_list: []\n"
|
||||
"general_settings: {}\n"
|
||||
"litellm_settings: {}\n"
|
||||
"router_settings:\n"
|
||||
" plugins:\n"
|
||||
" - rs_plugin.rs_plugin_instance\n"
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False)
|
||||
monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False)
|
||||
|
||||
router, _model_list, _general_settings = await ProxyConfig().load_config(
|
||||
router=None, config_file_path=str(f)
|
||||
)
|
||||
|
||||
assert len(router.routing_plugins) == 1
|
||||
assert type(router.routing_plugins[0]).__name__ == "_Plugin"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ProxyConfig_load_config_rejects_bad_router_settings_plugin(tmp_path, monkeypatch):
|
||||
plugin_file = tmp_path / "bad_rs_plugin.py"
|
||||
plugin_file.write_text("not_a_plugin = object()\n")
|
||||
f = tmp_path / "c.yaml"
|
||||
f.write_text(
|
||||
"model_list: []\n"
|
||||
"general_settings: {}\n"
|
||||
"litellm_settings: {}\n"
|
||||
"router_settings:\n"
|
||||
" plugins:\n"
|
||||
" - bad_rs_plugin.not_a_plugin\n"
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False)
|
||||
monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False)
|
||||
|
||||
with pytest.raises(ValueError, match="does not implement the RoutingPlugin interface"):
|
||||
await ProxyConfig().load_config(router=None, config_file_path=str(f))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ProxyConfig_load_config_wires_general_settings_url_validation(tmp_path, monkeypatch):
|
||||
"""Regression for #26599: SSRF settings in general_settings must reach litellm globals."""
|
||||
|
|
|
|||
|
|
@ -242,3 +242,88 @@ class TestRagIngestSSRFBlocked:
|
|||
assert response.status_code != 400, (
|
||||
f"Clean Bedrock ingest_options should not be rejected: {response.json()}"
|
||||
)
|
||||
|
||||
|
||||
def test_rag_query_returns_response_cost_header(client_internal_user):
|
||||
"""
|
||||
/v1/rag/query must surface the completion cost via the
|
||||
x-litellm-response-cost response header, like /v1/chat/completions does.
|
||||
"""
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
mock_response = ModelResponse(
|
||||
id="chatcmpl-test",
|
||||
choices=[
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "The codename is AZURE-FALCON-42."},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
model="gpt-4o-mini",
|
||||
usage={"prompt_tokens": 35, "completion_tokens": 14, "total_tokens": 49},
|
||||
)
|
||||
mock_response._hidden_params["response_cost"] = 3.45e-06
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.rag_endpoints.endpoints.litellm.aquery",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_response,
|
||||
), patch("litellm.vector_store_registry", None), patch(
|
||||
"litellm.proxy.proxy_server.prisma_client", None
|
||||
):
|
||||
response = client_internal_user.post(
|
||||
"/v1/rag/query",
|
||||
json={
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [{"role": "user", "content": "What is the codename?"}],
|
||||
"retrieval_config": {
|
||||
"vector_store_id": "vs_test_123",
|
||||
"custom_llm_provider": "openai",
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200, response.json()
|
||||
assert response.headers.get("x-litellm-response-cost") == "3.45e-06"
|
||||
|
||||
|
||||
def test_rag_query_stream_returns_event_stream(client_internal_user):
|
||||
"""
|
||||
A stream=true /v1/rag/query must return an SSE response. Returning the raw
|
||||
stream wrapper makes FastAPI try to serialize it, which raises and turns
|
||||
every streaming RAG query into a 500; the stream then never drains, so its
|
||||
single billing event (which carries the folded sub-call costs) never fires.
|
||||
"""
|
||||
import litellm as litellm_module
|
||||
|
||||
async def fake_aquery(**kwargs):
|
||||
return await litellm_module.acompletion(
|
||||
model="gpt-4o-mini",
|
||||
messages=[{"role": "user", "content": "What is the codename?"}],
|
||||
mock_response="The codename is AZURE-FALCON-42.",
|
||||
stream=True,
|
||||
api_key="test-key",
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.rag_endpoints.endpoints.litellm.aquery",
|
||||
new=AsyncMock(side_effect=fake_aquery),
|
||||
), patch("litellm.vector_store_registry", None), patch("litellm.proxy.proxy_server.prisma_client", None):
|
||||
response = client_internal_user.post(
|
||||
"/v1/rag/query",
|
||||
json={
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [{"role": "user", "content": "What is the codename?"}],
|
||||
"retrieval_config": {
|
||||
"vector_store_id": "vs_test_123",
|
||||
"custom_llm_provider": "openai",
|
||||
},
|
||||
"stream": True,
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.headers.get("content-type", "").startswith("text/event-stream")
|
||||
assert '"object":"chat.completion.chunk"' in response.text
|
||||
assert "data: [DONE]" in response.text
|
||||
|
|
|
|||
|
|
@ -64,6 +64,65 @@ def test_dotted_module_path_is_unaffected_by_gate():
|
|||
assert result == "loaded"
|
||||
|
||||
|
||||
def test_installed_package_resolved_when_local_file_absent(tmp_path, monkeypatch):
|
||||
# Regression: with config_file_path set (startup load path) but no local
|
||||
# module file next to it, get_instance_fn must fall back to importing the
|
||||
# dotted name as an installed package. Previously it raised ImportError
|
||||
# ("Could not find module file ..."), so plugins shipped as pip packages
|
||||
# (e.g. router_settings/complexity_router plugins) could not be referenced.
|
||||
pkg_dir = tmp_path / "site"
|
||||
pkg_dir.mkdir()
|
||||
(pkg_dir / "my_installed_plugin.py").write_text(
|
||||
"class _P:\n"
|
||||
" async def run(self, context):\n"
|
||||
" return context\n"
|
||||
"\n"
|
||||
"instance = _P()\n"
|
||||
)
|
||||
monkeypatch.syspath_prepend(str(pkg_dir))
|
||||
config_dir = tmp_path / "cfg"
|
||||
config_dir.mkdir()
|
||||
|
||||
result = get_instance_fn(
|
||||
value="my_installed_plugin.instance",
|
||||
config_file_path=str(config_dir / "config.yaml"),
|
||||
)
|
||||
|
||||
assert type(result).__name__ == "_P"
|
||||
|
||||
|
||||
def test_local_module_file_wins_over_installed_package(tmp_path, monkeypatch):
|
||||
# A local module file next to the config must still take precedence over an
|
||||
# installed package of the same dotted name -- the fallback only kicks in
|
||||
# when no local file exists.
|
||||
pkg_dir = tmp_path / "site"
|
||||
pkg_dir.mkdir()
|
||||
(pkg_dir / "shadowed_mod.py").write_text("value = 'from-installed'\n")
|
||||
monkeypatch.syspath_prepend(str(pkg_dir))
|
||||
config_dir = tmp_path / "cfg"
|
||||
config_dir.mkdir()
|
||||
(config_dir / "shadowed_mod.py").write_text("value = 'from-local-file'\n")
|
||||
|
||||
result = get_instance_fn(
|
||||
value="shadowed_mod.value",
|
||||
config_file_path=str(config_dir / "config.yaml"),
|
||||
)
|
||||
|
||||
assert result == "from-local-file"
|
||||
|
||||
|
||||
def test_missing_module_everywhere_raises_import_error(tmp_path):
|
||||
# Neither a local file nor an installed package: the fallback import must
|
||||
# surface a real ImportError rather than silently succeeding.
|
||||
config_dir = tmp_path / "cfg"
|
||||
config_dir.mkdir()
|
||||
with pytest.raises(ImportError):
|
||||
get_instance_fn(
|
||||
value="definitely_not_a_real_module_xyz.instance",
|
||||
config_file_path=str(config_dir / "config.yaml"),
|
||||
)
|
||||
|
||||
|
||||
def test_pass_through_route_threads_config_file_path():
|
||||
# ``create_pass_through_route`` must forward ``config_file_path`` so
|
||||
# an operator with ``custom_handler: s3://...`` declared in
|
||||
|
|
|
|||
0
tests/test_litellm/rag/__init__.py
Normal file
0
tests/test_litellm/rag/__init__.py
Normal file
266
tests/test_litellm/rag/test_main.py
Normal file
266
tests/test_litellm/rag/test_main.py
Normal file
|
|
@ -0,0 +1,266 @@
|
|||
"""
|
||||
Tests for the RAG query pipeline in litellm/rag/main.py.
|
||||
|
||||
The RAG pipeline forwards its kwargs (including the parent litellm_logging_obj)
|
||||
into @client-decorated sub-calls (vector store search, completion). Each logging
|
||||
object allows exactly one async_success event, so if sub-calls are not marked as
|
||||
internal, the vector store search consumes the slot first and the LLM
|
||||
completion's usage/cost is never logged (spend tracking and budget enforcement
|
||||
are bypassed). These tests pin the invariant that the single billing event for
|
||||
aquery carries the completion response with real usage and cost.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import is_internal_call
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.types.utils import CallTypes, ModelResponse
|
||||
|
||||
|
||||
class RecordingLogger(CustomLogger):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.success_events = []
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
self.success_events.append({"kwargs": kwargs, "response_obj": response_obj})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("use_router", [False, True])
|
||||
async def test_aquery_single_billing_event_carries_completion_usage_and_cost(use_router):
|
||||
"""
|
||||
litellm.aquery must produce exactly one success event, and that event must
|
||||
carry the LLM completion (a ModelResponse with non-zero usage and cost),
|
||||
not the vector store search response. The proxy always passes a router, so
|
||||
both the router and non-router completion branches are pinned.
|
||||
"""
|
||||
recording_logger = RecordingLogger()
|
||||
original_callbacks = litellm.callbacks
|
||||
litellm.callbacks = [recording_logger]
|
||||
|
||||
router_kwargs = {}
|
||||
if use_router:
|
||||
router_kwargs["router"] = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-4o-mini",
|
||||
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "test-key"},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
try:
|
||||
response = await litellm.aquery(
|
||||
model="gpt-4o-mini",
|
||||
messages=[{"role": "user", "content": "What is the secret project codename?"}],
|
||||
retrieval_config={"vector_store_id": "vs_test_123", "custom_llm_provider": "openai"},
|
||||
mock_response="The secret project codename is AZURE-FALCON-42.",
|
||||
**router_kwargs,
|
||||
)
|
||||
|
||||
assert isinstance(response, ModelResponse)
|
||||
assert is_internal_call.get() is False
|
||||
|
||||
for _ in range(50):
|
||||
if recording_logger.success_events:
|
||||
break
|
||||
await asyncio.sleep(0.1)
|
||||
await asyncio.sleep(0.5)
|
||||
finally:
|
||||
litellm.callbacks = original_callbacks
|
||||
|
||||
assert len(recording_logger.success_events) == 1
|
||||
event = recording_logger.success_events[0]
|
||||
|
||||
response_obj = event["response_obj"]
|
||||
assert isinstance(response_obj, ModelResponse)
|
||||
assert response_obj.usage.total_tokens > 0
|
||||
|
||||
standard_logging_object = event["kwargs"]["standard_logging_object"]
|
||||
assert standard_logging_object["call_type"] == "aquery"
|
||||
assert standard_logging_object["total_tokens"] > 0
|
||||
assert standard_logging_object["prompt_tokens"] > 0
|
||||
assert standard_logging_object["completion_tokens"] > 0
|
||||
assert standard_logging_object["response_cost"] > 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aquery_response_hidden_params_carry_completion_cost():
|
||||
"""
|
||||
The aquery response must expose the completion's response_cost via hidden
|
||||
params, so the proxy can return the x-litellm-response-cost header.
|
||||
"""
|
||||
response = await litellm.aquery(
|
||||
model="gpt-4o-mini",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
retrieval_config={"vector_store_id": "vs_test_123", "custom_llm_provider": "openai"},
|
||||
mock_response="hi there",
|
||||
)
|
||||
|
||||
assert isinstance(response, ModelResponse)
|
||||
response_cost = response._hidden_params.get("response_cost")
|
||||
assert response_cost is not None
|
||||
assert response_cost > 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aquery_billed_cost_includes_priced_vector_store_search():
|
||||
"""
|
||||
When the vector store provider prices search calls (e.g. per-query cost),
|
||||
that cost must be folded into the aquery billing instead of being dropped
|
||||
with the suppressed sub-call event.
|
||||
"""
|
||||
recording_logger = RecordingLogger()
|
||||
original_callbacks = litellm.callbacks
|
||||
litellm.callbacks = [recording_logger]
|
||||
|
||||
try:
|
||||
with patch("litellm.rag.main.vector_store_search_cost", return_value=(0.002, 0.0)):
|
||||
response = await litellm.aquery(
|
||||
model="gpt-4o-mini",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
retrieval_config={"vector_store_id": "vs_test_123", "custom_llm_provider": "openai"},
|
||||
mock_response="hi there",
|
||||
)
|
||||
|
||||
for _ in range(50):
|
||||
if recording_logger.success_events:
|
||||
break
|
||||
await asyncio.sleep(0.1)
|
||||
await asyncio.sleep(0.5)
|
||||
finally:
|
||||
litellm.callbacks = original_callbacks
|
||||
|
||||
assert isinstance(response, ModelResponse)
|
||||
total_cost = response._hidden_params.get("response_cost")
|
||||
assert total_cost is not None
|
||||
assert total_cost > 0.002
|
||||
|
||||
assert len(recording_logger.success_events) == 1
|
||||
standard_logging_object = recording_logger.success_events[0]["kwargs"]["standard_logging_object"]
|
||||
assert standard_logging_object["response_cost"] == total_cost
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aquery_with_rerank_bills_once_and_folds_rerank_cost():
|
||||
"""
|
||||
When rerank is enabled, its sub-call must run under the internal-call
|
||||
context (no standalone billing event) and its cost must be folded into
|
||||
the single aquery billing event.
|
||||
"""
|
||||
from litellm.types.rerank import RerankResponse
|
||||
|
||||
recording_logger = RecordingLogger()
|
||||
original_callbacks = litellm.callbacks
|
||||
litellm.callbacks = [recording_logger]
|
||||
rerank_seen = {}
|
||||
|
||||
async def fake_arerank(**kwargs):
|
||||
rerank_seen["internal"] = is_internal_call.get()
|
||||
rerank_result = RerankResponse(id="rr_1", results=[{"index": 0, "relevance_score": 0.9}], meta={})
|
||||
rerank_result._hidden_params["response_cost"] = 0.001
|
||||
return rerank_result
|
||||
|
||||
try:
|
||||
with patch("litellm.arerank", side_effect=fake_arerank):
|
||||
response = await litellm.aquery(
|
||||
model="gpt-4o-mini",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
retrieval_config={"vector_store_id": "vs_test_123", "custom_llm_provider": "openai"},
|
||||
rerank={"enabled": True, "model": "cohere/rerank-english-v3.0", "top_n": 1},
|
||||
mock_response="hi there",
|
||||
)
|
||||
|
||||
for _ in range(50):
|
||||
if recording_logger.success_events:
|
||||
break
|
||||
await asyncio.sleep(0.1)
|
||||
await asyncio.sleep(0.5)
|
||||
finally:
|
||||
litellm.callbacks = original_callbacks
|
||||
|
||||
assert rerank_seen["internal"] is True
|
||||
assert is_internal_call.get() is False
|
||||
|
||||
assert isinstance(response, ModelResponse)
|
||||
total_cost = response._hidden_params.get("response_cost")
|
||||
assert total_cost is not None
|
||||
assert total_cost > 0.001
|
||||
|
||||
assert len(recording_logger.success_events) == 1
|
||||
standard_logging_object = recording_logger.success_events[0]["kwargs"]["standard_logging_object"]
|
||||
assert standard_logging_object["call_type"] == "aquery"
|
||||
assert standard_logging_object["response_cost"] == total_cost
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aquery_streaming_bills_sub_call_costs_into_final_event():
|
||||
"""
|
||||
On the streaming path the response cost is computed from the assembled
|
||||
chunks after the pipeline returns, so there is no response object to fold
|
||||
sub-call costs into. The pipeline must instead carry the accumulated
|
||||
search and rerank cost through the logging object so the single streamed
|
||||
billing event includes it; otherwise a caller passing stream=true incurs
|
||||
priced vector search and rerank costs that never reach spend tracking.
|
||||
"""
|
||||
from litellm.types.rerank import RerankResponse
|
||||
|
||||
recording_logger = RecordingLogger()
|
||||
original_callbacks = litellm.callbacks
|
||||
litellm.callbacks = [recording_logger]
|
||||
rerank_seen = {}
|
||||
|
||||
async def fake_arerank(**kwargs):
|
||||
rerank_seen["internal"] = is_internal_call.get()
|
||||
rerank_result = RerankResponse(id="rr_1", results=[{"index": 0, "relevance_score": 0.9}], meta={})
|
||||
rerank_result._hidden_params["response_cost"] = 0.001
|
||||
return rerank_result
|
||||
|
||||
try:
|
||||
with (
|
||||
patch("litellm.rag.main.vector_store_search_cost", return_value=(0.002, 0.0)),
|
||||
patch("litellm.arerank", side_effect=fake_arerank),
|
||||
):
|
||||
response = await litellm.aquery(
|
||||
model="gpt-4o-mini",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
retrieval_config={"vector_store_id": "vs_test_123", "custom_llm_provider": "openai"},
|
||||
rerank={"enabled": True, "model": "cohere/rerank-english-v3.0", "top_n": 1},
|
||||
mock_response="hi there",
|
||||
stream=True,
|
||||
)
|
||||
async for _ in response:
|
||||
pass
|
||||
|
||||
for _ in range(50):
|
||||
if recording_logger.success_events:
|
||||
break
|
||||
await asyncio.sleep(0.1)
|
||||
await asyncio.sleep(0.5)
|
||||
finally:
|
||||
litellm.callbacks = original_callbacks
|
||||
|
||||
assert rerank_seen["internal"] is True
|
||||
assert is_internal_call.get() is False
|
||||
|
||||
assert len(recording_logger.success_events) == 1
|
||||
standard_logging_object = recording_logger.success_events[0]["kwargs"]["standard_logging_object"]
|
||||
assert standard_logging_object["call_type"] == "aquery"
|
||||
assert standard_logging_object["response_cost"] >= 0.003
|
||||
|
||||
|
||||
def test_rag_call_types_are_registered():
|
||||
"""
|
||||
query/aquery/ingest/aingest are @client-decorated entry points, so their
|
||||
function names must resolve to CallTypes members (deployment hooks and
|
||||
call-type driven logic silently no-op for unregistered call types).
|
||||
"""
|
||||
assert CallTypes("query") is CallTypes.query
|
||||
assert CallTypes("aquery") is CallTypes.aquery
|
||||
assert CallTypes("ingest") is CallTypes.ingest
|
||||
assert CallTypes("aingest") is CallTypes.aingest
|
||||
|
|
@ -3119,3 +3119,260 @@ class TestRoutingPlugins:
|
|||
assert first.model == "gpt-4o-mini"
|
||||
assert second.model == "gpt-4o-mini"
|
||||
assert spy.call_count == 2
|
||||
|
||||
|
||||
class TestEscalationKeywords:
|
||||
"""Test user-triggered escalation: a keyword in the prompt bumps the resolved tier
|
||||
one step higher so a user can force a stronger model when unhappy with results."""
|
||||
|
||||
@staticmethod
|
||||
def _request_kwargs(session_id: str) -> Dict:
|
||||
return {"metadata": {"session_id": session_id}}
|
||||
|
||||
def test_default_escalation_keyword(self, complexity_router):
|
||||
assert complexity_router.escalation_keywords == ["LITELLM ESCALATE"]
|
||||
|
||||
def test_escalation_triggered_is_case_sensitive(self, complexity_router):
|
||||
assert complexity_router._escalation_triggered("please LITELLM ESCALATE now") is True
|
||||
assert complexity_router._escalation_triggered("please litellm escalate now") is False
|
||||
assert complexity_router._escalation_triggered("how do I escalate this ticket") is False
|
||||
|
||||
def test_escalate_tier_bumps_one_step(self, complexity_router):
|
||||
assert complexity_router._escalate_tier(ComplexityTier.SIMPLE) == ComplexityTier.MEDIUM
|
||||
assert complexity_router._escalate_tier(ComplexityTier.MEDIUM) == ComplexityTier.COMPLEX
|
||||
assert complexity_router._escalate_tier(ComplexityTier.COMPLEX) == ComplexityTier.REASONING
|
||||
|
||||
def test_escalate_tier_caps_at_highest_configured(self, complexity_router):
|
||||
assert complexity_router._escalate_tier(ComplexityTier.REASONING) == ComplexityTier.REASONING
|
||||
|
||||
def test_escalate_tier_skips_unconfigured_intermediate(self, mock_router_instance):
|
||||
router = ComplexityRouter(
|
||||
model_name="test-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config={"tiers": {"SIMPLE": "gpt-4o-mini", "REASONING": "o1-preview"}},
|
||||
)
|
||||
assert router._escalate_tier(ComplexityTier.SIMPLE) == ComplexityTier.REASONING
|
||||
|
||||
def test_tier_for_model_returns_most_severe(self, mock_router_instance):
|
||||
router = ComplexityRouter(
|
||||
model_name="test-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config={
|
||||
"tiers": {"SIMPLE": "shared", "COMPLEX": "shared", "REASONING": "top"}
|
||||
},
|
||||
)
|
||||
assert router._tier_for_model("shared") == ComplexityTier.COMPLEX
|
||||
assert router._tier_for_model("top") == ComplexityTier.REASONING
|
||||
assert router._tier_for_model("unknown") is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_escalation_bumps_classified_tier(self, mock_router_instance, basic_config):
|
||||
router = ComplexityRouter(
|
||||
model_name="test-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config=basic_config,
|
||||
)
|
||||
# Baseline: this prompt classifies SIMPLE.
|
||||
baseline = await router.async_pre_routing_hook(
|
||||
model="test-model", request_kwargs={}, messages=[{"role": "user", "content": "Hello there!"}]
|
||||
)
|
||||
assert baseline.model == "gpt-4o-mini"
|
||||
|
||||
escalated = await router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={},
|
||||
messages=[{"role": "user", "content": "LITELLM ESCALATE Hello there!"}],
|
||||
)
|
||||
assert escalated.model == "gpt-4o" # SIMPLE bumped to MEDIUM
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lowercase_keyword_does_not_escalate(self, mock_router_instance, basic_config):
|
||||
router = ComplexityRouter(
|
||||
model_name="test-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config=basic_config,
|
||||
)
|
||||
result = await router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={},
|
||||
messages=[{"role": "user", "content": "litellm escalate Hello there!"}],
|
||||
)
|
||||
assert result.model == "gpt-4o-mini" # not escalated
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_custom_escalation_keyword(self, mock_router_instance, basic_config):
|
||||
router = ComplexityRouter(
|
||||
model_name="test-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config={**basic_config, "escalation_keywords": ["MAKE IT BETTER"]},
|
||||
)
|
||||
# The default keyword no longer triggers once a custom list is supplied.
|
||||
default = await router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={},
|
||||
messages=[{"role": "user", "content": "LITELLM ESCALATE Hello there!"}],
|
||||
)
|
||||
assert default.model == "gpt-4o-mini"
|
||||
|
||||
custom = await router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={},
|
||||
messages=[{"role": "user", "content": "MAKE IT BETTER Hello there!"}],
|
||||
)
|
||||
assert custom.model == "gpt-4o"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_keyword_list_disables_escalation(self, mock_router_instance, basic_config):
|
||||
router = ComplexityRouter(
|
||||
model_name="test-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config={**basic_config, "escalation_keywords": []},
|
||||
)
|
||||
result = await router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={},
|
||||
messages=[{"role": "user", "content": "LITELLM ESCALATE Hello there!"}],
|
||||
)
|
||||
assert result.model == "gpt-4o-mini"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_escalation_caps_at_highest_tier(self, mock_router_instance, basic_config):
|
||||
router = ComplexityRouter(
|
||||
model_name="test-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config=basic_config,
|
||||
)
|
||||
result = await router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={},
|
||||
messages=[
|
||||
{
|
||||
"role": "user",
|
||||
"content": "LITELLM ESCALATE Let's think step by step and reason through this carefully.",
|
||||
}
|
||||
],
|
||||
)
|
||||
assert result.model == "o1-preview" # already REASONING, stays there
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_escalation_bumps_keyword_tier_override(self, mock_router_instance, basic_config):
|
||||
router = ComplexityRouter(
|
||||
model_name="test-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config={
|
||||
**basic_config,
|
||||
"keyword_tier_rules": [{"keywords": ["billing"], "tier": "SIMPLE"}],
|
||||
},
|
||||
)
|
||||
baseline = await router.async_pre_routing_hook(
|
||||
model="test-model", request_kwargs={}, messages=[{"role": "user", "content": "a billing question"}]
|
||||
)
|
||||
assert baseline.model == "gpt-4o-mini"
|
||||
|
||||
escalated = await router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={},
|
||||
messages=[{"role": "user", "content": "LITELLM ESCALATE a billing question"}],
|
||||
)
|
||||
assert escalated.model == "gpt-4o" # override SIMPLE bumped to MEDIUM
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_escalation_overrides_session_pin_and_persists(self, mock_router_instance, basic_config):
|
||||
"""Mid-session escalation bumps relative to the pinned model (never below it) and
|
||||
the bumped model persists for later turns."""
|
||||
mock_router_instance.cache = DualCache()
|
||||
router = ComplexityRouter(
|
||||
model_name="test-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config={**basic_config, "session_affinity": True},
|
||||
)
|
||||
request_kwargs = self._request_kwargs("session-1")
|
||||
first = await router.async_pre_routing_hook(
|
||||
model="test-model", request_kwargs=request_kwargs, messages=[{"role": "user", "content": "Hello!"}]
|
||||
)
|
||||
assert first.model == "gpt-4o-mini" # pinned SIMPLE
|
||||
|
||||
with patch.object(router, "aclassify", wraps=router.aclassify) as spy_aclassify:
|
||||
escalated = await router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs=request_kwargs,
|
||||
messages=[{"role": "user", "content": "LITELLM ESCALATE"}],
|
||||
)
|
||||
spy_aclassify.assert_not_called()
|
||||
assert escalated.model == "gpt-4o" # bumped relative to the SIMPLE pin, not reclassified
|
||||
|
||||
# The bump persists: a later ordinary turn stays on the escalated model.
|
||||
later = await router.async_pre_routing_hook(
|
||||
model="test-model", request_kwargs=request_kwargs, messages=[{"role": "user", "content": "thanks"}]
|
||||
)
|
||||
assert later.model == "gpt-4o"
|
||||
|
||||
# Escalating again climbs one more tier.
|
||||
again = await router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs=request_kwargs,
|
||||
messages=[{"role": "user", "content": "LITELLM ESCALATE still not good"}],
|
||||
)
|
||||
assert again.model == "claude-sonnet-4-20250514" # MEDIUM bumped to COMPLEX
|
||||
|
||||
def test_blank_escalation_keywords_are_stripped(self):
|
||||
"""Blank/whitespace-only phrases are dropped so `"" in message` can't escalate
|
||||
every request; surrounding whitespace on real phrases is trimmed."""
|
||||
assert ComplexityRouterConfig(
|
||||
tiers={"SIMPLE": "gpt-4o-mini", "MEDIUM": "gpt-4o"},
|
||||
escalation_keywords=["", " "],
|
||||
).escalation_keywords == []
|
||||
assert ComplexityRouterConfig(
|
||||
tiers={"SIMPLE": "gpt-4o-mini", "MEDIUM": "gpt-4o"},
|
||||
escalation_keywords=[" LITELLM ESCALATE ", ""],
|
||||
).escalation_keywords == ["LITELLM ESCALATE"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_blank_escalation_keyword_does_not_escalate_everything(
|
||||
self, mock_router_instance, basic_config
|
||||
):
|
||||
router = ComplexityRouter(
|
||||
model_name="test-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config={**basic_config, "escalation_keywords": [""]},
|
||||
)
|
||||
assert router.escalation_keywords == []
|
||||
result = await router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs={},
|
||||
messages=[{"role": "user", "content": "Hello there!"}],
|
||||
)
|
||||
assert result.model == "gpt-4o-mini" # not escalated
|
||||
|
||||
def test_escalated_pin_stays_on_same_model_at_ceiling(self, mock_router_instance):
|
||||
"""At the highest configured tier escalation keeps the exact pinned model, even
|
||||
when that tier's pool has peers `get_model_for_tier` could randomly pick instead."""
|
||||
router = ComplexityRouter(
|
||||
model_name="test-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config={
|
||||
"tiers": {"SIMPLE": "gpt-4o-mini", "REASONING": ["o1-a", "o1-b", "o1-c"]}
|
||||
},
|
||||
)
|
||||
for pinned in ("o1-a", "o1-b", "o1-c"):
|
||||
assert router._escalated_pin(pinned) == pinned
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_escalation_at_ceiling_keeps_multi_model_pin(self, mock_router_instance):
|
||||
mock_router_instance.cache = DualCache()
|
||||
router = ComplexityRouter(
|
||||
model_name="test-router",
|
||||
litellm_router_instance=mock_router_instance,
|
||||
complexity_router_config={
|
||||
"tiers": {"SIMPLE": "gpt-4o-mini", "REASONING": ["o1-a", "o1-b", "o1-c"]},
|
||||
"session_affinity": True,
|
||||
},
|
||||
)
|
||||
cache_key = router._get_session_affinity_cache_key("session-top", {})
|
||||
await mock_router_instance.cache.async_set_cache(key=cache_key, value="o1-b")
|
||||
result = await router.async_pre_routing_hook(
|
||||
model="test-model",
|
||||
request_kwargs=self._request_kwargs("session-top"),
|
||||
messages=[{"role": "user", "content": "LITELLM ESCALATE do better"}],
|
||||
)
|
||||
assert result.model == "o1-b" # unchanged: no random hop to o1-a / o1-c
|
||||
|
|
|
|||
82
tests/test_litellm/test_gpt_realtime_mode.py
Normal file
82
tests/test_litellm/test_gpt_realtime_mode.py
Normal file
|
|
@ -0,0 +1,82 @@
|
|||
import json
|
||||
import typing
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.types.utils import ModelInfoBase
|
||||
|
||||
REALTIME_ONLY_GPT_MODELS = (
|
||||
"azure/gpt-realtime-2025-08-28",
|
||||
"azure/gpt-realtime-1.5-2026-02-23",
|
||||
"azure/gpt-realtime-mini-2025-10-06",
|
||||
"gpt-realtime",
|
||||
"gpt-realtime-1.5",
|
||||
"gpt-realtime-2",
|
||||
"gpt-realtime-2.1",
|
||||
"gpt-realtime-2.1-mini",
|
||||
"gpt-realtime-mini",
|
||||
"gpt-realtime-2025-08-28",
|
||||
"gpt-realtime-mini-2025-10-06",
|
||||
"gpt-realtime-mini-2025-12-15",
|
||||
)
|
||||
|
||||
REALTIME_ONLY_GPT_MODELS_WITHOUT_ENDPOINTS = (
|
||||
"azure/eu/gpt-4o-mini-realtime-preview-2024-12-17",
|
||||
"azure/eu/gpt-4o-realtime-preview-2024-10-01",
|
||||
"azure/eu/gpt-4o-realtime-preview-2024-12-17",
|
||||
"azure/gpt-4o-mini-realtime-preview-2024-12-17",
|
||||
"azure/gpt-4o-realtime-preview-2024-10-01",
|
||||
"azure/gpt-4o-realtime-preview-2024-12-17",
|
||||
"azure/us/gpt-4o-mini-realtime-preview-2024-12-17",
|
||||
"azure/us/gpt-4o-realtime-preview-2024-10-01",
|
||||
"azure/us/gpt-4o-realtime-preview-2024-12-17",
|
||||
"gpt-4o-mini-realtime-preview",
|
||||
"gpt-4o-mini-realtime-preview-2024-12-17",
|
||||
"gpt-4o-realtime-preview",
|
||||
"gpt-4o-realtime-preview-2024-12-17",
|
||||
"gpt-4o-realtime-preview-2025-06-03",
|
||||
)
|
||||
|
||||
ALL_REALTIME_ONLY_GPT_MODELS = REALTIME_ONLY_GPT_MODELS + REALTIME_ONLY_GPT_MODELS_WITHOUT_ENDPOINTS
|
||||
|
||||
|
||||
def _load_cost_map() -> dict:
|
||||
json_path = Path(__file__).parents[2] / "model_prices_and_context_window.json"
|
||||
with open(json_path) as f:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
def test_realtime_is_a_valid_mode_literal():
|
||||
hints = typing.get_type_hints(ModelInfoBase, include_extras=False)
|
||||
assert "realtime" in typing.get_args(hints["mode"])
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", REALTIME_ONLY_GPT_MODELS)
|
||||
def test_realtime_only_gpt_models_are_mode_realtime(model):
|
||||
"""These models only serve /v1/realtime and are rejected by /v1/chat/completions
|
||||
("This is not a chat model ..."), so they must not be tagged mode=chat."""
|
||||
info = _load_cost_map()[model]
|
||||
assert info["supported_endpoints"] == ["/v1/realtime"]
|
||||
assert info["mode"] == "realtime"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", REALTIME_ONLY_GPT_MODELS_WITHOUT_ENDPOINTS)
|
||||
def test_realtime_only_gpt_4o_models_are_mode_realtime(model):
|
||||
"""gpt-4o(-mini)-realtime-preview are realtime-only and must not be mode=chat."""
|
||||
assert _load_cost_map()[model]["mode"] == "realtime"
|
||||
|
||||
|
||||
def test_get_model_info_reports_realtime_mode():
|
||||
assert litellm.get_model_info("gpt-realtime-mini")["mode"] == "realtime"
|
||||
|
||||
|
||||
def test_backup_matches_main_for_realtime_models():
|
||||
repo_root = Path(__file__).parents[2]
|
||||
with open(repo_root / "model_prices_and_context_window.json") as f:
|
||||
main_cost = json.load(f)
|
||||
with open(repo_root / "litellm" / "model_prices_and_context_window_backup.json") as f:
|
||||
backup_cost = json.load(f)
|
||||
for model in ALL_REALTIME_ONLY_GPT_MODELS:
|
||||
assert backup_cost.get(model) == main_cost.get(model)
|
||||
|
|
@ -269,4 +269,22 @@ describe("ComplexityRouterConfig", () => {
|
|||
);
|
||||
expect(screen.getAllByText("This tier is required")).toHaveLength(1);
|
||||
});
|
||||
|
||||
it("renders the escalation keywords section with current keywords when the handler is provided", () => {
|
||||
renderWithProviders(
|
||||
<ComplexityRouterConfig
|
||||
{...baseProps}
|
||||
escalationKeywords={["LITELLM ESCALATE"]}
|
||||
onEscalationKeywordsChange={vi.fn()}
|
||||
/>,
|
||||
);
|
||||
fireEvent.click(screen.getByText("Advanced: Escalation Keywords"));
|
||||
expect(screen.getByText("Escalation Keywords")).toBeInTheDocument();
|
||||
expect(screen.getByText("LITELLM ESCALATE")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("hides the escalation keywords section when no handler is provided", () => {
|
||||
renderWithProviders(<ComplexityRouterConfig {...baseProps} />);
|
||||
expect(screen.queryByText("Advanced: Escalation Keywords")).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ import React from "react";
|
|||
import { ModelGroup } from "@/components/llm_calls/fetch_models";
|
||||
import AdaptiveRoutingConfig from "./AdaptiveRoutingConfig";
|
||||
import ClassificationMethodConfig from "./ClassificationMethodConfig";
|
||||
import EscalationKeywords from "./EscalationKeywords";
|
||||
import KeywordTierRules, { KeywordTierRule } from "./KeywordTierRules";
|
||||
import SemanticKeywordMatching from "./SemanticKeywordMatching";
|
||||
|
||||
|
|
@ -61,6 +62,8 @@ interface ComplexityRouterConfigProps {
|
|||
onEmbeddingModelChange?: (model: string) => void;
|
||||
matchThreshold?: number;
|
||||
onMatchThresholdChange?: (threshold: number) => void;
|
||||
escalationKeywords?: string[];
|
||||
onEscalationKeywordsChange?: (keywords: string[]) => void;
|
||||
showValidationErrors?: boolean;
|
||||
}
|
||||
|
||||
|
|
@ -101,6 +104,8 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
|
|||
onEmbeddingModelChange = () => {},
|
||||
matchThreshold = 0.5,
|
||||
onMatchThresholdChange = () => {},
|
||||
escalationKeywords = [],
|
||||
onEscalationKeywordsChange,
|
||||
showValidationErrors = false,
|
||||
}) => {
|
||||
// Embedding models can't serve a chat-completion role, so they're excluded here.
|
||||
|
|
@ -213,6 +218,19 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
|
|||
),
|
||||
children: <AdaptiveRoutingConfig value={value} onChange={onChange} />,
|
||||
},
|
||||
...(onEscalationKeywordsChange
|
||||
? [
|
||||
{
|
||||
key: "escalation",
|
||||
label: (
|
||||
<Text strong style={{ color: "#374151" }}>
|
||||
Advanced: Escalation Keywords
|
||||
</Text>
|
||||
),
|
||||
children: <EscalationKeywords keywords={escalationKeywords} onChange={onEscalationKeywordsChange} />,
|
||||
},
|
||||
]
|
||||
: []),
|
||||
...(onKeywordTierRulesChange || onSemanticMatchingEnabledChange
|
||||
? [
|
||||
{
|
||||
|
|
|
|||
|
|
@ -0,0 +1,45 @@
|
|||
import { InfoCircleOutlined } from "@ant-design/icons";
|
||||
import { Select as AntdSelect, Tooltip, Typography } from "antd";
|
||||
import React from "react";
|
||||
|
||||
const { Text } = Typography;
|
||||
|
||||
export const DEFAULT_ESCALATION_KEYWORDS = ["LITELLM ESCALATE"];
|
||||
|
||||
interface EscalationKeywordsProps {
|
||||
keywords: string[];
|
||||
onChange: (keywords: string[]) => void;
|
||||
}
|
||||
|
||||
const EscalationKeywords: React.FC<EscalationKeywordsProps> = ({ keywords, onChange }) => {
|
||||
return (
|
||||
<div className="w-full max-w-none">
|
||||
<div className="flex items-center gap-2 mb-1">
|
||||
<Typography.Title level={4} style={{ margin: 0 }}>
|
||||
Escalation Keywords
|
||||
</Typography.Title>
|
||||
<Tooltip title="Case-sensitive phrases a user can include in their message to force a bump to the next-higher complexity tier when they aren't happy with results. They can force a stronger model, but not choose which one.">
|
||||
<InfoCircleOutlined className="text-gray-400" />
|
||||
</Tooltip>
|
||||
</div>
|
||||
<Text type="secondary" style={{ display: "block", marginBottom: 8, fontSize: 12 }}>
|
||||
Optional: when a user message contains one of these phrases, the request is bumped one tier higher than it would
|
||||
otherwise route to. Matching is case-sensitive, so "LITELLM ESCALATE" only fires on the exact, shouted
|
||||
form. Leave empty to disable.
|
||||
</Text>
|
||||
<AntdSelect
|
||||
mode="tags"
|
||||
value={keywords}
|
||||
onChange={onChange}
|
||||
placeholder="e.g., LITELLM ESCALATE"
|
||||
tokenSeparators={[","]}
|
||||
open={false}
|
||||
suffixIcon={null}
|
||||
style={{ width: "100%" }}
|
||||
allowClear
|
||||
/>
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
export default EscalationKeywords;
|
||||
|
|
@ -14,6 +14,7 @@ import ComplexityRouterConfig, {
|
|||
DEFAULT_TIER_DISTANCE_PENALTY,
|
||||
} from "./ComplexityRouterConfig";
|
||||
import { KeywordTierRule } from "./KeywordTierRules";
|
||||
import { DEFAULT_ESCALATION_KEYWORDS } from "./EscalationKeywords";
|
||||
import { DEFAULT_MATCH_THRESHOLD } from "./SemanticKeywordMatching";
|
||||
import {
|
||||
buildComplexityRouterConfig,
|
||||
|
|
@ -52,6 +53,7 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
|
|||
const [semanticMatchingEnabled, setSemanticMatchingEnabled] = useState<boolean>(false);
|
||||
const [embeddingModel, setEmbeddingModel] = useState<string | undefined>(undefined);
|
||||
const [matchThreshold, setMatchThreshold] = useState<number>(DEFAULT_MATCH_THRESHOLD);
|
||||
const [escalationKeywords, setEscalationKeywords] = useState<string[]>(DEFAULT_ESCALATION_KEYWORDS);
|
||||
const [showValidationErrors, setShowValidationErrors] = useState<boolean>(false);
|
||||
|
||||
// Semantic router config (existing)
|
||||
|
|
@ -141,6 +143,7 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
|
|||
semanticMatchingEnabled,
|
||||
embeddingModel,
|
||||
matchThreshold,
|
||||
escalationKeywords,
|
||||
adaptive,
|
||||
adaptiveWeights,
|
||||
tierDistancePenalty,
|
||||
|
|
@ -316,6 +319,8 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({ form, handleOk, acc
|
|||
onEmbeddingModelChange={setEmbeddingModel}
|
||||
matchThreshold={matchThreshold}
|
||||
onMatchThresholdChange={setMatchThreshold}
|
||||
escalationKeywords={escalationKeywords}
|
||||
onEscalationKeywordsChange={setEscalationKeywords}
|
||||
showValidationErrors={showValidationErrors}
|
||||
/>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ const baseParams: BuildComplexityRouterConfigParams = {
|
|||
semanticMatchingEnabled: false,
|
||||
embeddingModel: undefined,
|
||||
matchThreshold: 0.5,
|
||||
escalationKeywords: ["LITELLM ESCALATE"],
|
||||
adaptive: false,
|
||||
adaptiveWeights: { quality: 0.3, cost: 0.7 },
|
||||
tierDistancePenalty: 0.5,
|
||||
|
|
@ -28,9 +29,22 @@ const baseParams: BuildComplexityRouterConfigParams = {
|
|||
};
|
||||
|
||||
describe("buildComplexityRouterConfig", () => {
|
||||
it("emits only tiers and classifier_type when nothing else is configured", () => {
|
||||
it("emits tiers, classifier_type, and escalation_keywords when nothing else is configured", () => {
|
||||
const config = buildComplexityRouterConfig(baseParams);
|
||||
expect(config).toEqual({ tiers, classifier_type: "heuristic" });
|
||||
expect(config).toEqual({ tiers, classifier_type: "heuristic", escalation_keywords: ["LITELLM ESCALATE"] });
|
||||
});
|
||||
|
||||
it("trims escalation keywords and drops blank entries", () => {
|
||||
const config = buildComplexityRouterConfig({
|
||||
...baseParams,
|
||||
escalationKeywords: [" LITELLM ESCALATE ", "", " ", "MAKE IT BETTER"],
|
||||
});
|
||||
expect(config.escalation_keywords).toEqual(["LITELLM ESCALATE", "MAKE IT BETTER"]);
|
||||
});
|
||||
|
||||
it("emits an empty escalation_keywords list so clearing the field disables escalation", () => {
|
||||
const config = buildComplexityRouterConfig({ ...baseParams, escalationKeywords: [] });
|
||||
expect(config.escalation_keywords).toEqual([]);
|
||||
});
|
||||
|
||||
it("passes through a tier configured with more than one model as a pool", () => {
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ export interface BuildComplexityRouterConfigParams {
|
|||
semanticMatchingEnabled: boolean;
|
||||
embeddingModel: string | undefined;
|
||||
matchThreshold: number;
|
||||
escalationKeywords: string[];
|
||||
adaptive: boolean;
|
||||
adaptiveWeights: AdaptiveRouterWeights;
|
||||
tierDistancePenalty: number;
|
||||
|
|
@ -31,6 +32,7 @@ export interface ComplexityRouterConfigPayload {
|
|||
semantic_keyword_matching?: boolean;
|
||||
embedding_model?: string;
|
||||
match_threshold?: number;
|
||||
escalation_keywords?: string[];
|
||||
adaptive?: boolean;
|
||||
adaptive_weights?: AdaptiveRouterWeights;
|
||||
tier_distance_penalty?: number;
|
||||
|
|
@ -69,11 +71,13 @@ export const buildComplexityRouterConfig = ({
|
|||
semanticMatchingEnabled,
|
||||
embeddingModel,
|
||||
matchThreshold,
|
||||
escalationKeywords,
|
||||
adaptive,
|
||||
adaptiveWeights,
|
||||
tierDistancePenalty,
|
||||
adaptiveEligible,
|
||||
}: BuildComplexityRouterConfigParams): ComplexityRouterConfigPayload => {
|
||||
const cleanedEscalationKeywords = escalationKeywords.map((keyword) => keyword.trim()).filter(Boolean);
|
||||
// Trim keywords and drop empty ones; drop any rule left with no keywords. Clicking
|
||||
// "Add keyword rule" seeds a rule with an empty keywords list, so without this an
|
||||
// unfilled row (common in the heuristic flow, where getSemanticConfigError doesn't run)
|
||||
|
|
@ -88,6 +92,7 @@ export const buildComplexityRouterConfig = ({
|
|||
...(classifierType === "llm" && classifierLlmConfig && { classifier_llm_config: classifierLlmConfig }),
|
||||
...(customTechnicalKeywords.length > 0 && { custom_technical_keywords: customTechnicalKeywords }),
|
||||
...(cleanedKeywordTierRules.length > 0 && { keyword_tier_rules: cleanedKeywordTierRules }),
|
||||
escalation_keywords: cleanedEscalationKeywords,
|
||||
...(semanticMatchingEnabled && {
|
||||
semantic_keyword_matching: true,
|
||||
embedding_model: embeddingModel,
|
||||
|
|
|
|||
7
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
7
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -21690,7 +21690,7 @@ export interface components {
|
|||
* CallTypes
|
||||
* @enum {string}
|
||||
*/
|
||||
CallTypes: "embedding" | "aembedding" | "completion" | "acompletion" | "atext_completion" | "text_completion" | "image_generation" | "aimage_generation" | "image_edit" | "aimage_edit" | "moderation" | "amoderation" | "atranscription" | "transcription" | "aspeech" | "speech" | "rerank" | "arerank" | "search" | "asearch" | "_arealtime" | "_aresponses_websocket" | "create_batch" | "acreate_batch" | "aretrieve_batch" | "retrieve_batch" | "acancel_batch" | "cancel_batch" | "pass_through_endpoint" | "anthropic_messages" | "aanthropic_messages" | "get_assistants" | "aget_assistants" | "create_assistants" | "acreate_assistants" | "delete_assistant" | "adelete_assistant" | "acreate_thread" | "create_thread" | "aget_thread" | "get_thread" | "a_add_message" | "add_message" | "aget_messages" | "get_messages" | "arun_thread" | "run_thread" | "arun_thread_stream" | "run_thread_stream" | "afile_retrieve" | "file_retrieve" | "afile_delete" | "file_delete" | "afile_list" | "file_list" | "acreate_file" | "create_file" | "afile_content" | "file_content" | "create_fine_tuning_job" | "acreate_fine_tuning_job" | "create_video" | "acreate_video" | "avideo_retrieve" | "video_retrieve" | "avideo_content" | "video_content" | "video_remix" | "avideo_remix" | "video_list" | "avideo_list" | "video_retrieve_job" | "avideo_retrieve_job" | "video_delete" | "avideo_delete" | "video_create_character" | "avideo_create_character" | "video_get_character" | "avideo_get_character" | "video_edit" | "avideo_edit" | "video_extension" | "avideo_extension" | "vector_store_file_create" | "avector_store_file_create" | "vector_store_file_list" | "avector_store_file_list" | "vector_store_file_retrieve" | "avector_store_file_retrieve" | "vector_store_file_content" | "avector_store_file_content" | "vector_store_file_update" | "avector_store_file_update" | "vector_store_file_delete" | "avector_store_file_delete" | "vector_store_create" | "avector_store_create" | "vector_store_search" | "avector_store_search" | "create_container" | "acreate_container" | "list_containers" | "alist_containers" | "retrieve_container" | "aretrieve_container" | "delete_container" | "adelete_container" | "list_container_files" | "alist_container_files" | "upload_container_file" | "aupload_container_file" | "create_sandbox" | "acreate_sandbox" | "delete_sandbox" | "adelete_sandbox" | "run_code" | "arun_code" | "code_interpreter_tool" | "acode_interpreter_tool" | "acancel_fine_tuning_job" | "cancel_fine_tuning_job" | "alist_fine_tuning_jobs" | "list_fine_tuning_jobs" | "aretrieve_fine_tuning_job" | "retrieve_fine_tuning_job" | "responses" | "aresponses" | "alist_input_items" | "llm_passthrough_route" | "allm_passthrough_route" | "generate_content" | "agenerate_content" | "generate_content_stream" | "agenerate_content_stream" | "ocr" | "aocr" | "call_mcp_tool" | "list_mcp_tools" | "asend_message" | "send_message" | "acreate_skill";
|
||||
CallTypes: "embedding" | "aembedding" | "completion" | "acompletion" | "atext_completion" | "text_completion" | "image_generation" | "aimage_generation" | "image_edit" | "aimage_edit" | "moderation" | "amoderation" | "atranscription" | "transcription" | "aspeech" | "speech" | "rerank" | "arerank" | "search" | "asearch" | "_arealtime" | "_aresponses_websocket" | "create_batch" | "acreate_batch" | "aretrieve_batch" | "retrieve_batch" | "acancel_batch" | "cancel_batch" | "pass_through_endpoint" | "anthropic_messages" | "aanthropic_messages" | "get_assistants" | "aget_assistants" | "create_assistants" | "acreate_assistants" | "delete_assistant" | "adelete_assistant" | "acreate_thread" | "create_thread" | "aget_thread" | "get_thread" | "a_add_message" | "add_message" | "aget_messages" | "get_messages" | "arun_thread" | "run_thread" | "arun_thread_stream" | "run_thread_stream" | "afile_retrieve" | "file_retrieve" | "afile_delete" | "file_delete" | "afile_list" | "file_list" | "acreate_file" | "create_file" | "afile_content" | "file_content" | "create_fine_tuning_job" | "acreate_fine_tuning_job" | "create_video" | "acreate_video" | "avideo_retrieve" | "video_retrieve" | "avideo_content" | "video_content" | "video_remix" | "avideo_remix" | "video_list" | "avideo_list" | "video_retrieve_job" | "avideo_retrieve_job" | "video_delete" | "avideo_delete" | "video_create_character" | "avideo_create_character" | "video_get_character" | "avideo_get_character" | "video_edit" | "avideo_edit" | "video_extension" | "avideo_extension" | "vector_store_file_create" | "avector_store_file_create" | "vector_store_file_list" | "avector_store_file_list" | "vector_store_file_retrieve" | "avector_store_file_retrieve" | "vector_store_file_content" | "avector_store_file_content" | "vector_store_file_update" | "avector_store_file_update" | "vector_store_file_delete" | "avector_store_file_delete" | "vector_store_create" | "avector_store_create" | "vector_store_search" | "avector_store_search" | "ingest" | "aingest" | "query" | "aquery" | "create_container" | "acreate_container" | "list_containers" | "alist_containers" | "retrieve_container" | "aretrieve_container" | "delete_container" | "adelete_container" | "list_container_files" | "alist_container_files" | "upload_container_file" | "aupload_container_file" | "create_sandbox" | "acreate_sandbox" | "delete_sandbox" | "adelete_sandbox" | "run_code" | "arun_code" | "code_interpreter_tool" | "acode_interpreter_tool" | "acancel_fine_tuning_job" | "cancel_fine_tuning_job" | "alist_fine_tuning_jobs" | "list_fine_tuning_jobs" | "aretrieve_fine_tuning_job" | "retrieve_fine_tuning_job" | "responses" | "aresponses" | "alist_input_items" | "llm_passthrough_route" | "allm_passthrough_route" | "generate_content" | "agenerate_content" | "generate_content_stream" | "agenerate_content_stream" | "ocr" | "aocr" | "call_mcp_tool" | "list_mcp_tools" | "asend_message" | "send_message" | "acreate_skill";
|
||||
/** CallbackDelete */
|
||||
CallbackDelete: {
|
||||
/** Callback Name */
|
||||
|
|
@ -21870,6 +21870,11 @@ export interface components {
|
|||
};
|
||||
/** ChatCompletionCachedContent */
|
||||
ChatCompletionCachedContent: {
|
||||
/**
|
||||
* Ttl
|
||||
* @enum {string}
|
||||
*/
|
||||
ttl?: "5m" | "1h";
|
||||
/**
|
||||
* Type
|
||||
* @constant
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue