mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
commit
9036b90a38
108 changed files with 4560 additions and 941 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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`.
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
119
litellm/proxy/management_helpers/access_group_model_sync.py
Normal file
119
litellm/proxy/management_helpers/access_group_model_sync.py
Normal 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)
|
||||
|
|
@ -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.",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -482,7 +482,6 @@ def search(
|
|||
timeout=timeout or request_timeout,
|
||||
_is_async=_is_async,
|
||||
client=kwargs.get("client"),
|
||||
router=router,
|
||||
)
|
||||
|
||||
return response
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 = {}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)."""
|
||||
|
|
|
|||
|
|
@ -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",))
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1399,9 +1399,6 @@
|
|||
},
|
||||
"prefer-const": {
|
||||
"count": 2
|
||||
},
|
||||
"react-hooks/set-state-in-effect": {
|
||||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/TeamsPage/teamTableColumns.tsx": {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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} />}
|
||||
|
|
|
|||
|
|
@ -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("");
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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([
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
||||
|
|
|
|||
|
|
@ -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 () => {
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
>
|
||||
|
|
|
|||
|
|
@ -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" },
|
||||
]);
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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];
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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() }),
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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", {}],
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}),
|
||||
};
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -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", () => {
|
||||
|
|
|
|||
|
|
@ -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" };
|
||||
|
|
|
|||
|
|
@ -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", () => {
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue