mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
feat(cli): fuzzy model picker and auto-route Claude Code to autorouter
Numbered-index selection didn't scale past a handful of models, so switch
the tier picker to InquirerPy's fzf-style fuzzy search. Also set
ANTHROPIC_DEFAULT_{SONNET,HAIKU,OPUS}_MODEL to "autorouter" in Claude
Code's settings, since Router resolves auto-router deployments by literal
model name with no wildcard support, so a "*" catch-all model_name would
never match real traffic.
This commit is contained in:
parent
159c7ec8da
commit
53f6fdd874
11 changed files with 306 additions and 110 deletions
|
|
@ -507,10 +507,12 @@ Lists the model groups your key can reach on the proxy, via `/model_group/info`,
|
|||
lite autoroute configure
|
||||
```
|
||||
|
||||
An interactive wizard. It runs the same model-group discovery as above, splits the results into chat-capable and embedding-capable pools, and asks you to assign one or more models from the chat pool to each of the four complexity tiers -- SIMPLE, MEDIUM, COMPLEX, REASONING (enter comma-separated indices to assign a pool of models to a tier instead of just one; complexity_router picks randomly among a tier's pool per request, and adaptive mode specifically depends on having more than one candidate to choose from). From there it optionally offers: classifying prompt complexity with an LLM (again picked from your discovered pool) instead of the free built-in heuristic scorer, semantic keyword matching for tier assignment (needs an embedding model from the pool), and adaptive (bandit-based) selection layered on top of tiering.
|
||||
An interactive wizard. It runs the same model-group discovery as above, splits the results into chat-capable and embedding-capable pools, and asks you to assign one or more models from the chat pool to each of the four complexity tiers -- SIMPLE, MEDIUM, COMPLEX, REASONING. Each tier's picker is a type-to-filter fuzzy search (fzf-style) rather than a scrollable numbered list, so it stays usable even with hundreds of model groups: type a substring to narrow the list, tab to toggle a model into the selection, enter to confirm (assigning more than one model to a tier is exactly when this matters -- complexity_router picks randomly among a tier's pool per request, and adaptive mode specifically depends on having more than one candidate to choose from). From there it optionally offers: classifying prompt complexity with an LLM (again picked from your discovered pool) instead of the free built-in heuristic scorer, semantic keyword matching for tier assignment (needs an embedding model from the pool), and adaptive (bandit-based) selection layered on top of tiering.
|
||||
|
||||
The wizard writes the result to `~/.litellm/autorouter/config.yaml` with `0600` permissions, since the file embeds your real proxy API key. Every model referenced anywhere in that config -- tier targets, the classifier model, the embedding model -- becomes its own `litellm_proxy/<model-name>` deployment whose `api_base` and `api_key` point back at your real proxy. That is the trick that keeps your real proxy's config untouched: every actual network call this generates, whether it is the routed completion, an LLM-classifier call, or an embedding call, forwards transparently through your real, already-running proxy with your real key.
|
||||
|
||||
You do not need to tell Claude Code to request `autorouter` by name yourself: `lite autoroute up` also sets `ANTHROPIC_DEFAULT_SONNET_MODEL`, `ANTHROPIC_DEFAULT_HAIKU_MODEL`, and `ANTHROPIC_DEFAULT_OPUS_MODEL` to `autorouter` in `~/.claude/settings.json`, so every one of Claude Code's own model tiers requests it directly regardless of `/model` or whatever it defaults to otherwise. (A bare `model_name: "*"` deployment looks like the obvious way to catch any request instead, but litellm's Router looks up auto-router deployments by the literal requested model string with no wildcard resolution, so a `"*"` entry would never actually match real traffic -- these env var overrides are what makes it work.)
|
||||
|
||||
You must run `configure` at least once before `up`; running `up` first fails with a clear error telling you to configure first.
|
||||
|
||||
#### Launch the Ephemeral Auto-Router Proxy
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ from typing import Dict, FrozenSet, List, Literal, Tuple, Union
|
|||
from pydantic import BaseModel, ConfigDict, Field, JsonValue, TypeAdapter
|
||||
|
||||
TIER_NAMES: Tuple[str, ...] = ("SIMPLE", "MEDIUM", "COMPLEX", "REASONING")
|
||||
AUTOROUTER_MODEL_NAME = "autorouter"
|
||||
|
||||
|
||||
class ConfigGenerationError(Exception):
|
||||
|
|
@ -184,14 +185,21 @@ def build_generated_model_list(config: AutorouteConfig) -> List[JsonValue]:
|
|||
if config.adaptive:
|
||||
complexity_router_config["adaptive"] = True
|
||||
|
||||
auto_router_deployment: Dict[str, JsonValue] = {
|
||||
"model_name": "autorouter",
|
||||
"litellm_params": {
|
||||
"model": "auto_router/complexity_router",
|
||||
"complexity_router_config": complexity_router_config,
|
||||
},
|
||||
auto_router_litellm_params: Dict[str, JsonValue] = {
|
||||
"model": "auto_router/complexity_router",
|
||||
"complexity_router_config": complexity_router_config,
|
||||
}
|
||||
return [*proxy_deployments, auto_router_deployment]
|
||||
# A bare "*" model_name looks like the obvious way to catch every request Claude Code
|
||||
# might send regardless of which model it thinks it's using, but Router's auto-router
|
||||
# registry is keyed by the literal requested model string (router.py:10711-10717), not
|
||||
# resolved through pattern/wildcard matching first -- so a "*" entry here would only ever
|
||||
# match a client that literally sends model="*", never an actual wildcard catch-all. Callers
|
||||
# instead need to make Claude Code request this "autorouter" name directly (see
|
||||
# ANTHROPIC_DEFAULT_*_MODEL in settings.py's merge_claude_settings_static_token).
|
||||
return [
|
||||
*proxy_deployments,
|
||||
{"model_name": AUTOROUTER_MODEL_NAME, "litellm_params": auto_router_litellm_params},
|
||||
]
|
||||
|
||||
|
||||
def build_generated_proxy_config(config: AutorouteConfig, master_key: str) -> Dict[str, JsonValue]:
|
||||
|
|
@ -210,6 +218,7 @@ def build_generated_proxy_config(config: AutorouteConfig, master_key: str) -> Di
|
|||
|
||||
__all__ = [
|
||||
"TIER_NAMES",
|
||||
"AUTOROUTER_MODEL_NAME",
|
||||
"ConfigGenerationError",
|
||||
"DiscoveredModel",
|
||||
"parse_discovered_models",
|
||||
|
|
|
|||
|
|
@ -2,11 +2,23 @@ from typing import Dict
|
|||
|
||||
from pydantic import JsonValue
|
||||
|
||||
from .config import AUTOROUTER_MODEL_NAME
|
||||
|
||||
ENV_KEY = "env"
|
||||
API_KEY_HELPER_KEY = "apiKeyHelper"
|
||||
ANTHROPIC_API_KEY_KEY = "ANTHROPIC_API_KEY"
|
||||
ANTHROPIC_AUTH_TOKEN_KEY = "ANTHROPIC_AUTH_TOKEN"
|
||||
ANTHROPIC_BASE_URL_KEY = "ANTHROPIC_BASE_URL"
|
||||
# Force every one of Claude Code's own model tiers to request the auto-router by name.
|
||||
# Router's auto-router registry is keyed by the literal requested model string
|
||||
# (litellm/router.py:10711-10717) with no wildcard/pattern resolution, so a bare "*"
|
||||
# model_name can never work as a catch-all -- these overrides are what actually makes
|
||||
# Claude Code send "autorouter" regardless of /model or its own version-specific defaults.
|
||||
ANTHROPIC_DEFAULT_MODEL_ENV_KEYS = (
|
||||
"ANTHROPIC_DEFAULT_SONNET_MODEL",
|
||||
"ANTHROPIC_DEFAULT_HAIKU_MODEL",
|
||||
"ANTHROPIC_DEFAULT_OPUS_MODEL",
|
||||
)
|
||||
|
||||
|
||||
def merge_claude_settings_static_token(
|
||||
|
|
@ -25,6 +37,7 @@ def merge_claude_settings_static_token(
|
|||
**base_env,
|
||||
ANTHROPIC_BASE_URL_KEY: base_url.rstrip("/"),
|
||||
ANTHROPIC_AUTH_TOKEN_KEY: auth_token,
|
||||
**{key: AUTOROUTER_MODEL_NAME for key in ANTHROPIC_DEFAULT_MODEL_ENV_KEYS},
|
||||
}
|
||||
env.pop(ANTHROPIC_API_KEY_KEY, None)
|
||||
merged: Dict[str, JsonValue] = {**settings, ENV_KEY: env}
|
||||
|
|
|
|||
|
|
@ -1,10 +1,11 @@
|
|||
import sys
|
||||
from pathlib import Path
|
||||
from typing import List, Optional, Tuple
|
||||
from typing import List, Tuple
|
||||
|
||||
import click
|
||||
import yaml
|
||||
from rich.console import Console
|
||||
from rich.table import Table
|
||||
from InquirerPy import inquirer
|
||||
from InquirerPy.base.control import Choice
|
||||
|
||||
from .... import Client
|
||||
from .config import (
|
||||
|
|
@ -25,52 +26,42 @@ from .config import (
|
|||
from .process import CONFIG_PATH
|
||||
|
||||
|
||||
def _render_model_table(models: Tuple[DiscoveredModel, ...], prompt_label: str) -> None:
|
||||
console = Console()
|
||||
table = Table(title=f"Pick model(s) for {prompt_label}")
|
||||
table.add_column("Index", style="cyan", no_wrap=True)
|
||||
table.add_column("Model", style="magenta")
|
||||
for i, model in enumerate(models):
|
||||
table.add_row(str(i + 1), model.name)
|
||||
console.print(table)
|
||||
def _is_interactive() -> bool:
|
||||
return sys.stdin.isatty()
|
||||
|
||||
|
||||
def _parse_indices(choice: str, count: int) -> Optional[Tuple[int, ...]]:
|
||||
raw_parts = [part.strip() for part in choice.split(",") if part.strip()]
|
||||
if not raw_parts:
|
||||
return None
|
||||
indices: List[int] = []
|
||||
for part in raw_parts:
|
||||
try:
|
||||
index = int(part) - 1
|
||||
except ValueError:
|
||||
return None
|
||||
if not (0 <= index < count):
|
||||
return None
|
||||
indices.append(index)
|
||||
return tuple(dict.fromkeys(indices))
|
||||
def _fuzzy_pick(models: Tuple[DiscoveredModel, ...], prompt_label: str, multiselect: bool) -> List[str]:
|
||||
"""Type-to-filter picker over a (possibly huge) model pool, using InquirerPy's fzf-style fuzzy prompt.
|
||||
|
||||
A plain numbered table + typed index does not scale past a handful of models -- proxies with
|
||||
hundreds of model groups made that interaction unusable. This lets the user narrow the pool by
|
||||
typing a substring instead of scrolling/counting.
|
||||
|
||||
Assumes the caller already checked interactivity (run_configure_wizard does, once, up front) --
|
||||
checking here too would check the wrong thing under test, where InquirerPy is driven through its
|
||||
own injected input/output rather than the real process stdin.
|
||||
"""
|
||||
choices = [Choice(value=model.name, name=model.name) for model in models]
|
||||
toggle_hint = "tab to toggle, " if multiselect else ""
|
||||
while True:
|
||||
result = inquirer.fuzzy(
|
||||
message=f"{prompt_label}: type to filter, {toggle_hint}enter to confirm",
|
||||
choices=choices,
|
||||
multiselect=multiselect,
|
||||
max_height="70%",
|
||||
).execute()
|
||||
selected = result if multiselect else [result]
|
||||
if selected:
|
||||
return selected
|
||||
click.echo("Select at least one model.")
|
||||
|
||||
|
||||
def _render_and_prompt_for_model(models: Tuple[DiscoveredModel, ...], prompt_label: str) -> str:
|
||||
_render_model_table(models, prompt_label)
|
||||
while True:
|
||||
choice = click.prompt(f"\nSelect a model for {prompt_label} by index", type=str).strip()
|
||||
indices = _parse_indices(choice, len(models))
|
||||
if indices is not None and len(indices) == 1:
|
||||
return models[indices[0]].name
|
||||
click.echo(f"Invalid selection. Please enter a single number between 1 and {len(models)}")
|
||||
return _fuzzy_pick(models, prompt_label, multiselect=False)[0]
|
||||
|
||||
|
||||
def _render_and_prompt_for_models(models: Tuple[DiscoveredModel, ...], prompt_label: str) -> Tuple[str, ...]:
|
||||
_render_model_table(models, prompt_label)
|
||||
while True:
|
||||
choice = click.prompt(
|
||||
f"\nSelect model(s) for {prompt_label} by index (comma-separated for multiple)", type=str
|
||||
).strip()
|
||||
indices = _parse_indices(choice, len(models))
|
||||
if indices is not None:
|
||||
return tuple(models[i].name for i in indices)
|
||||
click.echo(f"Invalid selection. Please enter number(s) between 1 and {len(models)}, comma-separated")
|
||||
return tuple(_fuzzy_pick(models, prompt_label, multiselect=True))
|
||||
|
||||
|
||||
def run_configure_wizard(ctx: click.Context) -> Path:
|
||||
|
|
@ -88,6 +79,9 @@ def run_configure_wizard(ctx: click.Context) -> Path:
|
|||
if not chat_pool:
|
||||
raise click.ClickException("Your key has no chat-capable models available on this proxy.")
|
||||
|
||||
if not _is_interactive():
|
||||
raise click.ClickException("`lite autoroute configure` requires an interactive terminal.")
|
||||
|
||||
click.echo("Assign model(s) to each complexity tier (from what your key can access):")
|
||||
tiers = {tier: _render_and_prompt_for_models(chat_pool, tier) for tier in TIER_NAMES}
|
||||
default_model = tiers["MEDIUM"][0]
|
||||
|
|
|
|||
|
|
@ -48,7 +48,10 @@ def load_json_or_empty(path: Path) -> Dict[str, JsonValue]:
|
|||
if not path.exists():
|
||||
return {}
|
||||
with open(path, "r") as f:
|
||||
return _SETTINGS_ADAPTER.validate_json(f.read())
|
||||
content = f.read()
|
||||
if not content.strip():
|
||||
return {}
|
||||
return _SETTINGS_ADAPTER.validate_json(content)
|
||||
|
||||
|
||||
def merge_claude_settings(
|
||||
|
|
|
|||
|
|
@ -66,6 +66,7 @@ proxy = [
|
|||
"litellm-enterprise==0.1.49",
|
||||
"RestrictedPython>=8.1,<9.0",
|
||||
"rich>=13.9.4,<14.0",
|
||||
"InquirerPy>=0.3.4,<1.0",
|
||||
"polars>=1.38.1,<2.0",
|
||||
"soundfile>=0.12.1,<1.0",
|
||||
"pyroscope-io>=0.8.16,<1.0; sys_platform != 'win32'",
|
||||
|
|
@ -74,11 +75,12 @@ proxy = [
|
|||
]
|
||||
# Thin client install for the `lite` CLI on developer laptops. The CLI's heavy
|
||||
# imports (fastapi, cryptography, ...) are all guarded, so it runs on the base
|
||||
# SDK plus just these three; none of the server runtime in `proxy` is pulled in.
|
||||
# SDK plus just these four; none of the server runtime in `proxy` is pulled in.
|
||||
cli = [
|
||||
"rich>=13.9.4,<14.0",
|
||||
"pyyaml>=6.0.3,<7.0",
|
||||
"requests>=2.32.0,<3.0",
|
||||
"InquirerPy>=0.3.4,<1.0",
|
||||
]
|
||||
extra_proxy = [
|
||||
"prisma>=0.11.0,<1.0",
|
||||
|
|
|
|||
|
|
@ -91,7 +91,7 @@ class TestBuildGeneratedModelList:
|
|||
def test_every_proxy_deployment_points_back_at_customer_proxy(self):
|
||||
config = _base_config()
|
||||
model_list = build_generated_model_list(config)
|
||||
proxy_entries = [m for m in model_list if m["model_name"] != "autorouter"]
|
||||
proxy_entries = [m for m in model_list if m["model_name"] not in ("autorouter", "*")]
|
||||
names = {m["model_name"] for m in proxy_entries}
|
||||
assert names == {"gpt-4o-mini", "gpt-4o", "o1"}
|
||||
for entry in proxy_entries:
|
||||
|
|
@ -99,6 +99,15 @@ class TestBuildGeneratedModelList:
|
|||
assert entry["litellm_params"]["api_base"] == config.base_url
|
||||
assert entry["litellm_params"]["api_key"] == config.api_key
|
||||
|
||||
def test_no_wildcard_deployment_is_generated(self):
|
||||
# A bare "*" model_name looks like the obvious catch-all, but Router's auto-router
|
||||
# registry is keyed by the literal requested model string with no wildcard resolution
|
||||
# (litellm/router.py:10711-10717), so a "*" entry here would silently never match real
|
||||
# traffic. Regression guard: don't reintroduce it.
|
||||
config = _base_config()
|
||||
model_list = build_generated_model_list(config)
|
||||
assert not any(m["model_name"] == "*" for m in model_list)
|
||||
|
||||
def test_complexity_router_config_reflects_llm_classifier(self):
|
||||
config = _base_config(classifier=LLMClassifier(model="gpt-4o", timeout_ms=1234))
|
||||
autorouter = next(m for m in build_generated_model_list(config) if m["model_name"] == "autorouter")
|
||||
|
|
|
|||
|
|
@ -1,4 +1,7 @@
|
|||
from litellm.proxy.client.cli.commands.autoroute.settings import merge_claude_settings_static_token
|
||||
from litellm.proxy.client.cli.commands.autoroute.settings import (
|
||||
ANTHROPIC_DEFAULT_MODEL_ENV_KEYS,
|
||||
merge_claude_settings_static_token,
|
||||
)
|
||||
|
||||
|
||||
def test_preserves_unrelated_top_level_keys():
|
||||
|
|
@ -34,3 +37,20 @@ def test_does_not_mutate_input():
|
|||
settings = {"env": {"FOO": "bar"}, "apiKeyHelper": "old-helper"}
|
||||
merge_claude_settings_static_token(settings, "http://127.0.0.1:4000", "token-abc")
|
||||
assert settings == {"env": {"FOO": "bar"}, "apiKeyHelper": "old-helper"}
|
||||
|
||||
|
||||
def test_forces_all_claude_code_default_model_tiers_to_the_autorouter():
|
||||
# A bare "*" model_name deployment looks like the obvious way to catch every request
|
||||
# regardless of which model Claude Code thinks it's using, but Router's auto-router
|
||||
# registry is keyed by the literal requested model string with no wildcard resolution
|
||||
# (litellm/router.py:10711-10717) -- so the only reliable way to make every one of Claude
|
||||
# Code's own tiers hit the auto-router is to override the env vars it reads per tier.
|
||||
merged = merge_claude_settings_static_token({}, "http://127.0.0.1:4000", "token-abc")
|
||||
for key in ANTHROPIC_DEFAULT_MODEL_ENV_KEYS:
|
||||
assert merged["env"][key] == "autorouter"
|
||||
|
||||
|
||||
def test_overrides_a_preexisting_default_model_env_var():
|
||||
settings = {"env": {"ANTHROPIC_DEFAULT_SONNET_MODEL": "claude-opus-4-8"}}
|
||||
merged = merge_claude_settings_static_token(settings, "http://127.0.0.1:4000", "token-abc")
|
||||
assert merged["env"]["ANTHROPIC_DEFAULT_SONNET_MODEL"] == "autorouter"
|
||||
|
|
|
|||
|
|
@ -1,16 +1,19 @@
|
|||
import asyncio
|
||||
from typing import Any, Dict, List, Tuple
|
||||
from unittest.mock import patch
|
||||
|
||||
import click
|
||||
import pytest
|
||||
import yaml
|
||||
from click.testing import CliRunner
|
||||
from InquirerPy.base.control import Choice
|
||||
from prompt_toolkit.application import create_app_session
|
||||
from prompt_toolkit.input import create_pipe_input
|
||||
from prompt_toolkit.output import DummyOutput
|
||||
|
||||
from litellm.proxy.client.cli.commands.autoroute import wizard as wizard_module
|
||||
from litellm.proxy.client.cli.commands.autoroute.config import DiscoveredModel
|
||||
from litellm.proxy.client.cli.commands.autoroute.wizard import (
|
||||
_render_and_prompt_for_model,
|
||||
run_configure_wizard,
|
||||
)
|
||||
from litellm.proxy.client.cli.commands.autoroute.wizard import run_configure_wizard
|
||||
|
||||
CHAT_AND_EMBEDDING_GROUPS: List[Dict[str, Any]] = [
|
||||
{"model_group": "gpt-4o-mini", "mode": "chat", "input_cost_per_token": 0.01, "output_cost_per_token": 0.02},
|
||||
|
|
@ -38,12 +41,37 @@ def _invoke_wizard(ctx: click.Context) -> None:
|
|||
run_configure_wizard(ctx)
|
||||
|
||||
|
||||
def _run(tmp_path, raw_groups: List[Dict[str, Any]], input_str: str):
|
||||
def _run(
|
||||
tmp_path,
|
||||
raw_groups: List[Dict[str, Any]],
|
||||
tier_picks: Dict[str, Tuple[str, ...]],
|
||||
input_str: str,
|
||||
classifier_pick: str = "",
|
||||
embedding_pick: str = "",
|
||||
):
|
||||
"""Drives run_configure_wizard's orchestration logic (discovery, validation, config writing,
|
||||
classifier/semantic/adaptive branching) by mocking the fuzzy picker itself, since that widget
|
||||
is a real prompt_toolkit application tested separately in TestFuzzyPickWidget. CliRunner's
|
||||
injected input still drives the plain click.confirm() y/n prompts."""
|
||||
config_path = tmp_path / "config.yaml"
|
||||
runner = CliRunner()
|
||||
|
||||
def _fake_prompt_for_models(models, prompt_label):
|
||||
return tier_picks[prompt_label]
|
||||
|
||||
def _fake_prompt_for_model(models, prompt_label):
|
||||
if prompt_label == "LLM classifier":
|
||||
return classifier_pick
|
||||
if prompt_label == "semantic embeddings":
|
||||
return embedding_pick
|
||||
raise AssertionError(f"unexpected single-pick prompt_label {prompt_label!r}")
|
||||
|
||||
with (
|
||||
patch.object(wizard_module, "Client") as mock_client_cls,
|
||||
patch.object(wizard_module, "CONFIG_PATH", config_path),
|
||||
patch.object(wizard_module, "_is_interactive", return_value=True),
|
||||
patch.object(wizard_module, "_render_and_prompt_for_models", side_effect=_fake_prompt_for_models),
|
||||
patch.object(wizard_module, "_render_and_prompt_for_model", side_effect=_fake_prompt_for_model),
|
||||
):
|
||||
mock_client_cls.return_value.model_groups.info.return_value = raw_groups
|
||||
result = runner.invoke(
|
||||
|
|
@ -60,13 +88,17 @@ def _router_config(config_path) -> Dict[str, Any]:
|
|||
return autorouter["litellm_params"]["complexity_router_config"]
|
||||
|
||||
|
||||
_SIMPLE_TIER_PICKS: Dict[str, Tuple[str, ...]] = {
|
||||
"SIMPLE": ("gpt-4o-mini",),
|
||||
"MEDIUM": ("gpt-4o",),
|
||||
"COMPLEX": ("claude-opus",),
|
||||
"REASONING": ("o1",),
|
||||
}
|
||||
|
||||
|
||||
class TestRunConfigureWizardHappyPath:
|
||||
def test_assigns_tiers_and_declines_everything(self, tmp_path):
|
||||
result, config_path = _run(
|
||||
tmp_path,
|
||||
CHAT_AND_EMBEDDING_GROUPS,
|
||||
input_str="1\n2\n3\n4\nn\nn\nn\n",
|
||||
)
|
||||
result, config_path = _run(tmp_path, CHAT_AND_EMBEDDING_GROUPS, _SIMPLE_TIER_PICKS, input_str="n\nn\nn\n")
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
router_config = _router_config(config_path)
|
||||
|
|
@ -83,11 +115,8 @@ class TestRunConfigureWizardHappyPath:
|
|||
assert "adaptive" not in router_config
|
||||
|
||||
def test_assigns_multiple_models_to_a_single_tier(self, tmp_path):
|
||||
result, config_path = _run(
|
||||
tmp_path,
|
||||
CHAT_AND_EMBEDDING_GROUPS,
|
||||
input_str="1,2\n2\n3\n4\nn\nn\nn\n",
|
||||
)
|
||||
tier_picks = {**_SIMPLE_TIER_PICKS, "SIMPLE": ("gpt-4o-mini", "gpt-4o")}
|
||||
result, config_path = _run(tmp_path, CHAT_AND_EMBEDDING_GROUPS, tier_picks, input_str="n\nn\nn\n")
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
router_config = _router_config(config_path)
|
||||
|
|
@ -95,22 +124,14 @@ class TestRunConfigureWizardHappyPath:
|
|||
assert router_config["default_model"] == "gpt-4o"
|
||||
|
||||
def test_writes_config_file_with_restricted_permissions(self, tmp_path):
|
||||
result, config_path = _run(
|
||||
tmp_path,
|
||||
CHAT_AND_EMBEDDING_GROUPS,
|
||||
input_str="1\n2\n3\n4\nn\nn\nn\n",
|
||||
)
|
||||
result, config_path = _run(tmp_path, CHAT_AND_EMBEDDING_GROUPS, _SIMPLE_TIER_PICKS, input_str="n\nn\nn\n")
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
assert config_path.exists()
|
||||
assert oct(config_path.stat().st_mode)[-3:] == "600"
|
||||
|
||||
def test_no_embedding_pool_skips_semantic_prompt_entirely(self, tmp_path):
|
||||
result, config_path = _run(
|
||||
tmp_path,
|
||||
CHAT_ONLY_GROUPS,
|
||||
input_str="1\n2\n3\n4\nn\nn\n",
|
||||
)
|
||||
result, config_path = _run(tmp_path, CHAT_ONLY_GROUPS, _SIMPLE_TIER_PICKS, input_str="n\nn\n")
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
router_config = _router_config(config_path)
|
||||
|
|
@ -120,9 +141,7 @@ class TestRunConfigureWizardHappyPath:
|
|||
class TestRunConfigureWizardLLMClassifier:
|
||||
def test_accepting_llm_classifier_records_chosen_model(self, tmp_path):
|
||||
result, config_path = _run(
|
||||
tmp_path,
|
||||
CHAT_AND_EMBEDDING_GROUPS,
|
||||
input_str="1\n2\n3\n4\ny\n2\nn\nn\n",
|
||||
tmp_path, CHAT_AND_EMBEDDING_GROUPS, _SIMPLE_TIER_PICKS, input_str="y\nn\nn\n", classifier_pick="gpt-4o"
|
||||
)
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
|
|
@ -136,7 +155,9 @@ class TestRunConfigureWizardSemanticMatching:
|
|||
result, config_path = _run(
|
||||
tmp_path,
|
||||
CHAT_AND_EMBEDDING_GROUPS,
|
||||
input_str="1\n2\n3\n4\nn\ny\n1\nn\n",
|
||||
_SIMPLE_TIER_PICKS,
|
||||
input_str="n\ny\nn\n",
|
||||
embedding_pick="text-embedding-3-small",
|
||||
)
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
|
|
@ -147,11 +168,7 @@ class TestRunConfigureWizardSemanticMatching:
|
|||
|
||||
class TestRunConfigureWizardAdaptive:
|
||||
def test_accepting_adaptive_sets_adaptive_flag(self, tmp_path):
|
||||
result, config_path = _run(
|
||||
tmp_path,
|
||||
CHAT_AND_EMBEDDING_GROUPS,
|
||||
input_str="1\n2\n3\n4\nn\nn\ny\n",
|
||||
)
|
||||
result, config_path = _run(tmp_path, CHAT_AND_EMBEDDING_GROUPS, _SIMPLE_TIER_PICKS, input_str="n\nn\ny\n")
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
router_config = _router_config(config_path)
|
||||
|
|
@ -160,41 +177,109 @@ class TestRunConfigureWizardAdaptive:
|
|||
|
||||
class TestRunConfigureWizardNoChatModels:
|
||||
def test_fails_cleanly_without_prompting_when_no_chat_models(self, tmp_path):
|
||||
result, config_path = _run(tmp_path, EMBEDDING_ONLY_GROUPS, input_str="")
|
||||
result, config_path = _run(tmp_path, EMBEDDING_ONLY_GROUPS, {}, input_str="")
|
||||
|
||||
assert result.exit_code != 0
|
||||
assert "no chat-capable models" in result.output.lower()
|
||||
assert not config_path.exists()
|
||||
|
||||
|
||||
class TestRenderAndPromptForModel:
|
||||
class TestRunConfigureWizardNotInteractive:
|
||||
def test_fails_cleanly_when_not_a_tty(self, tmp_path):
|
||||
config_path = tmp_path / "config.yaml"
|
||||
runner = CliRunner()
|
||||
with (
|
||||
patch.object(wizard_module, "Client") as mock_client_cls,
|
||||
patch.object(wizard_module, "CONFIG_PATH", config_path),
|
||||
patch.object(wizard_module, "_is_interactive", return_value=False),
|
||||
):
|
||||
mock_client_cls.return_value.model_groups.info.return_value = CHAT_AND_EMBEDDING_GROUPS
|
||||
result = runner.invoke(_invoke_wizard, obj={"base_url": "http://localhost:4000", "api_key": "sk-test"})
|
||||
|
||||
assert result.exit_code != 0
|
||||
assert "interactive terminal" in result.output
|
||||
assert not config_path.exists()
|
||||
|
||||
|
||||
def _drive_fuzzy_pick(
|
||||
models: Tuple[DiscoveredModel, ...],
|
||||
prompt_label: str,
|
||||
multiselect: bool,
|
||||
key_events: List[Tuple[str, float]],
|
||||
) -> List[str]:
|
||||
"""Drives the real InquirerPy fuzzy prompt through prompt_toolkit's own test input/output,
|
||||
exercising the actual widget (filtering, tab-to-toggle, enter-to-confirm) rather than mocking
|
||||
it away. asyncio.to_thread propagates the create_app_session context into the worker thread
|
||||
running _fuzzy_pick's synchronous .execute() call."""
|
||||
|
||||
async def _run() -> List[str]:
|
||||
with create_pipe_input() as pipe_input:
|
||||
with create_app_session(input=pipe_input, output=DummyOutput()):
|
||||
task = asyncio.ensure_future(
|
||||
asyncio.to_thread(wizard_module._fuzzy_pick, models, prompt_label, multiselect)
|
||||
)
|
||||
await asyncio.sleep(0.05)
|
||||
for text, delay in key_events:
|
||||
pipe_input.send_text(text)
|
||||
await asyncio.sleep(delay)
|
||||
return await task
|
||||
|
||||
return asyncio.run(_run())
|
||||
|
||||
|
||||
class TestFuzzyPickWidget:
|
||||
def _models(self) -> Tuple[DiscoveredModel, ...]:
|
||||
return (
|
||||
DiscoveredModel(name="model-a"),
|
||||
DiscoveredModel(name="model-b"),
|
||||
return tuple(DiscoveredModel(name=f"model-{i}") for i in range(20))
|
||||
|
||||
def test_single_select_filters_and_returns_highlighted_match(self):
|
||||
result = _drive_fuzzy_pick(
|
||||
self._models(), "test", multiselect=False, key_events=[("model-13", 0.3), ("\r", 0.1)]
|
||||
)
|
||||
assert result == ["model-13"]
|
||||
|
||||
def test_reprompts_on_non_numeric_input(self):
|
||||
with patch("click.prompt", side_effect=["not-a-number", "2"]):
|
||||
result = _render_and_prompt_for_model(self._models(), "test tier")
|
||||
def test_multiselect_requires_tab_to_toggle_before_enter(self):
|
||||
result = _drive_fuzzy_pick(
|
||||
self._models(), "test", multiselect=True, key_events=[("model-7", 0.3), ("\t", 0.1), ("\r", 0.1)]
|
||||
)
|
||||
assert result == ["model-7"]
|
||||
|
||||
assert result == "model-b"
|
||||
def test_multiselect_can_pick_more_than_one_across_filters(self):
|
||||
result = _drive_fuzzy_pick(
|
||||
self._models(),
|
||||
"test",
|
||||
multiselect=True,
|
||||
key_events=[
|
||||
("model-3", 0.3),
|
||||
("\t", 0.1),
|
||||
*[("\x7f", 0.02) for _ in range("model-3".__len__())],
|
||||
("model-15", 0.3),
|
||||
("\t", 0.1),
|
||||
("\r", 0.1),
|
||||
],
|
||||
)
|
||||
assert set(result) == {"model-3", "model-15"}
|
||||
|
||||
def test_reprompts_on_out_of_range_index(self):
|
||||
with patch("click.prompt", side_effect=["5", "1"]):
|
||||
result = _render_and_prompt_for_model(self._models(), "test tier")
|
||||
def test_choice_wraps_name_and_value_to_the_same_model_name(self):
|
||||
model = DiscoveredModel(name="only-model")
|
||||
choice = Choice(value=model.name, name=model.name)
|
||||
assert choice.value == choice.name == "only-model"
|
||||
|
||||
|
||||
class TestRenderAndPromptForModelWrappers:
|
||||
def test_single_pick_wrapper_returns_bare_string(self):
|
||||
with patch.object(wizard_module, "_fuzzy_pick", return_value=["model-a"]) as mock_pick:
|
||||
result = wizard_module._render_and_prompt_for_model((), "tier")
|
||||
assert result == "model-a"
|
||||
mock_pick.assert_called_once_with((), "tier", multiselect=False)
|
||||
|
||||
def test_reprompts_on_zero_index(self):
|
||||
with patch("click.prompt", side_effect=["0", "2"]):
|
||||
result = _render_and_prompt_for_model(self._models(), "test tier")
|
||||
def test_multi_pick_wrapper_returns_tuple(self):
|
||||
with patch.object(wizard_module, "_fuzzy_pick", return_value=["model-a", "model-b"]) as mock_pick:
|
||||
result = wizard_module._render_and_prompt_for_models((), "tier")
|
||||
assert result == ("model-a", "model-b")
|
||||
mock_pick.assert_called_once_with((), "tier", multiselect=True)
|
||||
|
||||
assert result == "model-b"
|
||||
|
||||
def test_valid_first_answer_returns_immediately(self):
|
||||
with patch("click.prompt", return_value="1") as mock_prompt:
|
||||
result = _render_and_prompt_for_model(self._models(), "test tier")
|
||||
|
||||
assert result == "model-a"
|
||||
mock_prompt.assert_called_once()
|
||||
@pytest.mark.parametrize("isatty_value", [True, False])
|
||||
def test_is_interactive_reflects_stdin_isatty(isatty_value):
|
||||
with patch.object(wizard_module.sys.stdin, "isatty", return_value=isatty_value):
|
||||
assert wizard_module._is_interactive() is isatty_value
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from litellm.proxy.client.cli.commands.up import (
|
|||
BackupRecord,
|
||||
UpError,
|
||||
down,
|
||||
load_json_or_empty,
|
||||
merge_claude_settings,
|
||||
read_backup,
|
||||
resolve_api_key_helper,
|
||||
|
|
@ -66,6 +67,26 @@ class TestMergeClaudeSettings:
|
|||
assert settings == {"env": {"FOO": "bar"}}
|
||||
|
||||
|
||||
class TestLoadJsonOrEmpty:
|
||||
def test_returns_empty_dict_when_file_does_not_exist(self, tmp_path):
|
||||
assert load_json_or_empty(tmp_path / "missing.json") == {}
|
||||
|
||||
def test_returns_empty_dict_when_file_is_empty(self, tmp_path):
|
||||
path = tmp_path / "settings.json"
|
||||
path.write_text("")
|
||||
assert load_json_or_empty(path) == {}
|
||||
|
||||
def test_returns_empty_dict_when_file_is_whitespace_only(self, tmp_path):
|
||||
path = tmp_path / "settings.json"
|
||||
path.write_text(" \n")
|
||||
assert load_json_or_empty(path) == {}
|
||||
|
||||
def test_parses_real_content(self, tmp_path):
|
||||
path = tmp_path / "settings.json"
|
||||
path.write_text(json.dumps({"theme": "dark"}))
|
||||
assert load_json_or_empty(path) == {"theme": "dark"}
|
||||
|
||||
|
||||
class TestBackupRoundTrip:
|
||||
def test_restores_original_content_when_file_existed(self, monkeypatch, tmp_path):
|
||||
settings_path, backup_path = _patch_paths(monkeypatch, tmp_path)
|
||||
|
|
|
|||
40
uv.lock
generated
40
uv.lock
generated
|
|
@ -9,7 +9,7 @@ resolution-markers = [
|
|||
]
|
||||
|
||||
[options]
|
||||
exclude-newer = "2026-07-10T16:47:58.286372Z"
|
||||
exclude-newer = "2026-07-11T19:28:02.260785Z"
|
||||
exclude-newer-span = "P3D"
|
||||
|
||||
[manifest]
|
||||
|
|
@ -2655,6 +2655,19 @@ wheels = [
|
|||
{ url = "https://files.pythonhosted.org/packages/cb/b1/3846dd7f199d53cb17f49cba7e651e9ce294d8497c8c150530ed11865bb8/iniconfig-2.3.0-py3-none-any.whl", hash = "sha256:f631c04d2c48c52b84d0d0549c99ff3859c98df65b3101406327ecc7d53fbf12", size = 7484, upload-time = "2025-10-18T21:55:41.639Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "inquirerpy"
|
||||
version = "0.3.4"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "pfzy" },
|
||||
{ name = "prompt-toolkit" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/64/73/7570847b9da026e07053da3bbe2ac7ea6cde6bb2cbd3c7a5a950fa0ae40b/InquirerPy-0.3.4.tar.gz", hash = "sha256:89d2ada0111f337483cb41ae31073108b2ec1e618a49d7110b0d7ade89fc197e", size = 44431, upload-time = "2022-06-27T23:11:20.598Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/ce/ff/3b59672c47c6284e8005b42e84ceba13864aa0f39f067c973d1af02f5d91/InquirerPy-0.3.4-py3-none-any.whl", hash = "sha256:c65fdfbac1fa00e3ee4fb10679f4d3ed7a012abf4833910e63c295827fe2a7d4", size = 67677, upload-time = "2022-06-27T23:11:17.723Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "isodate"
|
||||
version = "0.7.2"
|
||||
|
|
@ -3305,6 +3318,7 @@ caching = [
|
|||
{ name = "diskcache" },
|
||||
]
|
||||
cli = [
|
||||
{ name = "inquirerpy" },
|
||||
{ name = "pyyaml" },
|
||||
{ name = "requests" },
|
||||
{ name = "rich" },
|
||||
|
|
@ -3340,6 +3354,7 @@ proxy = [
|
|||
{ name = "fastapi-sso" },
|
||||
{ name = "granian" },
|
||||
{ name = "gunicorn" },
|
||||
{ name = "inquirerpy" },
|
||||
{ name = "litellm-enterprise" },
|
||||
{ name = "litellm-proxy-extras" },
|
||||
{ name = "mcp" },
|
||||
|
|
@ -3517,6 +3532,8 @@ requires-dist = [
|
|||
{ name = "gunicorn", marker = "extra == 'proxy'", specifier = ">=23.0.0,<24.0" },
|
||||
{ name = "httpx", specifier = ">=0.28.0,<1.0" },
|
||||
{ name = "importlib-metadata", specifier = ">=8.0.0,<9.0" },
|
||||
{ name = "inquirerpy", marker = "extra == 'cli'", specifier = ">=0.3.4,<1.0" },
|
||||
{ name = "inquirerpy", marker = "extra == 'proxy'", specifier = ">=0.3.4,<1.0" },
|
||||
{ name = "jinja2", specifier = ">=3.1.6,<4.0" },
|
||||
{ name = "jsonschema", specifier = ">=4.0.0,<5.0" },
|
||||
{ name = "langfuse", marker = "extra == 'proxy-runtime'", specifier = ">=2.59.7,<3.0" },
|
||||
|
|
@ -5277,6 +5294,15 @@ wheels = [
|
|||
{ url = "https://files.pythonhosted.org/packages/7d/eb/b6260b31b1a96386c0a880edebe26f89669098acea8e0318bff6adb378fd/pathable-0.4.4-py3-none-any.whl", hash = "sha256:5ae9e94793b6ef5a4cbe0a7ce9dbbefc1eec38df253763fd0aeeacf2762dbbc2", size = 9592, upload-time = "2025-01-10T18:43:11.88Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "pfzy"
|
||||
version = "0.3.4"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/d9/5a/32b50c077c86bfccc7bed4881c5a2b823518f5450a30e639db5d3711952e/pfzy-0.3.4.tar.gz", hash = "sha256:717ea765dd10b63618e7298b2d98efd819e0b30cd5905c9707223dceeb94b3f1", size = 8396, upload-time = "2022-01-28T02:26:17.946Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/8c/d7/8ff98376b1acc4503253b685ea09981697385ce344d4e3935c2af49e044d/pfzy-0.3.4-py3-none-any.whl", hash = "sha256:5f50d5b2b3207fa72e7ec0ef08372ef652685470974a107d0d4999fc5a903a96", size = 8537, upload-time = "2022-01-28T02:26:16.047Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "pillow"
|
||||
version = "12.3.0"
|
||||
|
|
@ -5469,6 +5495,18 @@ wheels = [
|
|||
{ url = "https://files.pythonhosted.org/packages/c7/98/745b810d822103adca2df8decd4c0bbe839ba7ad3511af3f0d09692fc0f0/prometheus_client-0.20.0-py3-none-any.whl", hash = "sha256:cde524a85bce83ca359cc837f28b8c0db5cac7aa653a588fd7e84ba061c329e7", size = 54474, upload-time = "2024-02-14T15:55:03.957Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "prompt-toolkit"
|
||||
version = "3.0.52"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "wcwidth" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/a1/96/06e01a7b38dce6fe1db213e061a4602dd6032a8a97ef6c1a862537732421/prompt_toolkit-3.0.52.tar.gz", hash = "sha256:28cde192929c8e7321de85de1ddbe736f1375148b02f2e17edd840042b1be855", size = 434198, upload-time = "2025-08-27T15:24:02.057Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/84/03/0d3ce49e2505ae70cf43bc5bb3033955d2fc9f932163e84dc0779cc47f48/prompt_toolkit-3.0.52-py3-none-any.whl", hash = "sha256:9aac639a3bbd33284347de5ad8d68ecc044b91a762dc39b7c21095fcd6a19955", size = 391431, upload-time = "2025-08-27T15:23:59.498Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "propcache"
|
||||
version = "0.5.2"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue