From 7b8cc0319e70d90dd096fbe86830b987a67ef491 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 2 Sep 2026 23:56:09 -0700 Subject: [PATCH 1/9] fix(proxy-extras): only spend a migrate-deploy attempt when a pass made no progress The v2 migration resolver gave `prisma migrate deploy` four attempts, and every recovery path ended in a bare `continue`, so each one burned an attempt. A database first brought up with `--use_prisma_db_push` has a full schema and no migrations ledger, so the baseline spent attempt one and the first three migrations whose objects already existed spent the rest. The proxy then exited before binding its port, and that database could never be moved onto the resolver. The retry budget now counts only attempts that got nowhere. Creating the baseline, and each migration newly marked applied, leaves the budget alone, so a push-created database works through its pre-existing objects one pass at a time. Timeouts, deadlock rollbacks, advisory-lock waits, and a repeat of a recovery that already ran still spend an attempt, so a run that stops making progress gives up exactly as before. --- .../litellm_proxy_extras/utils.py | 67 ++++++- .../test_litellm_proxy_extras_utils.py | 167 ++++++++++++++++++ 2 files changed, 226 insertions(+), 8 deletions(-) diff --git a/litellm-proxy-extras/litellm_proxy_extras/utils.py b/litellm-proxy-extras/litellm_proxy_extras/utils.py index d22484bc0e8..21f255b0666 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/utils.py +++ b/litellm-proxy-extras/litellm_proxy_extras/utils.py @@ -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,10 +804,11 @@ 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, ) + budget = budget.spend() time.sleep(random.randrange(5, 15)) continue @@ -781,7 +819,14 @@ class ProxyExtrasDBManager: logger.info( "Schema exists but no migrations ledger — creating baseline" ) - ProxyExtrasDBManager._create_baseline_migration(schema_path) + baselined = ProxyExtrasDBManager._create_baseline_migration( + schema_path + ) + budget = ( + budget.after_recovery("baseline") + if baselined + else budget.spend() + ) continue if "P3009" in stderr: @@ -818,6 +863,7 @@ class ProxyExtrasDBManager: f"intervention may be required.\n\n" f"Detail: {resolve_err}" ) from resolve_err + budget = budget.after_recovery(f"resolved:{name}") continue if migration_match: migration_name = migration_match.group(1) @@ -831,6 +877,7 @@ class ProxyExtrasDBManager: migration_name, ) ProxyExtrasDBManager._roll_back_migration_best_effort(migration_name) + budget = budget.spend() time.sleep(random.randrange(5, 15)) continue raise RuntimeError( @@ -876,6 +923,7 @@ class ProxyExtrasDBManager: f"intervention may be required.\n\n" f"Detail: {resolve_err}" ) from resolve_err + budget = budget.after_recovery(f"resolved:{name}") continue if migration_match and _MIGRATION_DEADLOCK_MARKER in stderr: @@ -888,6 +936,7 @@ class ProxyExtrasDBManager: ProxyExtrasDBManager._roll_back_migration_best_effort( migration_match.group(1) ) + budget = budget.spend() time.sleep(random.randrange(5, 15)) continue @@ -900,8 +949,9 @@ class ProxyExtrasDBManager: logger.info( "prisma migrate deploy attempt %s deadlocked against " "a concurrent migrate deploy, retrying", - attempt + 1, + budget.attempt_number, ) + budget = budget.spend() time.sleep(random.randrange(5, 15)) continue @@ -909,8 +959,9 @@ class ProxyExtrasDBManager: logger.info( "prisma migrate deploy attempt %s timed out waiting for " "the advisory lock a concurrent migrate deploy holds, retrying", - attempt + 1, + budget.attempt_number, ) + budget = budget.spend() time.sleep(random.randrange(5, 15)) continue @@ -920,9 +971,9 @@ class ProxyExtrasDBManager: ) from e 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." ) diff --git a/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py b/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py index b3d457707b8..22558faedad 100644 --- a/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py +++ b/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py @@ -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 == [] From 8572544b44662f00f2cea9a573bc6bf9977faa9d Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 3 Sep 2026 00:20:46 -0700 Subject: [PATCH 2/9] refactor(proxy-extras): pull the migrate deploy recovery branches into a budget helper _setup_database_v2 decided the next attempt budget inline in eight branches, each rebinding budget before continuing. The branches now live in _budget_after_deploy_failure, which returns the budget the next pass runs under, and the two identical idempotent-recovery blocks share _mark_migration_applied. The loop backs off whenever a pass spent an attempt, which is the same set of paths that slept before. --- .../litellm_proxy_extras/utils.py | 289 ++++++++---------- 1 file changed, 131 insertions(+), 158 deletions(-) diff --git a/litellm-proxy-extras/litellm_proxy_extras/utils.py b/litellm-proxy-extras/litellm_proxy_extras/utils.py index 21f255b0666..a4ea4789b49 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/utils.py +++ b/litellm-proxy-extras/litellm_proxy_extras/utils.py @@ -808,167 +808,16 @@ class ProxyExtrasDBManager: deploy_timeout, PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR, ) - budget = budget.spend() - 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" - ) - baselined = ProxyExtrasDBManager._create_baseline_migration( - schema_path - ) - budget = ( - budget.after_recovery("baseline") - if baselined - else budget.spend() - ) - 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 - budget = budget.after_recovery(f"resolved:{name}") - 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) - budget = budget.spend() - 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 - budget = budget.after_recovery(f"resolved:{name}") - 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) - ) - budget = budget.spend() - 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", - budget.attempt_number, - ) - budget = budget.spend() - 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", - budget.attempt_number, - ) - budget = budget.spend() - 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( f"Database migration failed after {MAX_MIGRATE_DEPLOY_ATTEMPTS} " @@ -980,6 +829,130 @@ class ProxyExtrasDBManager: 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: """ From e80e78d3ef9a65db5cd2382d4cb7edc7b2554773 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 3 Sep 2026 08:38:25 -0700 Subject: [PATCH 3/9] feat(cli): enable Claude Code gateway model discovery by default in lite claude (#39445) * feat(cli): enable Claude Code gateway model discovery by default in lite claude Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(cli): build agent env declaratively and document discovery key for lite up Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(cli): keep build_agent_env within LIT002 type-discipline budget Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/client/cli/README.md | 6 +++--- litellm/proxy/client/cli/commands/agents.py | 8 +++++++- litellm/proxy/client/cli/commands/claude_settings.py | 11 +++++++++-- tests/test_litellm/proxy/client/cli/test_agents.py | 11 +++++++++++ .../test_litellm/proxy/client/cli/test_up_commands.py | 7 +++++++ 5 files changed, 37 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/client/cli/README.md b/litellm/proxy/client/cli/README.md index 3ddce35b53d..47355f328dd 100644 --- a/litellm/proxy/client/cli/README.md +++ b/litellm/proxy/client/cli/README.md @@ -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`. diff --git a/litellm/proxy/client/cli/commands/agents.py b/litellm/proxy/client/cli/commands/agents.py index c591cbabee1..baa21996c7e 100644 --- a/litellm/proxy/client/cli/commands/agents.py +++ b/litellm/proxy/client/cli/commands/agents.py @@ -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 diff --git a/litellm/proxy/client/cli/commands/claude_settings.py b/litellm/proxy/client/cli/commands/claude_settings.py index 46af641636e..5e3ce95f088 100644 --- a/litellm/proxy/client/cli/commands/claude_settings.py +++ b/litellm/proxy/client/cli/commands/claude_settings.py @@ -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", diff --git a/tests/test_litellm/proxy/client/cli/test_agents.py b/tests/test_litellm/proxy/client/cli/test_agents.py index 0191dad3d94..62b94e948be 100644 --- a/tests/test_litellm/proxy/client/cli/test_agents.py +++ b/tests/test_litellm/proxy/client/cli/test_agents.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( diff --git a/tests/test_litellm/proxy/client/cli/test_up_commands.py b/tests/test_litellm/proxy/client/cli/test_up_commands.py index c78bdfa75b1..aead1764b0e 100644 --- a/tests/test_litellm/proxy/client/cli/test_up_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_up_commands.py @@ -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" From 1de960bce7a4ae0ddd4bba419efc25044511f749 Mon Sep 17 00:00:00 2001 From: Rakesh <51485939+rakeshrepository@users.noreply.github.com> Date: Thu, 3 Sep 2026 21:16:07 +0530 Subject: [PATCH 4/9] fix(docker): bump nginx runtime to 1.31.5-alpine3.24 and pin digest to resolve critical CVEs (#39561) --- ui/Dockerfile | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ui/Dockerfile b/ui/Dockerfile index 24140093270..f14b631685a 100644 --- a/ui/Dockerfile +++ b/ui/Dockerfile @@ -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 From 34d4f7f8aef2951ffcf5ee04f33bff69aab4fd9d Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 3 Sep 2026 08:58:48 -0700 Subject: [PATCH 5/9] fix: 1.99.0-rc2 UI bug batch (empty org on key create, session pagination, access group rename/delete) (#39436) * fix(ui): clearing the organization picker no longer sends organization_id="" on key create * fix(proxy): paginate Request Logs by conversation and aggregate session type counts and models server-side * fix(proxy): keep access groups in sync when a model is renamed or deleted * fix(proxy): cap the Request Logs conversation total like the row total * fix(proxy): judge access group backing by the database for db models A worker whose router has not polled the database yet still lists a sibling under its old name, so a delete or rename handled there kept the stale name in every access group. Only config-sourced deployments count as router backing now; db models are counted in the table. * fix(ui): keep the conversation badge when an MCP call represents a conversation A conversation that straddles the bounded page window can be represented by one of its MCP rows, which showed a plain MCP badge and hid the session counts. The badge now reads the server aggregates whenever the conversation has more than one call. * fix(proxy): list every model of a conversation in Request Logs and keep the conversation badge for MCP representatives Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): type session spend aggregates Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(ui): satisfy request logs lint budget Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): cap per-session model aggregation in request logs Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore: ratchet type-discipline budget after staging merge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(ui): send an explicit null when the key edit form clears the organization Clearing the Organization picker in the key edit form wrote undefined into the form value, and JSON.stringify drops undefined-valued keys, so /key/update never saw the field and the key kept its old organization. Writing null instead survives serialization, and the backend's model_dump(exclude_unset=True) preserves it, so the column is set to NULL. --------- Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_types.py | 11 +- .../model_management_endpoints.py | 54 ++++- .../access_group_model_sync.py | 119 +++++++++++ .../spend_management_endpoints.py | 116 +++++++---- .../test_key_management_endpoints.py | 10 + .../test_model_management_endpoints.py | 188 ++++++++++++++++++ .../test_access_group_model_sync.py | 170 ++++++++++++++++ .../test_spend_management_endpoints.py | 51 +++++ type-discipline-budget.json | 2 +- .../create_key_button.integration.test.tsx | 13 ++ .../organisms/create_key_button.tsx | 2 +- .../templates/key_edit_view.test.tsx | 26 +++ .../components/templates/key_edit_view.tsx | 4 +- .../RequestLogsTableColumns.test.tsx | 62 ++++++ .../view_logs/RequestLogsTableColumns.tsx | 24 ++- .../src/components/view_logs/columns.tsx | 2 + 16 files changed, 797 insertions(+), 57 deletions(-) create mode 100644 litellm/proxy/management_helpers/access_group_model_sync.py create mode 100644 tests/test_litellm/proxy/management_helpers/test_access_group_model_sync.py diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 849e54c65aa..c6ecb0be8b9 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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: diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index ca66640bf46..613e726f89d 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -68,6 +68,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, @@ -715,6 +719,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, @@ -741,6 +746,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() @@ -1673,6 +1692,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( @@ -2027,25 +2052,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() diff --git a/litellm/proxy/management_helpers/access_group_model_sync.py b/litellm/proxy/management_helpers/access_group_model_sync.py new file mode 100644 index 00000000000..b9d81f2981f --- /dev/null +++ b/litellm/proxy/management_helpers/access_group_model_sync.py @@ -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) diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index bb7dfafb297..dcc798b81cb 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -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: diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 7e2e680743f..68704f476b5 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -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).""" diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index 4661cc17dbc..5fa59a85c9d 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -1,5 +1,6 @@ import inspect import asyncio +import contextlib import json from typing import Dict, Optional from unittest.mock import AsyncMock, MagicMock, patch @@ -875,6 +876,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) @@ -936,6 +938,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) @@ -2079,6 +2082,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 @@ -2191,6 +2195,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 ) @@ -2273,6 +2278,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 ) @@ -2349,6 +2355,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 ) @@ -2434,6 +2441,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 ) @@ -2515,6 +2523,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 ) @@ -2584,6 +2593,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 ) @@ -2702,6 +2712,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 ) @@ -3844,6 +3855,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) @@ -4654,3 +4666,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",)) diff --git a/tests/test_litellm/proxy/management_helpers/test_access_group_model_sync.py b/tests/test_litellm/proxy/management_helpers/test_access_group_model_sync.py new file mode 100644 index 00000000000..65ef2d55cb8 --- /dev/null +++ b/tests/test_litellm/proxy/management_helpers/test_access_group_model_sync.py @@ -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() diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index 30b086bab61..c8b35e8a841 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -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(): """ diff --git a/type-discipline-budget.json b/type-discipline-budget.json index d5bf3883be4..cbcb5dca443 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,6 +1,6 @@ { "LIT001": { - "limit": 22330 + "limit": 22328 }, "LIT002": { "limit": 26763 diff --git a/ui/litellm-dashboard/src/components/organisms/create_key_button.integration.test.tsx b/ui/litellm-dashboard/src/components/organisms/create_key_button.integration.test.tsx index 8630be8548f..045dcfa3ceb 100644 --- a/ui/litellm-dashboard/src/components/organisms/create_key_button.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/organisms/create_key_button.integration.test.tsx @@ -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", () => { diff --git a/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx b/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx index d38a8995c1b..5749541dcea 100644 --- a/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx +++ b/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx @@ -588,7 +588,7 @@ const CreateKey: React.FC = ({ 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); diff --git a/ui/litellm-dashboard/src/components/templates/key_edit_view.test.tsx b/ui/litellm-dashboard/src/components/templates/key_edit_view.test.tsx index 97b09e00808..10b1983b54b 100644 --- a/ui/litellm-dashboard/src/components/templates/key_edit_view.test.tsx +++ b/ui/litellm-dashboard/src/components/templates/key_edit_view.test.tsx @@ -1457,6 +1457,32 @@ describe("KeyEditView", () => { expect(screen.getByLabelText("Organization")).toHaveValue("Engineering"); }); }); + + it("submits organization_id as null after the organization is cleared", async () => { + const onSubmit = vi.fn().mockResolvedValue(undefined); + renderWithProviders( + {}} + onSubmit={onSubmit} + accessToken="" + userID="" + userRole="Admin" + premiumUser={false} + />, + ); + + await waitFor(() => { + expect(screen.getByLabelText("Organization")).toHaveValue("Engineering"); + }); + await userEvent.click(screen.getByRole("button", { name: "Clear" })); + await userEvent.click(screen.getByRole("button", { name: /save changes/i })); + + await waitFor(() => { + expect(onSubmit).toHaveBeenCalledWith(expect.objectContaining({ organization_id: null })); + }); + expect(JSON.parse(JSON.stringify(onSubmit.mock.calls[0][0]))).toHaveProperty("organization_id", null); + }); }); describe("models dropdown team gating", () => { diff --git a/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx b/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx index df1af2ca8e9..3e772fd0e9b 100644 --- a/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx +++ b/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx @@ -303,8 +303,8 @@ export function KeyEditView({ } }; - const handleOrganizationChange = (setField: (value: string | undefined) => void, orgId: string | undefined) => { - setField(orgId); + const handleOrganizationChange = (setField: (value: string | null) => void, orgId: string | undefined) => { + setField(orgId || null); setSelectedOrganizationId(orgId || null); form.setValue("team_id", undefined); }; diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx index ce59c62f1c4..e3bacc0908a 100644 --- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx @@ -75,6 +75,68 @@ describe("Cost column", () => { }); }); +describe("Type column", () => { + it("shows the conversation badge and composition even when an MCP call represents the conversation", async () => { + const user = userEvent.setup(); + const mcpRepresentative = { + request_id: "req-mcp-rep", + call_type: "call_mcp_tool", + session_id: "sess-edge", + session_total_count: 3, + session_llm_count: 2, + mcp_tool_call_count: 1, + session_agent_count: 0, + }; + renderRows([logEntry(mcpRepresentative)]); + + expect(screen.queryByText("MCP")).not.toBeInTheDocument(); + await user.hover(screen.getByText("3")); + expect(await screen.findByText("2 LLM • 1 MCP")).toBeInTheDocument(); + }); + + it("keeps the plain MCP badge for a single MCP call", () => { + renderRows([logEntry({ request_id: "req-mcp-solo", call_type: "call_mcp_tool", session_total_count: 1 })]); + + expect(screen.getByText("MCP")).toBeInTheDocument(); + }); +}); + +describe("Model column", () => { + it("lists every model used across a conversation, not only the representative call's model", () => { + const conversationCall: Partial = { + request_id: "req-session", + model: "gpt-5.6", + session_id: "sess-1", + session_total_count: 3, + session_models: ["claude-sonnet-5", "gpt-5.6"], + }; + renderRows([logEntry(conversationCall)]); + + expect(screen.getByText("claude-sonnet-5, gpt-5.6")).toBeInTheDocument(); + expect(screen.queryByText("gpt-5.6")).not.toBeInTheDocument(); + }); + + it("marks a conversation whose model list was capped by the server", () => { + const cappedCall = { + request_id: "req-capped", + model: "gpt-5.6", + session_id: "sess-2", + session_total_count: 30, + session_models: ["claude-sonnet-5", "gpt-5.6"], + session_models_truncated: true, + }; + renderRows([logEntry(cappedCall)]); + + expect(screen.getByText("claude-sonnet-5, gpt-5.6, ...")).toBeInTheDocument(); + }); + + it("keeps a single call's own model", () => { + renderRows([logEntry({ request_id: "req-single", model: "gpt-5.6" })]); + + expect(screen.getByText("gpt-5.6")).toBeInTheDocument(); + }); +}); + describe("row action cells", () => { it("reports the key hash through the injected dependency rather than a row field", async () => { const user = userEvent.setup(); diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx index b9058d02a6b..8db0b106851 100644 --- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx @@ -63,9 +63,11 @@ export const getRequestLogsTableColumns = ({ const sessionAgentCount = log.session_agent_count ?? (isAgent ? sessionCount : 0); const sessionMcpCount = log.mcp_tool_call_count ?? (isMcp ? sessionCount : 0); - if (isMcp) return ; - if (isAgent && sessionCount <= 1) return ; - if (sessionCount <= 1) return ; + if (sessionCount <= 1) { + if (isMcp) return ; + if (isAgent) return ; + return ; + } const sessionTypeBadge = ( @@ -224,10 +226,13 @@ export const getRequestLogsTableColumns = ({ cell: ({ row }) => { const log = row.original; const provider = log.custom_llm_provider; - const modelName = log.model ?? ""; + const sessionModels = log.session_models ?? []; + const modelNames = sessionModels.length > 0 ? sessionModels : [log.model ?? ""]; + const modelLabel = log.session_models_truncated ? `${modelNames.join(", ")}, ...` : modelNames.join(", "); + const isSingleModel = modelNames.length === 1; return (
- {provider && ( + {provider && isSingleModel && ( )} - {modelName}} /> + + {modelLabel} + + } + />
); }, diff --git a/ui/litellm-dashboard/src/components/view_logs/columns.tsx b/ui/litellm-dashboard/src/components/view_logs/columns.tsx index 2f3a3681352..21e09faf454 100644 --- a/ui/litellm-dashboard/src/components/view_logs/columns.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/columns.tsx @@ -47,4 +47,6 @@ export type LogEntry = { mcp_tool_call_spend?: number; session_llm_count?: number; session_agent_count?: number; + session_models?: string[]; + session_models_truncated?: boolean; }; From 4990f06accc847a886e2dfc2236a89abd8bd3acc Mon Sep 17 00:00:00 2001 From: moe-berri Date: Thu, 3 Sep 2026 08:59:04 -0700 Subject: [PATCH 6/9] feat(auto-router): support classifier reasoning effort (#39372) * feat(auto-router): support classifier reasoning effort * fix(auto-router): harden classifier reasoning effort * fix(ui): satisfy classifier config lint limits * refactor(auto-router): simplify classifier effort support * fix(auto-router): clear frontend-lint and type-discipline gates, trim LOC --------- Co-authored-by: Tin Chi Lo --- litellm/router.py | 119 ++++++++++++++-- .../complexity_router/README.md | 7 +- .../complexity_router/complexity_router.py | 14 +- .../complexity_router/config.py | 8 ++ .../router_strategy/test_complexity_router.py | 51 ++++++- tests/test_litellm/test_router.py | 130 +++++++++++++++++- .../add_model/ClassificationMethodConfig.tsx | 35 ++++- .../ClassifierReasoningEffortSelect.tsx | 86 ++++++++++++ .../add_model/ComplexityRouterConfig.test.tsx | 112 ++++++++++++++- .../add_model/ComplexityRouterConfig.tsx | 17 +-- .../add_model/HeuristicScoringConfig.test.tsx | 6 +- .../add_model/add_auto_router_tab.tsx | 35 +++-- .../auto_router_connection_test.test.tsx | 8 +- .../add_model/auto_router_connection_test.tsx | 10 +- .../build_auto_router_test_targets.test.ts | 28 ++++ .../build_auto_router_test_targets.ts | 17 ++- .../build_complexity_router_config.test.ts | 51 ++++++- .../build_complexity_router_config.ts | 33 ++++- .../add_model/complexity_router_tiers.ts | 18 +++ ...d_updated_complexity_router_config.test.ts | 6 +- .../edit_auto_router_modal.tsx | 7 + .../llm_calls/fetch_models.test.tsx | 18 +++ .../src/components/llm_calls/fetch_models.tsx | 6 +- .../src/components/networking.test.ts | 9 ++ .../src/components/networking.tsx | 6 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 5 + 26 files changed, 774 insertions(+), 68 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/add_model/ClassifierReasoningEffortSelect.tsx diff --git a/litellm/router.py b/litellm/router.py index 303b22c9484..2b8b342d253 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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, ) @@ -231,6 +232,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, @@ -2168,6 +2170,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: @@ -2203,9 +2275,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: @@ -3216,6 +3295,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() @@ -3237,7 +3321,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: @@ -10255,6 +10341,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 @@ -10441,10 +10529,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: @@ -12308,10 +12409,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 diff --git a/litellm/router_strategy/complexity_router/README.md b/litellm/router_strategy/complexity_router/README.md index ee51add1ca1..e3da70f50fe 100644 --- a/litellm/router_strategy/complexity_router/README.md +++ b/litellm/router_strategy/complexity_router/README.md @@ -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. diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index a00ae6bee80..17e3d1256d0 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -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 diff --git a/litellm/router_strategy/complexity_router/config.py b/litellm/router_strategy/complexity_router/config.py index 9f2054dda01..70c1b281e31 100644 --- a/litellm/router_strategy/complexity_router/config.py +++ b/litellm/router_strategy/complexity_router/config.py @@ -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", diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index 3ecccb673f3..aa1b51afe10 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -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 diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index c843a66a1c1..3b4c80b6b4f 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -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(): @@ -9783,11 +9784,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=[ { @@ -9821,7 +9823,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 @@ -9958,11 +9960,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" @@ -9993,7 +9995,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: @@ -12252,6 +12254,120 @@ class TestTierParamsTheTargetAccepts: assert accepted == {"reasoning_effort": "max"} +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 diff --git a/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx b/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx index a29527a20fa..00ee7bd7d6e 100644 --- a/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx @@ -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; customTechnicalKeywords?: string[]; onCustomTechnicalKeywordsChange?: (keywords: string[]) => void; showValidationErrors?: boolean; @@ -236,6 +239,7 @@ const ClassificationMethodConfig: React.FC = ({ value, onChange, modelOptions, + effortOptionsByModel, customTechnicalKeywords, onCustomTechnicalKeywordsChange, showValidationErrors = false, @@ -251,6 +255,9 @@ const ClassificationMethodConfig: React.FC = ({ 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 = ({ }; 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 = ({ emptyText="No models found" allowClear={false} className={classifierModelMissing ? "border-destructive" : undefined} + aria-label="Classifier Model" /> {classifierModelMissing && A classifier model is required} +