Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_internal_copy_38013

# Conflicts:
#	litellm/proxy/credential_endpoints/endpoints.py
#	tests/test_litellm/proxy/credential_endpoints/test_endpoints.py
#	tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py
#	tests/test_litellm/test_router.py
This commit is contained in:
mateo-berri 2026-09-03 10:09:18 -07:00
commit 9036b90a38
108 changed files with 4560 additions and 941 deletions

View file

@ -11,17 +11,35 @@ import sys
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
import functools
import tempfile
from typing import Optional
from contextvars import ContextVar
from typing import TYPE_CHECKING, ClassVar, Literal, Optional
from litellm._logging import verbose_proxy_logger
from litellm.caching.caching import DualCache
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails._content_utils import walk_user_text
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
GUARDRAIL_NAME = "hide_secrets"
GUARDRAIL_PROVIDER = "hide-secrets"
# Per-invocation tally of redacted secrets by detect-secrets plugin type; None
# means the guardrail did not run, so _process_response records nothing.
_masked_entity_count: ContextVar[Optional[dict]] = ContextVar(
"hide_secrets_masked_entity_count", default=None
)
_custom_plugins_path = "file://" + os.path.join(
os.path.dirname(os.path.abspath(__file__)), "secrets_plugins"
)
@ -422,6 +440,10 @@ _default_detect_secrets_config = {
class _ENTERPRISE_SecretDetection(CustomGuardrail):
# Keeps proxied traffic on async_pre_call_hook (the unified apply_guardrail
# path skips should_run_check and never sees data["prompt"]).
use_native_lifecycle_hooks: ClassVar[bool] = True
def __init__(self, detect_secrets_config: Optional[dict] = None, **kwargs):
self.user_defined_detect_secrets_config = detect_secrets_config
super().__init__(**kwargs)
@ -455,6 +477,26 @@ class _ENTERPRISE_SecretDetection(CustomGuardrail):
return detected_secrets
def redact_text(self, text: str, source: str = "message") -> str:
"""Replace every detected secret in ``text`` with ``[REDACTED]`` and
tally the detected types into the per-invocation masked-entity count."""
detected_secrets = self.scan_message_for_secrets(text)
if not detected_secrets:
return text
counts = _masked_entity_count.get()
if counts is not None:
for secret in detected_secrets:
counts[secret["type"]] = counts.get(secret["type"], 0) + 1
secret_types = [secret["type"] for secret in detected_secrets]
verbose_proxy_logger.warning(
f"Detected and redacted secrets in {source}: {secret_types}"
)
return functools.reduce(
lambda redacted, secret: redacted.replace(secret["value"], "[REDACTED]"),
detected_secrets,
text,
)
async def should_run_check(self, user_api_key_dict: UserAPIKeyAuth) -> bool:
if user_api_key_dict.permissions is not None:
if GUARDRAIL_NAME in user_api_key_dict.permissions:
@ -463,7 +505,45 @@ class _ENTERPRISE_SecretDetection(CustomGuardrail):
return True
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional["LiteLLMLoggingObj"] = None,
) -> GenericGuardrailAPIInputs:
"""Unified-interface entrypoint, used by /guardrails/apply_guardrail
(the UI test playground). Proxied traffic keeps using
``async_pre_call_hook``, see ``use_native_lifecycle_hooks``."""
texts = inputs.get("texts")
if not texts or not any(texts):
return inputs
_masked_entity_count.set({})
return {**inputs, "texts": [self.redact_text(text) for text in texts]}
def _redact_prompt(self, data: dict) -> int:
"""Redact ``data["prompt"]`` (the text-completion shape, which
``walk_user_text`` does not cover) and return how many non-empty
strings were inspected."""
prompt = data.get("prompt")
if isinstance(prompt, str):
if not prompt:
return 0
data["prompt"] = self.redact_text(prompt, source="prompt")
return 1
if isinstance(prompt, list):
data["prompt"] = [ # mutable-ok: data["prompt"] is a list on the wire
self.redact_text(item, source="prompt")
if isinstance(item, str) and item
else item
for item in prompt
]
return sum(1 for item in prompt if isinstance(item, str) and item)
return 0
#### CALL HOOKS - proxy only ####
@log_guardrail_information
async def async_pre_call_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
@ -471,53 +551,84 @@ class _ENTERPRISE_SecretDetection(CustomGuardrail):
data: dict,
call_type: str, # "completion", "embeddings", "image_generation", "moderation"
):
_masked_entity_count.set(None)
if await self.should_run_check(user_api_key_dict) is False:
return
_masked_entity_count.set({})
# Covers multimodal list content + Responses-API input.
def _redact_message_text(text: str) -> str:
detected_secrets = self.scan_message_for_secrets(text)
for secret in detected_secrets:
text = text.replace(secret["value"], "[REDACTED]")
if detected_secrets:
secret_types = [secret["type"] for secret in detected_secrets]
verbose_proxy_logger.warning(
f"Detected and redacted secrets in message: {secret_types}"
)
return text
inspected = walk_user_text(data, self.redact_text) + self._redact_prompt(data)
walk_user_text(data, _redact_message_text)
if inspected == 0:
# Image-only, empty-text, and unsupported payloads inspected
# nothing, so recording "allow" would count a run that never
# looked at any content.
_masked_entity_count.set(None)
if "prompt" in data:
if isinstance(data["prompt"], str):
detected_secrets = self.scan_message_for_secrets(data["prompt"])
for secret in detected_secrets:
data["prompt"] = data["prompt"].replace(
secret["value"], "[REDACTED]"
)
if len(detected_secrets) > 0:
secret_types = [secret["type"] for secret in detected_secrets]
verbose_proxy_logger.warning(
f"Detected and redacted secrets in prompt: {secret_types}"
)
elif isinstance(data["prompt"], list):
# Index back into the list — assigning to ``item`` would only
# rebind the loop variable and leave ``data["prompt"]``
# carrying the unredacted secret.
for idx, item in enumerate(data["prompt"]):
if isinstance(item, str):
detected_secrets = self.scan_message_for_secrets(item)
for secret in detected_secrets:
item = item.replace(secret["value"], "[REDACTED]")
data["prompt"][idx] = item
if len(detected_secrets) > 0:
secret_types = [
secret["type"] for secret in detected_secrets
]
verbose_proxy_logger.warning(
f"Detected and redacted secrets in prompt: {secret_types}"
)
# ``data["input"]`` (Responses API and embeddings/moderation) is
# already covered by ``walk_user_text`` above.
return
def _process_response(
self,
response: Optional[dict],
request_data: dict,
start_time: Optional[float] = None,
end_time: Optional[float] = None,
duration: Optional[float] = None,
event_type: Optional[GuardrailEventHooks] = None,
original_inputs: Optional[dict] = None,
):
"""Record allow/mask plus the masked-entity tally for a completed run.
Records nothing when the guardrail inspected nothing (opted-out key,
empty inputs) or when the instance has no guardrail_name (legacy
``litellm_settings.callbacks`` deployments, which predate guardrail
telemetry and stay without it).
"""
counts = _masked_entity_count.get()
_masked_entity_count.set(None)
if counts is None or self.guardrail_name is None:
return response
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_json_response="mask" if counts else "allow",
request_data=request_data,
guardrail_status="success",
duration=duration,
start_time=start_time,
end_time=end_time,
event_type=event_type,
guardrail_provider=GUARDRAIL_PROVIDER,
masked_entity_count=counts,
)
return response
def _process_error(
self,
e: Exception,
request_data: dict,
start_time: Optional[float] = None,
end_time: Optional[float] = None,
duration: Optional[float] = None,
event_type: Optional[GuardrailEventHooks] = None,
):
"""Label the failed run with this guardrail's provider so error rows
group with the successful ones in the monitor. Nameless legacy
instances record nothing, matching ``_process_response``."""
_masked_entity_count.set(None)
if self.guardrail_name is None:
raise e
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_json_response=e,
request_data=request_data,
guardrail_status=(
"guardrail_intervened"
if self._is_guardrail_intervention(e)
else "guardrail_failed_to_respond"
),
duration=duration,
start_time=start_time,
end_time=end_time,
event_type=event_type,
guardrail_provider=GUARDRAIL_PROVIDER,
)
raise e

View file

