From 159c7ec8dac53883488bf0944caa26cfa8f5dc38 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Tue, 14 Jul 2026 12:06:22 -0700 Subject: [PATCH] feat(cli): allow multiple models per autoroute tier complexity_router already supports a pool of models per tier (randomly picked per request; adaptive mode specifically needs a pool to choose within), but the configure wizard only ever let you assign one. Tiers are now a tuple of model names; the wizard prompt accepts comma-separated indices to pick more than one per tier. --- litellm/proxy/client/cli/README.md | 2 +- .../client/cli/commands/autoroute/config.py | 16 +++-- .../client/cli/commands/autoroute/wizard.py | 59 ++++++++++++++----- .../proxy/client/cli/autoroute/test_config.py | 15 +++-- .../proxy/client/cli/autoroute/test_wizard.py | 20 +++++-- 5 files changed, 80 insertions(+), 32 deletions(-) diff --git a/litellm/proxy/client/cli/README.md b/litellm/proxy/client/cli/README.md index ecd5f3d48e2..86ad9687901 100644 --- a/litellm/proxy/client/cli/README.md +++ b/litellm/proxy/client/cli/README.md @@ -507,7 +507,7 @@ 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 a model from the chat pool to each of the four complexity tiers -- SIMPLE, MEDIUM, COMPLEX, REASONING. 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 (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. 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/` 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. diff --git a/litellm/proxy/client/cli/commands/autoroute/config.py b/litellm/proxy/client/cli/commands/autoroute/config.py index 85b5223524b..c0f5f778fd8 100644 --- a/litellm/proxy/client/cli/commands/autoroute/config.py +++ b/litellm/proxy/client/cli/commands/autoroute/config.py @@ -103,7 +103,9 @@ class AutorouteConfig(BaseModel): base_url: str api_key: str - tiers: Dict[str, str] + # Each tier maps to a pool of one or more models; complexity_router picks randomly among + # them per request (or, in adaptive mode, learns which to prefer within the pool). + tiers: Dict[str, Tuple[str, ...]] default_model: str classifier: ClassifierChoice = Field(default_factory=HeuristicClassifier) semantic_matching: SemanticMatchingChoice = Field(default_factory=NoSemanticMatching) @@ -115,9 +117,10 @@ def validate_config(config: AutorouteConfig, discovered: Tuple[DiscoveredModel, chat_names: FrozenSet[str] = frozenset(m.name for m in chat_models(discovered)) embedding_names: FrozenSet[str] = frozenset(m.name for m in embedding_models(discovered)) - for tier, model in config.tiers.items(): - if model not in chat_names: - raise ConfigGenerationError(f"Tier {tier} references unknown chat model '{model}'") + for tier, models in config.tiers.items(): + for model in models: + if model not in chat_names: + raise ConfigGenerationError(f"Tier {tier} references unknown chat model '{model}'") if config.default_model not in chat_names: raise ConfigGenerationError(f"default_model '{config.default_model}' is not a known chat model") @@ -152,7 +155,8 @@ def build_generated_model_list(config: AutorouteConfig) -> List[JsonValue]: to exactly one `litellm_proxy/` deployment forwarding to the customer's real proxy, plus one `auto_router/complexity_router` deployment tying the tiers together. """ - referenced_names = {*config.tiers.values(), config.default_model} + referenced_names = {model for models in config.tiers.values() for model in models} + referenced_names.add(config.default_model) if isinstance(config.classifier, LLMClassifier): referenced_names.add(config.classifier.model) if isinstance(config.semantic_matching, SemanticMatching): @@ -163,7 +167,7 @@ def build_generated_model_list(config: AutorouteConfig) -> List[JsonValue]: ] complexity_router_config: Dict[str, JsonValue] = { - "tiers": dict(config.tiers), + "tiers": {tier: list(models) for tier, models in config.tiers.items()}, "default_model": config.default_model, } if isinstance(config.classifier, LLMClassifier): diff --git a/litellm/proxy/client/cli/commands/autoroute/wizard.py b/litellm/proxy/client/cli/commands/autoroute/wizard.py index e2d1b85dc6f..c3ac2bdd637 100644 --- a/litellm/proxy/client/cli/commands/autoroute/wizard.py +++ b/litellm/proxy/client/cli/commands/autoroute/wizard.py @@ -1,5 +1,5 @@ from pathlib import Path -from typing import Tuple +from typing import List, Optional, Tuple import click import yaml @@ -25,25 +25,52 @@ from .config import ( from .process import CONFIG_PATH -def _render_and_prompt_for_model(models: Tuple[DiscoveredModel, ...], prompt_label: str) -> str: +def _render_model_table(models: Tuple[DiscoveredModel, ...], prompt_label: str) -> None: console = Console() - table = Table(title=f"Pick a model for {prompt_label}") + 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 _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 _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() - try: - index = int(choice) - 1 - except ValueError: - click.echo("Invalid input. Please enter a number.") - continue - if 0 <= index < len(models): - return models[index].name - click.echo(f"Invalid selection. Please enter a number between 1 and {len(models)}") + 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)}") + + +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") def run_configure_wizard(ctx: click.Context) -> Path: @@ -61,9 +88,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.") - click.echo("Assign a model to each complexity tier (from what your key can access):") - tiers = {tier: _render_and_prompt_for_model(chat_pool, tier) for tier in TIER_NAMES} - default_model = tiers["MEDIUM"] + 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] classifier = HeuristicClassifier() if click.confirm("\nUse an LLM classifier instead of the free heuristic scorer?", default=False): @@ -98,8 +125,8 @@ def run_configure_wizard(ctx: click.Context) -> Path: CONFIG_PATH.chmod(0o600) click.echo(f"\nWrote {CONFIG_PATH}") - for tier, model in tiers.items(): - click.echo(f" {tier}: {model}") + for tier, models in tiers.items(): + click.echo(f" {tier}: {', '.join(models)}") return CONFIG_PATH diff --git a/tests/test_litellm/proxy/client/cli/autoroute/test_config.py b/tests/test_litellm/proxy/client/cli/autoroute/test_config.py index a48db0c79e4..2335f07a823 100644 --- a/tests/test_litellm/proxy/client/cli/autoroute/test_config.py +++ b/tests/test_litellm/proxy/client/cli/autoroute/test_config.py @@ -31,10 +31,10 @@ def _base_config(**overrides: Any) -> AutorouteConfig: "base_url": "http://real-proxy.internal:4000", "api_key": "sk-real-key", "tiers": { - "SIMPLE": "gpt-4o-mini", - "MEDIUM": "gpt-4o", - "COMPLEX": "gpt-4o", - "REASONING": "o1", + "SIMPLE": ("gpt-4o-mini",), + "MEDIUM": ("gpt-4o",), + "COMPLEX": ("gpt-4o",), + "REASONING": ("o1",), }, "default_model": "gpt-4o", } @@ -146,7 +146,12 @@ class TestValidateConfig: def test_raises_for_tier_referencing_unknown_model(self): config = _base_config( - tiers={"SIMPLE": "unknown-model", "MEDIUM": "gpt-4o", "COMPLEX": "gpt-4o", "REASONING": "o1"} + tiers={ + "SIMPLE": ("unknown-model",), + "MEDIUM": ("gpt-4o",), + "COMPLEX": ("gpt-4o",), + "REASONING": ("o1",), + } ) with pytest.raises(ConfigGenerationError, match="unknown-model"): validate_config(config, DISCOVERED) diff --git a/tests/test_litellm/proxy/client/cli/autoroute/test_wizard.py b/tests/test_litellm/proxy/client/cli/autoroute/test_wizard.py index 81f54bf9b78..07ad1abd71f 100644 --- a/tests/test_litellm/proxy/client/cli/autoroute/test_wizard.py +++ b/tests/test_litellm/proxy/client/cli/autoroute/test_wizard.py @@ -71,10 +71,10 @@ class TestRunConfigureWizardHappyPath: assert result.exit_code == 0, result.output router_config = _router_config(config_path) assert router_config["tiers"] == { - "SIMPLE": "gpt-4o-mini", - "MEDIUM": "gpt-4o", - "COMPLEX": "claude-opus", - "REASONING": "o1", + "SIMPLE": ["gpt-4o-mini"], + "MEDIUM": ["gpt-4o"], + "COMPLEX": ["claude-opus"], + "REASONING": ["o1"], } assert router_config["default_model"] == "gpt-4o" assert "classifier_type" not in router_config @@ -82,6 +82,18 @@ class TestRunConfigureWizardHappyPath: assert "semantic_keyword_matching" not in router_config 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", + ) + + assert result.exit_code == 0, result.output + router_config = _router_config(config_path) + assert router_config["tiers"]["SIMPLE"] == ["gpt-4o-mini", "gpt-4o"] + assert router_config["default_model"] == "gpt-4o" + def test_writes_config_file_with_restricted_permissions(self, tmp_path): result, config_path = _run( tmp_path,