@ -6,6 +6,7 @@ import shutil
import subprocess
import tempfile
import time
from dataclasses import dataclass, replace
from pathlib import Path
from typing import Optional
@ -45,6 +46,38 @@ _MIGRATION_TS_RE = re.compile(r"^(\d{14})_")
_MIGRATION_DEADLOCK_MARKER = "deadlock detected"
MAX_MIGRATE_DEPLOY_ATTEMPTS = 4
@dataclass(frozen=True)
class _MigrateAttemptBudget:
"""Retries left, and the recoveries already run.
A recovery that lands something new costs nothing, so a database full of
objects `prisma db push` created works through them one per pass. Anything
that made no progress spends an attempt, so a stuck run still gives up.
"""
attempts_left: int
recoveries: frozenset[str] = frozenset()
@property
def exhausted(self) -> bool:
return self.attempts_left <= 0
@property
def attempt_number(self) -> int:
return MAX_MIGRATE_DEPLOY_ATTEMPTS - self.attempts_left + 1
def spend(self) -> "_MigrateAttemptBudget":
return replace(self, attempts_left=self.attempts_left - 1)
def after_recovery(self, recovery: str) -> "_MigrateAttemptBudget":
if recovery in self.recoveries:
return self.spend()
return replace(self, recoveries=self.recoveries | {recovery})
_SPEND_LOGS_ALTER_RE = re.compile(r'^ALTER\s+TABLE\s+"LiteLLM_SpendLogs"\s', re.IGNORECASE)
_SPEND_LOGS_ARTIFACT_DROP_RE = re.compile(
r'^DROP\s+TABLE\s+"LiteLLM_SpendLogs_[^"]*"', re.IGNORECASE
@ -716,6 +749,9 @@ class ProxyExtrasDBManager:
Ahead-of-HEAD state (DB has migrations newer than this build ships)
is logged as a warning, not a fatal error — users whose DBs got into
weird shapes from the old thrashing should still be able to start.
The retry budget only counts attempts that made no progress: see
_MigrateAttemptBudget.
"""
schema_path = ProxyExtrasDBManager._get_prisma_dir() + "/schema.prisma"
migrations_dir = ProxyExtrasDBManager._get_prisma_dir()
@ -749,8 +785,9 @@ class ProxyExtrasDBManager:
original_dir = os.getcwd()
os.chdir(migrations_dir)
deploy_timeout = prisma_migrate_deploy_timeout()
budget = _MigrateAttemptBudget(attempts_left=MAX_MIGRATE_DEPLOY_ATTEMPTS)
try:
for attempt in range(4):
while not budget.exhausted:
try:
result = subprocess.run(
[_get_prisma_command(), "migrate", "deploy"],
@ -767,168 +804,155 @@ class ProxyExtrasDBManager:
logger.warning(
"prisma migrate deploy attempt %s timed out after %ss, retrying. "
"Raise %s if this database needs longer to apply its pending migrations.",
attempt + 1,
budget.attempt_number,
deploy_timeout,
PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR,
)
time.sleep(random.randrange(5, 15))
continue
next_budget = budget.spend()
except subprocess.CalledProcessError as e:
stderr = e.stderr or ""
next_budget = ProxyExtrasDBManager._budget_after_deploy_failure(
e, budget, schema_path
)
if "P3005" in stderr and "database schema is not empty" in stderr:
logger.info(
"Schema exists but no migrations ledger — creating baseline"
)
ProxyExtrasDBManager._create_baseline_migration(schema_path)
continue
if "P3009" in stderr:
migration_match = re.search(r"`(\d+_\S+?)`", stderr)
if (
migration_match
and ProxyExtrasDBManager._is_idempotent_error(stderr)
):
name = migration_match.group(1)
logger.info(
f"Migration {name} failed idempotently — marking applied and retrying"
)
try:
ProxyExtrasDBManager._roll_back_migration(name)
except (
subprocess.CalledProcessError,
subprocess.TimeoutExpired,
):
pass # may already be rolled-back
try:
ProxyExtrasDBManager._resolve_specific_migration(name)
except (
subprocess.CalledProcessError,
subprocess.TimeoutExpired,
) as resolve_err:
# We're already inside the outer
# `except CalledProcessError` handler —
# re-raising CalledProcessError from here
# would escape as itself, bypassing
# proxy_cli.py's `except RuntimeError`.
raise RuntimeError(
f"Failed to mark migration {name} as applied "
f"after idempotent recovery. Manual "
f"intervention may be required.\n\n"
f"Detail: {resolve_err}"
) from resolve_err
continue
if migration_match:
migration_name = migration_match.group(1)
ledger_logs = ProxyExtrasDBManager._failed_migration_logs(migration_name)
if ledger_logs is not None and (
ledger_logs == "" or _MIGRATION_DEADLOCK_MARKER in ledger_logs
):
logger.info(
"Migration %s failed in a concurrent migrate deploy "
"deadlock race, rolling its ledger row back and retrying",
migration_name,
)
ProxyExtrasDBManager._roll_back_migration_best_effort(migration_name)
time.sleep(random.randrange(5, 15))
continue
raise RuntimeError(
"Database migration failed and cannot be auto-recovered. "
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
) from e
if "P3018" in stderr:
if ProxyExtrasDBManager._is_permission_error(stderr):
raise RuntimeError(
"Database migration failed due to insufficient "
"permissions. Please grant the required privileges "
f"and retry.\n\nPrisma error:\n{stderr}"
) from e
migration_match = re.search(
r"Migration name: (\d+_\S+)", stderr
)
if (
migration_match
and ProxyExtrasDBManager._is_idempotent_error(stderr)
):
name = migration_match.group(1)
logger.info(
f"Migration {name} SQL hit idempotent error — marking applied and retrying"
)
try:
ProxyExtrasDBManager._roll_back_migration(name)
except (
subprocess.CalledProcessError,
subprocess.TimeoutExpired,
):
pass # may already be rolled-back
try:
ProxyExtrasDBManager._resolve_specific_migration(name)
except (
subprocess.CalledProcessError,
subprocess.TimeoutExpired,
) as resolve_err:
raise RuntimeError(
f"Failed to mark migration {name} as applied "
f"after idempotent recovery. Manual "
f"intervention may be required.\n\n"
f"Detail: {resolve_err}"
) from resolve_err
continue
if migration_match and _MIGRATION_DEADLOCK_MARKER in stderr:
logger.info(
"Migration %s deadlocked against a concurrent "
"migrate deploy, rolling its ledger row back "
"and retrying",
migration_match.group(1),
)
ProxyExtrasDBManager._roll_back_migration_best_effort(
migration_match.group(1)
)
time.sleep(random.randrange(5, 15))
continue
raise RuntimeError(
"Database migration failed and cannot be auto-recovered. "
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
) from e
if _MIGRATION_DEADLOCK_MARKER in stderr:
logger.info(
"prisma migrate deploy attempt %s deadlocked against "
"a concurrent migrate deploy, retrying",
attempt + 1,
)
time.sleep(random.randrange(5, 15))
continue
if "P1002" in stderr and "advisory lock" in stderr:
logger.info(
"prisma migrate deploy attempt %s timed out waiting for "
"the advisory lock a concurrent migrate deploy holds, retrying",
attempt + 1,
)
time.sleep(random.randrange(5, 15))
continue
raise RuntimeError(
"Database migration failed and cannot be auto-recovered. "
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
) from e
if next_budget.attempts_left < budget.attempts_left:
time.sleep(random.randrange(5, 15))
budget = next_budget # rebind-ok: the loop carries the budget from one migrate deploy pass to the next
raise RuntimeError(
"Database migration failed after 4 attempts (retry loop "
"exhausted by timeouts, deadlock retries, or repeated "
"idempotent-recovery continues). Check database connectivity, "
f"Database migration failed after {MAX_MIGRATE_DEPLOY_ATTEMPTS} "
"attempts that made no progress (timeouts, deadlock retries, or a "
"recovery that had already run once). Check database connectivity, "
"load, and _prisma_migrations ledger state, and raise "
f"{PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR} if the attempts timed out."
)
finally:
os.chdir(original_dir)
@staticmethod
def _budget_after_deploy_failure(
error: subprocess.CalledProcessError,
budget: "_MigrateAttemptBudget",
schema_path: str,
) -> "_MigrateAttemptBudget":
"""Recover from one failed `prisma migrate deploy`, and price the pass.
Returns the budget the next pass runs under, or raises when the failure
is not one this resolver knows how to recover from.
"""
stderr = error.stderr or ""
if "P3005" in stderr and "database schema is not empty" in stderr:
logger.info("Schema exists but no migrations ledger — creating baseline")
if ProxyExtrasDBManager._create_baseline_migration(schema_path):
return budget.after_recovery("baseline")
return budget.spend()
if "P3009" in stderr:
migration_match = re.search(r"`(\d+_\S+?)`", stderr)
if migration_match and ProxyExtrasDBManager._is_idempotent_error(stderr):
name = migration_match.group(1)
logger.info(
f"Migration {name} failed idempotently — marking applied and retrying"
)
ProxyExtrasDBManager._mark_migration_applied(name)
return budget.after_recovery(f"resolved:{name}")
if migration_match:
migration_name = migration_match.group(1)
ledger_logs = ProxyExtrasDBManager._failed_migration_logs(migration_name)
if ledger_logs is not None and (
ledger_logs == "" or _MIGRATION_DEADLOCK_MARKER in ledger_logs
):
logger.info(
"Migration %s failed in a concurrent migrate deploy "
"deadlock race, rolling its ledger row back and retrying",
migration_name,
)
ProxyExtrasDBManager._roll_back_migration_best_effort(migration_name)
return budget.spend()
raise RuntimeError(
"Database migration failed and cannot be auto-recovered. "
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
) from error
if "P3018" in stderr:
if ProxyExtrasDBManager._is_permission_error(stderr):
raise RuntimeError(
"Database migration failed due to insufficient "
"permissions. Please grant the required privileges "
f"and retry.\n\nPrisma error:\n{stderr}"
) from error
migration_match = re.search(r"Migration name: (\d+_\S+)", stderr)
if migration_match and ProxyExtrasDBManager._is_idempotent_error(stderr):
name = migration_match.group(1)
logger.info(
f"Migration {name} SQL hit idempotent error — marking applied and retrying"
)
ProxyExtrasDBManager._mark_migration_applied(name)
return budget.after_recovery(f"resolved:{name}")
if migration_match and _MIGRATION_DEADLOCK_MARKER in stderr:
logger.info(
"Migration %s deadlocked against a concurrent "
"migrate deploy, rolling its ledger row back "
"and retrying",
migration_match.group(1),
)
ProxyExtrasDBManager._roll_back_migration_best_effort(
migration_match.group(1)
)
return budget.spend()
raise RuntimeError(
"Database migration failed and cannot be auto-recovered. "
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
) from error
if _MIGRATION_DEADLOCK_MARKER in stderr:
logger.info(
"prisma migrate deploy attempt %s deadlocked against "
"a concurrent migrate deploy, retrying",
budget.attempt_number,
)
return budget.spend()
if "P1002" in stderr and "advisory lock" in stderr:
logger.info(
"prisma migrate deploy attempt %s timed out waiting for "
"the advisory lock a concurrent migrate deploy holds, retrying",
budget.attempt_number,
)
return budget.spend()
raise RuntimeError(
"Database migration failed and cannot be auto-recovered. "
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
) from error
@staticmethod
def _mark_migration_applied(name: str) -> None:
"""Roll a failed ledger row back if it is still there, then mark it applied."""
try:
ProxyExtrasDBManager._roll_back_migration(name)
except (subprocess.CalledProcessError, subprocess.TimeoutExpired):
pass # may already be rolled-back
try:
ProxyExtrasDBManager._resolve_specific_migration(name)
except (
subprocess.CalledProcessError,
subprocess.TimeoutExpired,
) as resolve_err:
# We're called from inside an `except CalledProcessError` handler —
# re-raising CalledProcessError from here would escape as itself,
# bypassing proxy_cli.py's `except RuntimeError`.
raise RuntimeError(
f"Failed to mark migration {name} as applied "
f"after idempotent recovery. Manual "
f"intervention may be required.\n\n"
f"Detail: {resolve_err}"
) from resolve_err
@staticmethod
def apply_replica_identity_full_if_requested() -> bool:
"""

View file

@ -4,6 +4,7 @@ import logging
import time
from collections.abc import Mapping, Sequence
from functools import lru_cache
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, cast
from httpx import Response
@ -1164,6 +1165,12 @@ def _store_cost_breakdown_in_logging_obj(
# Don't fail the main cost calculation if breakdown storage fails
def _without_provider_stated_cost(usage: Usage | None) -> Usage | None:
if usage is None or getattr(usage, "cost", None) is None:
return usage
return usage.model_copy(update=MappingProxyType({"cost": None}))
def completion_cost(
completion_response: object | None = None,
model: str | None = None,
@ -1243,7 +1250,10 @@ def completion_cost(
cache_creation_input_tokens: int | None = None
cache_read_input_tokens: int | None = None
audio_transcription_file_duration: float = 0.0
cost_per_token_usage_object: Final[Usage | None] = _get_usage_object(completion_response=completion_response)
provider_usage_object: Final = _get_usage_object(completion_response=completion_response)
cost_per_token_usage_object: Final[Usage | None] = (
_without_provider_stated_cost(provider_usage_object) if custom_pricing else provider_usage_object
)
rerank_billed_units: RerankBilledUnits | None = None
# Extract service_tier from optional_params if not provided directly

View file

@ -1485,7 +1485,7 @@ def log_guardrail_information(func):
if func.__name__ == "apply_guardrail" and "inputs" in kwargs:
original_inputs = kwargs.get("inputs")
logging_obj: Final = kwargs.get("logging_obj")
logging_obj: Final = kwargs.get("logging_obj") or request_data.get("litellm_logging_obj")
self_recorded_token: Final = _guardrail_self_recorded.set(False)
try:
response: Final = await func(*args, **kwargs)
@ -1527,7 +1527,7 @@ def log_guardrail_information(func):
if func.__name__ == "apply_guardrail" and "inputs" in kwargs:
original_inputs = kwargs.get("inputs")
logging_obj: Final = kwargs.get("logging_obj")
logging_obj: Final = kwargs.get("logging_obj") or request_data.get("litellm_logging_obj")
self_recorded_token: Final = _guardrail_self_recorded.set(False)
try:
response: Final = func(*args, **kwargs)

View file

@ -54,6 +54,7 @@ FUNCTION_CALL_ATTRIBUTE: Final = "function_call"
_SYNC_ITER_EXHAUSTED: Final = object()
_GCHUNK_FIELDS: Final[frozenset] = frozenset(GChunk.__annotations__)
_USAGE_COST_HEADER_PROVIDERS: Final[frozenset[str]] = frozenset({LlmProviders.OPENROUTER.value})
def _next_sync_or_exhausted(it: Any) -> object:
@ -1886,8 +1887,8 @@ class CustomStreamWrapper:
@staticmethod
def _resolve_provider_reported_cost(usage_cost: object) -> float | None:
"""
Providers report usage.cost either as a number or, for Perplexity, as a
breakdown object whose total lives under ``total_cost``.
Providers report usage.cost either as a number or as a breakdown object
whose total lives under ``total_cost``.
"""
if isinstance(usage_cost, bool):
return None
@ -1900,12 +1901,10 @@ class CustomStreamWrapper:
@staticmethod
def _propagate_usage_cost_to_hidden_params(
response: "ModelResponse",
custom_llm_provider: str | None,
) -> None:
"""
If the assembled response carries a provider-reported cost on
usage.cost, copy it into _hidden_params so litellm's cost
calculator uses it instead of a token-based estimate.
"""
if custom_llm_provider not in _USAGE_COST_HEADER_PROVIDERS:
return
_usage: Final[Usage | None] = getattr(response, "usage", None)
_cost: Final = CustomStreamWrapper._resolve_provider_reported_cost(getattr(_usage, "cost", None))
if _cost is not None:
@ -2020,7 +2019,7 @@ class CustomStreamWrapper:
response = self.model_response_creator()
if complete_streaming_response is not None:
self._propagate_usage_cost_to_hidden_params(complete_streaming_response)
self._propagate_usage_cost_to_hidden_params(complete_streaming_response, self.custom_llm_provider)
setattr(
response,
@ -2270,7 +2269,7 @@ class CustomStreamWrapper:
response: Final = self.model_response_creator()
if complete_streaming_response is not None:
self._propagate_usage_cost_to_hidden_params(complete_streaming_response)
self._propagate_usage_cost_to_hidden_params(complete_streaming_response, self.custom_llm_provider)
setattr(
response,

View file

@ -24,7 +24,6 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
from litellm.router import Router
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
@ -121,11 +120,10 @@ class AzureAIVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig, BaseAzureLLM
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: Mapping[str, object],
extra_body: Mapping[str, object] | None = None,
router: Router | None = None,
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
) -> tuple[str, dict[str, object]]:
query_text: Final = self.query_text(query)
query_vector: Final = self.embed_query(query_text, litellm_params, embedding_executor, router)
query_vector: Final = self.embed_query(query_text, litellm_params, embedding_executor)
return self._search_request(
vector_store_id,
query_text,
@ -145,11 +143,10 @@ class AzureAIVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig, BaseAzureLLM
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: Mapping[str, object],
extra_body: Mapping[str, object] | None = None,
router: Router | None = None,
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
) -> tuple[str, dict[str, object]]:
query_text: Final = self.query_text(query)
query_vector: Final = await self.aembed_query(query_text, litellm_params, embedding_executor, router)
query_vector: Final = await self.aembed_query(query_text, litellm_params, embedding_executor)
return self._search_request(
vector_store_id,
query_text,

View file

@ -99,12 +99,9 @@ class RouterVectorStoreEmbeddingExecutor:
)
return bool(resolved) or model in deployment_models
def _embeds_through_sdk(self, model: str, configuration: Mapping[str, object]) -> bool:
return bool(configuration) and not self._router_serves(model)
def embed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse:
embedding_kwargs: Final = self._embedding_kwargs(configuration)
if self._embeds_through_sdk(model, configuration):
if not self._router_serves(model):
return LiteLLMVectorStoreEmbeddingExecutor().embed(model, query, embedding_kwargs)
return self.router.embedding( # pyright: ignore[reportUnknownMemberType] # Router embedding input retains a legacy untyped list
model=model,
@ -114,7 +111,7 @@ class RouterVectorStoreEmbeddingExecutor:
async def aembed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse:
embedding_kwargs: Final = self._embedding_kwargs(configuration)
if self._embeds_through_sdk(model, configuration):
if not self._router_serves(model):
return await LiteLLMVectorStoreEmbeddingExecutor().aembed(model, query, embedding_kwargs)
return await self.router.aembedding( # pyright: ignore[reportUnknownMemberType] # Router embedding input retains a legacy untyped list
model=model,
@ -153,7 +150,6 @@ class BaseVectorStoreConfig:
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: dict[str, Any] | None = None,
router: Router | None = None,
) -> tuple[str, dict]:
pass
@ -166,7 +162,6 @@ class BaseVectorStoreConfig:
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: dict[str, Any] | None = None,
router: Router | None = None,
) -> tuple[str, dict]:
"""
Optional async version of transform_search_vector_store_request.
@ -182,7 +177,6 @@ class BaseVectorStoreConfig:
litellm_logging_obj=litellm_logging_obj,
litellm_params=litellm_params,
extra_body=extra_body,
router=router,
)
@abstractmethod
@ -271,7 +265,6 @@ class BaseQueryEmbeddingVectorStoreConfig(BaseVectorStoreConfig):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: Mapping[str, object],
extra_body: Mapping[str, object] | None = None,
router: Router | None = None,
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
) -> tuple[str, dict[str, object]]:
pass
@ -285,7 +278,6 @@ class BaseQueryEmbeddingVectorStoreConfig(BaseVectorStoreConfig):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: Mapping[str, object],
extra_body: Mapping[str, object] | None = None,
router: Router | None = None,
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
) -> tuple[str, dict[str, object]]:
return self.transform_search_vector_store_request(
@ -296,7 +288,6 @@ class BaseQueryEmbeddingVectorStoreConfig(BaseVectorStoreConfig):
litellm_logging_obj=litellm_logging_obj,
litellm_params=litellm_params,
extra_body=extra_body,
router=router,
embedding_executor=embedding_executor,
)
@ -338,11 +329,10 @@ class BaseQueryEmbeddingVectorStoreConfig(BaseVectorStoreConfig):
query_text: str,
litellm_params: Mapping[str, object],
embedding_executor: VectorStoreEmbeddingExecutor | None,
router: Router | None = None,
) -> Sequence[float]:
model: Final = self.query_embedding_model(litellm_params)
configuration: Final = self.query_embedding_configuration(litellm_params)
executor: Final = self.query_embedding_executor(embedding_executor, router)
executor: Final = self.query_embedding_executor(embedding_executor, None)
try:
response: Final = executor.embed(model, query_text, configuration)
except Exception as e:
@ -354,11 +344,10 @@ class BaseQueryEmbeddingVectorStoreConfig(BaseVectorStoreConfig):
query_text: str,
litellm_params: Mapping[str, object],
embedding_executor: VectorStoreEmbeddingExecutor | None,
router: Router | None = None,
) -> Sequence[float]:
model: Final = self.query_embedding_model(litellm_params)
configuration: Final = self.query_embedding_configuration(litellm_params)
executor: Final = self.query_embedding_executor(embedding_executor, router)
executor: Final = self.query_embedding_executor(embedding_executor, None)
try:
response: Final = await executor.aembed(model, query_text, configuration)
except Exception as e:
@ -408,7 +397,6 @@ class BaseDirectVectorStoreConfig(BaseVectorStoreConfig):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: Mapping[str, object],
extra_body: Mapping[str, object] | None = None,
router: Router | None = None,
) -> NoReturn:
raise NotImplementedError("Direct vector store providers execute the search themselves; no HTTP request shape")

View file

@ -27,7 +27,6 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.router import Router
else:
LiteLLMLoggingObj = Any
@ -197,7 +196,6 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: dict[str, Any] | None = None,
router: "Router | None" = None,
) -> tuple[str, dict]:
if isinstance(query, list):
query = " ".join(query)

View file

@ -198,7 +198,6 @@ if TYPE_CHECKING:
AnthropicMessagesStreamingResponse,
)
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
from litellm.router import Router
from litellm.types.llms.openai_evals import (
CancelEvalResponse,
CancelRunResponse,
@ -9872,7 +9871,6 @@ class BaseLLMHTTPHandler:
timeout: float | httpx.Timeout | None = None,
client: HTTPHandler | AsyncHTTPHandler | None = None,
_is_async: bool = False,
router: "Router | None" = None,
) -> VectorStoreSearchResponse:
if isinstance(vector_store_provider_config, BaseDirectVectorStoreConfig):
self._pre_call_direct_vector_store_search(
@ -9923,7 +9921,6 @@ class BaseLLMHTTPHandler:
litellm_logging_obj=logging_obj,
litellm_params=dict(litellm_params),
extra_body=extra_body,
router=router,
embedding_executor=embedding_executor,
)
else:
@ -9938,7 +9935,6 @@ class BaseLLMHTTPHandler:
litellm_logging_obj=logging_obj,
litellm_params=dict(litellm_params),
extra_body=extra_body,
router=router,
)
all_optional_params: Final[dict[str, object]] = dict(litellm_params)
all_optional_params.update(vector_store_search_optional_params or {})
@ -9991,7 +9987,6 @@ class BaseLLMHTTPHandler:
timeout: float | httpx.Timeout | None = None,
client: HTTPHandler | AsyncHTTPHandler | None = None,
_is_async: bool = False,
router: "Router | None" = None,
) -> VectorStoreSearchResponse | Coroutine[object, object, VectorStoreSearchResponse]:
if _is_async:
return self.async_vector_store_search_handler(
@ -10007,7 +10002,6 @@ class BaseLLMHTTPHandler:
extra_body=extra_body,
timeout=timeout,
client=client,
router=router,
)
if isinstance(vector_store_provider_config, BaseDirectVectorStoreConfig):
@ -10056,7 +10050,6 @@ class BaseLLMHTTPHandler:
litellm_logging_obj=logging_obj,
litellm_params=dict(litellm_params),
extra_body=extra_body,
router=router,
embedding_executor=embedding_executor,
)
else:
@ -10071,7 +10064,6 @@ class BaseLLMHTTPHandler:
litellm_logging_obj=logging_obj,
litellm_params=dict(litellm_params),
extra_body=extra_body,
router=router,
)
all_optional_params: Final[dict[str, object]] = dict(litellm_params)

View file

@ -33,7 +33,6 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.router import Router
else:
LiteLLMLoggingObj = Any
@ -169,7 +168,6 @@ class GeminiVectorStoreConfig(BaseVectorStoreConfig):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Mapping[str, object] | None = None,
router: "Router | None" = None,
) -> tuple[str, dict]:
"""
Transform search request to Gemini's generateContent format.

View file

@ -24,7 +24,6 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
from litellm.router import Router
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
@ -129,11 +128,10 @@ class MilvusVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: Mapping[str, object],
extra_body: Mapping[str, object] | None = None,
router: Router | None = None,
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
) -> tuple[str, dict[str, object]]:
query_text: Final = self.query_text(query)
query_vector: Final = self.embed_query(query_text, litellm_params, embedding_executor, router)
query_vector: Final = self.embed_query(query_text, litellm_params, embedding_executor)
return self._search_request(
vector_store_id,
query_text,
@ -153,11 +151,10 @@ class MilvusVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: Mapping[str, object],
extra_body: Mapping[str, object] | None = None,
router: Router | None = None,
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
) -> tuple[str, dict[str, object]]:
query_text: Final = self.query_text(query)
query_vector: Final = await self.aembed_query(query_text, litellm_params, embedding_executor, router)
query_vector: Final = await self.aembed_query(query_text, litellm_params, embedding_executor)
return self._search_request(
vector_store_id,
query_text,

View file

@ -21,7 +21,6 @@ from litellm.utils import add_openai_metadata
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
from litellm.router import Router
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
@ -100,7 +99,6 @@ class OpenAIVectorStoreConfig(BaseVectorStoreConfig):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: dict[str, Any] | None = None,
router: "Router | None" = None,
) -> tuple[str, dict]:
encoded_vector_store_id: Final = encode_url_path_segment(vector_store_id, field_name="vector_store_id")
url: Final = f"{api_base}/{encoded_vector_store_id}/search"

View file

@ -8,7 +8,6 @@ from litellm.types.vector_stores import VectorStoreSearchOptionalRequestParams
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.router import Router
else:
LiteLLMLoggingObj = Any
@ -81,7 +80,6 @@ class PGVectorStoreConfig(OpenAIVectorStoreConfig):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: dict[str, Any] | None = None,
router: "Router | None" = None,
) -> tuple[str, dict]:
encoded_vector_store_id: Final = encode_url_path_segment(vector_store_id, field_name="vector_store_id")
url: Final = f"{api_base}/{encoded_vector_store_id}/search"

View file

@ -17,7 +17,6 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.router import Router
else:
LiteLLMLoggingObj = Any
@ -93,7 +92,6 @@ class RAGFlowVectorStoreConfig(BaseVectorStoreConfig):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: dict[str, Any] | None = None,
router: "Router | None" = None,
) -> tuple[str, dict]:
"""RAGFlow vector stores are management-only, search is not supported."""
raise NotImplementedError("RAGFlow vector stores support dataset management only, not search/retrieval")

View file

@ -1,9 +1,12 @@
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final
import httpx
from litellm.caching._embedding_router import resolve_embedding_router
from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig
from litellm.llms.base_llm.vector_store.transformation import (
BaseQueryEmbeddingVectorStoreConfig,
VectorStoreEmbeddingExecutor,
)
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.types.router import GenericLiteLLMParams
from litellm.types.vector_stores import (
@ -18,16 +21,18 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.router import Router
else:
LiteLLMLoggingObj = Any
_DEFAULT_QUERY_EMBEDDING_MODEL: Final = "text-embedding-3-small"
_DEFAULT_TOP_K: Final = 5
class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
class S3VectorsVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig, BaseAWSLLM):
"""Vector store configuration for AWS S3 Vectors."""
def __init__(self) -> None:
BaseVectorStoreConfig.__init__(self)
BaseQueryEmbeddingVectorStoreConfig.__init__(self)
BaseAWSLLM.__init__(self)
def get_auth_credentials(self, litellm_params: dict) -> BaseVectorStoreAuthCredentials:
@ -59,141 +64,94 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
return headers
def get_complete_url(self, api_base: str | None, litellm_params: dict) -> str:
# Resolve region the same way the ingestion path does:
# dynamic param -> AWS_REGION_NAME -> AWS_REGION -> default (us-west-2)
aws_region_name: Final = self.get_aws_region_name_for_non_llm_api_calls(litellm_params.get("aws_region_name"))
return f"https://s3vectors.{aws_region_name}.api.aws"
def _resolve_query_embedding_router(self, embedding_model: str, router: "Router | None") -> "Router | None":
"""Return the router iff it serves ``embedding_model`` as a deployment."""
if router is None:
return None
model_list: Final = [
dict(m) for m in (router.get_model_list() or ())
] # mutable-ok: resolve_embedding_router requires list[dict]
return resolve_embedding_router(embedding_model=embedding_model, llm_router=router, llm_model_list=model_list)
@staticmethod
def query_embedding_model(litellm_params: Mapping[str, object]) -> str:
configured: Final = litellm_params.get("litellm_embedding_model") or litellm_params.get("embedding_model")
return configured if isinstance(configured, str) and configured else _DEFAULT_QUERY_EMBEDDING_MODEL
@staticmethod
def _query_target(vector_store_id: str, litellm_params: Mapping[str, object]) -> tuple[str, str]:
if ":" in vector_store_id:
bucket_name, index_name = vector_store_id.split(":", 1)
return bucket_name, index_name
bucket_name_from_params: Final = litellm_params.get("vector_bucket_name")
if not isinstance(bucket_name_from_params, str) or not bucket_name_from_params:
raise ValueError(
"vector_store_id must be in format 'bucket_name:index_name' for S3 Vectors, "
"or vector_bucket_name must be provided in litellm_params"
)
return bucket_name_from_params, vector_store_id
@staticmethod
def _query_request(
bucket_name: str,
index_name: str,
query_text: str,
query_vector: Sequence[float],
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
api_base: str,
litellm_logging_obj: LiteLLMLoggingObj,
) -> tuple[str, dict[str, object]]:
litellm_logging_obj.model_call_details["query"] = query_text
return f"{api_base}/QueryVectors", {
"vectorBucketName": bucket_name,
"indexName": index_name,
"queryVector": {"float32": list(query_vector)},
"topK": vector_store_search_optional_params.get("max_num_results", _DEFAULT_TOP_K),
"returnDistance": True,
"returnMetadata": True,
}
def transform_search_vector_store_request(
self,
vector_store_id: str,
query: str | list[str],
query: str | Sequence[str],
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
api_base: str,
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: dict[str, Any] | None = None,
router: "Router | None" = None,
) -> tuple[str, dict]:
"""Sync version - generates embedding synchronously."""
# For S3 Vectors, vector_store_id should be in format: bucket_name:index_name
# If not in that format, try to construct it from litellm_params
bucket_name: str
index_name: str
if ":" in vector_store_id:
bucket_name, index_name = vector_store_id.split(":", 1)
else:
# Try to get bucket_name from litellm_params
bucket_name_from_params: Final = litellm_params.get("vector_bucket_name")
if not bucket_name_from_params or not isinstance(bucket_name_from_params, str):
raise ValueError(
"vector_store_id must be in format 'bucket_name:index_name' for S3 Vectors, "
"or vector_bucket_name must be provided in litellm_params"
)
bucket_name = bucket_name_from_params
index_name = vector_store_id
if isinstance(query, list):
query = " ".join(query)
# Generate embedding for the query
embedding_model: Final = litellm_params.get("embedding_model", "text-embedding-3-small")
embedding_router: Final = self._resolve_query_embedding_router(embedding_model=embedding_model, router=router)
import litellm as litellm_module
embedding_input: Final = [query] # mutable-ok: the embedding API takes list input
embedding_response: Final = (
embedding_router.embedding(model=embedding_model, input=embedding_input)
if embedding_router is not None
else litellm_module.embedding(model=embedding_model, input=embedding_input)
litellm_params: Mapping[str, object],
extra_body: Mapping[str, object] | None = None,
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
) -> tuple[str, dict[str, object]]:
bucket_name, index_name = self._query_target(vector_store_id, litellm_params)
query_text: Final = self.query_text(query)
query_vector: Final = self.embed_query(query_text, litellm_params, embedding_executor)
return self._query_request(
bucket_name,
index_name,
query_text,
query_vector,
vector_store_search_optional_params,
api_base,
litellm_logging_obj,
)
query_embedding: Final = embedding_response.data[0]["embedding"]
url: Final = f"{api_base}/QueryVectors"
request_body: Final[dict[str, Any]] = {
"vectorBucketName": bucket_name,
"indexName": index_name,
"queryVector": {"float32": query_embedding},
"topK": vector_store_search_optional_params.get("max_num_results", 5), # Default to 5
"returnDistance": True,
"returnMetadata": True,
}
litellm_logging_obj.model_call_details["query"] = query
return url, request_body
async def atransform_search_vector_store_request(
self,
vector_store_id: str,
query: str | list[str],
query: str | Sequence[str],
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
api_base: str,
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: dict[str, Any] | None = None,
router: "Router | None" = None,
) -> tuple[str, dict]:
"""Async version - generates embedding asynchronously."""
# For S3 Vectors, vector_store_id should be in format: bucket_name:index_name
# If not in that format, try to construct it from litellm_params
bucket_name: str
index_name: str
if ":" in vector_store_id:
bucket_name, index_name = vector_store_id.split(":", 1)
else:
# Try to get bucket_name from litellm_params
bucket_name_from_params: Final = litellm_params.get("vector_bucket_name")
if not bucket_name_from_params or not isinstance(bucket_name_from_params, str):
raise ValueError(
"vector_store_id must be in format 'bucket_name:index_name' for S3 Vectors, "
"or vector_bucket_name must be provided in litellm_params"
)
bucket_name = bucket_name_from_params
index_name = vector_store_id
if isinstance(query, list):
query = " ".join(query)
# Generate embedding for the query asynchronously
embedding_model: Final = litellm_params.get("embedding_model", "text-embedding-3-small")
embedding_router: Final = self._resolve_query_embedding_router(embedding_model=embedding_model, router=router)
import litellm as litellm_module
embedding_input: Final = [query] # mutable-ok: the embedding API takes list input
embedding_response: Final = (
await embedding_router.aembedding(model=embedding_model, input=embedding_input)
if embedding_router is not None
else await litellm_module.aembedding(model=embedding_model, input=embedding_input)
litellm_params: Mapping[str, object],
extra_body: Mapping[str, object] | None = None,
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
) -> tuple[str, dict[str, object]]:
bucket_name, index_name = self._query_target(vector_store_id, litellm_params)
query_text: Final = self.query_text(query)
query_vector: Final = await self.aembed_query(query_text, litellm_params, embedding_executor)
return self._query_request(
bucket_name,
index_name,
query_text,
query_vector,
vector_store_search_optional_params,
api_base,
litellm_logging_obj,
)
query_embedding: Final = embedding_response.data[0]["embedding"]
url: Final = f"{api_base}/QueryVectors"
request_body: Final[dict[str, Any]] = {
"vectorBucketName": bucket_name,
"indexName": index_name,
"queryVector": {"float32": query_embedding},
"topK": vector_store_search_optional_params.get("max_num_results", 5), # Default to 5
"returnDistance": True,
"returnMetadata": True,
}
litellm_logging_obj.model_call_details["query"] = query
return url, request_body
def sign_request(
self,
@ -226,21 +184,13 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
if not source_text:
continue
# Extract file information from metadata
chunk_index = metadata.get("chunk_index", "0")
file_id = f"s3-vectors-chunk-{chunk_index}"
filename = metadata.get("filename", f"document-{chunk_index}")
# S3 Vectors returns distance, convert to similarity score (0-1)
# Lower distance = higher similarity
# We'll normalize using 1 / (1 + distance) to get a 0-1 score
distance = item.get("distance")
score = None
if distance is not None:
# Convert distance to similarity score between 0 and 1
# For cosine distance: similarity = 1 - distance
# For euclidean: use 1 / (1 + distance)
# Assuming cosine distance here
score = max(0.0, min(1.0, 1.0 - float(distance)))
results.append(
@ -265,7 +215,6 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
headers=response.headers,
)
# Vector store creation is not yet implemented
def transform_create_vector_store_request(
self,
vector_store_create_optional_params,

View file

@ -21,7 +21,6 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
from litellm.router import Router
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
@ -162,7 +161,6 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Mapping[str, object] | None = None,
router: "Router | None" = None,
) -> tuple[str, dict[str, object]]:
"""
Transform search request for Vertex AI RAG API

View file

@ -25,7 +25,6 @@ from litellm.types.vector_stores import (
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
from litellm.router import Router
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
@ -246,7 +245,6 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
litellm_logging_obj: LiteLLMLoggingObj,
litellm_params: dict,
extra_body: Mapping[str, object] | None = None,
router: "Router | None" = None,
) -> tuple[str, dict[str, object]]:
"""
Transform a search request for the Vertex AI Search (Discovery Engine) API.

View file

@ -1,4 +1,5 @@
from collections.abc import AsyncIterator, Iterator, Mapping
from types import MappingProxyType
from typing import Any, Final
import httpx
@ -11,7 +12,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
filter_value_from_dict,
strip_name_from_messages,
)
from litellm.llms.xai.common_utils import XAIModelInfo
from litellm.llms.xai.common_utils import XAIModelInfo, xai_reported_cost_in_usd
from litellm.llms.xai.cost_calculator import (
apply_server_side_tool_usage_details_to_usage,
)
@ -30,6 +31,13 @@ from ...openai.chat.gpt_transformation import (
)
def _usage_restated_from_xai_ticks(usage: Usage | None) -> Usage | None:
reported_cost: Final = xai_reported_cost_in_usd(getattr(usage, "cost_in_usd_ticks", None))
if usage is None or reported_cost is None:
return None
return usage.model_copy(update=MappingProxyType({"cost": reported_cost}))
class XAIChatConfig(OpenAIGPTConfig):
@property
def custom_llm_provider(self) -> str | None:
@ -283,6 +291,9 @@ class XAIChatConfig(OpenAIGPTConfig):
self._fold_reasoning_tokens_into_completion(response)
self._normalize_openai_compatible_usage_totals(getattr(response, "usage", None))
restated_usage: Final = _usage_restated_from_xai_ticks(getattr(response, "usage", None))
if restated_usage is not None:
response.usage = restated_usage
return response
@staticmethod
@ -411,4 +422,8 @@ class XAIChatCompletionStreamingHandler(OpenAIChatCompletionStreamingHandler):
XAIChatConfig._fold_reasoning_tokens_into_completion(chunk["usage"])
XAIChatConfig._normalize_openai_compatible_usage_totals(chunk["usage"])
return super().chunk_parser(chunk)
parsed_chunk: Final = super().chunk_parser(chunk)
restated_usage: Final = _usage_restated_from_xai_ticks(getattr(parsed_chunk, "usage", None))
if restated_usage is not None:
parsed_chunk.usage = restated_usage
return parsed_chunk

View file

@ -8,6 +8,17 @@ from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ProviderSpecificModelInfo
USD_TICKS_PER_DOLLAR: Final = 10_000_000_000
def xai_reported_cost_in_usd(cost_in_usd_ticks: object) -> float | None:
"""xAI bills in ticks of a dollar: https://docs.x.ai/developers/cost-tracking"""
if not isinstance(cost_in_usd_ticks, int) or isinstance(cost_in_usd_ticks, bool):
return None
if cost_in_usd_ticks < 0:
return None
return cost_in_usd_ticks / USD_TICKS_PER_DOLLAR
class XAIModelInfo(BaseLLMModelInfo):
def get_provider_info(

View file

@ -1,9 +1,11 @@
"""
Helper util for handling XAI-specific cost calculation
- Prefers the cost xAI reports on the response over recomputing it locally
- Uses the generic cost calculator which already handles tiered pricing correctly
- Handles XAI-specific reasoning token billing (billed as part of completion tokens)
"""
import math
from collections.abc import Mapping
from typing import TYPE_CHECKING, Final
@ -36,6 +38,17 @@ def apply_server_side_tool_usage_details_to_usage(usage: Usage, details: Mapping
usage.prompt_tokens_details = prompt_tokens_details # rebind-ok: write details onto caller usage
def _cost_reported_by_xai(usage: "Usage") -> float | None:
reported_cost: Final[object] = getattr(usage, "cost", None)
if not isinstance(reported_cost, (int, float)) or isinstance(reported_cost, bool):
return None
if not math.isfinite(reported_cost):
return None
if reported_cost < 0:
return None
return float(reported_cost)
def cost_per_token(model: str, usage: Usage) -> tuple[float, float]:
"""
Calculates the cost per token for a given XAI model, prompt tokens, and completion tokens.
@ -48,6 +61,10 @@ def cost_per_token(model: str, usage: Usage) -> tuple[float, float]:
Returns:
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
"""
reported_cost: Final = _cost_reported_by_xai(usage)
if reported_cost is not None:
return 0.0, reported_cost
# XAI-specific completion cost: completion is billed as visible + reasoning
# tokens. Detect when the transformation layer already folded them so we
# don't double-count; fall back to raw xAI shape for callers that bypass
@ -112,6 +129,9 @@ def cost_per_web_search_request(usage: "Usage", model_info: "ModelInfo") -> floa
Per-call rate comes from model_info.search_context_cost_per_query when set,
otherwise the default xAI tools rate ($5 / 1k calls).
"""
if _cost_reported_by_xai(usage) is not None:
return 0.0
details: Final = getattr(usage, "server_side_tool_usage_details", None)
if not isinstance(details, Mapping):
return 0.0

View file

@ -1,17 +1,44 @@
from typing import Any, Final
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final
import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.constants import XAI_API_BASE
from litellm.exceptions import AuthenticationError
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
from litellm.llms.xai.common_utils import XAIModelInfo
from litellm.llms.xai.common_utils import XAIModelInfo, xai_reported_cost_in_usd
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams
from litellm.types.llms.openai import (
ResponseAPIUsage,
ResponseCompletedEvent,
ResponseFailedEvent,
ResponseIncompleteEvent,
ResponsesAPIOptionalRequestParams,
ResponsesAPIResponse,
ResponsesAPIStreamingResponse,
)
from litellm.types.llms.xai import XAIWebSearchTool, XAIXSearchTool
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import LlmProviders
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import (
Logging as _LiteLLMLoggingObj,
)
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
def _usage_restated_from_xai_ticks(usage: ResponseAPIUsage | None) -> ResponseAPIUsage | None:
reported_cost: Final = xai_reported_cost_in_usd(getattr(usage, "cost_in_usd_ticks", None))
if usage is None or reported_cost is None:
return None
return usage.model_copy(update=MappingProxyType({"cost": reported_cost}))
class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
"""
@ -250,6 +277,41 @@ class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
return f"{api_base}/responses"
def transform_response_api_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
) -> ResponsesAPIResponse:
response: Final = super().transform_response_api_response(
model=model,
raw_response=raw_response,
logging_obj=logging_obj,
)
restated_usage: Final = _usage_restated_from_xai_ticks(response.usage)
if restated_usage is not None:
response.usage = restated_usage
return response
def transform_streaming_response(
self,
model: str,
parsed_chunk: dict, # mutable-ok: overrides the base class signature
logging_obj: LiteLLMLoggingObj,
) -> ResponsesAPIStreamingResponse:
event: Final = super().transform_streaming_response(
model=model,
parsed_chunk=parsed_chunk,
logging_obj=logging_obj,
)
if not isinstance(event, (ResponseCompletedEvent, ResponseIncompleteEvent, ResponseFailedEvent)):
return event
restated_usage: Final = _usage_restated_from_xai_ticks(event.response.usage)
if restated_usage is not None:
event.response.usage = restated_usage
return event
def supports_native_websocket(self) -> bool:
"""XAI does not support native WebSocket for Responses API"""
return False

View file

@ -8595,9 +8595,19 @@ def stream_chunk_builder_text_completion(chunks: list, messages: list | None = N
return TextCompletionResponse(**response)
_CALCULATOR_PRICED_REPORTED_COST_PROVIDERS: Final = frozenset({LlmProviders.XAI.value})
def _reported_cost_is_priced_by_calculator(logging_obj: Optional["Logging"]) -> bool:
if logging_obj is None:
return False
provider: Final[object] = logging_obj.model_call_details.get("custom_llm_provider")
return provider in _CALCULATOR_PRICED_REPORTED_COST_PROVIDERS
def _stream_builder_response_cost(response: ModelResponse, logging_obj: Optional["Logging"]) -> float | None:
usage_cost: Final = getattr(getattr(response, "usage", None), "cost", None)
if isinstance(usage_cost, (int, float)):
if isinstance(usage_cost, (int, float)) and not _reported_cost_is_priced_by_calculator(logging_obj):
return float(usage_cost)
if logging_obj is not None:
return None

View file

@ -1216,9 +1216,9 @@ class GenerateKeyRequest(KeyRequestBase):
organization_id: str | None = None
project_id: str | None = None
@field_validator("team_id", mode="before")
@field_validator("team_id", "organization_id", mode="before")
@classmethod
def treat_cleared_team_id_as_unset(cls, v: object) -> object:
def treat_cleared_id_as_unset(cls, v: object) -> object:
if v == "":
return None
return v
@ -1278,6 +1278,13 @@ class UpdateKeyRequest(KeyRequestBase):
rotation_interval: str | None = None
organization_id: str | None = None
@field_validator("organization_id", mode="before")
@classmethod
def treat_cleared_organization_id_as_unset(cls, v: object) -> object:
if v == "":
return None
return v
@model_validator(mode="after")
def validate_temp_budget(self) -> "UpdateKeyRequest":
if self.temp_budget_increase is not None or self.temp_budget_expiry is not None:

View file

@ -489,7 +489,7 @@ lite codex exec "summarize the repo"
Each command resolves your LiteLLM key (logging in via SSO when none is stored and you are at a terminal; otherwise it expects `LITELLM_PROXY_API_KEY` or `--api-key`), checks the key against the proxy so bad credentials fail immediately instead of deep inside the agent, exports the environment variables the agent reads, then replaces itself with the agent process.
The right variables are picked per agent. Claude Code gets `ANTHROPIC_BASE_URL` (the proxy root, so it appends `/v1/messages`) and `ANTHROPIC_AUTH_TOKEN`, with any stray `ANTHROPIC_API_KEY` cleared so the proxy token wins, and `ENABLE_TOOL_SEARCH=true` (unless you already set it) so Claude Code keeps tool search on even though the base URL is a proxy rather than a first-party Anthropic host. Codex and OpenCode get `OPENAI_BASE_URL` (the proxy plus `/v1`) and `OPENAI_API_KEY`. Codex ignores `OPENAI_BASE_URL`, so it is additionally pointed at the proxy through a custom provider passed as `-c` config overrides (HTTP/SSE Responses transport, since the proxy does not speak the Responses WebSocket protocol).
The right variables are picked per agent. Claude Code gets `ANTHROPIC_BASE_URL` (the proxy root, so it appends `/v1/messages`) and `ANTHROPIC_AUTH_TOKEN`, with any stray `ANTHROPIC_API_KEY` cleared so the proxy token wins, and `ENABLE_TOOL_SEARCH=true` (unless you already set it) so Claude Code keeps tool search on even though the base URL is a proxy rather than a first-party Anthropic host. It also gets `CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY=1` (again unless you already set it) so Claude Code v2.1.129+ fills its `/model` picker from the proxy's `/v1/models`; Claude Code only lists entries whose id contains `claude` or `anthropic`, and older versions ignore the variable. Export it as `0` to turn discovery off. Codex and OpenCode get `OPENAI_BASE_URL` (the proxy plus `/v1`) and `OPENAI_API_KEY`. Codex ignores `OPENAI_BASE_URL`, so it is additionally pointed at the proxy through a custom provider passed as `-c` config overrides (HTTP/SSE Responses transport, since the proxy does not speak the Responses WebSocket protocol).
Options (these belong to the wrapper, so put them before the agent's own flags):
@ -505,7 +505,7 @@ The credential is short-lived by design (default 24h, configurable via `LITELLM_
### Route Every Claude Code Session Through the Proxy
`lite claude` wraps a single invocation, but `lite up` goes further: it patches `~/.claude/settings.json`, Claude Code's own config file, so that every Claude Code session started afterward -- from any terminal, launched normally with just `claude`, no wrapper needed -- routes through your LiteLLM proxy. It sets `env.ANTHROPIC_BASE_URL` to the proxy URL, `env.ENABLE_TOOL_SEARCH` to `true` when that key is missing, and `apiKeyHelper` to a `lite auth print-token` invocation, drops any stray static `ANTHROPIC_API_KEY` so the helper-issued token wins, and leaves every other setting in the file untouched. It backs up the original file before patching it.
`lite claude` wraps a single invocation, but `lite up` goes further: it patches `~/.claude/settings.json`, Claude Code's own config file, so that every Claude Code session started afterward -- from any terminal, launched normally with just `claude`, no wrapper needed -- routes through your LiteLLM proxy. It sets `env.ANTHROPIC_BASE_URL` to the proxy URL, `env.ENABLE_TOOL_SEARCH` to `true` and `env.CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY` to `1` when those keys are missing, and `apiKeyHelper` to a `lite auth print-token` invocation, drops any stray static `ANTHROPIC_API_KEY` so the helper-issued token wins, and leaves every other setting in the file untouched. It backs up the original file before patching it.
Two things need to already be true: you've run `lite login` (or `lite login --pkce`, whose key the helper renews on its own), since the apiKeyHelper depends on that stored token, and the proxy is already reachable, since `lite up` does not start one for you.
@ -529,7 +529,7 @@ Cursor is not supported: it has no equivalent file-based config to hot-patch thi
lite --base-url https://your-proxy.example.com login --config-claude
```
It writes the same settings `lite up` does, `env.ANTHROPIC_BASE_URL`, `env.ENABLE_TOOL_SEARCH`, and `apiKeyHelper`, but persistently: there is no backup, nothing to restore, and no foreground process to keep alive. Every other key in `~/.claude/settings.json` is preserved, the file is created if it does not exist, and it is written atomically with owner-only permissions. Plain `lite login` is unchanged; nothing happens to your Claude Code config unless you pass the flag.
It writes the same settings `lite up` does, `env.ANTHROPIC_BASE_URL`, `env.ENABLE_TOOL_SEARCH`, `env.CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY`, and `apiKeyHelper`, but persistently: there is no backup, nothing to restore, and no foreground process to keep alive. Every other key in `~/.claude/settings.json` is preserved, the file is created if it does not exist, and it is written atomically with owner-only permissions. Plain `lite login` is unchanged; nothing happens to your Claude Code config unless you pass the flag.
Because the credential is reached through `apiKeyHelper` rather than copied into the file, a later `lite login` refreshes it with no further action: Claude Code re-runs the helper on every request and picks up whatever token the most recent login stored. Nothing secret is written to `settings.json`.

View file

@ -16,6 +16,8 @@ ANTHROPIC_AUTH_TOKEN_ENV: Final = "ANTHROPIC_AUTH_TOKEN"
ANTHROPIC_API_KEY_ENV: Final = "ANTHROPIC_API_KEY"
ENABLE_TOOL_SEARCH_ENV: Final = "ENABLE_TOOL_SEARCH"
ENABLE_TOOL_SEARCH_VALUE: Final = "true"
ENABLE_GATEWAY_MODEL_DISCOVERY_ENV: Final = "CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY"
ENABLE_GATEWAY_MODEL_DISCOVERY_VALUE: Final = "1"
OPENAI_BASE_URL_ENV: Final = "OPENAI_BASE_URL"
OPENAI_API_KEY_ENV: Final = "OPENAI_API_KEY"
@ -67,7 +69,9 @@ def build_agent_env(
Anthropic key cannot win over the bearer token we set. ENABLE_TOOL_SEARCH
defaults to true because Claude Code turns tool search off when
ANTHROPIC_BASE_URL is not a first-party Anthropic host; a value already in
the environment is left alone.
the environment is left alone. CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY
defaults to 1 so Claude Code (v2.1.129+) fills its /model picker from the
proxy's /v1/models; likewise left alone when already set.
"""
env: Final = dict(base_env)
root: Final = base_url.rstrip("/")
@ -77,6 +81,8 @@ def build_agent_env(
env.pop(ANTHROPIC_API_KEY_ENV, None)
if ENABLE_TOOL_SEARCH_ENV not in env:
env[ENABLE_TOOL_SEARCH_ENV] = ENABLE_TOOL_SEARCH_VALUE
if ENABLE_GATEWAY_MODEL_DISCOVERY_ENV not in env:
env[ENABLE_GATEWAY_MODEL_DISCOVERY_ENV] = ENABLE_GATEWAY_MODEL_DISCOVERY_VALUE
if PROFILE_OPENAI in profiles:
env[OPENAI_BASE_URL_ENV] = root + "/v1"
env[OPENAI_API_KEY_ENV] = api_key

View file

@ -26,6 +26,8 @@ ANTHROPIC_BASE_URL_KEY: Final = "ANTHROPIC_BASE_URL"
ANTHROPIC_API_KEY_KEY: Final = "ANTHROPIC_API_KEY"
ENABLE_TOOL_SEARCH_KEY: Final = "ENABLE_TOOL_SEARCH"
ENABLE_TOOL_SEARCH_VALUE: Final = "true"
ENABLE_GATEWAY_MODEL_DISCOVERY_KEY: Final = "CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY"
ENABLE_GATEWAY_MODEL_DISCOVERY_VALUE: Final = "1"
CLAUDE_SETTINGS_PATH: Final = Path.home() / ".claude" / "settings.json"
BACKUP_PATH: Final = Path.home() / ".litellm" / "claude_settings_backup.json"
@ -77,13 +79,16 @@ def merge_claude_settings(
stray env.ANTHROPIC_API_KEY is dropped so it cannot outrank the helper-issued
token (same reasoning as build_agent_env in agents.py). ENABLE_TOOL_SEARCH
defaults to true because Claude Code turns tool search off when
ANTHROPIC_BASE_URL is not a first-party Anthropic host; an existing value is
left alone. Every other key is preserved untouched.
ANTHROPIC_BASE_URL is not a first-party Anthropic host, and
CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY defaults to 1 so the /model picker
is filled from the proxy's /v1/models; existing values of both are left
alone. Every other key is preserved untouched.
"""
raw_env: Final = settings.get(ENV_KEY, {})
base_env: Final = raw_env if isinstance(raw_env, dict) else {}
env: Final = {
ENABLE_TOOL_SEARCH_KEY: ENABLE_TOOL_SEARCH_VALUE,
ENABLE_GATEWAY_MODEL_DISCOVERY_KEY: ENABLE_GATEWAY_MODEL_DISCOVERY_VALUE,
**{key: value for key, value in base_env.items() if key != ANTHROPIC_API_KEY_KEY},
ANTHROPIC_BASE_URL_KEY: base_url.rstrip("/"),
}
@ -156,6 +161,8 @@ __all__ = (
"AUTOROUTE_BACKUP_PATH",
"BACKUP_PATH",
"CLAUDE_SETTINGS_PATH",
"ENABLE_GATEWAY_MODEL_DISCOVERY_KEY",
"ENABLE_GATEWAY_MODEL_DISCOVERY_VALUE",
"ENABLE_TOOL_SEARCH_KEY",
"ENABLE_TOOL_SEARCH_VALUE",
"ENV_KEY",

View file

@ -235,7 +235,7 @@ async def get_credentials(
]
return {"success": True, "credentials": masked_credentials}
except Exception as e:
return handle_exception_on_proxy(e)
raise handle_exception_on_proxy(e)
@router.get(
@ -406,7 +406,12 @@ async def delete_credential(
_reject_non_admin_wif_fields(
await named_credential_wif_fields(credential_name, prisma_client), user_api_key_dict
)
await CredentialsRepository(prisma_client).delete_by_name(credential_name)
deleted: Final = await CredentialsRepository(prisma_client).delete_by_name(credential_name)
if deleted is None:
raise HTTPException(
status_code=404,
detail="Credential not found. Got credential name: " + credential_name,
)
## DELETE FROM LITELLM ##
litellm.credential_list = [cred for cred in litellm.credential_list if cred.credential_name != credential_name]
@ -414,7 +419,7 @@ async def delete_credential(
except HTTPException:
raise
except Exception as e:
return handle_exception_on_proxy(e)
raise handle_exception_on_proxy(e)
def update_db_credential(

View file

@ -8,7 +8,7 @@ import json
import os
from collections.abc import Awaitable, Callable, Mapping, Sequence
from datetime import datetime, timezone
from types import UnionType
from types import MappingProxyType, UnionType
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypeVar, Union, cast, get_args, get_origin
from urllib.parse import urlparse
@ -51,6 +51,9 @@ from litellm.types.guardrails import (
SupportedGuardrailIntegrations,
ToolPermissionGuardrailConfigModel,
)
from litellm.types.proxy.guardrails.guardrail_hooks.hide_secrets import (
HideSecretsGuardrailConfigModel,
)
if TYPE_CHECKING:
from types import CodeType
@ -1401,7 +1404,11 @@ async def get_guardrail_ui_settings():
provider: [hook.value for hook in hooks]
for provider, guardrail_class in guardrail_class_registry.items()
if (hooks := guardrail_class.get_supported_event_hooks()) is not None
}
} | MappingProxyType(
# hide-secrets lives in the enterprise package, not in the registry
# above; it only runs on pre_call.
{SupportedGuardrailIntegrations.HIDE_SECRETS.value: [GuardrailEventHooks.pre_call.value]}
)
return GuardrailUIAddGuardrailSettings(
supported_entities=[entity.value for entity in PiiEntityType],
@ -1953,12 +1960,18 @@ async def get_provider_specific_params():
tool_permission_fields["ui_friendly_name"] = ToolPermissionGuardrailConfigModel.ui_friendly_name()
# hide-secrets lives in the enterprise package, not in the registry loop below.
hide_secrets_fields: Final = _get_fields_from_model(HideSecretsGuardrailConfigModel)
hide_secrets_fields["ui_friendly_name"] = HideSecretsGuardrailConfigModel.ui_friendly_name()
# Return the provider-specific parameters
provider_params: Final = {
SupportedGuardrailIntegrations.BEDROCK.value: bedrock_fields,
SupportedGuardrailIntegrations.PRESIDIO.value: presidio_fields,
SupportedGuardrailIntegrations.LAKERA_V2.value: lakera_v2_fields,
SupportedGuardrailIntegrations.TOOL_PERMISSION.value: tool_permission_fields,
SupportedGuardrailIntegrations.HIDE_SECRETS.value: hide_secrets_fields,
}
### get the config model for the guardrail - go through the registry and get the config model for the guardrail

View file

@ -74,6 +74,10 @@ from litellm.proxy.management_endpoints.team_endpoints import (
from litellm.proxy.management_endpoints.team_endpoints import (
update_team as _legacy_update_team,
)
from litellm.proxy.management_helpers.access_group_model_sync import (
sync_access_groups_for_deleted_model,
sync_access_groups_for_renamed_model,
)
from litellm.proxy.management_helpers.audit_logs import create_object_audit_log
from litellm.proxy.spend_tracking.ptu_feature_flag import (
PTU_COST_ATTRIBUTION_ENV_VAR,
@ -742,6 +746,7 @@ async def patch_model(
existing_params=db_model.litellm_params,
)
requested_model_name: Final = patch_data.model_name
# Handle team model updates with proper alias management
update_data: Final = await _update_team_model_in_db(
db_model=db_model,
@ -768,6 +773,20 @@ async def patch_model(
param=None,
)
stored_model_name: Final = update_data.get("model_name")
if (
stored_model_name is not None
and stored_model_name == requested_model_name
and stored_model_name != db_model.model_name
):
await sync_access_groups_for_renamed_model(
prisma_client=prisma_client,
model_id=model_id,
old_name=db_model.model_name,
new_name=stored_model_name,
llm_router=llm_router,
)
# Clear cache and reload models (uses config setting or defaults to preserving config models for DB updates)
live_before_reload: Final = live_model_ids_snapshot()
reload_outcome: Final = await clear_cache()
@ -1754,6 +1773,12 @@ async def delete_model(
proxy_logging_obj=proxy_logging_obj,
llm_router=llm_router,
)
await sync_access_groups_for_deleted_model(
prisma_client=prisma_client,
model_id=model_info.id,
model_name=model_params.model_name,
llm_router=llm_router,
)
## CREATE AUDIT LOG ##
asyncio.create_task(
@ -2222,25 +2247,36 @@ async def update_model(
model_params.litellm_params[k] = encrypted_value
### MERGE WITH EXISTING DATA ###
merged_dictionary: Final = {}
_mp: Final[dict[str, object]] = model_params.litellm_params.dict()
merged_dictionary: Final = {
key: _existing_litellm_params_dict[key] if value is None else value
for key, value in _mp.items()
if value is not None or _existing_litellm_params_dict.get(key) is not None
}
for key, value in _mp.items():
if value is not None:
merged_dictionary[key] = value
elif key in _existing_litellm_params_dict and _existing_litellm_params_dict[key] is not None:
merged_dictionary[key] = _existing_litellm_params_dict[key]
else:
pass
renamed_to: Final = (
model_params.model_name
if model_params.model_name not in (None, deployment.model_name)
and deployment.model_info.team_id is None
else None
)
_data: Final[dict[str, str]] = {
"litellm_params": json.dumps(merged_dictionary),
"updated_by": user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME,
**({} if renamed_to is None else {"model_name": renamed_to}),
}
model_response: Final = await _proxy_model_table(prisma_client).update(
where={"model_id": _model_id},
data=_data,
)
if renamed_to is not None:
await sync_access_groups_for_renamed_model(
prisma_client=prisma_client,
model_id=_model_id,
old_name=deployment.model_name,
new_name=renamed_to,
llm_router=llm_router,
)
# Clear cache and reload models (uses config setting or defaults to preserving config models for DB updates)
live_before_reload: Final = live_model_ids_snapshot()

View file

@ -443,7 +443,9 @@ class SAMLAuthHandler:
last_name: Final = SAMLAuthHandler._attribute_value(
attributes, "SAML_ATTRIBUTE_LAST_NAME", _LAST_NAME_ATTRIBUTE_CANDIDATES
)
role_value = SAMLAuthHandler._attribute_value(attributes, "SAML_ATTRIBUTE_ROLE", _ROLE_ATTRIBUTE_CANDIDATES)
role_values: Final = SAMLAuthHandler._attribute_values(
attributes, "SAML_ATTRIBUTE_ROLE", _ROLE_ATTRIBUTE_CANDIDATES
)
team_ids: Final = SAMLAuthHandler._attribute_values(
attributes, "SAML_ATTRIBUTE_TEAM_IDS", _TEAM_IDS_ATTRIBUTE_CANDIDATES
)
@ -464,7 +466,7 @@ class SAMLAuthHandler:
picture=None,
provider="saml",
team_ids=team_ids,
user_role=get_litellm_user_role(role_value) if role_value else None,
user_role=get_litellm_user_role(role_values),
)
except ValidationError as e:
raise HTTPException(

View file

@ -4,12 +4,44 @@ Types for the management endpoints
Might include fastapi/proxy requirements.txt related imports
"""
from collections.abc import Iterable, Sequence
from typing import Any, Final, cast
from fastapi_sso.sso.base import OpenID
from litellm.proxy._types import LitellmUserRoles
# Ordered highest to lowest privilege
LITELLM_USER_ROLE_HIERARCHY: Final = (
LitellmUserRoles.PROXY_ADMIN,
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
LitellmUserRoles.INTERNAL_USER,
LitellmUserRoles.INTERNAL_USER_VIEW_ONLY,
)
def highest_privilege_role(roles: Iterable[LitellmUserRoles]) -> LitellmUserRoles | None:
"""
Pick the highest privilege role out of the roles an IdP asserted for one user.
IdPs do not guarantee ordering within a multi-valued role claim, so a user holding
several roles resolves to the most privileged one rather than whichever came first.
Roles the hierarchy does not rank (org_admin, team, customer) resolve by name to stay
deterministic.
Args:
roles: The roles resolved from the claim
Returns:
The highest privilege role, or None if `roles` is empty
"""
resolved: Final = frozenset(roles)
if not resolved:
return None
ranked: Final = next((role for role in LITELLM_USER_ROLE_HIERARCHY if role in resolved), None)
return ranked if ranked is not None else min(resolved, key=lambda role: role.value)
def is_valid_litellm_user_role(role_str: str) -> bool:
"""
@ -28,12 +60,22 @@ def is_valid_litellm_user_role(role_str: str) -> bool:
return False
def get_litellm_user_role(role_str) -> LitellmUserRoles | None:
def _role_from_claim_value(role_str: object) -> LitellmUserRoles | None:
if not isinstance(role_str, str):
return None
# Use _value2member_map_ for O(1) lookup, case-insensitive
result: Final = LitellmUserRoles._value2member_map_.get(role_str.lower())
return cast(LitellmUserRoles | None, result)
def get_litellm_user_role(role_str: object) -> LitellmUserRoles | None:
"""
Convert a string (or list of strings) to a LitellmUserRoles enum if valid (case-insensitive).
Handles list inputs since some SSO providers (e.g., Keycloak) return roles
as arrays like ["proxy_admin"] instead of plain strings.
as arrays like ["proxy_admin"] instead of plain strings. A claim carrying several
roles resolves to the highest privilege one, so a user does not lose access just
because the IdP listed a weaker role first.
Args:
role_str: String or list to convert (e.g., "proxy_admin", ["proxy_admin"])
@ -41,16 +83,12 @@ def get_litellm_user_role(role_str) -> LitellmUserRoles | None:
Returns:
LitellmUserRoles enum if valid, None otherwise
"""
try:
if isinstance(role_str, list):
if len(role_str) == 0:
return None
role_str = role_str[0]
# Use _value2member_map_ for O(1) lookup, case-insensitive
result: Final = LitellmUserRoles._value2member_map_.get(role_str.lower())
return cast(LitellmUserRoles | None, result)
except Exception:
return None
if isinstance(role_str, (list, tuple)):
entries: Final = cast(Sequence[object], role_str) # cast-ok: isinstance narrows the claim, not its elements
return highest_privilege_role(
role for role in (_role_from_claim_value(entry) for entry in entries) if role is not None
)
return _role_from_claim_value(role_str)
class CustomOpenID(OpenID):

View file

@ -112,6 +112,7 @@ from litellm.proxy.management_endpoints.sso_helper_utils import (
)
from litellm.proxy.management_endpoints.team_endpoints import new_team, team_member_add
from litellm.proxy.management_endpoints.types import (
LITELLM_USER_ROLE_HIERARCHY,
CustomOpenID,
get_litellm_user_role,
is_valid_litellm_user_role,
@ -809,15 +810,6 @@ def normalize_email(email: str | None) -> str | None:
return email.lower() if isinstance(email, str) else email
# Ordered highest to lowest privilege
LITELLM_USER_ROLE_HIERARCHY: Final = (
LitellmUserRoles.PROXY_ADMIN,
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
LitellmUserRoles.INTERNAL_USER,
LitellmUserRoles.INTERNAL_USER_VIEW_ONLY,
)
def determine_role_from_groups(
user_groups: list[str],
role_mappings: "RoleMappings",
@ -4312,14 +4304,7 @@ class MicrosoftSSOHandler:
listed first. Roles the hierarchy does not rank (org_admin, team, customer)
resolve by name to stay deterministic
"""
resolved: Final = frozenset(
role for role in (get_litellm_user_role(role_str) for role_str in app_roles or ()) if role is not None
)
if not resolved:
return None
ranked: Final = next((role for role in LITELLM_USER_ROLE_HIERARCHY if role in resolved), None)
return ranked if ranked is not None else min(resolved, key=lambda role: role.value)
return get_litellm_user_role(tuple(app_roles or ()))
@staticmethod
def get_app_roles_from_id_token(id_token: str | None) -> list[str]:

View file

@ -0,0 +1,119 @@
"""
Keep `litellm_accessgrouptable.access_model_names` pointing at deployment names that still exist.
Unified access groups store model names, not ids, so a deployment rename or delete that leaves
the arrays alone strands every group on a name nothing serves any more.
"""
from collections.abc import Sequence
from typing import Final, Protocol
from pydantic import BaseModel
from litellm.proxy.db.routing_prisma_wrapper import WriterPinnedClient
from litellm.proxy.management_helpers.access_group_team_sync import invalidate_access_group_caches
from litellm.repositories.table_repositories import AccessGroupRepository
from litellm.router import Router
class _TouchedGroupRow(BaseModel):
access_group_id: str
class _DeploymentCountRow(BaseModel):
deployment_count: int
class _RawExecutor(Protocol):
async def query_raw(self, query: str, *args: str) -> Sequence[object]: ...
_BACKING_DEPLOYMENTS_SQL: Final = (
'SELECT COUNT(*)::int AS deployment_count FROM "LiteLLM_ProxyModelTable" WHERE "model_name" = $1'
)
_REPLACE_MODEL_NAME_SQL: Final = (
'UPDATE "LiteLLM_AccessGroupTable" '
'SET "access_model_names" = array_replace(array_remove("access_model_names", $2), $1, $2) '
'WHERE $1 = ANY("access_model_names") '
'RETURNING "access_group_id"'
)
_APPEND_MODEL_NAME_SQL: Final = (
'UPDATE "LiteLLM_AccessGroupTable" '
'SET "access_model_names" = array_append("access_model_names", $2) '
'WHERE $1 = ANY("access_model_names") AND NOT ($2 = ANY("access_model_names")) '
'RETURNING "access_group_id"'
)
_REMOVE_MODEL_NAME_SQL: Final = (
'UPDATE "LiteLLM_AccessGroupTable" '
'SET "access_model_names" = array_remove("access_model_names", $1) '
'WHERE $1 = ANY("access_model_names") '
'RETURNING "access_group_id"'
)
def _raw_executor(prisma_client: object) -> _RawExecutor:
db: Final = AccessGroupRepository(prisma_client).prisma_client.db # pyright: ignore[reportAny] # untyped Prisma client
return WriterPinnedClient(db).db # pyright: ignore[reportAny, reportReturnType] # untyped Prisma client behind the pin
def _config_sourced_sibling(llm_router: Router, deployment_id: str, model_id: str) -> bool:
if deployment_id == model_id:
return False
deployment: Final = llm_router.get_deployment(model_id=deployment_id)
return deployment is not None and not deployment.model_info.db_model
def _served_by_a_config_deployment(llm_router: Router | None, model_name: str, model_id: str) -> bool:
if llm_router is None:
return False
return any(
_config_sourced_sibling(llm_router, deployment_id, model_id)
for deployment_id in llm_router.get_model_ids(model_name=model_name)
)
async def _still_backed(executor: _RawExecutor, llm_router: Router | None, model_name: str, model_id: str) -> bool:
if _served_by_a_config_deployment(llm_router, model_name, model_id):
return True
count_rows: Final = await executor.query_raw(_BACKING_DEPLOYMENTS_SQL, model_name)
return any(_DeploymentCountRow.model_validate(row).deployment_count > 0 for row in count_rows)
async def _rewrite_groups(executor: _RawExecutor, sql: str, *names: str) -> None:
touched_rows: Final = await executor.query_raw(sql, *names)
await invalidate_access_group_caches(
tuple(_TouchedGroupRow.model_validate(row).access_group_id for row in touched_rows)
)
async def sync_access_groups_for_renamed_model(
prisma_client: object,
*,
model_id: str,
old_name: str,
new_name: str,
llm_router: Router | None,
) -> None:
if old_name == new_name:
return
executor: Final = _raw_executor(prisma_client)
old_name_still_backed: Final = await _still_backed(executor, llm_router, old_name, model_id)
await _rewrite_groups(
executor, _APPEND_MODEL_NAME_SQL if old_name_still_backed else _REPLACE_MODEL_NAME_SQL, old_name, new_name
)
async def sync_access_groups_for_deleted_model(
prisma_client: object,
*,
model_id: str,
model_name: str,
llm_router: Router | None,
) -> None:
executor: Final = _raw_executor(prisma_client)
if await _still_backed(executor, llm_router, model_name, model_id):
return
await _rewrite_groups(executor, _REMOVE_MODEL_NAME_SQL, model_name)

View file

@ -1,4 +1,26 @@
{
"1m_context": {
"label": "1M Context",
"description": "Routes across models with 1M-token context windows: Luna for simple queries, Terra for medium, Opus 5 for complex, Opus 5 at high thinking for reasoning.",
"complexity_router_config": {
"tiers": {
"SIMPLE": ["gpt-5.6-luna"],
"MEDIUM": ["gpt-5.6-terra"],
"COMPLEX": ["claude-opus-5"],
"REASONING": ["claude-opus-5"]
},
"tier_model_configs": {
"REASONING": [{ "model_name": "claude-opus-5", "litellm_params": { "reasoning_effort": "high" } }]
},
"classifier_type": "heuristic_v2",
"escalation_keywords": ["LITELLM ESCALATE"],
"classification_mode": "every_request",
"session_affinity": false,
"modality_routing": false,
"modality_pin_override": false,
"deployment_affinity": true
}
},
"anthropic_family": {
"label": "Anthropic Family",
"description": "Routes across the Claude model family: Haiku for simple queries, Sonnet for medium, Opus for complex, Opus at high thinking for reasoning.",

View file

@ -12,6 +12,7 @@ from typing import (
Literal,
NamedTuple,
Protocol,
TypeAlias,
TypedDict,
TypeVar,
cast, # noqa: TID251 # prisma group_by returns untyped aggregate mappings
@ -19,6 +20,7 @@ from typing import (
import fastapi
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
from pydantic import TypeAdapter
from typing_extensions import ReadOnly
import litellm
@ -158,6 +160,26 @@ class _SessionSpendRow(TypedDict):
session_cache_hit_count: ReadOnly[int]
session_llm_count: ReadOnly[int]
session_agent_count: ReadOnly[int]
session_models: ReadOnly[Sequence[str]]
_SESSION_MODELS_LIMIT: Final = 10
_SESSION_MODEL_NAME_MAX_LEN: Final = 256
class _SessionSpendStats(NamedTuple):
session_total_count: int
session_total_spend: float
mcp_tool_call_count: int
mcp_tool_call_spend: float
session_cache_hit_count: int
session_llm_count: int
session_agent_count: int
session_models: Sequence[str]
session_models_truncated: bool
_SessionSpendMap: TypeAlias = Mapping[tuple[str, str], _SessionSpendStats]
class _SpendSumAggregate(TypedDict, total=False):
@ -4121,7 +4143,7 @@ async def _build_ui_spend_logs_response(
}
)
session_spend_map: dict[tuple[str, str], dict[str, int | float]] = {}
session_spend_map: _SessionSpendMap = {}
if enrich_session_counts and session_ids:
from prisma.errors import PrismaError
@ -4139,40 +4161,60 @@ async def _build_ui_spend_logs_response(
rows: Final[Sequence[_SessionSpendRow]] = await _query_raw(
prisma_client,
f"""
SELECT session_id, api_key,
COUNT(*)::int AS session_total_count,
COALESCE(SUM(spend), 0)::double precision AS session_total_spend,
COUNT(*) FILTER (
WHERE call_type IN {_MCP_CALL_TYPES_SQL}
)::int AS mcp_tool_call_count,
COALESCE(SUM(spend) FILTER (
WHERE call_type IN {_MCP_CALL_TYPES_SQL}
), 0)::double precision AS mcp_tool_call_spend,
COUNT(*) FILTER (WHERE LOWER(cache_hit) = 'true')::int AS session_cache_hit_count,
COUNT(*) FILTER (
WHERE call_type NOT IN {_MCP_CALL_TYPES_SQL} AND call_type != {_AGENT_CALL_TYPE_SQL}
)::int AS session_llm_count,
COUNT(*) FILTER (WHERE call_type = {_AGENT_CALL_TYPE_SQL})::int AS session_agent_count
FROM "LiteLLM_SpendLogs"
WHERE session_id = ANY($1::text[])
AND api_key = ANY($2::text[])
GROUP BY session_id, api_key
SELECT s.*, COALESCE(m.session_models, ARRAY[]::text[]) AS session_models
FROM (
SELECT session_id, api_key,
COUNT(*)::int AS session_total_count,
COALESCE(SUM(spend), 0)::double precision AS session_total_spend,
COUNT(*) FILTER (
WHERE call_type IN {_MCP_CALL_TYPES_SQL}
)::int AS mcp_tool_call_count,
COALESCE(SUM(spend) FILTER (
WHERE call_type IN {_MCP_CALL_TYPES_SQL}
), 0)::double precision AS mcp_tool_call_spend,
COUNT(*) FILTER (WHERE LOWER(cache_hit) = 'true')::int AS session_cache_hit_count,
COUNT(*) FILTER (
WHERE call_type NOT IN {_MCP_CALL_TYPES_SQL} AND call_type != {_AGENT_CALL_TYPE_SQL}
)::int AS session_llm_count,
COUNT(*) FILTER (WHERE call_type = {_AGENT_CALL_TYPE_SQL})::int AS session_agent_count
FROM "LiteLLM_SpendLogs"
WHERE session_id = ANY($1::text[])
AND api_key = ANY($2::text[])
GROUP BY session_id, api_key
) s
LEFT JOIN LATERAL (
SELECT ARRAY_AGG(d.model ORDER BY d.model) AS session_models
FROM (
SELECT DISTINCT LEFT(model, $3::int) AS model
FROM "LiteLLM_SpendLogs"
WHERE session_id = s.session_id
AND api_key = s.api_key
AND model IS NOT NULL AND model <> ''
ORDER BY 1
LIMIT $4::int
) d
) m ON TRUE
""",
session_ids,
authorized_api_keys,
_SESSION_MODEL_NAME_MAX_LEN,
_SESSION_MODELS_LIMIT + 1,
)
session_spend_map = {
(row["session_id"], row["api_key"]): {
"session_total_count": int(row.get("session_total_count") or 0),
"session_total_spend": float(row.get("session_total_spend") or 0.0),
"mcp_tool_call_count": int(row.get("mcp_tool_call_count") or 0),
"mcp_tool_call_spend": float(row.get("mcp_tool_call_spend") or 0.0),
"session_cache_hit_count": int(row.get("session_cache_hit_count") or 0),
"session_llm_count": int(row.get("session_llm_count") or 0),
"session_agent_count": int(row.get("session_agent_count") or 0),
}
(row["session_id"], row["api_key"]): _SessionSpendStats(
session_total_count=int(row.get("session_total_count") or 0),
session_total_spend=float(row.get("session_total_spend") or 0.0),
mcp_tool_call_count=int(row.get("mcp_tool_call_count") or 0),
mcp_tool_call_spend=float(row.get("mcp_tool_call_spend") or 0.0),
session_cache_hit_count=int(row.get("session_cache_hit_count") or 0),
session_llm_count=int(row.get("session_llm_count") or 0),
session_agent_count=int(row.get("session_agent_count") or 0),
session_models=models[:_SESSION_MODELS_LIMIT],
session_models_truncated=len(models) > _SESSION_MODELS_LIMIT,
)
for row in rows
if row.get("session_id") and row.get("api_key") is not None
for models in (TypeAdapter(list[str]).validate_python(row.get("session_models") or ()),)
}
except PrismaError:
verbose_proxy_logger.debug(
@ -4187,15 +4229,17 @@ async def _build_ui_spend_logs_response(
sid = row_dict.get("session_id")
row_api_key = row_dict.get("api_key")
session_stats = session_spend_map.get((sid, row_api_key)) if sid and row_api_key is not None else None
row_dict["session_total_count"] = int(session_stats["session_total_count"]) if session_stats else 1
row_dict["session_total_count"] = session_stats.session_total_count if session_stats else 1
if session_stats:
row_dict["session_total_spend"] = session_stats["session_total_spend"]
if session_stats["mcp_tool_call_count"]:
row_dict["mcp_tool_call_count"] = session_stats["mcp_tool_call_count"]
row_dict["mcp_tool_call_spend"] = session_stats["mcp_tool_call_spend"]
row_dict["session_cache_hit_count"] = session_stats["session_cache_hit_count"]
row_dict["session_llm_count"] = session_stats["session_llm_count"]
row_dict["session_agent_count"] = session_stats["session_agent_count"]
row_dict["session_total_spend"] = session_stats.session_total_spend
if session_stats.mcp_tool_call_count:
row_dict["mcp_tool_call_count"] = session_stats.mcp_tool_call_count
row_dict["mcp_tool_call_spend"] = session_stats.mcp_tool_call_spend
row_dict["session_cache_hit_count"] = session_stats.session_cache_hit_count
row_dict["session_llm_count"] = session_stats.session_llm_count
row_dict["session_agent_count"] = session_stats.session_agent_count
row_dict["session_models"] = session_stats.session_models
row_dict["session_models_truncated"] = session_stats.session_models_truncated
enriched.append(row_dict)
response_data: list = enriched
else:

View file

@ -50,6 +50,7 @@ from litellm.constants import (
DEFAULT_HEALTH_CHECK_INTERVAL,
DEFAULT_HEALTH_CHECK_STALENESS_MULTIPLIER,
DEFAULT_MAX_LRU_CACHE_SIZE,
INTERNAL_CALL_ORIGIN_METADATA_KEY,
RUNTIME_UPDATABLE_ROUTER_SETTINGS,
SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY,
)
@ -232,6 +233,7 @@ from litellm.types.router import (
)
from litellm.types.services import ServiceTypes
from litellm.types.utils import (
AUTOROUTER_CLASSIFIER_CALL_ORIGIN,
PROMPT_QUOTING_ROUTING_DECISION_FIELDS,
CustomPricingLiteLLMParams,
GenericBudgetConfigType,
@ -2169,6 +2171,76 @@ class Router:
verbose_router_logger.debug("Error occurred while printing deployment - %s", e)
raise e
@staticmethod
def _deployment_params_with_request_reasoning_override(
deployment_params: Mapping[str, object], request_kwargs: Mapping[str, object]
) -> dict[str, object]: # mutable-ok: litellm's request pipeline consumes a mutable kwargs mapping
"""Return deployment params whose equivalent effort controls cannot outrank a request override.
Providers expose the same setting through several native carriers. A request-level
``reasoning_effort`` is the portable override, so a deployment's ``thinking`` or nested
``*.effort`` must not remain beside it and either win or trigger a conflicting-params 400.
Every changed mapping is copied so the Router's shared deployment config stays immutable.
"""
sanitized: Final = dict(deployment_params) # mutable-ok: request-local copy protects shared Router state
if request_kwargs.get("reasoning_effort") is None:
return sanitized
sanitized.pop("thinking", None)
Router._pop_effort_from_nested_carrier(sanitized, "output_config")
Router._pop_effort_from_nested_carrier(sanitized, "reasoning")
extra_body: Final = sanitized.get("extra_body")
if isinstance(extra_body, Mapping):
sanitized_extra_body: Final = dict(extra_body) # mutable-ok: request-local nested copy
sanitized_extra_body.pop("reasoning_effort", None)
sanitized_extra_body.pop("thinking", None)
Router._pop_effort_from_nested_carrier(sanitized_extra_body, "output_config")
Router._pop_effort_from_nested_carrier(sanitized_extra_body, "reasoning")
if sanitized_extra_body:
sanitized["extra_body"] = sanitized_extra_body
else:
sanitized.pop("extra_body", None)
return sanitized
@staticmethod
def _is_classifier_internal_call(kwargs: Mapping[str, object]) -> bool:
metadata: Final = kwargs.get("metadata")
litellm_metadata: Final = kwargs.get("litellm_metadata")
return any(
isinstance(candidate, Mapping)
and candidate.get(INTERNAL_CALL_ORIGIN_METADATA_KEY) == AUTOROUTER_CLASSIFIER_CALL_ORIGIN
for candidate in (metadata, litellm_metadata)
)
def _drop_unsupported_classifier_reasoning_effort(
self,
deployment: DeploymentTypedDict,
model: str,
kwargs: dict[str, object], # mutable-ok: fallback must update the active request and its log body together
) -> None:
"""Let a classifier fallback without reasoning support remain a usable fallback.
The dashboard only offers explicitly advertised levels, but an existing config can outlive
a model change and fallbacks can target a different group. Unknown capability fails open;
only a provider that explicitly rejects the parameter has it removed.
"""
if kwargs.get("reasoning_effort") is None or not self._is_classifier_internal_call(kwargs):
return
if self._deployment_accepts_param(deployment, model, "reasoning_effort"):
return
verbose_router_logger.warning(
"litellm.router.py: dropping classifier reasoning_effort for model=%s because the selected deployment does not support it",
model,
)
kwargs.pop("reasoning_effort", None)
proxy_server_request: Final = kwargs.get("proxy_server_request")
if not isinstance(proxy_server_request, dict):
return
body: Final = proxy_server_request.get("body")
if isinstance(body, dict):
body.pop("reasoning_effort", None)
### COMPLETION, EMBEDDING, IMG GENERATION FUNCTIONS
def completion(self, model: str, messages: list[dict[str, str]], **kwargs) -> ModelResponse | CustomStreamWrapper:
@ -2204,9 +2276,16 @@ class Router:
specific_deployment=kwargs.pop("specific_deployment", None),
request_kwargs=kwargs,
)
self._drop_unsupported_classifier_reasoning_effort(
deployment=cast(DeploymentTypedDict, deployment), # cast-ok: selection returns a router deployment
model=model,
kwargs=kwargs,
)
# Check for silent model experiment
# Make a local copy of litellm_params to avoid mutating the Router's state
litellm_params: Final = deployment["litellm_params"].copy()
litellm_params: Final = self._deployment_params_with_request_reasoning_override(
deployment["litellm_params"], kwargs
)
silent_model: Final = litellm_params.pop("silent_model", None)
if silent_model is not None:
@ -3217,6 +3296,11 @@ class Router:
specific_deployment=kwargs.pop("specific_deployment", None),
request_kwargs=kwargs,
)
self._drop_unsupported_classifier_reasoning_effort(
deployment=cast(DeploymentTypedDict, deployment), # cast-ok: selection returns a router deployment
model=model,
kwargs=kwargs,
)
_timeout_debug_deployment_dict = deployment
end_time: Final = time.time()
@ -3238,7 +3322,9 @@ class Router:
# Check for silent model experiment
# Make a local copy of litellm_params to avoid mutating the Router's state
litellm_params: Final = deployment["litellm_params"].copy()
litellm_params: Final = self._deployment_params_with_request_reasoning_override(
deployment["litellm_params"], kwargs
)
silent_model: Final = litellm_params.pop("silent_model", None)
if silent_model is not None:
@ -10253,6 +10339,8 @@ class Router:
total_itpm: int | None = None
total_otpm: int | None = None
configurable_clientside_auth_params: CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS = None
reasoning_efforts_initialized = False
reasoning_efforts_unknown = False
model_list: Final = self.get_model_list(model_name=model_group)
if model_list is None:
return None
@ -10439,10 +10527,23 @@ class Router:
if model_info.get("rpm", None) is not None and _deployment_rpm is None:
_deployment_rpm = model_info.get("rpm")
model_group_info.supported_reasoning_efforts = intersect_supported_reasoning_efforts(
model_group_info.supported_reasoning_efforts,
resolve_supported_reasoning_efforts(model_info, deployment_is_mapped=deployment_is_mapped),
deployment_reasoning_efforts = (
resolve_supported_reasoning_efforts( # rebind-ok: recalculated per deployment
model_info, deployment_is_mapped=deployment_is_mapped
)
)
if deployment_reasoning_efforts is None:
reasoning_efforts_unknown = True
model_group_info.supported_reasoning_efforts = None
elif not reasoning_efforts_initialized:
reasoning_efforts_initialized = True
if not reasoning_efforts_unknown:
model_group_info.supported_reasoning_efforts = deployment_reasoning_efforts
elif not reasoning_efforts_unknown:
model_group_info.supported_reasoning_efforts = intersect_supported_reasoning_efforts(
model_group_info.supported_reasoning_efforts,
deployment_reasoning_efforts,
)
if _deployment_tpm is not None:
if total_tpm is None:
@ -12306,10 +12407,12 @@ class Router:
@staticmethod
def _pop_effort_from_nested_carrier(request_kwargs: dict[str, object], carrier: str) -> None:
nested: Final = request_kwargs.get(carrier)
if not isinstance(nested, dict):
if not isinstance(nested, Mapping):
return
nested.pop("effort", None)
if not nested:
sanitized: Final = {key: value for key, value in nested.items() if key != "effort"}
if sanitized:
request_kwargs[carrier] = sanitized # rebind-ok: copy-on-write, so a shared nested carrier is never edited
else:
request_kwargs.pop(carrier, None)
@staticmethod

View file

@ -255,7 +255,8 @@ model_list:
classifier_type: heuristic_first
heuristic_first_max_tier: SIMPLE
classifier_llm_config:
model: gpt-4o-mini
model: gpt-5-mini
reasoning_effort: low
tiers:
SIMPLE: gpt-4o-mini
MEDIUM: gpt-4o
@ -263,6 +264,10 @@ model_list:
REASONING: o1-preview
```
`classifier_llm_config.reasoning_effort` applies only to the internal classifier call. Omit it to
keep the classifier deployment or provider default, or set a supported value such as `none` or
`low` to override that call.
A request short-circuits, meaning it routes on the scorer's own tier with no classifier call, when
two things hold: the scorer landed at or below `heuristic_first_max_tier`, and it produced at least
one signal. Everything else goes to the classifier, which then decides as it normally would.

View file

@ -28,6 +28,7 @@ from pydantic import BaseModel, create_model
from litellm._logging import verbose_router_logger
from litellm.constants import (
EMPTY_MAPPING,
INTERNAL_CALL_ORIGIN_METADATA_KEY,
RETURN_RAW_MODEL_NAME_METADATA_KEY,
SESSION_ID_GENERATED_METADATA_KEY,
)
@ -42,6 +43,7 @@ from litellm.router_strategy.complexity_router.tier_predictor import (
TierSuccessPredictor,
resolve_tier_artifact,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import (
AUTOROUTER_CLASSIFIER_CALL_ORIGIN,
ModelResponse,
@ -1668,20 +1670,27 @@ class ComplexityRouter(CustomLogger):
)
request_metadata = (request_kwargs or {}).get("litellm_metadata") or (request_kwargs or {}).get("metadata")
metadata: Final = forwarded_internal_call_metadata(request_metadata, AUTOROUTER_CLASSIFIER_CALL_ORIGIN)
metadata: Final = { # mutable-ok: SDK metadata kwarg is enriched by the request pipeline
**forwarded_internal_call_metadata(request_metadata, AUTOROUTER_CLASSIFIER_CALL_ORIGIN),
INTERNAL_CALL_ORIGIN_METADATA_KEY: AUTOROUTER_CLASSIFIER_CALL_ORIGIN,
}
turn_off_message_logging: Final = _effective_turn_off_message_logging(request_kwargs)
messages_for_call: Final = [
messages_for_call: Final[list[AllMessageValues]] = [ # mutable-ok: SDK request payload list is built once
{"role": "system", "content": classifier_system_prompt},
{"role": "user", "content": user_payload},
]
response_format: Final = classifier_response_format
classifier_call_params: Mapping[str, str] = EMPTY_MAPPING
if llm_config.reasoning_effort is not None:
classifier_call_params = MappingProxyType({"reasoning_effort": llm_config.reasoning_effort})
proxy_server_request: Final = {
"body": {
"model": llm_config.model,
"messages": messages_for_call,
"response_format": response_format,
**classifier_call_params,
}
}
@ -1693,6 +1702,7 @@ class ComplexityRouter(CustomLogger):
metadata=metadata,
proxy_server_request=proxy_server_request,
turn_off_message_logging=turn_off_message_logging,
**classifier_call_params,
**_parent_session_kwargs(request_kwargs),
)
content: Final = response.choices[0].message.content

View file

@ -12,6 +12,7 @@ from typing import Annotated, Final, Literal
from pydantic import BaseModel, ConfigDict, Field, SkipValidation, field_serializer, field_validator, model_validator
from litellm.types.llms.openai import REASONING_EFFORT
from litellm.types.router import AdaptiveRouterWeights, ClassifierPlugin, RoutingPlugin
from .tier_predictor import TrainedTierArtifact
@ -432,6 +433,13 @@ class ClassifierLLMConfig(BaseModel):
model: str = Field(
description="Model name (from the router's model_list) to call for classification",
)
reasoning_effort: REASONING_EFFORT | None = Field(
default=None,
description=(
"Reasoning effort override for classifier calls. Leave unset to use "
"the classifier deployment or provider default."
),
)
timeout_ms: int = Field(
default=3000,
description="Timeout budget for the classification call, in milliseconds",

View file

@ -0,0 +1,20 @@
"""Types for the Hide Secrets guardrail."""
from pydantic import Field
from .base import GuardrailConfigModel
class HideSecretsGuardrailConfigModel(GuardrailConfigModel):
"""Configuration for the Hide Secrets guardrail. Detection runs in-process
on the detect-secrets library; ``detect_secrets_config`` overrides the
bundled plugin set."""
detect_secrets_config: dict | None = Field( # mutable-ok: UI type derivation maps dict to "object"
default=None,
description="Optional detect-secrets configuration (plugins_used, filters_used) overriding the bundled plugin set",
)
@staticmethod
def ui_friendly_name() -> str:
return "Hide Secrets"

View file

@ -482,7 +482,6 @@ def search(
timeout=timeout or request_timeout,
_is_async=_is_async,
client=kwargs.get("client"),
router=router,
)
return response

View file

@ -1,6 +1,6 @@
{
"ANN001": {
"limit": 2985
"limit": 2984
},
"ANN002": {
"limit": 71
@ -57,7 +57,7 @@
"limit": 3
},
"BLE001": {
"limit": 2917
"limit": 2916
},
"C401": {
"limit": 8

View file

@ -703,3 +703,170 @@ class TestSpendLogsPartitionDetectionMissingPsycopg:
assert any(
"psycopg is not installed" in record.message for record in caplog.records
)
_ATTEMPT_BUDGET = 4
_P3005_STDERR = """Error: P3005
The database schema is not empty. Read more about how to baseline an existing production database: https://pris.ly/d/migrate-baseline
"""
def _p3018_stderr(migration_name):
return f"""Error: P3018
A migration failed to apply. New migrations cannot be applied before the error is recovered from.
Migration name: {migration_name}
Database error code: 42P07
Database error:
ERROR: relation "SomeTable" already exists
"""
class _MigrateDeployHarness:
"""Drives _setup_database_v2 with a scripted sequence of
`prisma migrate deploy` outcomes, with every recovery command faked out so
nothing touches a database or the packaged migrations directory."""
def __init__(self, monkeypatch, tmp_path, outcomes, repeat_last=False):
import subprocess as subprocess_module
import litellm_proxy_extras.utils as utils_module
self.deploy_calls = []
self.resolved = []
self.baselines = 0
self._outcomes = list(outcomes)
self._repeat_last = repeat_last
self._subprocess_module = subprocess_module
monkeypatch.delenv("DATABASE_URL", raising=False)
monkeypatch.setattr(
ProxyExtrasDBManager, "_get_prisma_dir", staticmethod(lambda: str(tmp_path))
)
monkeypatch.setattr(
ProxyExtrasDBManager,
"_create_baseline_migration",
staticmethod(self._fake_baseline),
)
monkeypatch.setattr(
ProxyExtrasDBManager,
"_roll_back_migration",
staticmethod(lambda name: None),
)
monkeypatch.setattr(
ProxyExtrasDBManager,
"_resolve_specific_migration",
staticmethod(self.resolved.append),
)
monkeypatch.setattr(utils_module.subprocess, "run", self._fake_run)
monkeypatch.setattr(utils_module.time, "sleep", lambda seconds: None)
self.baseline_succeeds = True
def _fake_baseline(self, *args, **kwargs):
self.baselines += 1
return self.baseline_succeeds
def _next_outcome(self):
if self._outcomes:
if self._repeat_last and len(self._outcomes) == 1:
return self._outcomes[0]
return self._outcomes.pop(0)
raise AssertionError("prisma migrate deploy called more times than scripted")
def _fake_run(self, cmd, **kwargs):
assert cmd[1:] == ["migrate", "deploy"], f"unexpected prisma command: {cmd}"
self.deploy_calls.append(cmd)
outcome = self._next_outcome()
if outcome == "ok":
return _FakeCompleted()
if outcome == "timeout":
raise self._subprocess_module.TimeoutExpired(cmd, 1)
raise self._subprocess_module.CalledProcessError(1, cmd, stderr=outcome)
def run(self):
return ProxyExtrasDBManager._setup_database_v2(use_migrate=True)
class TestMigrateDeployAttemptAccounting:
"""A `prisma db push` database has a full schema and no ledger, so the v2
resolver baselines it and then works through every migration whose objects
already exist. Those recoveries make progress, so they must not spend the
retry budget, which is there to stop a run that is getting nowhere."""
def test_a_push_created_database_finishes_bootstrapping(
self, monkeypatch, tmp_path
):
already_there = [
"20250329084805_new_cron_job_table",
"20250806095134_rename_alias_to_server_name_mcp_table",
"20260224203854_add_agent_object_permissions_table",
"20260301120000_fourth_table",
"20260302120000_fifth_table",
"20260303120000_sixth_table",
]
harness = _MigrateDeployHarness(
monkeypatch,
tmp_path,
[_P3005_STDERR]
+ [_p3018_stderr(name) for name in already_there]
+ ["ok"],
)
assert harness.run() is True
assert harness.baselines == 1
assert harness.resolved == already_there
assert len(harness.deploy_calls) == len(already_there) + 2
def test_repeated_recovery_of_one_migration_still_gives_up(
self, monkeypatch, tmp_path
):
harness = _MigrateDeployHarness(
monkeypatch,
tmp_path,
[_p3018_stderr("20250329084805_new_cron_job_table")],
repeat_last=True,
)
with pytest.raises(RuntimeError):
harness.run()
assert len(harness.deploy_calls) <= _ATTEMPT_BUDGET + 1
def test_timeouts_still_spend_the_budget(self, monkeypatch, tmp_path):
harness = _MigrateDeployHarness(
monkeypatch, tmp_path, ["timeout"], repeat_last=True
)
with pytest.raises(RuntimeError):
harness.run()
assert len(harness.deploy_calls) == _ATTEMPT_BUDGET
def test_a_baseline_that_never_lands_stops_after_the_budget(
self, monkeypatch, tmp_path
):
harness = _MigrateDeployHarness(
monkeypatch, tmp_path, [_P3005_STDERR], repeat_last=True
)
harness.baseline_succeeds = False
with pytest.raises(RuntimeError):
harness.run()
assert len(harness.deploy_calls) == _ATTEMPT_BUDGET
def test_an_unrecoverable_error_is_not_retried(self, monkeypatch, tmp_path):
harness = _MigrateDeployHarness(
monkeypatch,
tmp_path,
["Error: P3018\n\nMigration name: 20260101000000_x\n\nERROR: syntax error at or near \"SLECT\"\n"],
repeat_last=True,
)
with pytest.raises(RuntimeError):
harness.run()
assert len(harness.deploy_calls) == 1
assert harness.resolved == []

View file

@ -376,7 +376,6 @@ async def test_bedrock_kb_request_body_has_transformed_filters(
timeout=None,
client=None,
_is_async=False,
router: "litellm.Router | None" = None,
):
litellm_params_dict = (
litellm_params.model_dump(exclude_none=False)

View file

@ -187,7 +187,7 @@ class TestRouterEmbeddingIntegration:
assert _sent(store_route, 1) == ("Bearer store-key", "text-embedding-3-large", ["async query"])
@pytest.mark.asyncio
async def test_router_executor_rejects_unserved_models_without_explicit_config(
async def test_router_executor_embeds_unserved_models_through_the_sdk(
self, respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
):
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
@ -198,12 +198,13 @@ class TestRouterEmbeddingIntegration:
metadata={"user_api_key_team_id": "team-a"},
)
with pytest.raises(litellm.BadRequestError):
executor.embed("openai/text-embedding-3-large", "sync query", {})
with pytest.raises(litellm.BadRequestError):
await executor.aembed("openai/text-embedding-3-large", "async query", {})
sync_response = executor.embed("text-embedding-3-large", "sync query", {})
async_response = await executor.aembed("text-embedding-3-large", "async query", {})
assert openai_route.call_count == 0
assert sync_response.data[0]["embedding"] == QUERY_VECTOR
assert async_response.data[0]["embedding"] == QUERY_VECTOR
assert _sent(openai_route, 0) == ("Bearer env-key", "text-embedding-3-large", ["sync query"])
assert _sent(openai_route, 1) == ("Bearer env-key", "text-embedding-3-large", ["async query"])
def test_router_executor_routes_deployment_model_names_through_the_router(
self, respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch

View file

@ -0,0 +1,274 @@
"""Tests for the hide-secrets guardrail (LIT-3548).
Covers the three defects from the ticket:
- ``apply_guardrail`` (the UI test playground path) must redact, not echo.
- Guardrail runs must record ``standard_logging_guardrail_information`` so
Spend Logs / the guardrails monitor show activity, with hits ("mask" +
masked_entity_count) distinguishable from clean requests ("allow").
- Defining ``apply_guardrail`` must NOT reroute proxied traffic off the
native ``async_pre_call_hook`` (per-key opt-out and ``data["prompt"]``
handling live only on the native path).
"""
import pytest
from litellm_enterprise.enterprise_callbacks.secret_detection import (
_ENTERPRISE_SecretDetection,
)
from litellm.caching.caching import DualCache
from litellm.proxy._types import UserAPIKeyAuth
AWS_KEY = "AKIAIOSFODNN7EXAMPLE"
def _guardrail() -> _ENTERPRISE_SecretDetection:
return _ENTERPRISE_SecretDetection(
guardrail_name="hide-secrets", event_hook="pre_call", default_on=True
)
def _recorded(request_data: dict) -> dict:
entries = request_data["metadata"]["standard_logging_guardrail_information"]
assert len(entries) == 1
return entries[0]
@pytest.mark.asyncio
async def test_apply_guardrail_redacts_secrets():
"""Playground path: the returned texts must carry [REDACTED], not the secret."""
guardrail = _guardrail()
request_data: dict = {"metadata": {}}
result = await guardrail.apply_guardrail(
inputs={"texts": [f"my key is {AWS_KEY}, keep it safe"]},
request_data=request_data,
input_type="request",
)
assert result["texts"] == ["my key is [REDACTED], keep it safe"]
recorded = _recorded(request_data)
assert recorded["guardrail_status"] == "success"
assert recorded["guardrail_response"] == "mask"
assert recorded["guardrail_provider"] == "hide-secrets"
assert recorded["masked_entity_count"] == {"AWS Access Key": 1}
@pytest.mark.asyncio
async def test_apply_guardrail_clean_text_records_allow():
guardrail = _guardrail()
request_data: dict = {"metadata": {}}
result = await guardrail.apply_guardrail(
inputs={"texts": ["nothing sensitive here"]},
request_data=request_data,
input_type="request",
)
assert result["texts"] == ["nothing sensitive here"]
recorded = _recorded(request_data)
assert recorded["guardrail_status"] == "success"
assert recorded["guardrail_response"] == "allow"
assert recorded["masked_entity_count"] == {}
@pytest.mark.asyncio
async def test_pre_call_hook_records_mask_with_entity_count():
"""Live-traffic path: a redaction must be visible in spend-log telemetry."""
guardrail = _guardrail()
data = {
"messages": [{"role": "user", "content": f"use {AWS_KEY} for auth"}],
"metadata": {},
}
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(),
cache=DualCache(),
data=data,
call_type="completion",
)
assert data["messages"][0]["content"] == "use [REDACTED] for auth"
recorded = _recorded(data)
assert recorded["guardrail_status"] == "success"
assert recorded["guardrail_response"] == "mask"
assert recorded["guardrail_provider"] == "hide-secrets"
assert recorded["masked_entity_count"] == {"AWS Access Key": 1}
@pytest.mark.asyncio
async def test_pre_call_hook_clean_request_records_allow():
"""A request with no secrets must be distinguishable from a redacted one."""
guardrail = _guardrail()
data = {
"messages": [{"role": "user", "content": "what's the weather"}],
"metadata": {},
}
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(),
cache=DualCache(),
data=data,
call_type="completion",
)
recorded = _recorded(data)
assert recorded["guardrail_status"] == "success"
assert recorded["guardrail_response"] == "allow"
assert recorded["masked_entity_count"] == {}
@pytest.mark.asyncio
async def test_pre_call_hook_opt_out_records_nothing():
"""A key with permissions={"hide_secrets": False} skips redaction, so no
telemetry is recorded: every reader of a recorded entry (guardrail usage
tracking, compliance checks, the spend-log viewer) counts it as a run."""
guardrail = _guardrail()
content = f"my key is {AWS_KEY}"
data = {"messages": [{"role": "user", "content": content}], "metadata": {}}
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(permissions={"hide_secrets": False}),
cache=DualCache(),
data=data,
call_type="completion",
)
assert data["messages"][0]["content"] == content # untouched
assert "standard_logging_guardrail_information" not in data["metadata"]
@pytest.mark.asyncio
async def test_pre_call_hook_still_redacts_text_completion_prompt():
"""data["prompt"] (str and list) is a native-hook-only surface; it must
keep redacting now that the class also implements apply_guardrail."""
guardrail = _guardrail()
data = {"prompt": f"key {AWS_KEY} end", "metadata": {}}
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(),
cache=DualCache(),
data=data,
call_type="completion",
)
assert data["prompt"] == "key [REDACTED] end"
guardrail = _guardrail()
data = {"prompt": [f"key {AWS_KEY}", "clean"], "metadata": {}}
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(),
cache=DualCache(),
data=data,
call_type="completion",
)
assert data["prompt"] == ["key [REDACTED]", "clean"]
def test_proxied_traffic_stays_on_native_hooks():
"""Implementing apply_guardrail must not reroute proxied requests onto the
unified path: that path skips ``should_run_check`` (per-key opt-out) and
never sees ``data["prompt"]``."""
guardrail = _guardrail()
assert guardrail.uses_apply_guardrail_interface() is True
assert guardrail._deployment_pre_call_target() is guardrail
@pytest.mark.asyncio
async def test_apply_guardrail_without_texts_records_nothing():
"""No inputs means nothing was inspected, so no "allow" row is recorded.
Empty strings count as no input: there is no content to inspect."""
guardrail = _guardrail()
empty_variants: list[list[str]] = [[], ["", ""]]
for texts in empty_variants:
request_data: dict = {"metadata": {}}
result = await guardrail.apply_guardrail(
inputs={"texts": texts}, request_data=request_data, input_type="request"
)
assert result == {"texts": texts}
assert "standard_logging_guardrail_information" not in request_data["metadata"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"data",
[
pytest.param(
{
"messages": [
{
"role": "user",
"content": [
{"type": "image_url", "image_url": {"url": "https://x/y.png"}}
],
}
],
"metadata": {},
},
id="image_only",
),
pytest.param(
{"messages": [{"role": "user", "content": ""}], "metadata": {}},
id="empty_message",
),
pytest.param({"prompt": "", "metadata": {}}, id="empty_prompt"),
pytest.param({"prompt": ["", ""], "metadata": {}}, id="empty_prompt_list"),
],
)
async def test_pre_call_hook_without_inspectable_text_records_nothing(data: dict):
"""A payload the guardrail could not inspect (image-only content, empty
strings) must not record an "allow" run: monitoring would count a check
that never looked at any text."""
guardrail = _guardrail()
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(),
cache=DualCache(),
data=data,
call_type="completion",
)
assert "standard_logging_guardrail_information" not in data["metadata"]
@pytest.mark.asyncio
async def test_pre_call_hook_mixed_prompt_list_still_redacts_and_records():
"""A prompt list mixing empty and real strings is inspected, so the run is
recorded and the non-empty entry is still redacted."""
guardrail = _guardrail()
data = {"prompt": ["", f"key {AWS_KEY}"], "metadata": {}}
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(),
cache=DualCache(),
data=data,
call_type="completion",
)
assert data["prompt"] == ["", "key [REDACTED]"]
recorded = _recorded(data)
assert recorded["guardrail_response"] == "mask"
assert recorded["masked_entity_count"] == {"AWS Access Key": 1}
@pytest.mark.asyncio
async def test_legacy_nameless_instance_records_nothing():
"""``litellm_settings.callbacks: ["hide_secrets"]`` builds an arg-less
instance with no guardrail_name. It still redacts, but recording a nameless
entry would flip every spend row's guardrail status with nothing to join on."""
guardrail = _ENTERPRISE_SecretDetection()
data = {
"messages": [{"role": "user", "content": f"use {AWS_KEY} for auth"}],
"metadata": {},
}
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(),
cache=DualCache(),
data=data,
call_type="completion",
)
assert data["messages"][0]["content"] == "use [REDACTED] for auth"
assert "standard_logging_guardrail_information" not in data["metadata"]

View file

@ -21,6 +21,7 @@ from litellm.litellm_core_utils.streaming_handler import (
from litellm.types.utils import (
CompletionTokensDetailsWrapper,
Delta,
ModelResponse,
ModelResponseStream,
PromptTokensDetailsWrapper,
StandardLoggingPayload,
@ -1750,7 +1751,7 @@ def test_openrouter_streaming_cost_propagates_to_hidden_params():
assert complete_response.usage.cost == 0.00025
# Use the real propagation method from CustomStreamWrapper
CustomStreamWrapper._propagate_usage_cost_to_hidden_params(complete_response)
CustomStreamWrapper._propagate_usage_cost_to_hidden_params(complete_response, "openrouter")
assert "additional_headers" in complete_response._hidden_params
assert (
@ -1769,14 +1770,12 @@ def test_openrouter_streaming_cost_propagates_to_hidden_params():
assert provider_cost == 0.00025
def test_perplexity_streaming_dict_cost_propagates_to_hidden_params():
"""
Regression: Perplexity reports usage.cost as a breakdown object, which used to
blow up the end of the stream with
`float() argument must be a string or a real number, not 'dict'`.
"""
def test_perplexity_streaming_dict_cost_bills_through_its_own_calculator():
import litellm
from litellm.cost_calculator import get_response_cost_from_hidden_params
from litellm.cost_calculator import (
get_response_cost_from_hidden_params,
response_cost_calculator,
)
chunks = [
ModelResponseStream(
@ -1828,13 +1827,81 @@ def test_perplexity_streaming_dict_cost_propagates_to_hidden_params():
assert complete_response is not None
CustomStreamWrapper._propagate_usage_cost_to_hidden_params(complete_response)
CustomStreamWrapper._propagate_usage_cost_to_hidden_params(complete_response, "perplexity")
assert (
get_response_cost_from_hidden_params(complete_response._hidden_params)
== 0.00503
assert get_response_cost_from_hidden_params(complete_response._hidden_params) is None
assert response_cost_calculator(
response_object=complete_response,
model="perplexity/sonar",
custom_llm_provider="perplexity",
call_type="completion",
optional_params={},
) == pytest.approx(0.00503)
def test_openai_compatible_streaming_cost_is_priced_from_the_cost_map():
import litellm
from litellm.cost_calculator import (
get_response_cost_from_hidden_params,
response_cost_calculator,
)
model = "openai/streams-cost-in-nanodollars"
litellm.register_model(
{
model: {
"input_cost_per_token": 1e-6,
"output_cost_per_token": 2e-6,
"litellm_provider": "openai",
"mode": "chat",
}
}
)
complete_response = ModelResponse(
id="chatcmpl-openai-compatible",
model=model,
choices=[],
usage=Usage(completion_tokens=5, prompt_tokens=10, total_tokens=15, cost=3_144_000),
)
CustomStreamWrapper._propagate_usage_cost_to_hidden_params(complete_response, "openai")
assert get_response_cost_from_hidden_params(complete_response._hidden_params) is None
assert response_cost_calculator(
response_object=complete_response,
model=model,
custom_llm_provider="openai",
call_type="completion",
optional_params={},
) == pytest.approx(2e-5)
def test_xai_streaming_reported_cost_still_takes_the_margin(monkeypatch):
import litellm
from litellm.cost_calculator import (
get_response_cost_from_hidden_params,
response_cost_calculator,
)
complete_response = ModelResponse(
id="chatcmpl-xai",
model="grok-4-latest",
choices=[],
usage=Usage(completion_tokens=353, prompt_tokens=198, total_tokens=551, cost=0.0009956),
)
CustomStreamWrapper._propagate_usage_cost_to_hidden_params(complete_response, "xai")
assert get_response_cost_from_hidden_params(complete_response._hidden_params) is None
monkeypatch.setattr(litellm, "cost_margin_config", {"xai": 0.5})
assert response_cost_calculator(
response_object=complete_response,
model="xai/grok-4-latest",
custom_llm_provider="xai",
call_type="completion",
optional_params={},
) == pytest.approx(0.0009956 * 1.5)
def test_provider_reported_cost_ignores_unusable_shapes():
assert CustomStreamWrapper._resolve_provider_reported_cost(None) is None

View file

@ -1,40 +1,70 @@
from collections.abc import Mapping
from unittest.mock import AsyncMock, MagicMock, Mock, patch
import httpx
import pytest
from litellm.llms.base_llm.vector_store.transformation import (
RouterVectorStoreEmbeddingExecutor,
)
from litellm.llms.s3_vectors.vector_stores.transformation import (
S3VectorsVectorStoreConfig,
)
from litellm.types.utils import EmbeddingResponse
from litellm.types.vector_stores import VectorStoreSearchResponse
QUERY_VECTOR = [0.1, 0.2, 0.3]
def _mock_router(model_names, sync=False):
"""Router mock serving the given embedding model names."""
router = MagicMock()
router.get_model_list.return_value = [{"model_name": name} for name in model_names]
embedding_response = Mock(data=[{"embedding": [0.1, 0.2, 0.3]}])
if sync:
router.embedding = MagicMock(return_value=embedding_response)
else:
router.aembedding = AsyncMock(return_value=embedding_response)
return router
def _embedding_response(vector):
return EmbeddingResponse(data=[{"embedding": vector, "index": 0, "object": "embedding"}])
class _RecordingExecutor:
def __init__(self, vector=QUERY_VECTOR):
self.vector = vector
self.calls = []
def embed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse:
self.calls.append((model, query, dict(configuration)))
return _embedding_response(self.vector)
async def aembed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse:
self.calls.append((model, query, dict(configuration)))
return _embedding_response(self.vector)
def _logging_obj():
logging_obj = Mock()
logging_obj.model_call_details = {}
return logging_obj
def _search_kwargs(**overrides):
kwargs = {
"vector_store_id": "test-bucket:test-index",
"query": "test query",
"vector_store_search_optional_params": {},
"api_base": "https://s3vectors.us-west-2.api.aws",
"litellm_logging_obj": _logging_obj(),
"litellm_params": {},
"extra_body": None,
}
kwargs.update(overrides)
return kwargs
class TestS3VectorsVectorStoreConfig:
def test_init(self):
"""Test that S3VectorsVectorStoreConfig initializes correctly"""
config = S3VectorsVectorStoreConfig()
assert config is not None
def test_get_supported_openai_params(self):
"""Test that supported OpenAI params are returned"""
config = S3VectorsVectorStoreConfig()
params = config.get_supported_openai_params("test-model")
assert "max_num_results" in params
def test_get_complete_url(self):
"""Test URL generation for S3 Vectors"""
config = S3VectorsVectorStoreConfig()
litellm_params = {"aws_region_name": "us-west-2"}
url = config.get_complete_url(None, litellm_params)
@ -57,180 +87,170 @@ class TestS3VectorsVectorStoreConfig:
assert url == "https://s3vectors.eu-west-1.api.aws"
def test_get_complete_url_invalid_region_format(self):
"""Invalid region format raises"""
config = S3VectorsVectorStoreConfig()
with pytest.raises(ValueError, match="Invalid AWS region format"):
config.get_complete_url(None, {"aws_region_name": "Bad_Region!"})
def test_transform_search_request(self):
"""Full request-body transformation with a router-injected embedding"""
config = S3VectorsVectorStoreConfig()
mock_logging_obj = Mock()
mock_logging_obj.model_call_details = {}
router = _mock_router(["text-embedding-3-small"], sync=True)
logging_obj = _logging_obj()
executor = _RecordingExecutor()
url, request_body = config.transform_search_vector_store_request(
vector_store_id="test-bucket:test-index",
query="test query",
vector_store_search_optional_params={"max_num_results": 7},
api_base="https://s3vectors.us-west-2.api.aws",
litellm_logging_obj=mock_logging_obj,
litellm_params={},
extra_body=None,
router=router,
**_search_kwargs(
vector_store_search_optional_params={"max_num_results": 7},
litellm_logging_obj=logging_obj,
embedding_executor=executor,
)
)
assert url == "https://s3vectors.us-west-2.api.aws/QueryVectors"
assert request_body == {
"vectorBucketName": "test-bucket",
"indexName": "test-index",
"queryVector": {"float32": [0.1, 0.2, 0.3]},
"queryVector": {"float32": QUERY_VECTOR},
"topK": 7,
"returnDistance": True,
"returnMetadata": True,
}
assert mock_logging_obj.model_call_details["query"] == "test query"
assert executor.calls == [("text-embedding-3-small", "test query", {})]
assert logging_obj.model_call_details["query"] == "test query"
@pytest.mark.parametrize(
("litellm_params", "expected_model"),
[
({}, "text-embedding-3-small"),
({"embedding_model": ""}, "text-embedding-3-small"),
({"embedding_model": "my-embedding-model"}, "my-embedding-model"),
({"litellm_embedding_model": "shared-key-model"}, "shared-key-model"),
(
{"litellm_embedding_model": "shared-key-model", "embedding_model": "legacy-alias"},
"shared-key-model",
),
],
)
def test_query_embedding_model_accepts_embedding_model_alias(self, litellm_params, expected_model):
assert S3VectorsVectorStoreConfig.query_embedding_model(litellm_params) == expected_model
@pytest.mark.asyncio
async def test_atransform_search_uses_router_for_virtual_model(self):
"""Regression: router-served embedding models must resolve via the router,
not a bare litellm.aembedding call (which has no deployment credentials)."""
async def test_atransform_search_embeds_alias_and_store_config_through_executor(self):
config = S3VectorsVectorStoreConfig()
mock_logging_obj = Mock()
mock_logging_obj.model_call_details = {}
router = _mock_router(["my-embedding-model"])
executor = _RecordingExecutor(vector=[0.4, 0.5])
with patch("litellm.aembedding", new=AsyncMock()) as mock_bare_aembedding: # test-quality-ok: guards that the bare-embedding path is not taken; dispatch seam is the behavior under test
url, request_body = await config.atransform_search_vector_store_request(
vector_store_id="test-bucket:test-index",
query="test query",
vector_store_search_optional_params={},
api_base="https://s3vectors.us-west-2.api.aws",
litellm_logging_obj=mock_logging_obj,
litellm_params={"embedding_model": "my-embedding-model"},
extra_body=None,
router=router,
_, request_body = await config.atransform_search_vector_store_request(
**_search_kwargs(
query=["test", "query"],
litellm_params={
"embedding_model": "my-embedding-model",
"litellm_embedding_config": {"api_key": "store-key"},
},
embedding_executor=executor,
)
)
router.aembedding.assert_awaited_once_with(model="my-embedding-model", input=["test query"])
mock_bare_aembedding.assert_not_awaited()
assert request_body["queryVector"]["float32"] == [0.1, 0.2, 0.3]
assert request_body["topK"] == 5 # default
@pytest.mark.asyncio
async def test_atransform_search_falls_back_when_router_does_not_serve_model(self):
"""Router present but embedding_model is not a router deployment ->
bare litellm.aembedding keeps working (provider-prefixed + env creds stores)."""
config = S3VectorsVectorStoreConfig()
mock_logging_obj = Mock()
mock_logging_obj.model_call_details = {}
router = _mock_router(["some-other-model"])
mock_bare = AsyncMock(return_value=Mock(data=[{"embedding": [0.4, 0.5]}]))
with patch("litellm.aembedding", new=mock_bare): # test-quality-ok: stubs the bare-embedding fallback whose request body the test asserts on
_, request_body = await config.atransform_search_vector_store_request(
vector_store_id="test-bucket:test-index",
query="test query",
vector_store_search_optional_params={},
api_base="https://s3vectors.us-west-2.api.aws",
litellm_logging_obj=mock_logging_obj,
litellm_params={"embedding_model": "azure/text-embedding-3-small"},
extra_body=None,
router=router,
)
mock_bare.assert_awaited_once_with(model="azure/text-embedding-3-small", input=["test query"])
router.aembedding.assert_not_awaited()
assert executor.calls == [("my-embedding-model", "test query", {"api_key": "store-key"})]
assert request_body["queryVector"]["float32"] == [0.4, 0.5]
assert request_body["topK"] == 5
@pytest.mark.asyncio
async def test_atransform_search_without_router_uses_bare_embedding(self):
"""Backward compat: no router -> bare litellm.aembedding as before"""
async def test_atransform_search_router_executor_carries_request_metadata(self):
"""Regression (LIT-6750): a bare Router alias resolves through the Router with the request's
team metadata on the embedding call, so the embedding is attributed to the calling key and team."""
config = S3VectorsVectorStoreConfig()
mock_logging_obj = Mock()
mock_logging_obj.model_call_details = {}
router = MagicMock()
router.aembedding = AsyncMock(return_value=_embedding_response(QUERY_VECTOR))
request_metadata = {"user_api_key_team_id": "team-a", "user_api_key": "hashed-key"}
mock_bare = AsyncMock(return_value=Mock(data=[{"embedding": [0.6, 0.7]}]))
with patch("litellm.aembedding", new=mock_bare): # test-quality-ok: stubs the bare-embedding fallback whose request body the test asserts on
_, request_body = await config.atransform_search_vector_store_request(
vector_store_id="test-bucket:test-index",
query="test query",
vector_store_search_optional_params={},
api_base="https://s3vectors.us-west-2.api.aws",
litellm_logging_obj=mock_logging_obj,
litellm_params={},
extra_body=None,
_, request_body = await config.atransform_search_vector_store_request(
**_search_kwargs(
litellm_params={"embedding_model": "team-embeddings"},
embedding_executor=RouterVectorStoreEmbeddingExecutor(router=router, metadata=request_metadata),
)
)
router.aembedding.assert_awaited_once_with(
model="team-embeddings", input=["test query"], metadata=request_metadata
)
assert request_body["queryVector"]["float32"] == QUERY_VECTOR
@pytest.mark.asyncio
async def test_atransform_search_default_model_falls_back_to_the_sdk(self):
"""Regression (LIT-6750): a store that never named an embedding model keeps working on a proxy
whose model list has no text-embedding-3-small, embedding through the SDK instead of erroring."""
config = S3VectorsVectorStoreConfig()
router = MagicMock()
router.get_model_list.return_value = [
{"model_name": "team-embeddings", "litellm_params": {"model": "openai/text-embedding-3-small"}}
]
router.resolved_litellm_models.return_value = []
router.aembedding = AsyncMock(side_effect=AssertionError("unserved model must not reach the Router"))
request_metadata = {"user_api_key_team_id": "team-a"}
mock_bare = AsyncMock(return_value=_embedding_response(QUERY_VECTOR))
with patch("litellm.aembedding", new=mock_bare): # test-quality-ok: stubs the bare-embedding fallback whose call the test asserts on
_, request_body = await config.atransform_search_vector_store_request(
**_search_kwargs(
embedding_executor=RouterVectorStoreEmbeddingExecutor(router=router, metadata=request_metadata)
)
)
mock_bare.assert_awaited_once_with(
model="text-embedding-3-small", input=["test query"], metadata=request_metadata
)
assert request_body["queryVector"]["float32"] == QUERY_VECTOR
@pytest.mark.asyncio
async def test_atransform_search_without_executor_uses_bare_embedding(self):
config = S3VectorsVectorStoreConfig()
mock_bare = AsyncMock(return_value=_embedding_response([0.6, 0.7]))
with patch("litellm.aembedding", new=mock_bare): # test-quality-ok: stubs the bare-embedding fallback whose request body the test asserts on
_, request_body = await config.atransform_search_vector_store_request(**_search_kwargs())
mock_bare.assert_awaited_once_with(model="text-embedding-3-small", input=["test query"])
assert request_body["queryVector"]["float32"] == [0.6, 0.7]
def test_transform_search_uses_router_for_virtual_model_sync(self):
"""Sync twin: router-served embedding model resolves via router.embedding"""
def test_transform_search_without_executor_uses_bare_embedding_sync(self):
config = S3VectorsVectorStoreConfig()
mock_logging_obj = Mock()
mock_logging_obj.model_call_details = {}
router = _mock_router(["my-embedding-model"], sync=True)
with patch("litellm.embedding", new=MagicMock()) as mock_bare_embedding: # test-quality-ok: guards that the bare-embedding path is not taken; dispatch seam is the behavior under test
_, request_body = config.transform_search_vector_store_request(
vector_store_id="test-bucket:test-index",
query="test query",
vector_store_search_optional_params={},
api_base="https://s3vectors.us-west-2.api.aws",
litellm_logging_obj=mock_logging_obj,
litellm_params={"embedding_model": "my-embedding-model"},
extra_body=None,
router=router,
)
router.embedding.assert_called_once_with(model="my-embedding-model", input=["test query"])
mock_bare_embedding.assert_not_called()
assert request_body["queryVector"]["float32"] == [0.1, 0.2, 0.3]
def test_transform_search_without_router_uses_bare_embedding_sync(self):
"""Sync twin: no router -> bare litellm.embedding as before"""
config = S3VectorsVectorStoreConfig()
mock_logging_obj = Mock()
mock_logging_obj.model_call_details = {}
mock_bare = MagicMock(return_value=Mock(data=[{"embedding": [0.8, 0.9]}]))
mock_bare = MagicMock(return_value=_embedding_response([0.8, 0.9]))
with patch("litellm.embedding", new=mock_bare): # test-quality-ok: stubs the bare-embedding fallback whose request body the test asserts on
_, request_body = config.transform_search_vector_store_request(
vector_store_id="test-bucket:test-index",
query="test query",
vector_store_search_optional_params={},
api_base="https://s3vectors.us-west-2.api.aws",
litellm_logging_obj=mock_logging_obj,
litellm_params={},
extra_body=None,
**_search_kwargs(litellm_params={"embedding_model": "my-embedding-model"})
)
mock_bare.assert_called_once_with(model="text-embedding-3-small", input=["test query"])
mock_bare.assert_called_once_with(model="my-embedding-model", input=["test query"])
assert request_body["queryVector"]["float32"] == [0.8, 0.9]
def test_transform_search_request_invalid_vector_store_id(self):
"""Test that invalid vector_store_id format raises error"""
config = S3VectorsVectorStoreConfig()
mock_logging_obj = Mock()
mock_logging_obj.model_call_details = {}
executor = _RecordingExecutor()
with pytest.raises(
ValueError,
match="vector_store_id must be in format 'bucket_name:index_name'",
):
config.transform_search_vector_store_request(
vector_store_id="invalid-format",
query="test query",
vector_store_search_optional_params={},
api_base="https://s3vectors.us-west-2.api.aws",
litellm_logging_obj=mock_logging_obj,
litellm_params={},
extra_body=None,
**_search_kwargs(vector_store_id="invalid-format", embedding_executor=executor)
)
assert executor.calls == []
def test_transform_search_request_bucket_from_litellm_params(self):
config = S3VectorsVectorStoreConfig()
_, request_body = config.transform_search_vector_store_request(
**_search_kwargs(
vector_store_id="only-index",
litellm_params={"vector_bucket_name": "params-bucket"},
embedding_executor=_RecordingExecutor(),
)
)
assert request_body["vectorBucketName"] == "params-bucket"
assert request_body["indexName"] == "only-index"
def test_transform_search_response(self):
"""Test search response transformation"""
config = S3VectorsVectorStoreConfig()
mock_logging_obj = Mock()
mock_logging_obj.model_call_details = {"query": "test query"}
@ -239,7 +259,7 @@ class TestS3VectorsVectorStoreConfig:
mock_response.json.return_value = {
"vectors": [
{
"distance": 0.05, # S3 Vectors returns distance, not score
"distance": 0.05,
"metadata": {
"source_text": "This is test content",
"chunk_index": "0",
@ -258,23 +278,18 @@ class TestS3VectorsVectorStoreConfig:
mock_response.status_code = 200
mock_response.headers = {}
result = config.transform_search_vector_store_response(
mock_response, mock_logging_obj
)
result = config.transform_search_vector_store_response(mock_response, mock_logging_obj)
# VectorStoreSearchResponse is a TypedDict, so check structure instead of isinstance
assert result["object"] == "vector_store.search_results.page"
assert result["search_query"] == "test query"
assert len(result["data"]) == 2
# Score should be 1 - distance (cosine similarity)
assert result["data"][0]["score"] == 0.95 # 1 - 0.05
assert result["data"][0]["score"] == 0.95
assert result["data"][0]["content"][0]["text"] == "This is test content"
assert result["data"][0]["filename"] == "test.pdf"
assert result["data"][1]["score"] == 0.85 # 1 - 0.15
assert result["data"][1]["score"] == 0.85
assert result["data"][1]["content"][0]["text"] == "More test content"
def test_map_openai_params(self):
"""Test OpenAI parameter mapping"""
config = S3VectorsVectorStoreConfig()
non_default_params = {"max_num_results": 5}
optional_params = {}

View file

@ -7,12 +7,13 @@ transformations for the Responses API.
Source: litellm/llms/xai/responses/transformation.py
"""
from unittest.mock import MagicMock
from unittest.mock import MagicMock, Mock
import httpx
import pytest
import litellm
from litellm.llms.xai.cost_calculator import cost_per_token
from litellm.llms.xai.responses.transformation import XAIResponsesAPIConfig
from litellm.responses.utils import ResponseAPILoggingUtils
from litellm.types.llms.openai import (
@ -400,3 +401,94 @@ class TestXAIResponsesWebSearchBilling:
bridged = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(event.response.usage)
assert getattr(bridged, "server_side_tool_usage_details") == self._TOOL_DETAILS
class TestXAIResponsesReportedCost:
"""xAI reports what it charged; the transformation moves it to where litellm bills from.
``ResponseAPILoggingUtils`` copies ``usage.cost`` onto the chat Usage that cost
tracking prices, so restating ``cost_in_usd_ticks`` there is what makes /v1/responses
bill the reported figure. At 10^10 ticks to the dollar, 37756000 ticks is $0.0037756.
"""
@staticmethod
def _response_body(usage: dict) -> dict:
return {
"id": "resp_xai",
"object": "response",
"created_at": 0,
"model": "grok-4-latest",
"status": "completed",
"output": [],
"parallel_tool_calls": False,
"tool_choice": "auto",
"tools": [],
"usage": usage,
}
def _transformed_usage(self, usage: dict) -> ResponseAPIUsage | None:
raw_response = httpx.Response(status_code=200, json=self._response_body(usage))
response = XAIResponsesAPIConfig().transform_response_api_response(
model="grok-4-latest",
raw_response=raw_response,
logging_obj=Mock(),
)
return response.usage
def test_reported_cost_reaches_the_cost_calculator(self):
usage = self._transformed_usage(
{
"input_tokens": 100,
"output_tokens": 200,
"total_tokens": 300,
"cost_in_usd_ticks": 37756000,
}
)
assert usage.cost == 0.0037756
chat_usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
assert cost_per_token(model="grok-4-latest", usage=chat_usage) == (0.0, 0.0037756)
def test_streamed_reported_cost_reaches_the_cost_calculator(self):
event = XAIResponsesAPIConfig().transform_streaming_response(
model="grok-4-latest",
parsed_chunk={
"type": "response.completed",
"sequence_number": 7,
"response": self._response_body(
{
"input_tokens": 100,
"output_tokens": 200,
"total_tokens": 300,
"cost_in_usd_ticks": 37756000,
}
),
},
logging_obj=Mock(),
)
assert isinstance(event, ResponseCompletedEvent)
chat_usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(event.response.usage)
assert cost_per_token(model="grok-4-latest", usage=chat_usage) == (0.0, 0.0037756)
def test_usage_without_a_reported_cost_is_left_alone(self):
usage = self._transformed_usage(
{"input_tokens": 100, "output_tokens": 200, "total_tokens": 300}
)
assert usage.cost is None
def test_negative_reported_cost_is_not_carried(self):
"""A caller who can set api_base must not be able to report negative spend."""
usage = self._transformed_usage(
{
"input_tokens": 100,
"output_tokens": 200,
"total_tokens": 300,
"cost_in_usd_ticks": -37756000,
}
)
assert usage.cost is None

View file

@ -1,9 +1,14 @@
from unittest.mock import Mock
import httpx
import pytest
import litellm
from litellm.llms.xai.chat.transformation import XAIChatConfig
from litellm.llms.xai.chat.transformation import (
XAIChatCompletionStreamingHandler,
XAIChatConfig,
)
from litellm.llms.xai.cost_calculator import cost_per_token
from litellm.types.utils import (
CompletionTokensDetailsWrapper,
ModelResponse,
@ -195,3 +200,113 @@ class TestXAIChatWebSearchBilling:
)
assert with_search - without_search == pytest.approx(3 * 5.0 / 1000.0)
class TestXAIReportedCost:
"""xAI reports what it charged; the transformation moves it to where litellm bills from.
``cost`` is the field litellm already carries a provider stated cost in, so restating
``cost_in_usd_ticks`` there is what lets ``llms/xai/cost_calculator.py`` bill the
reported figure. At 10^10 ticks to the dollar, 37756000 ticks is $0.0037756.
"""
@staticmethod
def _transformed_usage(usage: dict) -> Usage:
raw_response = httpx.Response(
status_code=200,
json={
"id": "chatcmpl-xai",
"object": "chat.completion",
"created": 0,
"model": "grok-4-latest",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "hi"},
"finish_reason": "stop",
}
],
"usage": usage,
},
)
response = XAIChatConfig().transform_response(
model="grok-4-latest",
raw_response=raw_response,
model_response=ModelResponse(),
logging_obj=Mock(),
request_data={},
messages=[{"role": "user", "content": "hi"}],
optional_params={},
litellm_params={},
encoding=None,
)
return response.usage
def test_reported_cost_reaches_the_cost_calculator(self):
usage = self._transformed_usage(
{
"prompt_tokens": 100,
"completion_tokens": 200,
"total_tokens": 300,
"cost_in_usd_ticks": 37756000,
}
)
assert usage.cost == 0.0037756
assert cost_per_token(model="grok-4-latest", usage=usage) == (0.0, 0.0037756)
def test_usage_without_a_reported_cost_is_left_alone(self):
usage = self._transformed_usage(
{"prompt_tokens": 100, "completion_tokens": 200, "total_tokens": 300}
)
assert getattr(usage, "cost", None) is None
def test_negative_reported_cost_is_not_carried(self):
"""A caller who can set api_base must not be able to report negative spend."""
usage = self._transformed_usage(
{
"prompt_tokens": 100,
"completion_tokens": 200,
"total_tokens": 300,
"cost_in_usd_ticks": -37756000,
}
)
assert getattr(usage, "cost", None) is None
def test_streamed_reported_cost_survives_chunk_aggregation(self):
"""Streamed spend only matches if the conversion happens on the chunk.
Chunk aggregation rebuilds usage from the fields it models plus ``cost``, so a
chunk still carrying only ``cost_in_usd_ticks`` loses the reported amount.
"""
handler = XAIChatCompletionStreamingHandler(
streaming_response=iter([]), sync_stream=True
)
parsed = handler.chunk_parser(
{
"id": "chatcmpl-xai",
"object": "chat.completion.chunk",
"created": 0,
"model": "grok-4-latest",
"choices": [],
"usage": {
"prompt_tokens": 100,
"completion_tokens": 200,
"total_tokens": 300,
"cost_in_usd_ticks": 37756000,
},
}
)
assert parsed.usage.cost == 0.0037756
assembled = litellm.stream_chunk_builder(chunks=[parsed])
assert assembled.usage.cost == 0.0037756
assert cost_per_token(model="grok-4-latest", usage=assembled.usage) == (
0.0,
0.0037756,
)

View file

@ -7,7 +7,10 @@ import os
import litellm
from litellm.types.utils import (
Choices,
CompletionTokensDetailsWrapper,
Message,
ModelResponse,
PromptTokensDetailsWrapper,
Usage,
)
@ -361,6 +364,145 @@ class TestXAICostCalculator:
response_object=object(), usage=usage
)
def test_reported_cost_is_preferred_over_token_math(self):
"""The amount xAI reported, carried on usage.cost by the transformation, is billed.
It lands entirely on completion cost because xAI does not split its total by
direction, the same shape the perplexity calculator returns.
"""
usage = Usage(
prompt_tokens=100,
completion_tokens=200,
total_tokens=300,
cost=0.0037756,
)
prompt_cost, completion_cost = cost_per_token(model="grok-4-latest", usage=usage)
assert prompt_cost == 0.0
assert math.isclose(completion_cost, 0.0037756, rel_tol=1e-10)
def test_reported_cost_suppresses_web_search_surcharge(self):
"""The reported total already covers server-side tool calls.
Without the suppression these 3 searches would be billed a second time on
top of the total xAI already charged.
"""
usage = Usage(
prompt_tokens=100,
completion_tokens=50,
total_tokens=150,
prompt_tokens_details=PromptTokensDetailsWrapper(
text_tokens=100,
web_search_requests=3,
),
cost=0.0037756,
)
assert cost_per_web_search_request(usage=usage, model_info={}) == 0.0
def test_web_search_surcharge_suppressed_through_the_dispatcher(self):
"""The suppression has to hold on the path cost tracking actually uses.
Legacy behaviour stays intact when xAI reports no cost.
"""
from litellm.llms import get_cost_for_web_search_request
usage = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150)
setattr(usage, "server_side_tool_usage_details", {"web_search_calls": 3})
assert get_cost_for_web_search_request("xai", usage, {}) > 0.0
reported = Usage(
prompt_tokens=100, completion_tokens=50, total_tokens=150, cost=0.0037756
)
setattr(reported, "server_side_tool_usage_details", {"web_search_calls": 3})
assert get_cost_for_web_search_request("xai", reported, {}) == 0.0
def test_no_reported_cost_falls_back_to_token_math(self):
"""Absent the provider figure, nothing changes for existing callers."""
usage = Usage(prompt_tokens=100, completion_tokens=200, total_tokens=300)
prompt_cost, completion_cost = cost_per_token(model="grok-4-latest", usage=usage)
assert prompt_cost > 0.0
assert completion_cost > 0.0
def test_malformed_reported_cost_falls_back_to_token_math(self):
"""A junk value must not fail the request, fall back to calculating."""
usage = Usage(prompt_tokens=100, completion_tokens=200, total_tokens=300)
setattr(usage, "cost", "not-a-number")
prompt_cost, completion_cost = cost_per_token(model="grok-4-latest", usage=usage)
assert prompt_cost > 0.0
assert completion_cost > 0.0
def test_boolean_reported_cost_falls_back_to_token_math(self):
"""True is an int in python and would otherwise be billed as $1."""
usage = Usage(prompt_tokens=100, completion_tokens=200, total_tokens=300)
setattr(usage, "cost", True)
prompt_cost, completion_cost = cost_per_token(model="grok-4-latest", usage=usage)
assert prompt_cost > 0.0
assert completion_cost > 0.0
assert completion_cost != 1.0
def test_negative_reported_cost_is_rejected(self):
"""A negative amount must never reach spend tracking.
A caller who can set api_base controls the response body, so trusting a
negative figure would let them subtract from their own recorded spend and
slip past a budget. Fall back to token pricing instead, and keep charging
the web search surcharge, since no trustworthy total was reported.
"""
usage = Usage(
prompt_tokens=100,
completion_tokens=200,
total_tokens=300,
cost=-0.0037756,
)
setattr(usage, "server_side_tool_usage_details", {"web_search_calls": 3})
prompt_cost, completion_cost = cost_per_token(model="grok-4-latest", usage=usage)
assert prompt_cost > 0.0
assert completion_cost > 0.0
assert cost_per_web_search_request(usage=usage, model_info={}) > 0.0
def test_non_finite_reported_cost_is_rejected(self):
"""NaN compares false against every budget threshold.
Usage stores a provider supplied cost without validating it, so a caller who
controls the response body could report NaN and leave spend >= max_budget
false for the life of the key rather than mispricing one request. The
infinities are refused alongside it. Fall back to token pricing and keep
charging the web search surcharge, since no trustworthy total was reported.
"""
for reported_cost in (float("nan"), float("inf"), float("-inf")):
usage = Usage(
prompt_tokens=100,
completion_tokens=200,
total_tokens=300,
cost=reported_cost,
)
setattr(usage, "server_side_tool_usage_details", {"web_search_calls": 3})
prompt_cost, completion_cost = cost_per_token(model="grok-4-latest", usage=usage)
assert math.isfinite(prompt_cost), reported_cost
assert math.isfinite(completion_cost), reported_cost
assert prompt_cost > 0.0, reported_cost
assert completion_cost > 0.0, reported_cost
assert cost_per_web_search_request(usage=usage, model_info={}) > 0.0, reported_cost
def test_zero_reported_cost_is_honoured(self):
"""A reported zero is a real answer, not a missing value."""
usage = Usage(prompt_tokens=100, completion_tokens=200, total_tokens=300, cost=0.0)
assert cost_per_token(model="grok-4-latest", usage=usage) == (0.0, 0.0)
def test_grok_4_20_beta_reasoning_cost_calculation(self):
"""Test cost calculation for grok-4.20-beta-0309-reasoning model."""
usage = Usage(prompt_tokens=100, completion_tokens=200, total_tokens=300)
@ -437,6 +579,48 @@ class TestXAICostCalculator:
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
def test_custom_pricing_beats_the_reported_cost(self):
response = ModelResponse(
id="chatcmpl-xai",
model="grok-4-latest",
choices=[Choices(index=0, message=Message(role="assistant", content="x"), finish_reason="stop")],
usage=Usage(prompt_tokens=198, completion_tokens=353, total_tokens=551, cost=0.0009956),
)
billed = litellm.completion_cost(
completion_response=response,
model="xai/grok-4-latest",
custom_llm_provider="xai",
custom_cost_per_token={"input_cost_per_token": 0.001, "output_cost_per_token": 0.001},
custom_pricing=True,
)
assert math.isclose(billed, 0.551, rel_tol=1e-10)
def test_deployment_custom_pricing_beats_the_reported_cost(self, monkeypatch):
deployment_id = "xai-deployment-priced-by-the-operator"
monkeypatch.setitem(
litellm.model_cost,
deployment_id,
{"input_cost_per_token": 0.001, "output_cost_per_token": 0.001, "litellm_provider": "xai", "mode": "chat"},
)
response = ModelResponse(
id="chatcmpl-xai",
model="grok-4-latest",
choices=[Choices(index=0, message=Message(role="assistant", content="x"), finish_reason="stop")],
usage=Usage(prompt_tokens=198, completion_tokens=353, total_tokens=551, cost=0.0009956),
)
billed = litellm.completion_cost(
completion_response=response,
model="xai/grok-4-latest",
custom_llm_provider="xai",
custom_pricing=True,
router_model_id=deployment_id,
)
assert math.isclose(billed, 0.551, rel_tol=1e-10)
class TestXAIWebSearchCostHelpers:
"""Focused coverage for web_search / tool-usage helpers in cost_calculator.py."""

View file

@ -78,9 +78,19 @@ class TestBuildAgentEnv:
assert env["ANTHROPIC_BASE_URL"] == "http://localhost:4000"
assert env["ANTHROPIC_AUTH_TOKEN"] == "sk-key"
assert env["ENABLE_TOOL_SEARCH"] == "true"
assert env["CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY"] == "1"
assert "OPENAI_BASE_URL" not in env
assert "OPENAI_API_KEY" not in env
def test_anthropic_profile_preserves_existing_gateway_model_discovery(self):
env = build_agent_env(
{"CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY": "0"},
"http://localhost:4000",
"sk-key",
frozenset({"anthropic"}),
)
assert env["CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY"] == "0"
def test_anthropic_profile_preserves_existing_tool_search(self):
env = build_agent_env(
{"ENABLE_TOOL_SEARCH": "false"},
@ -107,6 +117,7 @@ class TestBuildAgentEnv:
assert env["OPENAI_API_KEY"] == "sk-key"
assert "ANTHROPIC_BASE_URL" not in env
assert "ENABLE_TOOL_SEARCH" not in env
assert "CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY" not in env
def test_both_profiles_set_everything(self):
env = build_agent_env(

View file

@ -56,8 +56,14 @@ class TestMergeClaudeSettings:
merged = merge_claude_settings(settings, "http://localhost:4000/", "new-helper")
assert merged["env"]["ANTHROPIC_BASE_URL"] == "http://localhost:4000"
assert merged["env"]["ENABLE_TOOL_SEARCH"] == "true"
assert merged["env"]["CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY"] == "1"
assert merged["apiKeyHelper"] == "new-helper"
def test_preserves_existing_gateway_model_discovery(self):
settings = {"env": {"CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY": "0"}}
merged = merge_claude_settings(settings, "http://localhost:4000", "helper")
assert merged["env"]["CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY"] == "0"
def test_preserves_existing_tool_search(self):
settings = {"env": {"ENABLE_TOOL_SEARCH": "false"}}
merged = merge_claude_settings(settings, "http://localhost:4000", "helper")
@ -73,6 +79,7 @@ class TestMergeClaudeSettings:
assert merged["env"] == {
"ANTHROPIC_BASE_URL": "http://localhost:4000",
"ENABLE_TOOL_SEARCH": "true",
"CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY": "1",
}
assert merged["apiKeyHelper"] == "helper"

View file

@ -26,16 +26,12 @@ def _as_non_admin():
return UserAPIKeyAuth(api_key="test-key", user_role="internal_user")
def _patch_credential(name: str, body: dict, auth=_as_admin):
def _call_as(method: str, path: str, json_body: dict | None = None, auth=_as_admin):
missing = object()
previous_override = app.dependency_overrides.get(user_api_key_auth, missing)
app.dependency_overrides[user_api_key_auth] = auth
try:
return client.patch(
f"/credentials/{name}",
json=body,
headers={"Authorization": "Bearer test-key"},
)
return client.request(method, path, json=json_body, headers={"Authorization": "Bearer test-key"})
finally:
if previous_override is missing:
app.dependency_overrides.pop(user_api_key_auth, None)
@ -43,17 +39,49 @@ def _patch_credential(name: str, body: dict, auth=_as_admin):
app.dependency_overrides[user_api_key_auth] = previous_override
def _patch_credential(name: str, body: dict, auth=_as_admin):
return _call_as("PATCH", f"/credentials/{name}", body, auth)
def _post_credential(body: dict, auth=_as_admin):
missing = object()
previous_override = app.dependency_overrides.get(user_api_key_auth, missing)
app.dependency_overrides[user_api_key_auth] = auth
try:
return client.post("/credentials", json=body, headers={"Authorization": "Bearer test-key"})
finally:
if previous_override is missing:
app.dependency_overrides.pop(user_api_key_auth, None)
else:
app.dependency_overrides[user_api_key_auth] = previous_override
return _call_as("POST", "/credentials", body, auth)
def _delete_credential(name: str, auth=_as_admin):
return _call_as("DELETE", f"/credentials/{name}", auth=auth)
def _list_credentials():
return _call_as("GET", "/credentials")
def _prisma_without_credential_rows() -> MagicMock:
prisma_client = MagicMock()
prisma_client.db.litellm_credentialstable.find_unique = AsyncMock(return_value=None)
return prisma_client
@pytest.fixture
def credential_store():
"""Stands the credential store up for one test: whether the database is reachable, what
the proxy is already serving from memory, and what each repository call hands back."""
def install(
*,
connected: bool = True,
in_memory: tuple[object, ...] = (),
**repository_calls: AsyncMock,
) -> None:
patch("litellm.proxy.proxy_server.prisma_client", _prisma_without_credential_rows() if connected else None).start()
patch("litellm.proxy.proxy_server.master_key", "sk-test-master").start()
patch.object(litellm, "credential_list", list(in_memory)).start()
repository = patch("litellm.proxy.credential_endpoints.endpoints.CredentialsRepository").start()
repository.return_value.find_by_name = AsyncMock(return_value=None)
for call_name, result in repository_calls.items():
setattr(repository.return_value, call_name, result)
yield install
patch.stopall()
@contextmanager
@ -79,7 +107,7 @@ def _repository_holding(stored: CredentialItem | None):
repository.return_value.find_by_name = AsyncMock(return_value=stored)
repository.return_value.create = AsyncMock(return_value=None)
repository.return_value.update_by_name = AsyncMock(return_value=None)
repository.return_value.delete_by_name = AsyncMock(return_value=None)
repository.return_value.delete_by_name = AsyncMock(return_value=stored)
yield repository.return_value
@ -102,46 +130,35 @@ def test_create_credential_write_omits_the_patch_only_deletion_field(restore_cre
assert "credential_values_to_delete" not in written_data
def test_update_credential_answers_404_when_the_credential_does_not_exist():
def test_update_credential_answers_404_when_the_credential_does_not_exist(credential_store):
"""Regression: the handler used to ``return handle_exception_on_proxy(e)``, which makes
the exception the response body and lets FastAPI answer 200, so a write the handler
rejected read as a success to every caller that checks the status. The dashboard's API
client branches on the status, so it reported a failed edit as applied."""
with (
patch( # test-quality-ok: the proxy wiring under test is what this patches
"litellm.proxy.proxy_server.prisma_client", MagicMock()
),
patch( # test-quality-ok: the proxy wiring under test is what this patches
"litellm.proxy.credential_endpoints.endpoints.CredentialsRepository"
) as repository, # test-quality-ok: the proxy wiring under test is what this patches
):
repository.return_value.find_by_name = AsyncMock(return_value=None)
credential_store(find_by_name=AsyncMock(return_value=None))
response = _patch_credential(
"definitely-not-there",
{
"credential_name": "definitely-not-there",
"credential_values": {"api_key": "sk-x"},
"credential_info": {},
},
)
response = _patch_credential(
"definitely-not-there",
{"credential_name": "definitely-not-there", "credential_values": {"api_key": "sk-x"}, "credential_info": {}},
)
assert response.status_code == 404, f"rejected write answered {response.status_code}: {response.text}"
assert "error" in response.json()
def test_update_credential_answers_500_when_the_database_is_not_connected():
def test_update_credential_answers_500_when_the_database_is_not_connected(credential_store):
"""The other rejection this handler raises must carry its own status too."""
with patch("litellm.proxy.proxy_server.prisma_client", None):
response = _patch_credential(
"any-name",
{"credential_name": "any-name", "credential_values": {"api_key": "sk-x"}, "credential_info": {}},
)
credential_store(connected=False)
response = _patch_credential(
"any-name",
{"credential_name": "any-name", "credential_values": {"api_key": "sk-x"}, "credential_info": {}},
)
assert response.status_code == 500, f"rejected write answered {response.status_code}: {response.text}"
def test_update_credential_still_answers_200_on_a_successful_write():
def test_update_credential_still_answers_200_on_a_successful_write(credential_store):
"""The fix must not turn a legitimate update into an error; the dashboard and the
Playwright credentials spec both assert the success path."""
stored = CredentialItem(
@ -149,40 +166,19 @@ def test_update_credential_still_answers_200_on_a_successful_write():
credential_values={"api_key": "sk-old"},
credential_info={"custom_llm_provider": "openai"},
)
with (
patch( # test-quality-ok: the proxy wiring under test is what this patches
"litellm.proxy.proxy_server.prisma_client", MagicMock()
),
patch( # test-quality-ok: the proxy wiring under test is what this patches
"litellm.proxy.proxy_server.master_key", "sk-test-master"
),
patch( # test-quality-ok: the proxy wiring under test is what this patches
"litellm.proxy.credential_endpoints.endpoints.CredentialsRepository"
) as repository, # test-quality-ok: the proxy wiring under test is what this patches
):
repository.return_value.find_by_name = AsyncMock(return_value=stored)
repository.return_value.update_by_name = AsyncMock(return_value=None)
credential_store(find_by_name=AsyncMock(return_value=stored), update_by_name=AsyncMock(return_value=None))
response = _patch_credential(
"existing",
{"credential_name": "existing", "credential_values": {"api_key": "sk-new"}, "credential_info": {}},
)
response = _patch_credential(
"existing",
{"credential_name": "existing", "credential_values": {"api_key": "sk-new"}, "credential_info": {}},
)
assert response.status_code == 200, response.text
assert response.json()["success"] is True
def _get_jwks(name: str):
missing = object()
previous_override = app.dependency_overrides.get(user_api_key_auth, missing)
app.dependency_overrides[user_api_key_auth] = _as_admin
try:
return client.get(f"/credentials/{name}/jwks", headers={"Authorization": "Bearer test-key"})
finally:
if previous_override is missing:
app.dependency_overrides.pop(user_api_key_auth, None)
else:
app.dependency_overrides[user_api_key_auth] = previous_override
return _call_as("GET", f"/credentials/{name}/jwks")
@pytest.fixture
@ -572,19 +568,6 @@ class TestNonAdminCannotPersistWifFieldsOnCredential:
update_mock.assert_awaited_once()
def _delete_credential(name: str, auth=_as_admin):
missing = object()
previous_override = app.dependency_overrides.get(user_api_key_auth, missing)
app.dependency_overrides[user_api_key_auth] = auth
try:
return client.delete(f"/credentials/{name}", headers={"Authorization": "Bearer test-key"})
finally:
if previous_override is missing:
app.dependency_overrides.pop(user_api_key_auth, None)
else:
app.dependency_overrides[user_api_key_auth] = previous_override
def _wif_credential(name: str = "federated-cred") -> CredentialItem:
return CredentialItem(
credential_name=name,
@ -803,13 +786,17 @@ class TestNonAdminCannotTouchAStoredWifCredential:
assert litellm.credential_list == [config_credential]
def test_proxy_admin_can_delete_a_config_only_wif_credential(self, restore_credential_list, monkeypatch):
monkeypatch.setattr(litellm, "credential_list", [_wif_credential("config-wif")])
"""The gate lets the admin through to the row delete. The 404 that follows is the rule for
every config-only credential (no row to delete, the entry is back on the next boot), so the
in-memory entry stays put too."""
config_credential = _wif_credential("config-wif")
monkeypatch.setattr(litellm, "credential_list", [config_credential])
with _repository_holding(None) as repository:
response = _delete_credential("config-wif", auth=_as_admin)
assert response.status_code == 200, response.text
assert response.status_code == 404, response.text
repository.delete_by_name.assert_awaited_once_with("config-wif")
assert litellm.credential_list == []
assert litellm.credential_list == [config_credential]
def test_non_admin_cannot_shadow_a_config_only_wif_credential(self, restore_credential_list, monkeypatch):
"""POST with the same name carries no WIF field and collides with no DB row, yet
@ -996,3 +983,87 @@ class TestManagementReadsTheStoredCredential:
assert resolved is not None
assert resolved.credential_values["anthropic_issuer_url"] == "https://configured.example.com"
def test_delete_credential_answers_404_when_the_credential_does_not_exist(credential_store):
"""Regression: prisma's ``delete`` hands back None when the ``where`` clause matched no row
instead of raising, and the handler never looked. Deleting a name that was never stored
answered 200 "Credential deleted successfully", so an operator scripting cleanup could not
tell a real deletion from a typo."""
credential_store(delete_by_name=AsyncMock(return_value=None))
response = _delete_credential("definitely-not-there")
assert response.status_code == 404, f"delete of a missing credential answered {response.status_code}: {response.text}"
assert "definitely-not-there" in response.text
def test_delete_credential_still_answers_200_and_drops_the_credential_from_memory(credential_store):
"""The fix must not turn a real deletion into an error, and the deleted credential must
stop being served from the in-memory list the proxy routes on."""
stored = CredentialItem(
credential_name="doomed",
credential_values={"api_key": "sk-old"},
credential_info={"custom_llm_provider": "openai"},
)
survivor = CredentialItem(
credential_name="keeper",
credential_values={"api_key": "sk-keep"},
credential_info={},
)
credential_store(in_memory=(stored, survivor), delete_by_name=AsyncMock(return_value=MagicMock()))
response = _delete_credential("doomed")
assert response.status_code == 200, response.text
assert response.json()["success"] is True
assert [credential.credential_name for credential in litellm.credential_list] == ["keeper"]
def test_delete_credential_leaves_a_credential_that_only_exists_in_memory_in_place(credential_store):
"""A credential declared in the config yaml is never written to the table, so the delete
matches no row. Reporting success would be the same lie: it comes straight back on the next
proxy boot. ``PATCH /credentials/{name}`` already answers 404 for that credential."""
config_only = CredentialItem(
credential_name="from-config-yaml",
credential_values={"api_key": "sk-config"},
credential_info={},
)
credential_store(in_memory=(config_only,), delete_by_name=AsyncMock(return_value=None))
response = _delete_credential("from-config-yaml")
assert response.status_code == 404, response.text
assert [credential.credential_name for credential in litellm.credential_list] == ["from-config-yaml"]
def test_delete_credential_answers_500_when_the_database_is_not_connected(credential_store):
"""The handler used to ``return handle_exception_on_proxy(e)``, which makes the exception the
response body and lets FastAPI answer 200. A DB-less proxy answered its own 500 as a success."""
credential_store(connected=False)
response = _delete_credential("any-name")
assert response.status_code == 500, f"rejected delete answered {response.status_code}: {response.text}"
class _CredentialThatCannotBeMasked:
"""Stands in for anything that fails while ``GET /credentials`` builds its response."""
credential_name = "unreadable"
credential_info: dict = {}
@property
def credential_values(self):
raise RuntimeError("credential store unreadable")
def test_get_credentials_answers_an_error_status_when_the_listing_fails(credential_store):
"""Same ``return`` instead of ``raise`` on the list route: a failed listing was serialized as
a 200 whose body happened to be an error, so a caller reading the status saw an empty success."""
credential_store(in_memory=(_CredentialThatCannotBeMasked(),))
response = _list_credentials()
assert response.status_code == 500, f"failed listing answered {response.status_code}: {response.text}"
assert response.json().get("success") is not True

View file

@ -670,6 +670,37 @@ def test_get_provider_specific_params():
) # Literal type should be select
@pytest.mark.asyncio
async def test_provider_specific_params_includes_hide_secrets():
"""hide-secrets lives in the enterprise package so it is not in
guardrail_class_registry; the endpoint must still advertise it or the
Add Guardrail UI dropdown never offers it (LIT-3548)."""
from litellm.proxy.guardrails.guardrail_endpoints import (
get_provider_specific_params,
)
provider_params = await get_provider_specific_params()
assert "hide-secrets" in provider_params
# populateGuardrailProviders() in the dashboard only lists providers whose
# entry carries a ui_friendly_name.
assert provider_params["hide-secrets"]["ui_friendly_name"] == "Hide Secrets"
assert provider_params["hide-secrets"]["detect_secrets_config"]["required"] is False
@pytest.mark.asyncio
async def test_add_guardrail_settings_restricts_hide_secrets_to_pre_call():
"""hide-secrets only implements async_pre_call_hook, so offering the other
modes in the UI would create configs that boot clean and never run."""
from litellm.proxy.guardrails.guardrail_endpoints import (
get_guardrail_ui_settings,
)
settings = await get_guardrail_ui_settings()
assert settings.supported_modes_by_provider["hide-secrets"] == ["pre_call"]
def test_optional_params_not_returned_when_not_overridden():
"""Test that optional_params is not returned when the config model doesn't override it"""
from typing import Optional

View file

@ -17505,6 +17505,16 @@ def test_generate_key_request_blank_team_id_is_personal():
assert GenerateKeyRequest(team_id="team-1").team_id == "team-1"
def test_key_request_blank_organization_id_is_unset():
from litellm.proxy._types import RegenerateKeyRequest, UpdateKeyRequest
assert GenerateKeyRequest(organization_id="").organization_id is None
assert RegenerateKeyRequest(organization_id="").organization_id is None
assert UpdateKeyRequest(key="sk-1", organization_id="").organization_id is None
assert GenerateKeyRequest(organization_id="org-1").organization_id == "org-1"
assert UpdateKeyRequest(key="sk-1", organization_id="org-1").organization_id == "org-1"
def test_key_generation_check_blank_team_id_uses_personal_permissions(monkeypatch):
"""key_generation_check with team_id="" must take the personal-key path instead
of failing the team lookup with "Unable to find team object" (LIT-3925)."""

View file

@ -1,5 +1,6 @@
import inspect
import asyncio
import contextlib
import json
from typing import Dict, Optional
from unittest.mock import AsyncMock, MagicMock, patch
@ -861,6 +862,7 @@ class TestDeleteModelClearsRouterRegistry:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row)
mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row)
@ -922,6 +924,7 @@ class TestDeleteModelClearsRouterRegistry:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row)
mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row)
@ -2005,6 +2008,7 @@ class TestAddAndDeleteModelLifecycle:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=db_row)
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row)
mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row)
@ -2110,6 +2114,7 @@ class TestDeleteTeamBYOKModelGhost:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row)
mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row)
# After the row delete no team deployment remains -> nothing backs the public name.
@ -2186,6 +2191,7 @@ class TestDeleteTeamBYOKModelGhost:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row)
mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row)
mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
@ -2258,6 +2264,7 @@ class TestDeleteTeamBYOKModelGhost:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=deleted_row)
mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=deleted_row)
# After the deleted replica's row is gone, the sibling still backs the public name.
@ -2333,6 +2340,7 @@ class TestDeleteTeamBYOKModelGhost:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row)
mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row)
mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
@ -2408,6 +2416,7 @@ class TestDeleteTeamBYOKModelGhost:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row)
mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row)
mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
@ -2473,6 +2482,7 @@ class TestDeleteModelTeamAuth:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row)
mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row)
mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
@ -2585,6 +2595,7 @@ class TestDeleteModelTeamAuth:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row)
mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row)
mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
@ -3308,18 +3319,10 @@ class TestPatchModelRowDeletedBeforeWrite:
mock_prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=None)
with (
patch(
"litellm.proxy.proxy_server.prisma_client", mock_prisma
), # test-quality-ok: proxy_server module global is the endpoint's only injection point
patch(
"litellm.proxy.proxy_server.llm_router", MagicMock(**{"get_model_ids.return_value": ["m1"]})
), # test-quality-ok: proxy_server module global is the endpoint's only injection point
patch(
"litellm.proxy.proxy_server.store_model_in_db", True
), # test-quality-ok: proxy_server module global is the endpoint's only injection point
patch(
"litellm.proxy.proxy_server.premium_user", True
), # test-quality-ok: proxy_server module global is the endpoint's only injection point
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: proxy_server module global is the endpoint's only injection point
patch("litellm.proxy.proxy_server.llm_router", MagicMock(**{"get_model_ids.return_value": ["m1"]})), # test-quality-ok: proxy_server module global is the endpoint's only injection point
patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: proxy_server module global is the endpoint's only injection point
patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: proxy_server module global is the endpoint's only injection point
patch( # test-quality-ok: stubs the auth gate so the test exercises the not-found branch under test
"litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call",
new=AsyncMock(return_value=None),
@ -3684,6 +3687,7 @@ class TestDeleteEvictionsHoldTheReconcileLock:
prisma = MagicMock()
prisma.db.litellm_proxymodeltable = table
prisma.db.query_raw = AsyncMock(return_value=[])
router = MagicMock()
router.delete_deployment = MagicMock(return_value=True)
@ -5410,3 +5414,179 @@ class TestBlockModelResponseSerialization:
assert body["model_id"] == "m-block-1"
assert body["blocked"] is blocked
assert body["litellm_params"] == {"model": "openai/gpt-4o-mini", "api_key": "encrypted-value"}
class TestAccessGroupModelSync:
"""A rename or delete of a deployment must land in every unified access group that names it."""
_PS = "litellm.proxy.proxy_server"
_MOD = "litellm.proxy.management_endpoints.model_management_endpoints"
_INVALIDATE = "litellm.proxy.management_helpers.access_group_model_sync.invalidate_access_group_caches"
@staticmethod
def _admin():
return UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin")
@staticmethod
def _prisma_with_row(model_id: str, model_name: str, deployment_count: int):
row = LiteLLM_ProxyModelTable(
model_id=model_id,
model_name=model_name,
litellm_params={"model": "openai/gpt-5.6"},
model_info={"id": model_id},
created_by="admin",
updated_by="admin",
)
async def query_raw(sql, *params):
if sql.startswith("SELECT COUNT(*)"):
return [{"deployment_count": deployment_count}]
return [{"access_group_id": "ag-1"}]
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.query_raw = AsyncMock(side_effect=query_raw)
mock_prisma.db.litellm_proxymodeltable = AsyncMock()
mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=row)
mock_prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=row)
mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=row)
return mock_prisma
@staticmethod
def _access_group_updates(mock_prisma):
return [
call
for call in mock_prisma.db.query_raw.await_args_list
if call.args[0].startswith('UPDATE "LiteLLM_AccessGroupTable"')
]
@contextlib.contextmanager
def _endpoint_env(self, mock_prisma, router):
with contextlib.ExitStack() as stack:
for target in (
patch(f"{self._PS}.prisma_client", mock_prisma),
patch(f"{self._PS}.llm_router", router),
patch(f"{self._PS}.store_model_in_db", True),
patch(f"{self._PS}.premium_user", True),
patch(f"{self._PS}.proxy_logging_obj", MagicMock()),
patch(f"{self._PS}.user_api_key_cache", MagicMock()),
patch(f"{self._MOD}.ModelManagementAuthChecks.can_user_make_model_call", new=AsyncMock(return_value=None)),
patch(
f"{self._MOD}.clear_cache",
new=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None)),
),
patch(f"{self._MOD}.encrypt_value_helper", side_effect=lambda value, **kwargs: value),
):
stack.enter_context(target)
yield stack.enter_context(patch(self._INVALIDATE, new=AsyncMock()))
@pytest.mark.asyncio
async def test_patch_model_rename_rewrites_the_groups_that_named_the_model(self):
from litellm.proxy.management_endpoints.model_management_endpoints import patch_model
mock_prisma = self._prisma_with_row("m-rename", "gpt-5.6", deployment_count=0)
router = MagicMock()
router.get_model_ids.return_value = ["m-rename"]
with self._endpoint_env(mock_prisma, router) as invalidate:
await patch_model(
model_id="m-rename",
patch_data=updateDeployment(model_name="gpt-5.6-eu"),
user_api_key_dict=self._admin(),
)
written = mock_prisma.db.litellm_proxymodeltable.update.await_args.kwargs["data"]
assert written["model_name"] == "gpt-5.6-eu"
(update_call,) = self._access_group_updates(mock_prisma)
assert "array_replace" in update_call.args[0]
assert update_call.args[1:] == ("gpt-5.6", "gpt-5.6-eu")
invalidate.assert_awaited_once_with(("ag-1",))
@pytest.mark.asyncio
async def test_patch_model_rename_appends_when_a_sibling_deployment_keeps_the_old_name(self):
from litellm.proxy.management_endpoints.model_management_endpoints import patch_model
mock_prisma = self._prisma_with_row("m-rename", "gpt-5.6", deployment_count=1)
router = MagicMock()
router.get_model_ids.return_value = ["m-rename"]
with self._endpoint_env(mock_prisma, router):
await patch_model(
model_id="m-rename",
patch_data=updateDeployment(model_name="gpt-5.6-eu"),
user_api_key_dict=self._admin(),
)
(update_call,) = self._access_group_updates(mock_prisma)
assert "array_append" in update_call.args[0]
assert update_call.args[1:] == ("gpt-5.6", "gpt-5.6-eu")
@pytest.mark.asyncio
async def test_patch_model_without_a_rename_leaves_access_groups_alone(self):
from litellm.proxy.management_endpoints.model_management_endpoints import patch_model
mock_prisma = self._prisma_with_row("m-same", "gpt-5.6", deployment_count=0)
router = MagicMock()
router.get_model_ids.return_value = ["m-same"]
with self._endpoint_env(mock_prisma, router) as invalidate:
await patch_model(model_id="m-same", patch_data=updateDeployment(blocked=True), user_api_key_dict=self._admin())
mock_prisma.db.query_raw.assert_not_awaited()
invalidate.assert_not_awaited()
@pytest.mark.asyncio
async def test_delete_model_drops_the_name_from_groups_when_nothing_backs_it(self):
from litellm.proxy.management_endpoints.model_management_endpoints import ModelInfoDelete, delete_model
mock_prisma = self._prisma_with_row("m-doomed", "gpt-5.6", deployment_count=0)
router = MagicMock()
router.get_model_ids.return_value = []
with self._endpoint_env(mock_prisma, router) as invalidate:
await delete_model(model_info=ModelInfoDelete(id="m-doomed"), user_api_key_dict=self._admin())
(update_call,) = self._access_group_updates(mock_prisma)
assert "array_remove" in update_call.args[0]
assert update_call.args[1:] == ("gpt-5.6",)
invalidate.assert_awaited_once_with(("ag-1",))
@pytest.mark.asyncio
async def test_delete_model_keeps_the_name_while_a_sibling_deployment_backs_it(self):
from litellm.proxy.management_endpoints.model_management_endpoints import ModelInfoDelete, delete_model
mock_prisma = self._prisma_with_row("m-doomed", "gpt-5.6", deployment_count=1)
router = MagicMock()
router.get_model_ids.return_value = []
with self._endpoint_env(mock_prisma, router) as invalidate:
await delete_model(model_info=ModelInfoDelete(id="m-doomed"), user_api_key_dict=self._admin())
assert self._access_group_updates(mock_prisma) == []
invalidate.assert_not_awaited()
@pytest.mark.asyncio
async def test_update_model_persists_a_new_model_name_and_rewrites_the_groups(self):
from litellm.proxy.management_endpoints.model_management_endpoints import update_model
from litellm.types.router import ModelInfo, updateLiteLLMParams
mock_prisma = self._prisma_with_row("m-terraform", "gpt-5.6", deployment_count=0)
router = MagicMock()
router.get_model_ids.return_value = ["m-terraform"]
with self._endpoint_env(mock_prisma, router) as invalidate:
await update_model(
model_params=updateDeployment(
model_name="gpt-5.6-eu",
litellm_params=updateLiteLLMParams(model="openai/gpt-5.6"),
model_info=ModelInfo(id="m-terraform"),
),
user_api_key_dict=self._admin(),
)
written = mock_prisma.db.litellm_proxymodeltable.update.await_args.kwargs["data"]
assert written["model_name"] == "gpt-5.6-eu"
(update_call,) = self._access_group_updates(mock_prisma)
assert "array_replace" in update_call.args[0]
assert update_call.args[1:] == ("gpt-5.6", "gpt-5.6-eu")
invalidate.assert_awaited_once_with(("ag-1",))

View file

@ -516,6 +516,39 @@ async def test_team_ids_extracted_from_groups_attribute(saml_env_idp_initiated):
assert result.team_ids == ["team-a", "team-b"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"roles",
[
["internal_user", "proxy_admin_viewer"],
["proxy_admin_viewer", "internal_user"],
],
)
async def test_multi_valued_role_attribute_resolves_to_highest_privilege(saml_env_idp_initiated, roles):
"""An assertion carrying several roles must not depend on the order the IdP emitted them in."""
key_pem, cert_pem = saml_env_idp_initiated
resp = _build_signed_response(
key_pem,
cert_pem,
attributes={
"email": ["dave@example.com"],
"role": roles,
},
)
result = await _acs(_b64(resp), _shared_cache())
assert result.user_role == LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY
@pytest.mark.asyncio
async def test_assertion_without_role_attribute_has_no_user_role(saml_env_idp_initiated):
key_pem, cert_pem = saml_env_idp_initiated
resp = _build_signed_response(key_pem, cert_pem, attributes={"email": ["erin@example.com"]})
result = await _acs(_b64(resp), _shared_cache())
assert result.user_role is None
@pytest.mark.asyncio
async def test_build_login_redirect_targets_idp_and_caches_request_id(saml_env):
cache = DualCache()

View file

@ -6598,13 +6598,94 @@ def test_get_litellm_user_role_with_invalid_role():
assert result is None
def test_get_litellm_user_role_with_list_multiple_roles():
"""Test that get_litellm_user_role takes the first element from a multi-element list."""
@pytest.mark.parametrize(
"role_claim",
[
["proxy_admin", "internal_user"],
["internal_user", "proxy_admin"],
],
)
def test_get_litellm_user_role_picks_highest_privilege_regardless_of_order(role_claim):
"""A multi-valued role claim resolves to the most privileged role, not the first one listed."""
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.types import get_litellm_user_role
result = get_litellm_user_role(["proxy_admin", "internal_user"])
assert result == LitellmUserRoles.PROXY_ADMIN
assert get_litellm_user_role(role_claim) == LitellmUserRoles.PROXY_ADMIN
@pytest.mark.parametrize(
"role_claim",
[
["proxy_admin_viewer", "internal_user"],
["internal_user", "proxy_admin_viewer"],
],
)
def test_get_litellm_user_role_keeps_org_spend_visibility_for_mixed_roles(role_claim):
"""
Regression for LIT-6077: a user holding both proxy_admin_viewer and internal_user kept
losing org-level spend visibility whenever the IdP happened to list internal_user first.
"""
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.types import get_litellm_user_role
assert get_litellm_user_role(role_claim) == LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY
def test_get_litellm_user_role_ignores_unrecognised_entries():
"""Roles LiteLLM does not know about are skipped rather than swallowing the whole claim."""
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.types import get_litellm_user_role
assert get_litellm_user_role(["some_idp_group", "internal_user"]) == LitellmUserRoles.INTERNAL_USER
assert get_litellm_user_role(["some_idp_group", "another_group"]) is None
def test_get_litellm_user_role_list_lookup_is_case_insensitive():
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.types import get_litellm_user_role
assert get_litellm_user_role(["INTERNAL_USER", "Proxy_Admin"]) == LitellmUserRoles.PROXY_ADMIN
@pytest.mark.parametrize(
"role_claim",
[
["org_admin", "team"],
["team", "org_admin"],
],
)
def test_get_litellm_user_role_is_deterministic_for_unranked_roles(role_claim):
"""Roles outside the privilege hierarchy still resolve the same way in either claim order."""
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.types import get_litellm_user_role
assert get_litellm_user_role(role_claim) == LitellmUserRoles.ORG_ADMIN
@pytest.mark.parametrize(
"role_claim",
[
["org_admin", "internal_user"],
["internal_user", "org_admin"],
],
)
def test_get_litellm_user_role_prefers_a_ranked_role_over_an_unranked_one(role_claim):
"""
org_admin, team and customer sit outside the privilege ladder, so a claim mixing one of
them with a ranked role settles on the ranked role in either order. Same rule the Entra
app_roles and role_mappings paths already follow.
"""
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.types import get_litellm_user_role
assert get_litellm_user_role(role_claim) == LitellmUserRoles.INTERNAL_USER
def test_get_litellm_user_role_returns_none_for_non_string_claims():
from litellm.proxy.management_endpoints.types import get_litellm_user_role
assert get_litellm_user_role(None) is None
assert get_litellm_user_role({"role": "proxy_admin"}) is None
# ============================================================================
@ -6654,6 +6735,46 @@ def test_process_sso_jwt_access_token_extracts_role_from_access_token():
assert result.user_role == LitellmUserRoles.PROXY_ADMIN
@pytest.mark.parametrize(
"role_claim",
[
["internal_user", "proxy_admin_viewer"],
["proxy_admin_viewer", "internal_user"],
],
)
def test_process_sso_jwt_access_token_resolves_highest_privilege_role(role_claim):
"""
The generic SSO access-token path must land on the same role for a user whose role
claim holds several roles, whichever order the IdP emitted them in.
"""
import jwt as pyjwt
from litellm.proxy._types import LitellmUserRoles
access_token_str = pyjwt.encode(
{"sub": "user-123", "email": "mixed@test.com", "litellm_role": role_claim},
"secret",
algorithm="HS256",
)
result = CustomOpenID(
id="user-123",
email="mixed@test.com",
display_name="Mixed Role User",
team_ids=[],
user_role=None,
)
with patch.dict(os.environ, {"GENERIC_USER_ROLE_ATTRIBUTE": "litellm_role"}):
process_sso_jwt_access_token(
access_token_str=access_token_str,
sso_jwt_handler=None,
result=result,
role_mappings=None,
)
assert result.user_role == LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY
def test_process_sso_jwt_access_token_does_not_override_existing_role():
"""
Test that process_sso_jwt_access_token does NOT override a role that was

View file

@ -0,0 +1,170 @@
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.proxy.db.prisma_client import PrismaWrapper
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
from litellm.proxy.management_helpers.access_group_model_sync import (
sync_access_groups_for_deleted_model,
sync_access_groups_for_renamed_model,
)
_INVALIDATE = "litellm.proxy.management_helpers.access_group_model_sync.invalidate_access_group_caches"
def _routed_prisma_client(deployment_count: int):
async def query_raw(sql, *params):
if sql.startswith("SELECT COUNT(*)"):
return [{"deployment_count": deployment_count}]
return [{"access_group_id": "ag-1"}, {"access_group_id": "ag-2"}]
writer_inner = MagicMock(name="writer_prisma")
reader_inner = MagicMock(name="reader_prisma")
writer_inner.query_raw = AsyncMock(side_effect=query_raw)
reader_inner.query_raw = AsyncMock(side_effect=query_raw)
writer = PrismaWrapper(original_prisma=writer_inner, iam_token_db_auth=False)
reader = PrismaWrapper(original_prisma=reader_inner, iam_token_db_auth=False)
routing = RoutingPrismaWrapper(writer=writer, reader=reader)
return SimpleNamespace(db=routing), writer_inner, reader_inner
def _access_group_updates(writer_inner):
return [
call
for call in writer_inner.query_raw.await_args_list
if call.args[0].startswith('UPDATE "LiteLLM_AccessGroupTable"')
]
@pytest.mark.asyncio
async def test_rename_replaces_the_old_name_when_no_other_deployment_carries_it():
prisma_client, writer_inner, reader_inner = _routed_prisma_client(deployment_count=0)
with patch(_INVALIDATE, new=AsyncMock()) as invalidate:
await sync_access_groups_for_renamed_model(
prisma_client, model_id="m-1", old_name="gpt-5.6", new_name="gpt-5.6-eu", llm_router=None
)
(update_call,) = _access_group_updates(writer_inner)
assert "array_replace" in update_call.args[0]
assert update_call.args[1:] == ("gpt-5.6", "gpt-5.6-eu")
invalidate.assert_awaited_once_with(("ag-1", "ag-2"))
reader_inner.query_raw.assert_not_awaited()
@pytest.mark.asyncio
async def test_rename_appends_the_new_name_when_a_sibling_row_keeps_the_old_one():
prisma_client, writer_inner, _ = _routed_prisma_client(deployment_count=1)
with patch(_INVALIDATE, new=AsyncMock()) as invalidate:
await sync_access_groups_for_renamed_model(
prisma_client, model_id="m-1", old_name="gpt-5.6", new_name="gpt-5.6-eu", llm_router=None
)
(update_call,) = _access_group_updates(writer_inner)
assert "array_append" in update_call.args[0]
assert "array_replace" not in update_call.args[0]
assert update_call.args[1:] == ("gpt-5.6", "gpt-5.6-eu")
invalidate.assert_awaited_once_with(("ag-1", "ag-2"))
def _router_serving(db_model_by_deployment_id: dict[str, bool]):
llm_router = MagicMock()
llm_router.get_model_ids.return_value = list(db_model_by_deployment_id)
llm_router.get_deployment.side_effect = lambda model_id: SimpleNamespace(
model_info=SimpleNamespace(db_model=db_model_by_deployment_id[model_id])
)
return llm_router
@pytest.mark.asyncio
@pytest.mark.parametrize(
"db_model_by_deployment_id, expected_write",
[
({"m-1": True}, "array_replace"),
({"m-1": True, "m-from-config": False}, "array_append"),
({"m-1": True, "m-db-sibling-this-worker-has-not-refreshed": True}, "array_replace"),
],
)
async def test_rename_counts_only_config_deployments_with_another_id_as_backing_the_old_name(
db_model_by_deployment_id, expected_write
):
prisma_client, writer_inner, _ = _routed_prisma_client(deployment_count=0)
llm_router = _router_serving(db_model_by_deployment_id)
with patch(_INVALIDATE, new=AsyncMock()):
await sync_access_groups_for_renamed_model(
prisma_client, model_id="m-1", old_name="gpt-5.6", new_name="gpt-5.6-eu", llm_router=llm_router
)
llm_router.get_model_ids.assert_called_once_with(model_name="gpt-5.6")
(update_call,) = _access_group_updates(writer_inner)
assert expected_write in update_call.args[0]
@pytest.mark.asyncio
async def test_delete_ignores_a_db_sibling_this_worker_has_not_refreshed_yet():
prisma_client, writer_inner, _ = _routed_prisma_client(deployment_count=0)
llm_router = _router_serving({"m-1": True, "m-renamed-elsewhere": True})
with patch(_INVALIDATE, new=AsyncMock()) as invalidate:
await sync_access_groups_for_deleted_model(
prisma_client, model_id="m-1", model_name="gpt-5.6", llm_router=llm_router
)
(update_call,) = _access_group_updates(writer_inner)
assert "array_remove" in update_call.args[0]
invalidate.assert_awaited_once_with(("ag-1", "ag-2"))
@pytest.mark.asyncio
async def test_delete_keeps_the_name_while_a_config_deployment_still_serves_it():
prisma_client, writer_inner, _ = _routed_prisma_client(deployment_count=0)
llm_router = _router_serving({"m-1": True, "m-from-config": False})
with patch(_INVALIDATE, new=AsyncMock()) as invalidate:
await sync_access_groups_for_deleted_model(
prisma_client, model_id="m-1", model_name="gpt-5.6", llm_router=llm_router
)
assert _access_group_updates(writer_inner) == []
invalidate.assert_not_awaited()
@pytest.mark.asyncio
async def test_rename_to_the_same_name_writes_nothing():
prisma_client, writer_inner, _ = _routed_prisma_client(deployment_count=0)
with patch(_INVALIDATE, new=AsyncMock()) as invalidate:
await sync_access_groups_for_renamed_model(
prisma_client, model_id="m-1", old_name="gpt-5.6", new_name="gpt-5.6", llm_router=None
)
assert _access_group_updates(writer_inner) == []
invalidate.assert_not_awaited()
@pytest.mark.asyncio
async def test_delete_removes_the_name_when_no_row_backs_it_any_more():
prisma_client, writer_inner, reader_inner = _routed_prisma_client(deployment_count=0)
with patch(_INVALIDATE, new=AsyncMock()) as invalidate:
await sync_access_groups_for_deleted_model(prisma_client, model_id="m-1", model_name="gpt-5.6", llm_router=None)
(update_call,) = _access_group_updates(writer_inner)
assert "array_remove" in update_call.args[0]
assert update_call.args[1:] == ("gpt-5.6",)
invalidate.assert_awaited_once_with(("ag-1", "ag-2"))
reader_inner.query_raw.assert_not_awaited()
@pytest.mark.asyncio
async def test_delete_keeps_the_name_while_a_sibling_row_still_backs_it():
prisma_client, writer_inner, _ = _routed_prisma_client(deployment_count=2)
with patch(_INVALIDATE, new=AsyncMock()) as invalidate:
await sync_access_groups_for_deleted_model(prisma_client, model_id="m-1", model_name="gpt-5.6", llm_router=None)
assert _access_group_updates(writer_inner) == []
invalidate.assert_not_awaited()

View file

@ -4922,19 +4922,21 @@ class TestAnthropicProxyRouteCallerAuthHeaders:
@pytest.mark.asyncio
async def test_wif_credential_drops_configured_custom_key_header(self, monkeypatch):
from litellm.proxy import proxy_server
self._clear_anthropic_env(monkeypatch)
self._enable_wif(monkeypatch)
monkeypatch.setitem(proxy_server.general_settings, "litellm_key_header_name", "X-Tenant-Key")
with patch.dict("litellm.proxy.proxy_server.general_settings", {"litellm_key_header_name": "X-Tenant-Key"}):
sent: Final = await self._upstream_headers(
self._request(
{
"content-type": "application/json",
"x-tenant-key": "sk-caller-virtual-key",
"x-tenant-region": "eu",
}
)
sent: Final = await self._upstream_headers(
self._request(
{
"content-type": "application/json",
"x-tenant-key": "sk-caller-virtual-key",
"x-tenant-region": "eu",
}
)
)
assert sent["authorization"] == f"Bearer {self._MINTED}"
assert "x-tenant-key" not in sent

View file

@ -1271,6 +1271,16 @@ def test_get_autorouter_presets_local_mode_serves_bundled_catalog(
assert response.status_code == 200
payload = response.json()
assert "anthropic_family" in payload
assert payload["1m_context"]["complexity_router_config"]["classifier_type"] == "heuristic_v2"
assert payload["1m_context"]["complexity_router_config"]["tiers"] == {
"SIMPLE": ["gpt-5.6-luna"],
"MEDIUM": ["gpt-5.6-terra"],
"COMPLEX": ["claude-opus-5"],
"REASONING": ["claude-opus-5"],
}
assert payload["1m_context"]["complexity_router_config"]["tier_model_configs"] == {
"REASONING": [{"model_name": "claude-opus-5", "litellm_params": {"reasoning_effort": "high"}}]
}
for preset in payload.values():
assert isinstance(preset["label"], str)
assert isinstance(preset["description"], str)

View file

@ -3971,6 +3971,7 @@ async def test_build_ui_spend_logs_response_dict_rows_session_counts():
"mcp_tool_call_spend": 10.0,
"session_llm_count": 1,
"session_agent_count": 0,
"session_models": ["claude-haiku-4-5", "gpt-5.4-nano"],
}
]
)
@ -3997,6 +3998,7 @@ async def test_build_ui_spend_logs_response_dict_rows_session_counts():
assert rows[1]["mcp_tool_call_spend"] == 10.0
assert rows[0]["session_llm_count"] == 1
assert rows[0]["session_agent_count"] == 0
assert rows[0]["session_models"] == ["claude-haiku-4-5", "gpt-5.4-nano"]
# Every row in the session carries the full session spend, not just its own
assert rows[0]["session_total_spend"] == 15.0
@ -4004,11 +4006,60 @@ async def test_build_ui_spend_logs_response_dict_rows_session_counts():
# Row without a session_id defaults to 1
assert rows[2]["session_total_count"] == 1
assert "session_models" not in rows[2]
# The count is folded into the single aggregate query; no separate group_by call.
mock_prisma.db.litellm_spendlogs.group_by.assert_not_called()
@pytest.mark.asyncio
async def test_build_ui_spend_logs_response_caps_session_models():
"""The per-session model list is bounded server-side and flags when it was cut."""
from litellm.proxy.spend_tracking.spend_management_endpoints import (
_SESSION_MODELS_LIMIT,
_build_ui_spend_logs_response,
)
session_id = "sess-many-models"
api_key = "hashed-key-xyz"
over_limit_models = [f"model-{i:02d}" for i in range(_SESSION_MODELS_LIMIT + 1)]
mock_prisma = MagicMock()
mock_prisma.db.query_raw = AsyncMock(
return_value=[
{
"session_id": session_id,
"api_key": api_key,
"session_total_count": len(over_limit_models),
"session_total_spend": 1.0,
"mcp_tool_call_count": 0,
"mcp_tool_call_spend": 0.0,
"session_llm_count": len(over_limit_models),
"session_agent_count": 0,
"session_models": over_limit_models,
}
]
)
result = await _build_ui_spend_logs_response(
prisma_client=mock_prisma,
data=[{"request_id": "req-1", "session_id": session_id, "call_type": "completion", "api_key": api_key}],
total_records=1,
page=1,
page_size=50,
total_pages=1,
enrich_session_counts=True,
)
row = result["data"][0]
assert row["session_models"] == over_limit_models[:_SESSION_MODELS_LIMIT]
assert row["session_models_truncated"] is True
sql, *params = mock_prisma.db.query_raw.await_args.args
assert "LIMIT $4" in sql
assert params[3] == _SESSION_MODELS_LIMIT + 1
@pytest.mark.asyncio
async def test_build_ui_spend_logs_response_key_split_session_gets_per_key_aggregates():
"""

View file

@ -1558,6 +1558,14 @@ class TestLLMClassifierConfig:
assert config.classifier_type == "heuristic"
assert config.classifier_llm_config is None
@pytest.mark.parametrize("reasoning_effort", ["", "ultra"])
def test_classifier_reasoning_effort_rejects_unsupported_values(self, reasoning_effort):
with pytest.raises(ValidationError):
ComplexityRouterConfig(
classifier_type="llm",
classifier_llm_config={"model": "haiku-classifier", "reasoning_effort": reasoning_effort},
)
CUSTOM_TIER_LABELS: Dict[str, str] = {
"SIMPLE": "Cheap",
@ -1871,6 +1879,19 @@ class TestLLMClassifier:
call_kwargs = mock_router_instance.acompletion.call_args.kwargs
assert call_kwargs["metadata"] == {**request_metadata, "internal_call_origin": "autorouter_classifier"}
@pytest.mark.asyncio
async def test_aclassify_stamps_internal_origin_without_caller_metadata(
self, llm_complexity_router, mock_router_instance
):
"""Fallback handling must still recognize the classifier when an SDK caller supplied no metadata."""
mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "SIMPLE"}'))
await llm_complexity_router.aclassify("hi")
assert mock_router_instance.acompletion.call_args.kwargs["metadata"] == {
"internal_call_origin": "autorouter_classifier"
}
@pytest.mark.asyncio
@pytest.mark.parametrize(
"request_kwargs",
@ -1948,6 +1969,33 @@ class TestLLMClassifier:
"REASONING",
]
@pytest.mark.asyncio
@pytest.mark.parametrize("reasoning_effort", [None, "none", "low"], ids=["omitted", "none", "low"])
async def test_classifier_reasoning_effort_reaches_only_classifier_call(
self, mock_router_instance, llm_classifier_config, reasoning_effort
):
classifier_llm_config = {
**llm_classifier_config["classifier_llm_config"],
**({"reasoning_effort": reasoning_effort} if reasoning_effort is not None else {}),
}
router = ComplexityRouter(
model_name="test-complexity-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={**llm_classifier_config, "classifier_llm_config": classifier_llm_config},
)
mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "COMPLEX"}'))
await router.aclassify("explain quantum tunneling in depth")
call_kwargs = mock_router_instance.acompletion.call_args.kwargs
body = call_kwargs["proxy_server_request"]["body"]
if reasoning_effort is None:
assert "reasoning_effort" not in call_kwargs
assert "reasoning_effort" not in body
else:
assert call_kwargs["reasoning_effort"] == reasoning_effort
assert body["reasoning_effort"] == reasoning_effort
@pytest.mark.asyncio
async def test_aclassify_propagates_top_level_turn_off_message_logging(
self, llm_complexity_router, mock_router_instance
@ -8735,9 +8783,10 @@ class TestClassificationRubrics:
[
{"model": "haiku-classifier", "system_prompt": "Grade the data sensitivity of the request."},
{"model": "haiku-classifier", "classification_rubric": "chat"},
{"model": "haiku-classifier", "reasoning_effort": "low"},
{"model": "haiku-classifier"},
],
ids=["custom-prompt", "chat-preset", "neither"],
ids=["custom-prompt", "chat-preset", "reasoning-effort", "neither"],
)
def test_config_survives_a_dump_and_rebuild(self, classifier_llm_config):
"""/auto_router/test_routing dumps this config and hands the dict straight back to

View file

@ -3131,9 +3131,9 @@ def test_stream_chunk_builder_prices_proxy_alias_via_model_map():
assert response._hidden_params["response_cost"] == pytest.approx(expected_cost)
def _stream_builder_logging_obj() -> LiteLLMLogging:
def _stream_builder_logging_obj(model: str = "gpt-4o", custom_llm_provider: str = "openai") -> LiteLLMLogging:
logging_obj: Final = LiteLLMLogging(
model="gpt-4o",
model=model,
messages=[{"role": "user", "content": "hi"}],
stream=True,
call_type="completion",
@ -3142,10 +3142,11 @@ def _stream_builder_logging_obj() -> LiteLLMLogging:
function_id="test-function-id",
)
logging_obj.update_environment_variables(
model="gpt-4o",
model=model,
user=None,
optional_params={},
litellm_params={"custom_llm_provider": "openai"},
litellm_params={"custom_llm_provider": custom_llm_provider},
custom_llm_provider=custom_llm_provider,
)
return logging_obj
@ -3237,3 +3238,24 @@ def test_stream_chunk_builder_prices_alias_from_openai_sdk_usage_chunk():
assert response.usage.completion_tokens == 60
assert getattr(response.usage, "cost", None) == pytest.approx(0.000704)
assert response._hidden_params["response_cost"] == pytest.approx(0.000704)
def test_stream_chunk_builder_leaves_xai_reported_cost_to_the_calculator(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(litellm, "cost_margin_config", {"xai": 0.5})
usage_chunk: Final = _stream_builder_text_chunk("grok-4", "")
usage_chunk.usage = Usage(prompt_tokens=5, completion_tokens=2, total_tokens=7, cost=0.42)
chunks: Final = [
_stream_builder_text_chunk("grok-4", "Hello "),
_stream_builder_text_chunk("grok-4", "world.", finish_reason="stop"),
usage_chunk,
]
logging_obj: Final = _stream_builder_logging_obj(model="grok-4", custom_llm_provider="xai")
response: Final = litellm.stream_chunk_builder(
chunks=chunks, messages=[{"role": "user", "content": "hi"}], logging_obj=logging_obj
)
assert response is not None
assert getattr(response.usage, "cost", None) == pytest.approx(0.42)
assert response._hidden_params.get("response_cost") is None
assert logging_obj._response_cost_calculator(result=response) == pytest.approx(0.63)

View file

@ -35,6 +35,7 @@ from litellm.router import (
_anthropic_stream_should_drop_pre_content_ping,
_is_retriable_anthropic_status,
)
from litellm.types.router import DeploymentTypedDict
def test_update_kwargs_does_not_mutate_defaults_and_merges_metadata():
@ -9817,11 +9818,12 @@ def test_model_group_info_intersects_supported_reasoning_efforts():
assert result.supported_reasoning_efforts == ("minimal", "low", "medium", "high")
def test_model_group_info_reasoning_efforts_ignore_a_deployment_off_the_map():
def test_model_group_info_reasoning_efforts_are_unknown_when_any_deployment_is_off_the_map():
"""The router fills every ModelInfo key, so a deployment absent from the model map arrives with
supports_reasoning None rather than with the key missing. Its synthesized entry carries no mode,
which is what separates it from a mapped non-reasoning model, and nothing being known about it is
no reason to drop the levels the rest of the group agrees on."""
no evidence that the unknown deployment accepts levels its mapped sibling supports. The group
therefore reports unknown instead of advertising a value routing might send to either one."""
router = litellm.Router(
model_list=[
{
@ -9855,7 +9857,7 @@ def test_model_group_info_reasoning_efforts_ignore_a_deployment_off_the_map():
)
assert result is not None
assert result.supported_reasoning_efforts == ("none", "minimal", "low", "medium", "high", "max")
assert result.supported_reasoning_efforts is None
@ -9992,11 +9994,11 @@ def test_model_group_info_survives_a_junk_typed_operator_effort_value():
assert result.supported_reasoning_efforts == ("none", "minimal", "low", "medium", "high")
def test_model_group_info_reasoning_efforts_ignore_a_mode_the_operator_declared():
def test_model_group_info_reasoning_efforts_are_unknown_for_an_operator_declared_mode():
"""A deployment is registered in the cost map under its own id with whatever model_info the
operator wrote, so a mode they set themselves reads back exactly like one the map supplied. Only
a mode the map supplied marks the deployment as known, or an off-map deployment carrying any
mode empties the group it sits in."""
a mode the map supplied marks the deployment as known. An off-map deployment carrying an
operator mode remains unknown and must keep the whole group's level support unknown."""
from litellm.router_utils.reasoning_effort_capability import resolve_supported_reasoning_efforts
mapped_model = "openai/gpt-5.6-sol"
@ -10027,7 +10029,7 @@ def test_model_group_info_reasoning_efforts_ignore_a_mode_the_operator_declared(
)
assert result is not None
assert result.supported_reasoning_efforts == expected
assert result.supported_reasoning_efforts is None
class TestAddDeploymentApiBaseProviderResolution:
@ -12311,6 +12313,120 @@ def test_router_keeps_wif_secret_pointers_unresolved(monkeypatch):
assert litellm_params["anthropic_keycloak_client_secret_ref"] == "os.environ/WIF_TEST_KC_SECRET"
class TestRequestReasoningEffortOverride:
def test_drop_effort_from_nested_carrier_preserves_other_nested_values(self):
params: dict[str, object] = {"output_config": {"effort": "high", "format": "json"}}
litellm.Router._pop_effort_from_nested_carrier(params, "output_config")
assert params == {"output_config": {"format": "json"}}
@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"])
def test_is_classifier_internal_call_recognizes_both_metadata_carriers(self, metadata_key):
kwargs = {metadata_key: {"internal_call_origin": "autorouter_classifier"}}
assert litellm.Router._is_classifier_internal_call(kwargs) is True
assert litellm.Router._is_classifier_internal_call({metadata_key: {}}) is False
def test_removes_every_deployment_native_effort_carrier_without_mutating_shared_config(self):
extra_body: dict[str, object] = {
"reasoning_effort": "high",
"thinking": {"type": "enabled"},
"output_config": {"effort": "high", "format": "json"},
"reasoning": {"effort": "high", "summary": "detailed"},
"provider_option": True,
}
deployment_params: dict[str, object] = {
"model": "bedrock/converse/anthropic.claude-3-7-sonnet",
"thinking": {"type": "enabled", "budget_tokens": 2048},
"output_config": {"effort": "high", "format": {"type": "json_schema"}},
"reasoning": {"effort": "high", "summary": "auto"},
"extra_body": extra_body,
}
sanitized = litellm.Router._deployment_params_with_request_reasoning_override(
deployment_params, {"reasoning_effort": "low"}
)
assert sanitized == {
"model": "bedrock/converse/anthropic.claude-3-7-sonnet",
"output_config": {"format": {"type": "json_schema"}},
"reasoning": {"summary": "auto"},
"extra_body": {
"output_config": {"format": "json"},
"reasoning": {"summary": "detailed"},
"provider_option": True,
},
}
assert deployment_params["thinking"] == {"type": "enabled", "budget_tokens": 2048}
assert deployment_params["output_config"] == {"effort": "high", "format": {"type": "json_schema"}}
assert extra_body["reasoning_effort"] == "high"
@pytest.mark.parametrize("request_kwargs", [{}, {"reasoning_effort": None}])
def test_omitted_override_preserves_deployment_defaults(self, request_kwargs):
deployment_params = {
"model": "deepseek/deepseek-reasoner",
"thinking": {"type": "enabled"},
"output_config": {"effort": "high"},
}
assert (
litellm.Router._deployment_params_with_request_reasoning_override(deployment_params, request_kwargs)
== deployment_params
)
@pytest.mark.asyncio
async def test_280_concurrent_overrides_never_mutate_or_leak_through_shared_deployment_params(self):
deployment_params = {
"model": "fireworks_ai/accounts/fireworks/models/kimi-k2-thinking",
"thinking": {"type": "enabled"},
"output_config": {"effort": "high", "format": "json"},
"extra_body": {"reasoning_effort": "high", "tenant": "shared"},
}
efforts = ("none", "minimal", "low", "medium", "high", "xhigh", "max")
results = await asyncio.gather(
*(
asyncio.to_thread(
litellm.Router._deployment_params_with_request_reasoning_override,
deployment_params,
{"reasoning_effort": efforts[index % len(efforts)]},
)
for index in range(280)
)
)
assert all("thinking" not in result for result in results)
assert all(result["output_config"] == {"format": "json"} for result in results)
assert all(result["extra_body"] == {"tenant": "shared"} for result in results)
assert deployment_params["thinking"] == {"type": "enabled"}
assert deployment_params["output_config"] == {"effort": "high", "format": "json"}
assert deployment_params["extra_body"] == {"reasoning_effort": "high", "tenant": "shared"}
@pytest.mark.parametrize(
("metadata", "should_drop"),
[({"internal_call_origin": "autorouter_classifier"}, True), ({}, False)],
ids=["classifier", "ordinary-request"],
)
def test_only_classifier_calls_drop_effort_for_an_unsupported_fallback(self, metadata, should_drop):
router = litellm.Router(model_list=[])
body: dict[str, object] = {"model": "classifier", "reasoning_effort": "low"}
kwargs: dict[str, object] = {
"reasoning_effort": "low",
"metadata": metadata,
"proxy_server_request": {"body": body},
}
deployment: DeploymentTypedDict = {
"model_name": "fallback",
"litellm_params": {"model": "openai/gpt-4o-mini"},
}
router._drop_unsupported_classifier_reasoning_effort(deployment, "fallback", kwargs)
assert ("reasoning_effort" not in kwargs) is should_drop
assert ("reasoning_effort" not in body) is should_drop
class TestPreRoutingTierDrivesFallbacks:
"""#38832: a complexity/auto router picks a tier behind the router name, but fallback
lookup stayed on the router name, so the tier's configured chain never ran and a

View file

@ -2,14 +2,17 @@
Tests for litellm/vector_stores/main.py.
Pins the router threading contract for vector store search: the router is an
explicit named parameter that reaches the HTTP handler, and it must never leak
into litellm_params/kwargs where logging would model_dump() it (the #19550
serialization trap).
explicit named parameter that reaches the HTTP handler wrapped in the embedding
executor, and it must never leak into litellm_params/kwargs where logging would
model_dump() it (the #19550 serialization trap).
"""
from unittest.mock import MagicMock, patch
import litellm.vector_stores.main as vector_stores_main
from litellm.llms.base_llm.vector_store.transformation import (
RouterVectorStoreEmbeddingExecutor,
)
from litellm.vector_stores.main import search
MOCK_SEARCH_RESPONSE = {
@ -19,17 +22,18 @@ MOCK_SEARCH_RESPONSE = {
}
def test_search_threads_router_to_handler():
"""search() must pass its router param through to the HTTP handler"""
def test_search_wraps_router_into_the_handler_embedding_executor():
"""search() hands the HTTP handler a Router-backed embedding executor carrying the
request metadata, and no bare router kwarg (LIT-6750)"""
mock_router = MagicMock()
logger = MagicMock()
with (
patch( # test-quality-ok: stubs provider config resolution; the seam under test is the router kwarg threading
patch( # test-quality-ok: stubs provider config resolution; the seam under test is the executor threading
"litellm.vector_stores.main.ProviderConfigManager.get_provider_vector_stores_config",
return_value=MagicMock(),
),
patch.object( # test-quality-ok: the handler call is the observable boundary for the router kwarg contract
patch.object( # test-quality-ok: the handler call is the observable boundary for the executor contract
vector_stores_main.base_llm_http_handler,
"vector_store_search_handler",
return_value=MOCK_SEARCH_RESPONSE,
@ -41,11 +45,16 @@ def test_search_threads_router_to_handler():
custom_llm_provider="s3_vectors",
router=mock_router,
litellm_logging_obj=logger,
litellm_metadata={"user_api_key_team_id": "team-a"},
)
assert response == MOCK_SEARCH_RESPONSE
mock_handler.assert_called_once()
assert mock_handler.call_args.kwargs["router"] is mock_router
assert "router" not in mock_handler.call_args.kwargs
executor = mock_handler.call_args.kwargs["embedding_executor"]
assert isinstance(executor, RouterVectorStoreEmbeddingExecutor)
assert executor.router is mock_router
assert dict(executor.metadata) == {"user_api_key_team_id": "team-a"}
def test_search_router_not_in_litellm_params():

View file

@ -1,6 +1,6 @@
{
"LIT001": {
"limit": 22330
"limit": 22328
},
"LIT002": {
"limit": 26763
@ -27,10 +27,10 @@
"limit": 0
},
"LIT010": {
"limit": 16478
"limit": 16477
},
"LIT011": {
"limit": 5520
"limit": 5519
},
"LIT012": {
"limit": 4489

View file

@ -3,7 +3,7 @@
# UI container — Next.js static export served by nginx.
ARG UI_BUILD_IMAGE=node:24.19-alpine3.24@sha256:d32cdf619f63fe0471182d08996dd516c6275bb5fd31ae06e55a570bd9e1ad43
ARG NGINX_VERSION=1.31-alpine
ARG NGINX_VERSION=1.31.5-alpine3.24@sha256:34f40471dea485273c5e2a04dd5e97a682332ceb4a9adecd67de450dcb2fb390
# ---------- builder ----------
FROM ${UI_BUILD_IMAGE} AS builder

View file

@ -1399,9 +1399,6 @@
},
"prefer-const": {
"count": 2
},
"react-hooks/set-state-in-effect": {
"count": 1
}
},
"src/components/TeamsPage/teamTableColumns.tsx": {

View file

@ -201,6 +201,7 @@ export const guardrailLogoMap = {
XecGuard: xecguardLogo.src,
"LiteLLM Content Filter": litellmLogo.src,
"LiteLLM LLM as a Judge": litellmLogo.src,
"Hide Secrets": litellmLogo.src,
Akto: aktoLogo.src,
"DeepKeep AI Firewall": deepkeepLogo.src,
"Qostodian Nexus": qohashLogo.src,

View file

@ -0,0 +1,99 @@
import React from "react";
import { fireEvent, screen, waitFor } from "@testing-library/react";
import { useForm } from "react-hook-form";
import { renderWithProviders } from "@/../tests/test-utils";
import { beforeEach, describe, expect, it, vi } from "vitest";
import GuardrailProviderFields from "./guardrail_provider_fields";
import { populateGuardrailProviderMap } from "./guardrail_info_helpers";
import type { GuardrailFormValues } from "./GuardrailFormField";
vi.mock("@/lib/toast", () => ({ toast: { error: vi.fn() } }));
const HIDE_SECRETS_PARAMS = {
"hide-secrets": {
ui_friendly_name: "Hide Secrets",
detect_secrets_config: {
param: "detect_secrets_config",
description: "Optional detect-secrets configuration",
required: false,
type: "object",
},
},
};
const Harness: React.FC<{ onValid: (values: GuardrailFormValues) => void }> = ({ onValid }) => {
const form = useForm<GuardrailFormValues>();
return (
<form onSubmit={form.handleSubmit(onValid)}>
<GuardrailProviderFields
selectedProvider="Hide-secrets"
control={form.control}
providerParams={HIDE_SECRETS_PARAMS}
/>
<button type="submit">save</button>
</form>
);
};
const renderHarness = () => {
populateGuardrailProviderMap(HIDE_SECRETS_PARAMS);
const onValid = vi.fn();
renderWithProviders(<Harness onValid={onValid} />);
const textarea = screen.getByLabelText(/detect_secrets_config/) as HTMLTextAreaElement;
return { onValid, textarea };
};
describe("GuardrailProviderFields object field", () => {
beforeEach(() => {
vi.clearAllMocks();
});
it("commits a valid JSON object as a parsed dict", async () => {
const { onValid, textarea } = renderHarness();
fireEvent.change(textarea, { target: { value: '{"plugins_used": [{"name": "AWSKeyDetector"}]}' } });
fireEvent.blur(textarea);
fireEvent.click(screen.getByRole("button", { name: "save" }));
await waitFor(() => expect(onValid).toHaveBeenCalledTimes(1));
expect(onValid.mock.calls[0][0].detect_secrets_config).toEqual({
plugins_used: [{ name: "AWSKeyDetector" }],
});
});
it("blocks submission while the field holds malformed JSON", async () => {
const { onValid, textarea } = renderHarness();
fireEvent.change(textarea, { target: { value: "{not json" } });
fireEvent.blur(textarea);
fireEvent.click(screen.getByRole("button", { name: "save" }));
await screen.findByText("detect_secrets_config must be a valid JSON object");
expect(onValid).not.toHaveBeenCalled();
expect(textarea.value).toBe("{not json");
});
it.each(['["array"]', '"scalar"', "null", "42"])("blocks non-object JSON %s", async (raw) => {
const { onValid, textarea } = renderHarness();
fireEvent.change(textarea, { target: { value: raw } });
fireEvent.blur(textarea);
fireEvent.click(screen.getByRole("button", { name: "save" }));
await screen.findByText("detect_secrets_config must be a valid JSON object");
expect(onValid).not.toHaveBeenCalled();
});
it("treats a cleared field as unset and submits", async () => {
const { onValid, textarea } = renderHarness();
fireEvent.change(textarea, { target: { value: '{"a": 1}' } });
fireEvent.blur(textarea);
fireEvent.change(textarea, { target: { value: "" } });
fireEvent.blur(textarea);
fireEvent.click(screen.getByRole("button", { name: "save" }));
await waitFor(() => expect(onValid).toHaveBeenCalledTimes(1));
expect(onValid.mock.calls[0][0].detect_secrets_config).toBeUndefined();
});
});

View file

@ -13,6 +13,8 @@ import { FieldGroup } from "@/components/ui/field";
import { Input } from "@/components/ui/input";
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
import { Slider } from "@/components/ui/slider";
import { Textarea } from "@/components/ui/textarea";
import { toast } from "@/lib/toast";
import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner";
import {
asStringArray,
@ -22,6 +24,7 @@ import {
readRecord,
requiredRule,
type GuardrailFieldControlProps,
type GuardrailFieldRules,
type GuardrailFormControl,
} from "./GuardrailFormField";
@ -60,6 +63,44 @@ const BOOLEAN_ITEMS = [
const isSecretKey = (fieldKey: string): boolean =>
fieldKey.includes("password") || fieldKey.includes("secret") || fieldKey.includes("key");
const isPlainObject = (value: unknown): value is Record<string, unknown> =>
typeof value === "object" && value !== null && !Array.isArray(value);
// Object fields hold the raw text while the user types, so submission must be
// blocked until the value parses to a plain JSON object (or is cleared).
const jsonObjectRule = (fieldKey: string): GuardrailFieldRules => ({
validate: (value: unknown) =>
value === undefined || isPlainObject(value) ? true : `${fieldKey} must be a valid JSON object`,
});
// Commits a parsed object (or undefined for a cleared field) to the form on
// blur; anything else stays as raw text so jsonObjectRule blocks submission.
const commitObjectField = (raw: string, onChange: (value: unknown) => void): void => {
const next = raw.trim();
if (next === "") {
onChange(undefined);
return;
}
let parsed: unknown;
try {
parsed = JSON.parse(next);
} catch {
parsed = next;
}
if (isPlainObject(parsed)) {
onChange(parsed);
} else {
toast.error("Enter a valid JSON object for this configuration");
}
};
const fieldRules = (field: ProviderParam, fieldKey: string): GuardrailFieldRules | undefined => {
if (field.type === "object") {
return jsonObjectRule(fieldKey);
}
return field.required ? requiredRule(`${fieldKey} is required`) : undefined;
};
interface ProviderFieldInputProps {
descriptor: ProviderParam;
fieldKey: string;
@ -141,6 +182,25 @@ const ProviderFieldInput: React.FC<ProviderFieldInputProps> = ({ descriptor, fie
);
}
if (descriptor.type === "object") {
const objectValue = typeof value === "object" && value !== null ? JSON.stringify(value, null, 2) : asText(value);
return (
<Textarea
id={id}
name={name}
ref={ref}
placeholder={descriptor.description}
value={objectValue}
onChange={(event) => onChange(event.target.value)}
onBlur={(event) => {
commitObjectField(event.target.value, onChange);
onBlur();
}}
{...aria}
/>
);
}
if (descriptor.type === "number") {
return (
<NumericalInput
@ -316,7 +376,7 @@ const GuardrailProviderFields: React.FC<GuardrailProviderFieldsProps> = ({
control={control}
name={fullFieldKey}
label={labelWithHint(fieldKey, field.description)}
rules={field.required ? requiredRule(`${fieldKey} is required`) : undefined}
rules={fieldRules(field, fieldKey)}
defaultValue={resolvedInitialValue}
>
{(fieldControl) => <ProviderFieldInput descriptor={field} fieldKey={fieldKey} control={fieldControl} />}

View file

@ -15,6 +15,7 @@ import {
teamCreateCall,
} from "./networking";
import Teams from "./Teams";
import { chooseSelectOption } from "../../tests/test-utils";
const can = vi.fn();
vi.mock("@/app/(dashboard)/hooks/useCan", () => ({
@ -1488,3 +1489,204 @@ describe("Teams - the exact bytes the create call sends", () => {
expect(teamCreateCall).not.toHaveBeenCalled();
});
});
describe("Teams - the create form keeps the organization and models picks while it is open", () => {
const ORGS = [
{ organization_id: "org-1", organization_alias: "Org 1", models: [], members: [] },
{ organization_id: "org-2", organization_alias: "Org 2", models: [], members: [] },
];
const orgField = () => screen.getByRole("combobox", { name: /organization/i });
const modelsField = () => screen.getByTestId("create-team-models-select");
const openCreateModal = async () => {
act(() => {
fireEvent.click(screen.getAllByRole("button", { name: /create team/i })[0]);
});
await screen.findByLabelText(/team name/i);
};
beforeEach(() => {
vi.clearAllMocks();
mockTeamInfoView.mockClear();
vi.mocked(fetchAvailableModelsForTeamOrKey).mockResolvedValue(["gpt-4", "gpt-3.5-turbo"]);
vi.mocked(fetchMCPAccessGroups).mockResolvedValue([]);
vi.mocked(getGuardrailsList).mockResolvedValue({ guardrails: [] });
vi.mocked(getDefaultTeamSettings).mockResolvedValue({ values: {} });
mockUseOrganizations.mockReturnValue({ data: ORGS });
});
it("keeps both picks when the organizations list comes back changed from a refetch", async () => {
const user = userEvent.setup();
renderWithQueryClient(<Teams accessToken="test-token" userID="user-123" userRole="Admin" />);
await openCreateModal();
await chooseSelectOption(user, orgField(), /Org 1/);
fireEvent.change(modelsField(), { target: { value: "gpt-4" } });
mockUseOrganizations.mockReturnValue({ data: ORGS.map((org) => ({ ...org, spend: 1 })) });
fireEvent.click(screen.getByText("Additional Settings"));
expect(orgField()).toHaveValue("Org 1");
expect(modelsField()).toHaveValue("gpt-4");
});
it("keeps models picked before the available models finish loading", async () => {
let resolveModels: (models: string[]) => void = () => {};
vi.mocked(fetchAvailableModelsForTeamOrKey).mockReturnValue(
new Promise<string[]>((resolve) => {
resolveModels = resolve;
}),
);
renderWithQueryClient(<Teams accessToken="test-token" userID="user-123" userRole="Admin" />);
await openCreateModal();
fireEvent.change(modelsField(), { target: { value: "gpt-4" } });
await act(async () => {
resolveModels(["gpt-4", "gpt-3.5-turbo"]);
});
expect(modelsField()).toHaveValue("gpt-4");
});
it("clears the models pick when the organization is changed, since models are org scoped", async () => {
const user = userEvent.setup();
renderWithQueryClient(<Teams accessToken="test-token" userID="user-123" userRole="Admin" />);
await openCreateModal();
await chooseSelectOption(user, orgField(), /Org 1/);
fireEvent.change(modelsField(), { target: { value: "gpt-4" } });
await chooseSelectOption(user, orgField(), /Org 2/);
await waitFor(() => expect(orgField()).toHaveValue("Org 2"));
expect(modelsField()).toHaveValue("");
});
it("keeps the models pick when the same organization is chosen again", async () => {
const user = userEvent.setup();
renderWithQueryClient(<Teams accessToken="test-token" userID="user-123" userRole="Admin" />);
await openCreateModal();
await chooseSelectOption(user, orgField(), /Org 1/);
fireEvent.change(modelsField(), { target: { value: "gpt-4" } });
await chooseSelectOption(user, orgField(), /Org 1/);
expect(orgField()).toHaveValue("Org 1");
expect(modelsField()).toHaveValue("gpt-4");
});
it("still preselects the only organization an org admin can create teams in", async () => {
mockUseOrganizations.mockReturnValue({
data: [
{
organization_id: "org-1",
organization_alias: "Org 1",
models: [],
members: [{ user_id: "user-123", user_role: "org_admin" }],
},
],
});
renderWithQueryClient(<Teams accessToken="test-token" userID="user-123" userRole="Internal User" />);
await openCreateModal();
expect(orgField()).toHaveValue("Org 1");
expect(orgField()).toBeDisabled();
});
it("leaves an org admin able to pick when their admin orgs narrow to one while the form is open", async () => {
const orgAdminOrgs = [
{
organization_id: "org-1",
organization_alias: "Org 1",
models: [],
members: [{ user_id: "user-123", user_role: "org_admin" }],
},
{
organization_id: "org-2",
organization_alias: "Org 2",
models: [],
members: [{ user_id: "user-123", user_role: "org_admin" }],
},
];
mockUseOrganizations.mockReturnValue({ data: orgAdminOrgs });
renderWithQueryClient(<Teams accessToken="test-token" userID="user-123" userRole="Internal User" />);
await openCreateModal();
expect(orgField()).toHaveValue("");
mockUseOrganizations.mockReturnValue({ data: [orgAdminOrgs[0]] });
fireEvent.click(screen.getByText("Additional Settings"));
expect(orgField()).toBeEnabled();
});
it("refuses to create the team in an organization the admin has lost access to", async () => {
const user = userEvent.setup();
const orgAdminOrgs = ORGS.map((org) => ({ ...org, members: [{ user_id: "user-123", user_role: "org_admin" }] }));
mockUseOrganizations.mockReturnValue({ data: orgAdminOrgs });
renderWithQueryClient(<Teams accessToken="test-token" userID="user-123" userRole="Internal User" />);
await openCreateModal();
fireEvent.change(screen.getByTestId("team-name-input"), { target: { value: "Revoked Team" } });
await chooseSelectOption(user, orgField(), /Org 1/);
mockUseOrganizations.mockReturnValue({ data: [orgAdminOrgs[1]] });
fireEvent.click(screen.getByText("Additional Settings"));
const submitButtons = screen.getAllByRole("button", { name: /create team/i });
fireEvent.click(submitButtons[submitButtons.length - 1]);
await screen.findByText(/no longer create teams in this organization/i);
expect(teamCreateCall).not.toHaveBeenCalled();
});
it("lets the admin switch to the one organization left after losing access to their pick", async () => {
const user = userEvent.setup();
const orgAdminOrgs = ORGS.map((org) => ({ ...org, members: [{ user_id: "user-123", user_role: "org_admin" }] }));
mockUseOrganizations.mockReturnValue({ data: orgAdminOrgs });
const createdTeam = {
team_id: "new-team-1",
team_alias: "Recovered Team",
models: [],
organization_id: "org-2",
keys: [],
members_with_roles: [],
spend: 0,
};
vi.mocked(teamCreateCall).mockResolvedValue(createdTeam);
renderWithQueryClient(<Teams accessToken="test-token" userID="user-123" userRole="Internal User" />);
await openCreateModal();
fireEvent.change(screen.getByTestId("team-name-input"), { target: { value: "Recovered Team" } });
await chooseSelectOption(user, orgField(), /Org 1/);
mockUseOrganizations.mockReturnValue({ data: [orgAdminOrgs[1]] });
fireEvent.click(screen.getByText("Additional Settings"));
expect(orgField()).toBeEnabled();
await chooseSelectOption(user, orgField(), /Org 2/);
const submitButtons = screen.getAllByRole("button", { name: /create team/i });
fireEvent.click(submitButtons[submitButtons.length - 1]);
await waitFor(() =>
expect(teamCreateCall).toHaveBeenCalledWith(
"test-token",
expect.objectContaining({ team_alias: "Recovered Team", organization_id: "org-2" }),
),
);
});
it("starts the form clean again when the modal is closed and reopened", async () => {
const user = userEvent.setup();
renderWithQueryClient(<Teams accessToken="test-token" userID="user-123" userRole="Admin" />);
await openCreateModal();
await chooseSelectOption(user, orgField(), /Org 1/);
fireEvent.change(modelsField(), { target: { value: "gpt-4" } });
fireEvent.click(screen.getByRole("button", { name: /^close$/i }));
await waitFor(() => expect(screen.queryByLabelText(/team name/i)).not.toBeInTheDocument());
await openCreateModal();
expect(orgField()).toHaveValue("");
expect(modelsField()).toHaveValue("");
});
});

View file

@ -208,7 +208,6 @@ const Teams: React.FC<TeamProps> = ({ accessToken, userID, userRole, premiumUser
const queryClient = useQueryClient();
const refreshTeams = () => queryClient.invalidateQueries({ queryKey: teamsTableKeys.all });
const [currentOrg] = useState<Organization | null>(null);
const [currentOrgForCreateTeam, setCurrentOrgForCreateTeam] = useState<Organization | null>(null);
const isOrgAdmin = userRole !== "Admin";
const [additionalSettingsOpen, setAdditionalSettingsOpen] = useState(false);
@ -216,17 +215,33 @@ const Teams: React.FC<TeamProps> = ({ accessToken, userID, userRole, premiumUser
const [agentSettingsOpen, setAgentSettingsOpen] = useState(false);
const [searchToolSettingsOpen, setSearchToolSettingsOpen] = useState(false);
const adminOrgs = useMemo(
() => getAdminOrganizations(userRole, userID, organizations),
[userRole, userID, organizations],
);
const teamCreateSchema = useMemo(
() =>
teamCreateFieldsSchema.superRefine((values, ctx) => {
if (isOrgAdmin && !values.organization_id) {
ctx.addIssue({ code: "custom", message: SUPPRESSED_BY_DESCRIPTION, path: ["organization_id"] });
}
const organizationIsStillPickable =
values.organization_id == null ||
organizations == null ||
adminOrgs.some((org) => org.organization_id === values.organization_id);
if (!organizationIsStillPickable) {
ctx.addIssue({
code: "custom",
message: "You can no longer create teams in this organization",
path: ["organization_id"],
});
}
if (additionalSettingsOpen && !isParsableJson(values.secret_manager_settings)) {
ctx.addIssue({ code: "custom", message: SUPPRESSED_BY_DESCRIPTION, path: ["secret_manager_settings"] });
}
}),
[isOrgAdmin, additionalSettingsOpen],
[isOrgAdmin, additionalSettingsOpen, adminOrgs, organizations],
);
const form = useZodForm(teamCreateSchema, { defaultValues: EMPTY_TEAM_CREATE_VALUES });
@ -264,28 +279,6 @@ const Teams: React.FC<TeamProps> = ({ accessToken, userID, userRole, premiumUser
? `Default: ${getBudgetDurationLabel(defaultBudgetDuration)} (${defaultBudgetDuration})`
: "n/a";
useEffect(() => {
form.setValue("models", []);
}, [currentOrgForCreateTeam, userModels]);
// Handle organization preselection when modal opens
useEffect(() => {
if (isTeamModalVisible) {
const adminOrgs = getAdminOrganizations(userRole, userID, organizations);
// Org admins must scope a team to an org, so with exactly one we preselect it.
// Proxy admins can create org-less teams, so the field stays optional regardless of org count.
if (isOrgAdmin && adminOrgs.length === 1) {
const org = adminOrgs[0];
form.setValue("organization_id", org.organization_id);
setCurrentOrgForCreateTeam(org);
} else {
form.setValue("organization_id", currentOrg?.organization_id || null);
setCurrentOrgForCreateTeam(currentOrg);
}
}
}, [isTeamModalVisible, isOrgAdmin, userRole, userID, organizations, currentOrg]);
// Add this useEffect to fetch guardrails
useEffect(() => {
const fetchGuardrails = async () => {
@ -320,6 +313,26 @@ const Teams: React.FC<TeamProps> = ({ accessToken, userID, userRole, premiumUser
if (canViewPolicies) fetchPolicies();
}, [accessToken, canViewPolicies]);
const openCreateTeamModal = () => {
// Org admins must scope a team to an org, so with exactly one we preselect it.
// Proxy admins can create org-less teams, so the field stays optional regardless of org count.
if (isOrgAdmin && adminOrgs.length === 1) {
form.setValue("organization_id", adminOrgs[0].organization_id);
}
setIsTeamModalVisible(true);
};
const selectCreateTeamOrganization = (
next: string,
currentOrganizationId: string | null,
onChange: (organizationId: string | null) => void,
) => {
const nextOrganizationId = next === "" ? null : next;
if (nextOrganizationId === currentOrganizationId) return;
onChange(nextOrganizationId);
form.setValue("models", []);
};
const resetCreateForm = () => {
form.reset(EMPTY_TEAM_CREATE_VALUES);
setAdditionalSettingsOpen(false);
@ -636,7 +649,7 @@ const Teams: React.FC<TeamProps> = ({ accessToken, userID, userRole, premiumUser
subtitle="Manage teams, members, and their access to models and budgets"
primaryAction={
canCreateOrManageTeams(userRole, userID, organizations) ? (
<UIButton onClick={() => setIsTeamModalVisible(true)} data-testid="create-team-button">
<UIButton onClick={openCreateTeamModal} data-testid="create-team-button">
<Plus className="size-4" />
Create Team
</UIButton>
@ -683,9 +696,9 @@ const Teams: React.FC<TeamProps> = ({ accessToken, userID, userRole, premiumUser
)}
</FormField>
{(() => {
const adminOrgs = getAdminOrganizations(userRole, userID, organizations);
const isSingleOrg = adminOrgs.length === 1;
const hasNoOrgs = adminOrgs.length === 0;
const soleOrganizationId = isSingleOrg ? adminOrgs[0].organization_id ?? null : null;
return (
<>
@ -715,18 +728,13 @@ const Teams: React.FC<TeamProps> = ({ accessToken, userID, userRole, premiumUser
label: org.organization_alias ?? "",
sublabel: org.organization_id ?? "",
}))}
disabled={isOrgAdmin && isSingleOrg}
disabled={isOrgAdmin && soleOrganizationId !== null && value === soleOrganizationId}
allowClear={!isOrgAdmin}
placeholder={
hasNoOrgs ? "No organizations available" : "Search or select an Organization"
}
emptyText="No organizations available"
onValueChange={(next) => {
onChange(next === "" ? null : next);
setCurrentOrgForCreateTeam(
adminOrgs.find((org) => org.organization_id === next) ?? null,
);
}}
onValueChange={(next) => selectCreateTeamOrganization(next, value ?? null, onChange)}
/>
)}
</FormField>

View file

@ -4,6 +4,7 @@ import type { OnUrlUpdateFunction } from "nuqs/adapters/testing";
import { vi, it, expect, beforeEach, describe, Mock, MockedFunction } from "vitest";
import { chooseSelectOption, renderWithProviders } from "../../../tests/test-utils";
import { VirtualKeysTable } from "./VirtualKeysTable";
import { KEY_TABLE_SORT_FIELDS } from "./keyTableColumns";
import { KeyResponse, Team } from "../key_team_helpers/key_list";
import { useKeyInfo } from "@/app/(dashboard)/hooks/keys/useKeyInfo";
import { KeysResponse, useKeys } from "@/app/(dashboard)/hooks/keys/useKeys";
@ -176,8 +177,10 @@ const keysResult = (keys: KeyResponse[], data: Partial<KeysResponse> = {}, extra
const openFilters = () => fireEvent.click(screen.getByRole("button", { name: "Filters" }));
const lastKeyParam = (onUrlUpdate: Mock<OnUrlUpdateFunction>) =>
onUrlUpdate.mock.calls.at(-1)?.[0].searchParams.get("key");
const lastSearchParam = (onUrlUpdate: Mock<OnUrlUpdateFunction>, name: string) =>
onUrlUpdate.mock.calls.at(-1)?.[0].searchParams.get(name);
const lastKeyParam = (onUrlUpdate: Mock<OnUrlUpdateFunction>) => lastSearchParam(onUrlUpdate, "key");
beforeEach(() => {
vi.clearAllMocks();
@ -510,7 +513,8 @@ describe("server-side filtering – the LIT-4080 regression guard", () => {
});
it("drops the filter from the useKeys query when it is cleared", async () => {
renderWithProviders(<VirtualKeysTable />);
const onUrlUpdate = vi.fn<OnUrlUpdateFunction>();
renderWithProviders(<VirtualKeysTable />, { onUrlUpdate });
openFilters();
const userIdInput = await screen.findByPlaceholderText(/Enter User ID/);
@ -520,6 +524,12 @@ describe("server-side filtering – the LIT-4080 regression guard", () => {
await waitFor(() => {
expect(mockUseKeys).toHaveBeenLastCalledWith(1, 50, expect.objectContaining({ userID: "user-42" }));
});
// Let the filter reach the URL before clearing it: NuqsTestingAdapter runs
// resetUrlUpdateQueueOnMount on every render, so a still-queued write can be
// aborted by the re-render its own predecessor triggers.
await waitFor(() => {
expect(lastSearchParam(onUrlUpdate, "filter_user")).toBe("user-42");
});
fireEvent.click(screen.getByTestId("datatable-clear-filters"));
@ -638,3 +648,183 @@ describe("Status column reflects blocked / expiry / scim metadata", () => {
});
});
});
describe("table state lives in the URL so it survives leaving and returning to the page", () => {
it("restores the search term, sort and pagination from the URL on mount", async () => {
renderWithProviders(<VirtualKeysTable />, {
searchParams: { key_search: "prod", sort_by: "spend", sort_order: "asc", page: "3", page_size: "25" },
});
await waitFor(() => {
expect(mockUseKeys).toHaveBeenLastCalledWith(
3,
25,
expect.objectContaining({ selectedKeyAlias: "prod", sortBy: "spend", sortOrder: "asc" }),
);
});
expect(screen.getByPlaceholderText(/Search by key alias/)).toHaveValue("prod");
});
it("restores the drawer filters from the URL on mount", async () => {
renderWithProviders(<VirtualKeysTable />, { searchParams: { filter_team: "team-1", filter_user: "user-42" } });
await waitFor(() => {
expect(mockUseKeys).toHaveBeenLastCalledWith(
1,
50,
expect.objectContaining({ teamID: "team-1", userID: "user-42" }),
);
});
expect(screen.getByTestId("filter-chip-team_id")).toHaveTextContent("Test Team");
});
it("writes the search term to the URL", async () => {
const onUrlUpdate = vi.fn<OnUrlUpdateFunction>();
renderWithProviders(<VirtualKeysTable />, { onUrlUpdate });
fireEvent.change(screen.getByPlaceholderText(/Search by key alias/), { target: { value: "prod" } });
await waitFor(() => {
expect(lastSearchParam(onUrlUpdate, "key_search")).toBe("prod");
});
});
it("writes the sort field and direction to the URL", async () => {
const onUrlUpdate = vi.fn<OnUrlUpdateFunction>();
renderWithProviders(<VirtualKeysTable />, { onUrlUpdate });
fireEvent.click(screen.getByRole("button", { name: "Key" }));
await waitFor(() => {
expect(lastSearchParam(onUrlUpdate, "sort_by")).toBe("key_alias");
});
expect(lastSearchParam(onUrlUpdate, "sort_order")).toBe("asc");
});
it("writes an applied drawer filter to the URL and clears it again", async () => {
const onUrlUpdate = vi.fn<OnUrlUpdateFunction>();
renderWithProviders(<VirtualKeysTable />, { onUrlUpdate });
openFilters();
fireEvent.change(await screen.findByPlaceholderText(/Enter User ID/), { target: { value: "user-42" } });
fireEvent.click(screen.getByTestId("filter-drawer-apply"));
await waitFor(() => {
expect(lastSearchParam(onUrlUpdate, "filter_user")).toBe("user-42");
});
fireEvent.click(screen.getByTestId("datatable-clear-filters"));
await waitFor(() => {
expect(lastSearchParam(onUrlUpdate, "filter_user")).toBeNull();
});
expect(screen.queryByTestId("filter-chip-user_id")).not.toBeInTheDocument();
});
it("returns to page 1 when the search term changes", async () => {
const onUrlUpdate = vi.fn<OnUrlUpdateFunction>();
renderWithProviders(<VirtualKeysTable />, { searchParams: { page: "3" }, onUrlUpdate });
await waitFor(() => {
expect(mockUseKeys).toHaveBeenLastCalledWith(3, 50, expect.anything());
});
fireEvent.change(screen.getByPlaceholderText(/Search by key alias/), { target: { value: "prod" } });
await waitFor(() => {
expect(mockUseKeys).toHaveBeenLastCalledWith(1, 50, expect.objectContaining({ selectedKeyAlias: "prod" }));
});
await waitFor(() => {
expect(lastSearchParam(onUrlUpdate, "page")).toBeNull();
});
});
it("leaves the create-key deep link's team_id alone instead of filtering the list with it", async () => {
renderWithProviders(<VirtualKeysTable />, { searchParams: { create: "true", team_id: "team-1" } });
await waitFor(() => {
expect(mockUseKeys).toHaveBeenLastCalledWith(1, 50, expect.objectContaining({ teamID: undefined }));
});
expect(screen.queryByTestId("filter-chip-team_id")).not.toBeInTheDocument();
});
it.each([
["0", 1],
["-3", 1],
])("clamps a hand-edited page of %s up to the first page", async (page, expected) => {
renderWithProviders(<VirtualKeysTable />, { searchParams: { page } });
await waitFor(() => {
expect(mockUseKeys).toHaveBeenLastCalledWith(expected, 50, expect.anything());
});
});
it.each([
["0", 1],
["1000", 100],
])("clamps a hand-edited page_size of %s into the range /key/list accepts", async (pageSize, expected) => {
renderWithProviders(<VirtualKeysTable />, { searchParams: { page_size: pageSize } });
await waitFor(() => {
expect(mockUseKeys).toHaveBeenLastCalledWith(1, expected, expect.anything());
});
});
it("trims whitespace off a filter that arrived from the URL", async () => {
renderWithProviders(<VirtualKeysTable />, { searchParams: { filter_user: " user-42 " } });
await waitFor(() => {
expect(mockUseKeys).toHaveBeenLastCalledWith(1, 50, expect.objectContaining({ userID: "user-42" }));
});
});
it("falls back to the default sort when the URL names a column the table cannot sort by", async () => {
renderWithProviders(<VirtualKeysTable />, { searchParams: { sort_by: "totally_unknown_field" } });
await waitFor(() => {
expect(mockUseKeys).toHaveBeenLastCalledWith(
1,
50,
expect.objectContaining({ sortBy: "created_at", sortOrder: "desc" }),
);
});
expect(screen.getByText("Test Key Alias")).toBeInTheDocument();
});
it.each(KEY_TABLE_SORT_FIELDS)("round-trips a %s sort from the URL", async (field) => {
renderWithProviders(<VirtualKeysTable />, { searchParams: { sort_by: field, sort_order: "asc" } });
await waitFor(() => {
expect(mockUseKeys).toHaveBeenLastCalledWith(1, 50, expect.objectContaining({ sortBy: field, sortOrder: "asc" }));
});
});
it("clears sort_by from the URL when the Spend / Budget sort is reset", async () => {
const user = userEvent.setup();
const onUrlUpdate = vi.fn<OnUrlUpdateFunction>();
renderWithProviders(<VirtualKeysTable />, { onUrlUpdate });
await chooseSelectOption(user, screen.getByTestId("sort-trigger-spend"), "Spend ascending", "menuitem");
await waitFor(() => {
expect(lastSearchParam(onUrlUpdate, "sort_by")).toBe("spend");
});
await chooseSelectOption(user, screen.getByTestId("sort-trigger-spend"), "Reset", "menuitem");
await waitFor(() => {
expect(lastSearchParam(onUrlUpdate, "sort_by")).toBeNull();
});
expect(lastSearchParam(onUrlUpdate, "sort_order")).toBeNull();
});
it("drops the search param back out of the URL when the search box is cleared", async () => {
const onUrlUpdate = vi.fn<OnUrlUpdateFunction>();
renderWithProviders(<VirtualKeysTable />, { searchParams: { key_search: "prod" }, onUrlUpdate });
fireEvent.change(screen.getByPlaceholderText(/Search by key alias/), { target: { value: "" } });
await waitFor(() => {
expect(lastSearchParam(onUrlUpdate, "key_search")).toBeNull();
});
});
});

View file

@ -15,34 +15,65 @@ import { SearchSelect } from "@/components/shared/SearchSelect";
import { PageHeader } from "@/components/shared/PageHeader";
import { Input } from "@/components/ui/input";
import { useDebouncedValue } from "@tanstack/react-pacer/debouncer";
import { ColumnFiltersState, OnChangeFn, PaginationState, SortingState } from "@tanstack/react-table";
import { ColumnFiltersState, functionalUpdate, OnChangeFn, PaginationState, SortingState } from "@tanstack/react-table";
import { KeyRound } from "lucide-react";
import { parseAsString, useQueryState } from "nuqs";
import { createParser, parseAsInteger, parseAsString, parseAsStringLiteral, useQueryState, useQueryStates } from "nuqs";
import React, { useCallback, useMemo, useState } from "react";
import { KeyResponse, Team } from "../key_team_helpers/key_list";
import KeyInfoView from "../templates/key_info_view";
import { getKeyTableColumns, KEY_TABLE_HIDDEN_COLUMNS } from "./keyTableColumns";
import { getKeyTableColumns, KEY_TABLE_HIDDEN_COLUMNS, KEY_TABLE_SORT_FIELDS } from "./keyTableColumns";
interface VirtualKeysTableProps {
headerActions?: React.ReactNode;
}
const DEFAULT_SORTING: SortingState = [{ id: "created_at", desc: true }];
const FILTER_COLUMNS = ["team_id", "org_id", "user_id", "key_hash"] as const;
type FilterColumn = (typeof FILTER_COLUMNS)[number];
const toSortOrder = (sorting: SortingState): "asc" | "desc" | undefined => {
const active = sorting[0];
if (!active) return undefined;
return active.desc ? "desc" : "asc";
};
const FILTER_LABELS: Record<string, string> = {
const FILTER_LABELS: Record<FilterColumn, string> = {
team_id: "Team",
org_id: "Organization",
user_id: "User ID",
key_hash: "Key ID",
};
const DEFAULT_SORT_BY = "created_at";
const DEFAULT_SORT_ORDER = "desc";
const DEFAULT_PAGE_SIZE = 50;
const MAX_PAGE_SIZE = 100;
const MAX_PAGE = 100_000;
const boundedInteger = (min: number, max: number, fallback: number) =>
createParser({
parse: (value: string) => {
const parsed = parseAsInteger.parse(value);
return parsed === null ? null : Math.min(Math.max(parsed, min), max);
},
serialize: String,
}).withDefault(fallback);
// The filters carry a prefix because /api-keys also takes team_id, key_alias and key_type
// as create-key prefills; an unprefixed filter would hijack those deep links.
const TABLE_STATE = {
key_search: parseAsString.withDefault(""),
sort_by: parseAsString.withDefault(DEFAULT_SORT_BY),
sort_order: parseAsStringLiteral(["asc", "desc"] as const).withDefault(DEFAULT_SORT_ORDER),
page: boundedInteger(1, MAX_PAGE, 1),
page_size: boundedInteger(1, MAX_PAGE_SIZE, DEFAULT_PAGE_SIZE),
filter_team: parseAsString.withDefault(""),
filter_org: parseAsString.withDefault(""),
filter_user: parseAsString.withDefault(""),
filter_key_id: parseAsString.withDefault(""),
};
const toSortOrder = (active: SortingState[number]): "asc" | "desc" => (active.desc ? "desc" : "asc");
const filterValue = (filters: ColumnFiltersState, column: FilterColumn): string | null => {
const value = filters.find((filter) => filter.id === column)?.value;
return (typeof value === "string" ? value.trim() : "") || null;
};
export function VirtualKeysTable({ headerActions }: VirtualKeysTableProps) {
const { data: fetchedOrganizations } = useOrganizations();
const organizations = useMemo(() => fetchedOrganizations ?? [], [fetchedOrganizations]);
@ -50,32 +81,48 @@ export function VirtualKeysTable({ headerActions }: VirtualKeysTableProps) {
const allTeams = useMemo<Team[]>(() => fetchedTeams ?? [], [fetchedTeams]);
const [selectedKeyId, setSelectedKeyId] = useQueryState("key", parseAsString.withOptions({ history: "push" }));
const [sorting, setSorting] = useState<SortingState>(DEFAULT_SORTING);
const [tablePagination, setTablePagination] = useState<PaginationState>({ pageIndex: 0, pageSize: 50 });
const [columnFilters, setColumnFilters] = useState<ColumnFiltersState>([]);
const [tableState, setTableState] = useQueryStates(TABLE_STATE);
const [filtersOpen, setFiltersOpen] = useState(false);
const [searchInput, setSearchInput] = useState("");
const searchInput = tableState.key_search;
const [searchQuery] = useDebouncedValue(searchInput, { wait: DEBOUNCE_WAIT_MS });
const getFilterValue = useCallback(
(columnId: string): string | undefined => {
const entry = columnFilters.find((filter) => filter.id === columnId);
return typeof entry?.value === "string" && entry.value.trim() ? entry.value.trim() : undefined;
},
[columnFilters],
// A hand-edited sort_by the table cannot sort by would 400 at /key/list and leave the page loading.
const sortBy = KEY_TABLE_SORT_FIELDS.includes(tableState.sort_by) ? tableState.sort_by : DEFAULT_SORT_BY;
const sorting = useMemo<SortingState>(
() => [{ id: sortBy, desc: tableState.sort_order === "desc" }],
[sortBy, tableState.sort_order],
);
const tablePagination = useMemo<PaginationState>(
() => ({ pageIndex: tableState.page - 1, pageSize: tableState.page_size }),
[tableState.page, tableState.page_size],
);
const { filter_team, filter_org, filter_user, filter_key_id } = tableState;
const appliedFilters = useMemo(
() => ({
team_id: filter_team.trim(),
org_id: filter_org.trim(),
user_id: filter_user.trim(),
key_hash: filter_key_id.trim(),
}),
[filter_team, filter_org, filter_user, filter_key_id],
);
const columnFilters = useMemo<ColumnFiltersState>(
() =>
FILTER_COLUMNS.filter((column) => appliedFilters[column]).map((column) => ({
id: column,
value: appliedFilters[column],
})),
[appliedFilters],
);
const sortBy = sorting[0]?.id;
const sortOrder = toSortOrder(sorting);
const keyListOptions = {
teamID: getFilterValue("team_id"),
organizationID: getFilterValue("org_id"),
teamID: appliedFilters.team_id || undefined,
organizationID: appliedFilters.org_id || undefined,
selectedKeyAlias: searchQuery.trim() || undefined,
userID: getFilterValue("user_id"),
keyHash: getFilterValue("key_hash"),
userID: appliedFilters.user_id || undefined,
keyHash: appliedFilters.key_hash || undefined,
sortBy,
sortOrder,
sortOrder: tableState.sort_order,
expand: "user",
};
@ -89,20 +136,47 @@ export function VirtualKeysTable({ headerActions }: VirtualKeysTableProps) {
const keyList = useMemo(() => keys?.keys ?? [], [keys]);
const rowCount = keys?.total_count ?? 0;
const handleSearchChange = useCallback((value: string) => {
setSearchInput(value);
setTablePagination((prev) => ({ ...prev, pageIndex: 0 }));
}, []);
const handleSearchChange = useCallback(
(value: string) => {
void setTableState({ key_search: value || null, page: null });
},
[setTableState],
);
const handleSortingChange = useCallback<OnChangeFn<SortingState>>((updaterOrValue) => {
setSorting(updaterOrValue);
setTablePagination((prev) => ({ ...prev, pageIndex: 0 }));
}, []);
const handleSortingChange = useCallback<OnChangeFn<SortingState>>(
(updaterOrValue) => {
const active = functionalUpdate(updaterOrValue, sorting)[0];
void setTableState({
sort_by: active?.id ?? null,
sort_order: active ? toSortOrder(active) : null,
page: null,
});
},
[sorting, setTableState],
);
const handleColumnFiltersChange = useCallback<OnChangeFn<ColumnFiltersState>>((updaterOrValue) => {
setColumnFilters(updaterOrValue);
setTablePagination((prev) => ({ ...prev, pageIndex: 0 }));
}, []);
const handleColumnFiltersChange = useCallback<OnChangeFn<ColumnFiltersState>>(
(updaterOrValue) => {
const next = functionalUpdate(updaterOrValue, columnFilters);
const nextFilters = {
filter_team: filterValue(next, "team_id"),
filter_org: filterValue(next, "org_id"),
filter_user: filterValue(next, "user_id"),
filter_key_id: filterValue(next, "key_hash"),
page: null,
};
void setTableState(nextFilters);
},
[columnFilters, setTableState],
);
const handlePaginationChange = useCallback<OnChangeFn<PaginationState>>(
(updaterOrValue) => {
const next = functionalUpdate(updaterOrValue, tablePagination);
void setTableState({ page: next.pageIndex + 1, page_size: next.pageSize });
},
[tablePagination, setTableState],
);
const columns = useMemo(
() => getKeyTableColumns({ allTeams, organizations, onSelectKey: (key) => void setSelectedKeyId(key.token) }),
@ -199,7 +273,7 @@ export function VirtualKeysTable({ headerActions }: VirtualKeysTableProps) {
onSortingChange={handleSortingChange}
paginationMode="server"
pagination={tablePagination}
onPaginationChange={setTablePagination}
onPaginationChange={handlePaginationChange}
rowCount={rowCount}
filterMode="server"
columnFilters={columnFilters}

View file

@ -32,6 +32,14 @@ const SPEND_BUDGET_SORT_FIELDS: DataTableSortField[] = [
{ id: "max_budget", label: "Budget" },
];
export const KEY_TABLE_SORT_FIELDS: readonly string[] = [
"key_alias",
"token",
"created_at",
"updated_at",
...SPEND_BUDGET_SORT_FIELDS.map((field) => field.id),
];
const getKeyStatus = (key: KeyResponse): KeyStatus => {
if (key.blocked === true) {
const isScimBlocked = (key.metadata as Record<string, unknown> | null | undefined)?.scim_blocked === true;

View file

@ -13,6 +13,8 @@ import ClassifierPromptEditor from "./ClassifierPromptEditor";
import CustomTierPromptEditor from "./CustomTierPromptEditor";
import { RestrictedSection, restrictedBy } from "./TierRestrictions";
import HeuristicScoringConfig from "./HeuristicScoringConfig";
import ClassifierReasoningEffortSelect from "./ClassifierReasoningEffortSelect";
import type { ReasoningEffort } from "./complexity_router_tiers";
import { useComplexityScorerDefaults } from "@/app/(dashboard)/hooks/autoRouter/useComplexityScorerDefaults";
import {
ClassificationFrequency,
@ -154,6 +156,7 @@ interface ClassificationMethodConfigProps {
value: ComplexityRouterConfigValue;
onChange: (value: ComplexityRouterConfigValue) => void;
modelOptions: { value: string; label: string }[];
effortOptionsByModel: Record<string, string[] | null | undefined>;
customTechnicalKeywords?: string[];
onCustomTechnicalKeywordsChange?: (keywords: string[]) => void;
showValidationErrors?: boolean;
@ -236,6 +239,7 @@ const ClassificationMethodConfig: React.FC<ClassificationMethodConfigProps> = ({
value,
onChange,
modelOptions,
effortOptionsByModel,
customTechnicalKeywords,
onCustomTechnicalKeywordsChange,
showValidationErrors = false,
@ -251,6 +255,9 @@ const ClassificationMethodConfig: React.FC<ClassificationMethodConfigProps> = ({
const contextBudget = value.classifier_context_budget_chars ?? DEFAULT_CLASSIFIER_CONTEXT_BUDGET_CHARS;
const contextBudgetQuotesNothing = contextBudget > 0 && contextBudget < MIN_QUOTED_CONTEXT_TURN_CHARS;
const classificationRubric = value.classifier_llm_config?.classification_rubric ?? DEFAULT_CLASSIFICATION_RUBRIC;
const classifierModel = value.classifier_llm_config?.model ?? "";
const classifierReasoningEffort = value.classifier_llm_config?.reasoning_effort;
const explicitlySupportedClassifierEfforts = effortOptionsByModel[classifierModel];
const handleClassifierTypeChange = (classifierType: ClassifierType) => {
const nextValue: ComplexityRouterConfigValue = {
@ -299,16 +306,33 @@ const ClassificationMethodConfig: React.FC<ClassificationMethodConfigProps> = ({
};
const handleClassifierModelChange = (model: string) => {
if (model === value.classifier_llm_config?.model) return;
const { reasoning_effort: _reasoningEffort, ...classifierLlmConfig } = value.classifier_llm_config ?? {
model: "",
timeout_ms: DEFAULT_CLASSIFIER_TIMEOUT_MS,
};
onChange({
...value,
classifier_llm_config: {
...value.classifier_llm_config,
...classifierLlmConfig,
model,
timeout_ms: value.classifier_llm_config?.timeout_ms ?? DEFAULT_CLASSIFIER_TIMEOUT_MS,
timeout_ms: classifierLlmConfig.timeout_ms,
},
});
};
const handleClassifierReasoningEffortChange = (reasoningEffort: ReasoningEffort | undefined) => {
if (!value.classifier_llm_config) return;
const { reasoning_effort: _reasoningEffort, ...classifierLlmConfig } = value.classifier_llm_config;
onChange({
...value,
classifier_llm_config:
reasoningEffort === undefined
? classifierLlmConfig
: { ...classifierLlmConfig, reasoning_effort: reasoningEffort },
});
};
const handleClassifierTimeoutChange = (timeoutMs: number) => {
onChange({
...value,
@ -493,9 +517,16 @@ const ClassificationMethodConfig: React.FC<ClassificationMethodConfigProps> = ({
emptyText="No models found"
allowClear={false}
className={classifierModelMissing ? "border-destructive" : undefined}
aria-label="Classifier Model"
/>
{classifierModelMissing && <span className="text-xs text-destructive">A classifier model is required</span>}
</div>
<ClassifierReasoningEffortSelect
model={classifierModel}
value={classifierReasoningEffort}
explicitlySupported={explicitlySupportedClassifierEfforts}
onChange={handleClassifierReasoningEffortChange}
/>
<div>
<Label htmlFor={CLASSIFIER_TIMEOUT_ID} className="block mb-1 font-semibold">
Timeout (ms)

View file

@ -0,0 +1,86 @@
import { Info } from "lucide-react";
import { SimpleTooltip } from "@/components/ui/tooltip";
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
import type { ReasoningEffort } from "./complexity_router_tiers";
const PROVIDER_DEFAULT = "__classifier_provider_default__";
type EffortStatus = "supported" | "unsupported" | "unverified" | undefined;
const effortStatusFor = (
effort: string | undefined,
explicitlySupported: string[] | null | undefined,
): EffortStatus => {
if (effort === undefined) return undefined;
if (!Array.isArray(explicitlySupported)) return "unverified";
return explicitlySupported.includes(effort) ? "supported" : "unsupported";
};
interface ClassifierReasoningEffortSelectProps {
model: string;
value: ReasoningEffort | undefined;
explicitlySupported: string[] | null | undefined;
onChange: (value: ReasoningEffort | undefined) => void;
}
const ClassifierReasoningEffortSelect = ({
model,
value,
explicitlySupported,
onChange,
}: ClassifierReasoningEffortSelectProps) => {
const status = effortStatusFor(value, explicitlySupported);
const options = Array.from(new Set([...(explicitlySupported ?? []), ...(value ? [value] : [])]));
if (!model || options.length === 0) return null;
const optionLabel = (effort: string): string =>
effort === value && status !== "supported" ? `${effort} (${status})` : effort;
return (
<div>
<div className="flex items-center gap-2 mb-1">
<strong className="font-semibold">Reasoning Effort</strong>
<SimpleTooltip content="Sent only to the classifier call. Default leaves the classifier deployment or provider setting unchanged.">
<Info className="size-4 text-muted-foreground" />
</SimpleTooltip>
</div>
<Select
items={[
{ value: PROVIDER_DEFAULT, label: "Default" },
...options.map((effort) => ({ value: effort, label: optionLabel(effort) })),
]}
value={value ?? PROVIDER_DEFAULT}
onValueChange={(effort: string | null) =>
effort && onChange(effort === PROVIDER_DEFAULT ? undefined : (effort as ReasoningEffort))
}
>
<SelectTrigger aria-label={`Reasoning effort for classifier model ${model}`} className="w-full">
<SelectValue />
</SelectTrigger>
<SelectContent>
<SelectItem value={PROVIDER_DEFAULT}>Default</SelectItem>
{options.map((effort) => (
<SelectItem key={effort} value={effort}>
{optionLabel(effort)}
</SelectItem>
))}
</SelectContent>
</Select>
{status === "unverified" && (
<p className="mt-1 text-xs text-amber-700 dark:text-amber-400">
This saved effort cannot be verified for the selected model. Choose Default unless you have confirmed provider
support.
</p>
)}
{status === "unsupported" && (
<p className="mt-1 text-xs text-destructive">
This saved effort is not supported by every deployment in the selected model group. Choose Default or a
supported value before saving.
</p>
)}
</div>
);
};
export default ClassifierReasoningEffortSelect;

View file

@ -1124,6 +1124,109 @@ describe("ComplexityRouterConfig per-model reasoning effort", () => {
});
});
describe("ComplexityRouterConfig classifier reasoning effort", () => {
const llmValue: ComplexityRouterConfigValue = {
...defaultValue,
classifier_type: "llm",
classifier_llm_config: { model: "gpt-4", timeout_ms: 3000 },
};
const renderClassifier = (value: ComplexityRouterConfigValue = llmValue, onChange = vi.fn()) => {
renderWithProviders(<ComplexityRouterConfig {...baseProps} value={value} onChange={onChange} />);
fireEvent.click(screen.getByText("Advanced: Classification Method"));
return onChange;
};
it("defaults to the classifier provider setting and offers only supported efforts", async () => {
renderClassifier();
const user = userEvent.setup();
const select = screen.getByRole("combobox", { name: "Reasoning effort for classifier model gpt-4" });
expect(select).toHaveTextContent("Default");
await user.click(select);
expect((await screen.findAllByRole("option")).map((option) => option.textContent)).toEqual([
"Default",
"medium",
"high",
"xhigh",
]);
});
it("stores an explicit effort on the classifier config", async () => {
const onChange = renderClassifier();
const user = userEvent.setup();
await user.click(screen.getByRole("combobox", { name: "Reasoning effort for classifier model gpt-4" }));
await user.click(await screen.findByRole("option", { name: "high" }));
expect(onChange).toHaveBeenCalledWith({
...llmValue,
classifier_llm_config: { model: "gpt-4", timeout_ms: 3000, reasoning_effort: "high" },
});
});
it("removes the effort override when Default is selected", async () => {
const onChange = renderClassifier({
...llmValue,
classifier_llm_config: { model: "gpt-4", timeout_ms: 3000, reasoning_effort: "high" },
});
const user = userEvent.setup();
await user.click(screen.getByRole("combobox", { name: "Reasoning effort for classifier model gpt-4" }));
await user.click(await screen.findByRole("option", { name: "Default" }));
expect(onChange).toHaveBeenCalledWith(llmValue);
});
it("clears the old effort when the classifier model changes", async () => {
const onChange = renderClassifier({
...llmValue,
classifier_llm_config: { model: "gpt-4", timeout_ms: 3000, reasoning_effort: "high" },
});
const user = userEvent.setup();
await user.click(screen.getByRole("combobox", { name: "Classifier Model" }));
await user.click(await screen.findByRole("option", { name: "gpt-3.5-turbo" }));
expect(onChange).toHaveBeenCalledWith({
...llmValue,
classifier_llm_config: { model: "gpt-3.5-turbo", timeout_ms: 3000 },
});
});
it.each(["click", "enter"] as const)(
"keeps the effort when the selected model is confirmed by %s",
async (action) => {
const onChange = renderClassifier({
...llmValue,
classifier_llm_config: { model: "gpt-4", timeout_ms: 3000, reasoning_effort: "high" },
});
const user = userEvent.setup();
await user.click(screen.getByRole("combobox", { name: "Classifier Model" }));
if (action === "click") await user.click(await screen.findByRole("option", { name: "gpt-4" }));
else await user.keyboard("{Enter}");
expect(onChange).not.toHaveBeenCalled();
},
);
it.each([
["gpt-4", "max", "max (unsupported)", /not supported by every deployment/],
["claude-3-opus", "low", "low (unverified)", /cannot be verified/],
])("keeps a saved exceptional value visible for %s", (model, effort, label, warning) => {
renderClassifier({
...llmValue,
classifier_llm_config: { model, timeout_ms: 3000, reasoning_effort: effort },
});
expect(screen.getByRole("combobox", { name: `Reasoning effort for classifier model ${model}` })).toHaveTextContent(
label,
);
expect(screen.getByText(warning)).toBeInTheDocument();
});
it.each(["claude-3-opus", "gpt-3.5-turbo"])("hides the effort control when %s has no advertised options", (model) => {
renderClassifier({
...llmValue,
classifier_llm_config: { model, timeout_ms: 3000 },
});
expect(
screen.queryByRole("combobox", { name: `Reasoning effort for classifier model ${model}` }),
).not.toBeInTheDocument();
});
});
describe("ComplexityRouterConfig reasoning effort gating", () => {
it("offers no effort select for a model group without reasoning support", () => {
renderWithProviders(<ComplexityRouterConfig {...baseProps} />);
@ -1232,12 +1335,13 @@ describe("ComplexityRouterConfig custom technical keywords", () => {
});
it("hides the keywords when the scorer never runs, so they cannot imply an effect they have none", () => {
openClassificationPanel({
const llmWithDefaultFallback = {
...defaultValue,
classifier_type: "llm",
classifier_type: "llm" as const,
classifier_llm_config: llmConfig,
classifier_fallback: "default_model",
});
classifier_fallback: "default_model" as const,
};
openClassificationPanel(llmWithDefaultFallback);
expect(screen.queryByText("Custom Technical Keywords")).not.toBeInTheDocument();
});
});

View file

@ -36,10 +36,11 @@ import ContextWindowEscalationConfig from "./ContextWindowEscalationConfig";
import { Restricted, restrictedBy } from "./TierRestrictions";
import { type TierSetAction, applyTierSetAction, setFallbackTier } from "./tier_set_actions";
import {
REASONING_EFFORT_OPTIONS,
ReasoningEffort,
TierModelParamsByTier,
classifierEffortOptionsForModels,
setTierModelReasoningEffort,
tierEffortOptionsForModels,
tierRowLabel,
} from "./complexity_router_tiers";
import TierModelEffortRows from "./TierModelEffortRows";
@ -124,6 +125,7 @@ export const CLASSIFICATION_RUBRIC_KEYS = Object.keys(CLASSIFICATION_RUBRIC_DESC
export interface ClassifierLLMConfig {
model: string;
timeout_ms: number;
reasoning_effort?: ReasoningEffort;
classification_rubric?: ClassificationRubric;
system_prompt?: string;
}
@ -632,14 +634,8 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
const removeTierRow = (id: string) => dispatch({ kind: "remove", id });
const exitToBuiltInTiers = () => dispatch({ kind: "restore" });
// An absent list means the proxy does not send the field yet, so every level is offered as before.
// An empty list is the group's own answer that its deployments share no level, and is left empty.
const effortOptionsByModel: Record<string, string[]> = Object.fromEntries(
modelInfo.map((model) => [
model.model_group,
model.supported_reasoning_efforts ?? (model.supports_reasoning ? [...REASONING_EFFORT_OPTIONS] : []),
]),
);
const tierEffortOptionsByModel = tierEffortOptionsForModels(modelInfo);
const classifierEffortOptionsByModel = classifierEffortOptionsForModels(modelInfo);
// Embedding models can't serve a chat-completion role, so they're excluded here.
const modelOptions = modelInfo
@ -746,7 +742,7 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
<TierModelEffortRows
tierLabel={label}
models={row.models}
effortOptionsByModel={effortOptionsByModel}
effortOptionsByModel={tierEffortOptionsByModel}
paramsByModel={row.params}
onEffortChange={(model, effort) => handleTierModelEffortChange(row.id, model, effort)}
/>
@ -818,6 +814,7 @@ const ComplexityRouterConfig: React.FC<ComplexityRouterConfigProps> = ({
value={value}
onChange={onChange}
modelOptions={modelOptions}
effortOptionsByModel={classifierEffortOptionsByModel}
customTechnicalKeywords={customTechnicalKeywords}
onCustomTechnicalKeywordsChange={onCustomTechnicalKeywordsChange}
showValidationErrors={showValidationErrors}

View file

@ -155,7 +155,11 @@ describe("HeuristicScoringConfig", () => {
});
describe("ClassificationMethodConfig scorer gating", () => {
const props = { onChange: vi.fn(), modelOptions: [{ value: "gpt-4o-mini", label: "gpt-4o-mini" }] };
const props = {
onChange: vi.fn(),
modelOptions: [{ value: "gpt-4o-mini", label: "gpt-4o-mini" }],
effortOptionsByModel: {},
};
const withClassifier = (type: ClassifierType, fallback?: ClassifierFallback): ComplexityRouterConfigValue => ({
...BASE,
classifier_type: type,

View file

@ -636,7 +636,14 @@ describe("AddAutoRouterTab", () => {
const labels = visibleOptions().map((option) => option.querySelector(".font-medium")?.textContent);
expect(labels).toEqual(["Anthropic Family", "Gemini Family", "Lite", "OpenAI Family", "Custom Configuration"]);
expect(labels).toEqual([
"1M Context",
"Anthropic Family",
"Gemini Family",
"Lite",
"OpenAI Family",
"Custom Configuration",
]);
});
describe("routing test", () => {
@ -1060,7 +1067,14 @@ describe("AddAutoRouterTab", () => {
expect(isOptionDisabled(optionByLabel("Anthropic Family")!)).toBe(false);
});
const labels = visibleOptions().map((option) => option.querySelector(".font-medium")?.textContent);
expect(labels).toEqual(["Anthropic Family", "Gemini Family", "Lite", "OpenAI Family", "Custom Configuration"]);
expect(labels).toEqual([
"Anthropic Family",
"1M Context",
"Gemini Family",
"Lite",
"OpenAI Family",
"Custom Configuration",
]);
});
it.each([

View file

@ -19,11 +19,12 @@ import { all_admin_roles } from "@/utils/roles";
import { type ModelWriteScope } from "@/utils/modelPermissions";
import TeamDropdown from "../common_components/team_dropdown";
import { type AddAutoRouterValues, handleAddAutoRouterSubmit } from "./handle_add_auto_router_submit";
import { fetchAvailableModels } from "@/components/llm_calls/fetch_models";
import { fetchAvailableModels, type ModelGroup } from "@/components/llm_calls/fetch_models";
import { autoRouterListKey, fetchAllModelDeployments } from "@/app/(dashboard)/hooks/models/useModels";
import ComplexityRouterConfig, {
ComplexityRouterConfigValue,
effectiveClassifierType,
usesLlmClassifier,
DEFAULT_ADAPTIVE_WEIGHTS,
DEFAULT_SESSION_AFFINITY,
DEFAULT_DEPLOYMENT_AFFINITY,
@ -37,6 +38,7 @@ import {
buildComplexityRouterConfig,
getKeywordTierRulesError,
getClassifierModelError,
getClassifierReasoningEffortError,
getMissingTiersError,
getPlanModeTierError,
getSemanticConfigError,
@ -120,14 +122,21 @@ export const getSubmitBlockedReason = (
config: ComplexityRouterConfigValue,
keywordTierRules: KeywordTierRule[],
referencedModelsParams: Parameters<typeof getReferencedModelsError>[0],
availability: ModelAvailability,
): string | null =>
(config.custom_tier_set ? getCustomTierRowsError(config.custom_tier_set) : getTierLabelsError(config.tier_labels)) ??
getMissingTiersError(activeTierRows(config)) ??
getPlanModeTierError(config.plan_mode_min_tier, activeTierRows(config)) ??
getKeywordTierRulesError(keywordTierRules, activeTierRows(config)) ??
getClassifierModelError(config) ??
getReferencedModelsError(referencedModelsParams, availability);
...capabilities: [availability: ModelAvailability, modelInfo?: readonly ModelGroup[]]
): string | null => {
const [availability, modelInfo = []] = capabilities;
return (
(config.custom_tier_set
? getCustomTierRowsError(config.custom_tier_set)
: getTierLabelsError(config.tier_labels)) ??
getMissingTiersError(activeTierRows(config)) ??
getPlanModeTierError(config.plan_mode_min_tier, activeTierRows(config)) ??
getKeywordTierRulesError(keywordTierRules, activeTierRows(config)) ??
getClassifierModelError(config) ??
getClassifierReasoningEffortError(config, modelInfo) ??
getReferencedModelsError(referencedModelsParams, availability)
);
};
const autoRouterSchema = (requiresTeamScope: boolean) =>
z.object({
@ -338,6 +347,7 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({
keywordTierRules,
referencedModelsParams,
groupsOnlyAvailability,
modelInfo,
);
const complexityRouterConfigParams: BuildComplexityRouterConfigParams = {
@ -390,6 +400,7 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({
keywordTierRules,
referencedModelsParams,
groupsOnlyAvailability,
modelInfo,
) ?? getSemanticConfigError({ semanticMatchingEnabled, embeddingModel, keywordTierRules });
if (blockedReason) {
setShowValidationErrors(true);
@ -462,6 +473,12 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({
semanticMatchingEnabled,
embeddingModel,
defaultModel: resolveComplexityDefaultModel(complexityRouterConfig, complexityRouterConfig.default_model),
classifier: usesLlmClassifier(effectiveClassifierType(complexityRouterConfig))
? {
model: complexityRouterConfig.classifier_llm_config?.model ?? "",
reasoningEffort: complexityRouterConfig.classifier_llm_config?.reasoning_effort,
}
: undefined,
};
const targets = buildAutoRouterTestTargets(testTargetParams);

View file

@ -17,6 +17,7 @@ const targets: AutoRouterTestTarget[] = [
{ labels: ["SIMPLE"], modelGroup: "gpt-4o-mini", mode: "chat" },
{ labels: ["MEDIUM", "COMPLEX"], modelGroup: "claude-sonnet-4", mode: "chat" },
{ labels: ["Embedding"], modelGroup: "voyage-3-5", mode: "embedding" },
{ labels: ["Classifier"], modelGroup: "gpt-5-mini", mode: "chat", requestParams: { reasoning_effort: "low" } },
];
describe("AutoRouterConnectionTest", () => {
@ -30,11 +31,12 @@ describe("AutoRouterConnectionTest", () => {
renderWithProviders(<AutoRouterConnectionTest accessToken="sk-test" targets={targets} />);
await waitFor(() => expect(mock).toHaveBeenCalledTimes(3));
await waitFor(() => expect(mock).toHaveBeenCalledTimes(4));
expect(mock).toHaveBeenCalledWith("sk-test", "gpt-4o-mini", "chat");
expect(mock).toHaveBeenCalledWith("sk-test", "claude-sonnet-4", "chat");
expect(mock).toHaveBeenCalledWith("sk-test", "voyage-3-5", "embedding");
expect(mock).toHaveBeenCalledWith("sk-test", "gpt-5-mini", "chat", { reasoning_effort: "low" });
});
it("shows a success indicator per target when the routing probe passes", async () => {
@ -43,7 +45,7 @@ describe("AutoRouterConnectionTest", () => {
renderWithProviders(<AutoRouterConnectionTest accessToken="sk-test" targets={targets} />);
await waitFor(() => expect(screen.getAllByTestId("test-status-success")).toHaveLength(3));
await waitFor(() => expect(screen.getAllByTestId("test-status-success")).toHaveLength(4));
expect(screen.queryByTestId("test-status-error")).not.toBeInTheDocument();
expect(screen.getByText("MEDIUM, COMPLEX")).toBeInTheDocument();
});
@ -63,7 +65,7 @@ describe("AutoRouterConnectionTest", () => {
expect(await screen.findByTestId("test-error-message")).toBeInTheDocument();
expect(screen.getByTestId("test-error-message")).toHaveTextContent("invalid api key");
expect(screen.getByTestId("test-error-message")).not.toHaveTextContent("litellm.AuthenticationError");
expect(screen.getAllByTestId("test-status-success")).toHaveLength(2);
expect(screen.getAllByTestId("test-status-success")).toHaveLength(3);
});
it("renders a non-litellm error string verbatim", async () => {

View file

@ -29,7 +29,9 @@ const AutoRouterConnectionTest: React.FC<AutoRouterConnectionTestProps> = ({
const run = async () => {
await Promise.all(
targets.map(async (target, index) => {
const result = await testModelGroupConnection(accessToken, target.modelGroup, target.mode);
const result = target.requestParams
? await testModelGroupConnection(accessToken, target.modelGroup, target.mode, target.requestParams)
: await testModelGroupConnection(accessToken, target.modelGroup, target.mode);
if (cancelled) return;
const cleaned: TargetResult =
result.status === "error" ? { status: "error", error: cleanErrorMessage(result.error) } : result;
@ -56,14 +58,14 @@ const AutoRouterConnectionTest: React.FC<AutoRouterConnectionTestProps> = ({
return (
<div className="space-y-3">
<p className="mb-2 text-sm text-muted-foreground">
Each configured tier routes to a saved model group. Test Connection sends a minimal request through the proxy to
each one, exactly as the auto router would.
Test Connection sends a minimal request to every configured tier, classifier, default, and embedding model. The
classifier probe includes its reasoning effort override.
</p>
{targets.map((target, index) => {
const result = results[index] ?? { status: "pending" };
return (
<div
key={`${target.modelGroup}-${target.mode}`}
key={`${target.labels.join("-")}-${target.modelGroup}-${target.mode}`}
data-testid="auto-router-test-row"
className="flex items-start gap-3 rounded-lg border p-3"
>

View file

@ -126,4 +126,32 @@ describe("buildAutoRouterTestTargets", () => {
});
expect(targets).toEqual([{ labels: ["SIMPLE"], modelGroup: "gpt-4o-mini", mode: "chat" }]);
});
it("adds a distinct classifier probe with its reasoning effort", () => {
const targets = buildAutoRouterTestTargets({
tiers: tierEntries(["gpt-5-mini"]),
semanticMatchingEnabled: false,
embeddingModel: undefined,
classifier: { model: "gpt-5-mini", reasoningEffort: "low" },
});
expect(targets).toEqual([
{ labels: ["SIMPLE"], modelGroup: "gpt-5-mini", mode: "chat" },
{
labels: ["Classifier"],
modelGroup: "gpt-5-mini",
mode: "chat",
requestParams: { reasoning_effort: "low" },
},
]);
});
it("omits an empty classifier and omits params when the classifier uses provider defaults", () => {
const base = { tiers: tierEntries([]), semanticMatchingEnabled: false, embeddingModel: undefined };
const emptyClassifier = { ...base, classifier: { model: " " } };
const providerDefaultClassifier = { ...base, classifier: { model: "gpt-5-mini" } };
expect(buildAutoRouterTestTargets(emptyClassifier)).toEqual([]);
expect(buildAutoRouterTestTargets(providerDefaultClassifier)).toEqual([
{ labels: ["Classifier"], modelGroup: "gpt-5-mini", mode: "chat" },
]);
});
});

View file

@ -4,6 +4,7 @@ export interface AutoRouterTestTarget {
labels: string[];
modelGroup: string;
mode: AutoRouterTestMode;
requestParams?: Record<string, unknown>;
}
export interface BuildAutoRouterTestTargetsParams {
@ -14,6 +15,7 @@ export interface BuildAutoRouterTestTargetsParams {
/** The resolved default model - see resolveComplexityDefaultModel. A live fallback destination,
* so it is probed even when no tier lists it. */
defaultModel?: string;
classifier?: { model: string; reasoningEffort?: string };
}
export const buildAutoRouterTestTargets = ({
@ -21,6 +23,7 @@ export const buildAutoRouterTestTargets = ({
semanticMatchingEnabled,
embeddingModel,
defaultModel,
classifier,
}: BuildAutoRouterTestTargetsParams): AutoRouterTestTarget[] => {
const tieredByModel = tiers.reduce<Record<string, string[]>>((acc, [tier, models]) => {
return models.reduce((tierAcc, rawModel) => {
@ -50,5 +53,17 @@ export const buildAutoRouterTestTargets = ({
? [{ labels: ["Embedding"], modelGroup: embeddingModel.trim(), mode: "embedding" as const }]
: [];
return [...tierTargets, ...embeddingTarget];
const classifierModel = classifier?.model.trim();
const classifierTarget: AutoRouterTestTarget[] = classifierModel
? [
{
labels: ["Classifier"],
modelGroup: classifierModel,
mode: "chat",
...(classifier?.reasoningEffort && { requestParams: { reasoning_effort: classifier.reasoningEffort } }),
},
]
: [];
return [...tierTargets, ...embeddingTarget, ...classifierTarget];
};

View file

@ -4,6 +4,7 @@ import {
normalizeClassifierLlmConfig,
getKeywordTierRulesError,
getClassifierModelError,
getClassifierReasoningEffortError,
getMissingTiersError,
hydrateCustomTierSet,
getSemanticConfigError,
@ -506,6 +507,19 @@ describe("classifier prompt and fallback", () => {
expect(buildComplexityRouterConfig(llmParams)).not.toHaveProperty("classifier_fallback");
});
it("keeps an explicit classifier reasoning effort", () => {
const config = buildComplexityRouterConfig({
...llmParams,
classifierLlmConfig: { model: "haiku-classifier", timeout_ms: 400, reasoning_effort: "low" },
});
expect(config.classifier_llm_config?.reasoning_effort).toBe("low");
});
it("omits classifier reasoning effort when the provider default is selected", () => {
const config = buildComplexityRouterConfig(llmParams);
expect(config.classifier_llm_config).not.toHaveProperty("reasoning_effort");
});
it("sends the chat preset the operator picked", () => {
const config = buildComplexityRouterConfig({
...llmParams,
@ -540,12 +554,16 @@ describe("classifier prompt and fallback", () => {
});
it("normalizeClassifierLlmConfig leaves a real prompt untouched and strips an empty one", () => {
expect(normalizeClassifierLlmConfig({ model: "m", timeout_ms: 1, system_prompt: "x" })).toEqual({
const customPromptConfig = { model: "m", timeout_ms: 1, reasoning_effort: "none" as const, system_prompt: "x" };
const emptyPromptConfig = { model: "m", timeout_ms: 1, system_prompt: "" };
const expectedCustomPromptConfig = {
model: "m",
timeout_ms: 1,
reasoning_effort: "none",
system_prompt: "x",
});
expect(normalizeClassifierLlmConfig({ model: "m", timeout_ms: 1, system_prompt: "" })).toEqual({
};
expect(normalizeClassifierLlmConfig(customPromptConfig)).toEqual(expectedCustomPromptConfig);
expect(normalizeClassifierLlmConfig(emptyPromptConfig)).toEqual({
model: "m",
timeout_ms: 1,
});
@ -767,6 +785,26 @@ describe("getClassifierModelError", () => {
});
});
describe("getClassifierReasoningEffortError", () => {
const classifier = {
classifier_type: "llm" as const,
classifier_llm_config: { model: "classifier", timeout_ms: 3000, reasoning_effort: "low" },
};
it.each([
[["low", "medium"], null],
[["medium", "high"], "low reasoning effort is not supported"],
[null, null],
[undefined, null],
])("validates capability levels %o", (supportedReasoningEfforts, expectedError) => {
const error = getClassifierReasoningEffortError(classifier, [
{ model_group: "classifier", supported_reasoning_efforts: supportedReasoningEfforts },
]);
if (expectedError) expect(error).toContain(expectedError);
else expect(error).toBeNull();
});
});
describe("getKeywordTierRulesError orphaned tiers", () => {
const rows = activeTierRows({ tiers });
@ -944,11 +982,16 @@ describe("buildComplexityRouterConfig with an edited tier set", () => {
classifierLlmConfig: {
model: "gpt-4o-mini",
timeout_ms: 3000,
reasoning_effort: "low",
system_prompt: "replace the whole rubric",
classification_rubric: "agentic",
},
});
expect(payload.classifier_llm_config).toEqual({ model: "gpt-4o-mini", timeout_ms: 3000 });
expect(payload.classifier_llm_config).toEqual({
model: "gpt-4o-mini",
timeout_ms: 3000,
reasoning_effort: "low",
});
});
it.each(

View file

@ -1,4 +1,5 @@
import { KeywordTierRule } from "./KeywordTierRules";
import type { ModelGroup } from "../llm_calls/fetch_models";
import {
type CustomTierSet,
type TierRow,
@ -55,12 +56,18 @@ import {
export const normalizeClassifierLlmConfig = ({
model,
timeout_ms,
reasoning_effort,
classification_rubric,
system_prompt,
}: ClassifierLLMConfig): ClassifierLLMConfig =>
system_prompt?.trim()
? { model, timeout_ms, system_prompt }
: { model, timeout_ms, ...(classification_rubric && { classification_rubric }) };
? { model, timeout_ms, ...(reasoning_effort && { reasoning_effort }), system_prompt }
: {
model,
timeout_ms,
...(reasoning_effort && { reasoning_effort }),
...(classification_rubric && { classification_rubric }),
};
interface ScorerKnobInputs {
classifierType: ClassifierType;
@ -268,6 +275,20 @@ export const getClassifierModelError = (
: "Please select a classifier model, or switch back to Heuristic";
};
export const getClassifierReasoningEffortError = (
config: Pick<ComplexityRouterConfigValue, "custom_tier_set" | "classifier_type" | "classifier_llm_config">,
modelInfo: readonly ModelGroup[],
): string | null => {
if (!usesLlmClassifier(effectiveClassifierType(config))) return null;
const classifierConfig = config.classifier_llm_config;
if (!classifierConfig?.model || !classifierConfig.reasoning_effort) return null;
const supported = modelInfo.find(
(model) => model.model_group === classifierConfig.model,
)?.supported_reasoning_efforts;
if (!Array.isArray(supported) || supported.includes(classifierConfig.reasoning_effort)) return null;
return `${classifierConfig.reasoning_effort} reasoning effort is not supported by every deployment in ${classifierConfig.model}. Choose Default or a supported value.`;
};
export const getSemanticConfigError = ({
semanticMatchingEnabled,
embeddingModel,
@ -295,11 +316,15 @@ export const customTierWireFields = (
tier_definitions: tierDefinitionsFromRows(rows),
...(fallback && { fallback_tier: activeTierName(fallback) }),
classifier_type: "llm",
// Rebuilt from the two fields an edited tier set allows. The backend rejects system_prompt and
// Rebuilt from the fields an edited tier set allows. The backend rejects system_prompt and
// classification_rubric beside tier_definitions, and both live inside this object rather than at
// the top level the omit list covers. The opening instructions ride classification_prompt below.
...(classifierLlmConfig && {
classifier_llm_config: { model: classifierLlmConfig.model, timeout_ms: classifierLlmConfig.timeout_ms },
classifier_llm_config: {
model: classifierLlmConfig.model,
timeout_ms: classifierLlmConfig.timeout_ms,
...(classifierLlmConfig.reasoning_effort && { reasoning_effort: classifierLlmConfig.reasoning_effort }),
},
}),
session_affinity: false,
...(classificationPrompt?.trim() && { classification_prompt: classificationPrompt.trim() }),

View file

@ -1,4 +1,5 @@
import type { ComplexityTier } from "./KeywordTierRules";
import type { ModelGroup } from "@/components/llm_calls/fetch_models";
import { TIER_ORDER } from "./tier_rows";
export type TierModelParams = Record<string, unknown>;
@ -18,6 +19,23 @@ export const REASONING_EFFORT_OPTIONS = ["none", "minimal", "low", "medium", "hi
*/
export type ReasoningEffort = (typeof REASONING_EFFORT_OPTIONS)[number] | (string & {});
export const tierEffortOptionsForModels = (modelInfo: ModelGroup[]): Record<string, string[]> =>
Object.fromEntries(
modelInfo.map((model) => [
model.model_group,
model.supported_reasoning_efforts ?? (model.supports_reasoning ? [...REASONING_EFFORT_OPTIONS] : []),
]),
);
/**
* Stricter than the tier variant on purpose: classifier overrides are new in this release, so an
* unknown capability list stays unknown instead of inventing provider levels.
*/
export const classifierEffortOptionsForModels = (
modelInfo: ModelGroup[],
): Record<string, string[] | null | undefined> =>
Object.fromEntries(modelInfo.map((model) => [model.model_group, model.supported_reasoning_efforts]));
const asRecord = (raw: unknown): Record<string, unknown> | undefined =>
typeof raw === "object" && raw !== null && !Array.isArray(raw) ? (raw as Record<string, unknown>) : undefined;

View file

@ -16,7 +16,7 @@ const storedCustomConfig = (overrides: Record<string, unknown> = {}) => ({
],
fallback_tier: "CASUAL",
classifier_type: "llm",
classifier_llm_config: { model: "gpt-4o-mini", timeout_ms: 3000 },
classifier_llm_config: { model: "gpt-4o-mini", timeout_ms: 3000, reasoning_effort: "low" },
...overrides,
});
@ -111,7 +111,7 @@ describe("buildUpdatedComplexityRouterConfig keyword matching", () => {
const STORED_LLM = {
tiers: { SIMPLE: ["gpt-4o-mini"], MEDIUM: [], COMPLEX: [], REASONING: [] },
classifier_type: "llm",
classifier_llm_config: { model: "gpt-4o-mini", timeout_ms: 3000 },
classifier_llm_config: { model: "gpt-4o-mini", timeout_ms: 3000, reasoning_effort: "low" },
classifier_context_window_size: 5,
classifier_context_per_turn_chars: 300,
};
@ -522,7 +522,7 @@ describe("managed keys survive an untouched open-and-save", () => {
tier_labels: { SIMPLE: "Cheap" },
classifier_type: "heuristic_first",
heuristic_first_max_tier: "SIMPLE",
classifier_llm_config: { model: "gpt-4o-mini", timeout_ms: 3000 },
classifier_llm_config: { model: "gpt-4o-mini", timeout_ms: 3000, reasoning_effort: "low" },
classifier_context_window_size: 5,
classifier_context_budget_chars: 4000,
classifier_context_include_assistant_turns: true,

View file

@ -28,6 +28,7 @@ import {
type BuildComplexityRouterConfigParams,
buildComplexityRouterConfig,
getClassifierModelError,
getClassifierReasoningEffortError,
getKeywordTierRulesError,
getMissingTiersError,
getSemanticConfigError,
@ -560,6 +561,12 @@ const EditAutoRouterModal: React.FC<EditAutoRouterModalProps> = ({
toast.fromError(classifierError);
return;
}
const classifierEffortError = getClassifierReasoningEffortError(complexityRouterConfig, modelInfo);
if (classifierEffortError) {
setShowValidationErrors(true);
toast.fromError(classifierEffortError);
return;
}
// Same guards the create form applies (add_auto_router_tab.tsx). The backend rejects a
// keyword rule with no keyword, and semantic_keyword_matching without an embedding model
// or keyword rules (complexity_router/config.py), so without these a save fails as a raw

View file

@ -52,6 +52,24 @@ describe("fetchAvailableModels", () => {
]);
});
it("preserves absent, unknown, empty, and explicit effort capability states", async () => {
modelHubCallMock.mockResolvedValue({
data: [
{ model_group: "absent", supports_reasoning: true },
{ model_group: "unknown", supports_reasoning: true, supported_reasoning_efforts: null },
{ model_group: "empty", supports_reasoning: true, supported_reasoning_efforts: [] },
{ model_group: "known", supports_reasoning: true, supported_reasoning_efforts: ["low"] },
],
});
expect(await fetchAvailableModels("token")).toEqual([
{ model_group: "absent", supports_reasoning: true },
{ model_group: "empty", supports_reasoning: true, supported_reasoning_efforts: [] },
{ model_group: "known", supports_reasoning: true, supported_reasoning_efforts: ["low"] },
{ model_group: "unknown", supports_reasoning: true, supported_reasoning_efforts: null },
]);
});
it.each([
["an error payload in place of the list", { data: { error: "no access" } }],
["a missing data key", {}],

View file

@ -7,7 +7,7 @@ export interface ModelGroup {
model_group: string;
mode?: string;
supports_reasoning?: boolean;
supported_reasoning_efforts?: string[];
supported_reasoning_efforts?: string[] | null;
}
interface AvailableModel {
@ -25,7 +25,9 @@ const toModelGroup = (item: AvailableModel): ModelGroup => {
model_group: groupName,
...(item.mode && { mode: item.mode }),
...(item.supports_reasoning === true && { supports_reasoning: true }),
...(item.supported_reasoning_efforts && { supported_reasoning_efforts: item.supported_reasoning_efforts }),
...(item.supported_reasoning_efforts !== undefined && {
supported_reasoning_efforts: item.supported_reasoning_efforts,
}),
};
};

View file

@ -617,6 +617,15 @@ describe("buildModelGroupTestRequest", () => {
expect(path).toBe("/v1/embeddings");
expect(body).toEqual({ model: "text-embedding-3-small", input: "test from litellm" });
});
it("adds classifier request parameters to a chat probe", () => {
const { body } = Networking.buildModelGroupTestRequest("gpt-5-mini", "chat", { reasoning_effort: "low" });
expect(body).toEqual({
model: "gpt-5-mini",
messages: [{ role: "user", content: "test from litellm" }],
reasoning_effort: "low",
});
});
});
describe("testMCPToolsListRequest auth headers", () => {

View file

@ -2459,20 +2459,22 @@ export type ModelGroupConnectionResult = { status: "success" } | { status: "erro
export const buildModelGroupTestRequest = (
modelGroup: string,
mode: "chat" | "embedding",
requestParams: Record<string, unknown> = {},
): { path: string; body: Record<string, unknown> } =>
mode === "embedding"
? { path: "/v1/embeddings", body: { model: modelGroup, input: "test from litellm" } }
: {
path: "/v1/chat/completions",
body: { model: modelGroup, messages: [{ role: "user", content: "test from litellm" }] },
body: { ...requestParams, model: modelGroup, messages: [{ role: "user", content: "test from litellm" }] },
};
export const testModelGroupConnection = async (
accessToken: string,
modelGroup: string,
mode: "chat" | "embedding",
requestParams?: Record<string, unknown>,
): Promise<ModelGroupConnectionResult> => {
const { path, body } = buildModelGroupTestRequest(modelGroup, mode);
const { path, body } = buildModelGroupTestRequest(modelGroup, mode, requestParams);
try {
await apiClient.post(path, { accessToken, body });
return { status: "success" };

View file

@ -750,6 +750,19 @@ describe("CreateKey", () => {
expect((await createdPayload()).organization_id).toBe("org-1");
});
it("drops organization_id when the chosen organization is cleared again", async () => {
state.organizations = [{ organization_id: "org-1", organization_alias: "Engineering" }];
await openModal();
await nameTheKey();
await userEvent.click(await screen.findByLabelText("Organization"));
await userEvent.click(await screen.findByRole("option", { name: /Engineering/ }));
await userEvent.click(await screen.findByRole("button", { name: "Clear" }));
await submit();
expect((await createdPayload()).organization_id).toBeUndefined();
});
});
describe("policy and prompt fields", () => {

View file

@ -588,7 +588,7 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey, autoOp
};
const changeOrganization = (write: FieldWrite) => (orgId: string) => {
write(orgId);
write(orgId || undefined);
setSelectedOrganizationId(orgId || null);
// Clear team and project when org changes
setSelectedCreateKeyTeam(null);

Some files were not shown because too many files have changed in this diff Show more