From cd85a3b2c7e8179b2050f5a402becbe72dd27302 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Sat, 14 Mar 2026 13:25:06 -0700 Subject: [PATCH] refactor(setup_wizard): class with static methods, use check_valid_key from litellm.utils --- CLAUDE.md | 5 + litellm/setup_wizard.py | 705 ++++++++++++++++++---------------------- 2 files changed, 326 insertions(+), 384 deletions(-) diff --git a/CLAUDE.md b/CLAUDE.md index d9061b5e2be..f0478120181 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -140,6 +140,11 @@ LiteLLM is a unified interface for 100+ LLM providers with two main components: - **Check index coverage.** For new or modified queries, check `schema.prisma` for a supporting index. Prefer extending an existing index (e.g. `@@index([a])` → `@@index([a, b])`) over adding a new one, unless it's a `@@unique`. Only add indexes for large/frequent queries. - **Keep schema files in sync.** Apply schema changes to all `schema.prisma` copies (`schema.prisma`, `litellm/proxy/`, `litellm-proxy-extras/`, `litellm-js/spend-logs/` for SpendLogs) with a migration under `litellm-proxy-extras/litellm_proxy_extras/migrations/`. +### Setup Wizard (`litellm/setup_wizard.py`) +- The wizard is implemented as a single `SetupWizard` class with `@staticmethod` methods — keep it that way. No module-level functions except `run_setup_wizard()` (the public entrypoint) and pure helpers (color, ANSI). +- Use `litellm.utils.check_valid_key(model, api_key)` for credential validation — never roll a custom completion call. +- Do not hardcode provider env-key names or model lists that already exist in the codebase. Add a `test_model` field to each provider entry to drive `check_valid_key`; set it to `None` for providers that can't be validated with a single API key (Azure, Bedrock, Ollama). + ### Enterprise Features - Enterprise-specific code in `enterprise/` directory - Optional features enabled via environment variables diff --git a/litellm/setup_wizard.py b/litellm/setup_wizard.py index 813591053a4..9971328931d 100644 --- a/litellm/setup_wizard.py +++ b/litellm/setup_wizard.py @@ -17,17 +17,26 @@ import tty from pathlib import Path from typing import Dict, List, Optional, Set +from litellm.utils import check_valid_key + # --------------------------------------------------------------------------- # Provider definitions # --------------------------------------------------------------------------- +# Each entry describes one provider card shown in the wizard. +# `env_key` — primary env var name (None = no key needed, e.g. Ollama) +# `test_model` — model passed to check_valid_key for credential validation +# (None = skip validation, e.g. Azure needs a deployment name) +# `models` — default models written into the generated config +# --------------------------------------------------------------------------- -PROVIDERS = [ +PROVIDERS: List[Dict] = [ { "id": "openai", "name": "OpenAI", "description": "GPT-4o, GPT-4o-mini, o3-mini", "env_key": "OPENAI_API_KEY", "key_hint": "sk-...", + "test_model": "gpt-4o-mini", "models": ["gpt-4o", "gpt-4o-mini"], }, { @@ -36,6 +45,7 @@ PROVIDERS = [ "description": "Claude Opus 4.6, Sonnet 4.6, Haiku 4.5", "env_key": "ANTHROPIC_API_KEY", "key_hint": "sk-ant-...", + "test_model": "claude-haiku-4-5-20251001", "models": ["claude-opus-4-6", "claude-sonnet-4-6", "claude-haiku-4-5-20251001"], }, { @@ -44,6 +54,7 @@ PROVIDERS = [ "description": "Gemini 2.0 Flash, Gemini 2.5 Pro", "env_key": "GEMINI_API_KEY", "key_hint": "AIza...", + "test_model": "gemini/gemini-2.0-flash", "models": ["gemini/gemini-2.0-flash", "gemini/gemini-2.5-pro"], }, { @@ -52,6 +63,7 @@ PROVIDERS = [ "description": "GPT-4o via Azure", "env_key": "AZURE_API_KEY", "key_hint": "your-azure-key", + "test_model": None, # needs deployment name — skip validation "models": [], "needs_api_base": True, "api_base_hint": "https://.openai.azure.com/", @@ -63,8 +75,8 @@ PROVIDERS = [ "description": "Claude 3.5, Llama 3 via AWS", "env_key": "AWS_ACCESS_KEY_ID", "key_hint": "AKIA...", + "test_model": None, # multi-key auth — skip validation "models": ["bedrock/anthropic.claude-3-5-sonnet-20241022-v2:0"], - "needs_extra": True, "extra_keys": ["AWS_SECRET_ACCESS_KEY", "AWS_REGION_NAME"], "extra_hints": ["your-secret-key", "us-east-1"], }, @@ -74,6 +86,7 @@ PROVIDERS = [ "description": "Local models (llama3.2, mistral, etc.)", "env_key": None, "key_hint": None, + "test_model": None, # local — no remote validation "models": ["ollama/llama3.2", "ollama/mistral"], "api_base": "http://localhost:11434", }, @@ -81,7 +94,7 @@ PROVIDERS = [ # --------------------------------------------------------------------------- -# ANSI colour helpers (no external deps needed) +# ANSI colour helpers # --------------------------------------------------------------------------- _ORANGE = "\033[38;2;215;119;87m" @@ -92,6 +105,7 @@ _BLUE = "\033[38;2;177;185;249m" _GREY = "\033[38;2;153;153;153m" _RESET = "\033[0m" _CHECK = "✔" +_CROSS = "✘" _CURSOR_HIDE = "\033[?25l" _CURSOR_SHOW = "\033[?25h" @@ -103,9 +117,7 @@ def _supports_color() -> bool: def _c(code: str, text: str) -> str: - if _supports_color(): - return f"{code}{text}{_RESET}" - return text + return f"{code}{text}{_RESET}" if _supports_color() else text def orange(t: str) -> str: @@ -133,7 +145,7 @@ def dim(t: str) -> str: # --------------------------------------------------------------------------- -# ASCII art +# Layout constants # --------------------------------------------------------------------------- LITELLM_ASCII = r""" @@ -148,410 +160,340 @@ LITELLM_ASCII = r""" DIVIDER = dim(" " + "╌" * 74) -def _print_welcome() -> None: - try: - version = importlib.metadata.version("litellm") - except Exception: - version = "unknown" - - print() - print(orange(LITELLM_ASCII.rstrip("\n"))) - print(f" {orange('Welcome')} to {bold('LiteLLM')} {grey('v' + version)}") - print() - print(DIVIDER) - print() - - # --------------------------------------------------------------------------- -# Arrow-key provider selector +# Setup wizard # --------------------------------------------------------------------------- -def _read_key() -> str: - """Read one keypress from /dev/tty in raw mode.""" - with open("/dev/tty", "rb") as tty_fh: - fd = tty_fh.fileno() - old = termios.tcgetattr(fd) +class SetupWizard: + """ + Interactive onboarding wizard: provider selection → API keys → config file. + + All methods are static — the class is purely a namespace with clear + single-responsibility sections. Entry point: SetupWizard.run(). + """ + + # ── entry point ───────────────────────────────────────────────────────── + + @staticmethod + def run() -> None: try: - tty.setraw(fd) - ch = tty_fh.read(1) - if ch == b"\x1b": - ch2 = tty_fh.read(1) - if ch2 == b"[": - ch3 = tty_fh.read(1) - return "\x1b[" + ch3.decode("utf-8", errors="replace") - return "\x1b" + ch2.decode("utf-8", errors="replace") - return ch.decode("utf-8", errors="replace") - finally: - termios.tcsetattr(fd, termios.TCSADRAIN, old) + SetupWizard._wizard() + except (KeyboardInterrupt, EOFError): + print(f"\n\n {grey('Setup cancelled.')}\n") + # ── wizard steps ──────────────────────────────────────────────────────── -def _render_selector(cursor: int, selected: Set[int], first_render: bool) -> int: - """Draw (or redraw) the provider list. Returns number of lines printed.""" - lines = [] - lines.append(f"\n {bold('Add your first model')}\n") - lines.append(grey(" ↑↓ to navigate · Space to select · Enter to confirm") + "\n") - lines.append("\n") + @staticmethod + def _wizard() -> None: + SetupWizard._print_welcome() + print(f" {bold('Lets get started.')}") + print() - for i, p in enumerate(PROVIDERS): - arrow = blue("❯") if i == cursor else " " - bullet = green("◉") if i in selected else grey("○") - name_str = bold(p["name"]) if i == cursor else p["name"] - desc_str = grey(p["description"]) - lines.append(f" {arrow} {bullet} {name_str} {desc_str}\n") + providers = SetupWizard._select_providers() + env_vars = SetupWizard._collect_keys(providers) + port, master_key = SetupWizard._proxy_settings() - lines.append("\n") - content = "".join(lines) - line_count = content.count("\n") + config_path = Path(os.getcwd()) / "litellm_config.yaml" + config_path.write_text( + SetupWizard._build_config(providers, env_vars, port, master_key) + ) - if not first_render and _supports_color(): - sys.stdout.write(_MOVE_UP.format(line_count)) + SetupWizard._print_success(config_path, port, master_key) + SetupWizard._offer_start(config_path, port, master_key) - sys.stdout.write(content) - sys.stdout.flush() - return line_count + # ── welcome ───────────────────────────────────────────────────────────── + @staticmethod + def _print_welcome() -> None: + try: + version = importlib.metadata.version("litellm") + except Exception: + version = "unknown" + print() + print(orange(LITELLM_ASCII.rstrip("\n"))) + print(f" {orange('Welcome')} to {bold('LiteLLM')} {grey('v' + version)}") + print() + print(DIVIDER) + print() -def _select_providers() -> List[Dict]: - """Arrow-key multi-select. Falls back to number input if /dev/tty unavailable.""" - try: - return _select_providers_interactive() - except (OSError, termios.error): - return _select_providers_fallback() + # ── provider selector ─────────────────────────────────────────────────── + @staticmethod + def _select_providers() -> List[Dict]: + """Arrow-key multi-select. Falls back to number input if /dev/tty unavailable.""" + try: + return SetupWizard._select_interactive() + except (OSError, termios.error): + return SetupWizard._select_fallback() -def _select_providers_interactive() -> List[Dict]: - cursor = 0 - selected: Set[int] = set() + @staticmethod + def _read_key() -> str: + """Read one keypress from /dev/tty in raw mode.""" + with open("/dev/tty", "rb") as tty_fh: + fd = tty_fh.fileno() + old = termios.tcgetattr(fd) + try: + tty.setraw(fd) + ch = tty_fh.read(1) + if ch == b"\x1b": + ch2 = tty_fh.read(1) + if ch2 == b"[": + ch3 = tty_fh.read(1) + return "\x1b[" + ch3.decode("utf-8", errors="replace") + return "\x1b" + ch2.decode("utf-8", errors="replace") + return ch.decode("utf-8", errors="replace") + finally: + termios.tcsetattr(fd, termios.TCSADRAIN, old) - if _supports_color(): - sys.stdout.write(_CURSOR_HIDE) + @staticmethod + def _render_selector(cursor: int, selected: Set[int], first_render: bool) -> int: + """Draw or redraw the provider list. Returns the number of lines printed.""" + lines = [ + f"\n {bold('Add your first model')}\n", + grey(" ↑↓ to navigate · Space to select · Enter to confirm") + "\n", + "\n", + ] + for i, p in enumerate(PROVIDERS): + arrow = blue("❯") if i == cursor else " " + bullet = green("◉") if i in selected else grey("○") + name_str = bold(p["name"]) if i == cursor else p["name"] + lines.append(f" {arrow} {bullet} {name_str} {grey(p['description'])}\n") + lines.append("\n") + + content = "".join(lines) + if not first_render and _supports_color(): + sys.stdout.write(_MOVE_UP.format(content.count("\n"))) + sys.stdout.write(content) sys.stdout.flush() + return content.count("\n") - try: - _render_selector(cursor, selected, first_render=True) + @staticmethod + def _select_interactive() -> List[Dict]: + cursor, selected = 0, set() + + if _supports_color(): + sys.stdout.write(_CURSOR_HIDE) + sys.stdout.flush() + try: + SetupWizard._render_selector(cursor, selected, first_render=True) + while True: + key = SetupWizard._read_key() + if key == "\x1b[A": + cursor = (cursor - 1) % len(PROVIDERS) + elif key == "\x1b[B": + cursor = (cursor + 1) % len(PROVIDERS) + elif key == " ": + selected.symmetric_difference_update({cursor}) + elif key in ("\r", "\n"): + if not selected: + selected.add(cursor) + break + elif key in ("\x03", "\x04"): + raise KeyboardInterrupt + SetupWizard._render_selector(cursor, selected, first_render=False) + finally: + if _supports_color(): + sys.stdout.write(_CURSOR_SHOW) + sys.stdout.flush() + + return [PROVIDERS[i] for i in sorted(selected)] + + @staticmethod + def _select_fallback() -> List[Dict]: + """Number-based fallback when raw terminal input is unavailable.""" + print() + print(f" {bold('Add your first model')}") + print(grey(" Enter numbers separated by commas (e.g. 1,2). Press Enter to confirm.")) + print() + for i, p in enumerate(PROVIDERS, 1): + print(f" {grey(str(i) + '.')} {bold(p['name'])} {grey(p['description'])}") + print() while True: - key = _read_key() - - if key == "\x1b[A": # Up - cursor = (cursor - 1) % len(PROVIDERS) - elif key == "\x1b[B": # Down - cursor = (cursor + 1) % len(PROVIDERS) - elif key == " ": # Space — toggle - if cursor in selected: - selected.discard(cursor) - else: - selected.add(cursor) - elif key in ("\r", "\n"): # Enter — confirm - if not selected: - selected.add(cursor) # select highlighted item if nothing chosen - break - elif key in ("\x03", "\x04"): # Ctrl+C / Ctrl+D - raise KeyboardInterrupt - - _render_selector(cursor, selected, first_render=False) - finally: - if _supports_color(): - sys.stdout.write(_CURSOR_SHOW) - sys.stdout.flush() - - return [PROVIDERS[i] for i in sorted(selected)] - - -def _select_providers_fallback() -> List[Dict]: - """Number-based fallback when raw terminal input is unavailable.""" - print() - print(f" {bold('Add your first model')}") - print(grey(" Enter numbers separated by commas (e.g. 1,2). Press Enter to confirm.")) - print() - for i, p in enumerate(PROVIDERS, 1): - print(f" {grey(str(i) + '.')} {bold(p['name'])} {grey(p['description'])}") - print() - - selected_nums: List[int] = [] - while True: - raw = input(f" {blue('❯')} Provider(s): ").strip() - if not raw: - if not selected_nums: + raw = input(f" {blue('❯')} Provider(s): ").strip() + if not raw: print(grey(" Please select at least one provider.")) continue - break - try: - nums = [int(x.strip()) for x in raw.replace(" ", ",").split(",") if x.strip()] - valid = [n for n in nums if 1 <= n <= len(PROVIDERS)] - if not valid: - print(grey(f" Enter numbers between 1 and {len(PROVIDERS)}.")) + try: + nums = [int(x.strip()) for x in raw.replace(" ", ",").split(",") if x.strip()] + valid = sorted({n for n in nums if 1 <= n <= len(PROVIDERS)}) + if not valid: + print(grey(f" Enter numbers between 1 and {len(PROVIDERS)}.")) + continue + return [PROVIDERS[i - 1] for i in valid] + except ValueError: + print(grey(" Enter numbers separated by commas, e.g. 1,3")) + + # ── key collection ─────────────────────────────────────────────────────── + + @staticmethod + def _collect_keys(providers: List[Dict]) -> Dict[str, str]: + env_vars: Dict[str, str] = {} + print() + print(DIVIDER) + print() + print(f" {bold('Enter your API keys')}") + print(grey(" Keys are stored only in the generated config file.")) + print() + + for p in providers: + if p["env_key"] is None: + print(f" {green(p['name'])}: {grey('no key needed (uses local Ollama)')}") continue - selected_nums = sorted(set(valid)) - break - except ValueError: - print(grey(" Enter numbers separated by commas, e.g. 1,3")) - return [PROVIDERS[i - 1] for i in selected_nums] - - -# --------------------------------------------------------------------------- -# Credential validation -# --------------------------------------------------------------------------- - -def _test_provider_key(provider: Dict, env_vars: Dict[str, str]) -> Optional[str]: - """ - Make a minimal test completion to validate credentials. - Returns None on success, or an error string on failure. - Skipped for providers without models (Azure) or local providers (Ollama). - """ - models = provider.get("models", []) - if not models or provider["id"] in ("azure", "ollama"): - return None - - import litellm # noqa: PLC0415 - - litellm.suppress_debug_info = True - litellm.set_verbose = False - - # Temporarily inject the keys so litellm can pick them up - saved: Dict[str, Optional[str]] = {} - for k, v in env_vars.items(): - saved[k] = os.environ.get(k) - os.environ[k] = v - - try: - litellm.completion( - model=models[0], - messages=[{"role": "user", "content": "hi"}], - max_tokens=1, - ) - return None - except Exception as exc: - # Surface a short, readable reason - msg = str(exc) - for fragment in ("AuthenticationError", "InvalidAPIKey", "Unauthorized", "401", "403"): - if fragment in msg: - return "invalid API key" - return msg[:120] - finally: - for k, old_v in saved.items(): - if old_v is None: - os.environ.pop(k, None) - else: - os.environ[k] = old_v - - -# --------------------------------------------------------------------------- -# API key collection -# --------------------------------------------------------------------------- - -def _collect_keys(providers: List[Dict]) -> Dict[str, str]: - env_vars: Dict[str, str] = {} - print() - print(DIVIDER) - print() - print(f" {bold('Enter your API keys')}") - print(grey(" Keys are stored only in the generated config file.")) - print() - - for p in providers: - if p["env_key"] is None: - # Ollama — no key needed - print(f" {green(p['name'])}: {grey('no key needed (uses local Ollama)')}") - continue - - hint = grey(p.get("key_hint", "")) - key = "" - while not key: - key = input(f" {blue('❯')} {bold(p['name'])} API key {hint}: ").strip() + key = SetupWizard._prompt_key(p) if not key: - print(grey(" Key is required. Leave blank to skip this provider.")) - skip = input(grey(" Skip? (y/N): ")).strip().lower() - if skip == "y": - break + continue - if not key: - continue + env_vars[p["env_key"]] = key - env_vars[p["env_key"]] = key - - # Extra keys (e.g. AWS secret + region) - if p.get("needs_extra"): for extra_key, extra_hint in zip( p.get("extra_keys", []), p.get("extra_hints", []) ): - val = input( - f" {blue('❯')} {extra_key} {grey(extra_hint)}: " - ).strip() + val = input(f" {blue('❯')} {extra_key} {grey(extra_hint)}: ").strip() if val: env_vars[extra_key] = val - # API base for Azure - if p.get("needs_api_base"): - api_base = input( - f" {blue('❯')} Azure endpoint URL {grey(p.get('api_base_hint', ''))}: " + if p.get("needs_api_base"): + api_base = input( + f" {blue('❯')} Azure endpoint URL {grey(p.get('api_base_hint', ''))}: " + ).strip() + if api_base: + env_vars[f"_LITELLM_AZURE_API_BASE_{p['id'].upper()}"] = api_base + + SetupWizard._validate_and_report(p, key) + + return env_vars + + @staticmethod + def _prompt_key(provider: Dict) -> str: + """Prompt for a provider's API key, with skip option. Returns the key or ''.""" + hint = grey(provider.get("key_hint", "")) + while True: + key = input(f" {blue('❯')} {bold(provider['name'])} API key {hint}: ").strip() + if key: + return key + print(grey(" Key is required. Leave blank to skip this provider.")) + if input(grey(" Skip? (y/N): ")).strip().lower() == "y": + return "" + + @staticmethod + def _validate_and_report(provider: Dict, api_key: str) -> None: + """ + Validate credentials using litellm.utils.check_valid_key. + Offers a re-entry loop on failure. + """ + test_model: Optional[str] = provider.get("test_model") + if not test_model: + return # Azure / Bedrock / Ollama — skip + + while True: + print(f" {grey('Testing credentials…')}", end="", flush=True) + valid = check_valid_key(model=test_model, api_key=api_key) + if valid: + print(f"\r {green(_CHECK + ' ' + provider['name'])} credentials valid ") + return + + print(f"\r {_c(_BOLD, _CROSS)} {bold(provider['name'])} {grey('invalid API key')}") + if input(f" {blue('❯')} Re-enter key? {grey('(y/N)')}: ").strip().lower() != "y": + return + + hint = grey(provider.get("key_hint", "")) + new_key = input( + f" {blue('❯')} {bold(provider['name'])} API key {hint}: " ).strip() - if api_base: - env_vars[f"_LITELLM_AZURE_API_BASE_{p['id'].upper()}"] = api_base + if not new_key: + return + api_key = new_key - # Validate credentials - print(f" {grey('Testing credentials…')}", end="", flush=True) - error = _test_provider_key(p, env_vars) - if error is None: - print(f"\r {green(_CHECK + ' ' + p['name'])} credentials valid ") - else: - print(f"\r {_c(_BOLD, '✘')} {bold(p['name'])} {grey(error)}") - retry = input(f" {blue('❯')} Re-enter key? {grey('(y/N)')}: ").strip().lower() - if retry == "y": - del env_vars[p["env_key"]] - # Re-prompt by looping — restart key collection for this provider - while True: - key = input(f" {blue('❯')} {bold(p['name'])} API key {hint}: ").strip() - if not key: - break - env_vars[p["env_key"]] = key - print(f" {grey('Testing credentials…')}", end="", flush=True) - error = _test_provider_key(p, env_vars) - if error is None: - print(f"\r {green(_CHECK + ' ' + p['name'])} credentials valid ") - break - print(f"\r {_c(_BOLD, '✘')} {bold(p['name'])} {grey(error)}") + # ── proxy settings ─────────────────────────────────────────────────────── - return env_vars + @staticmethod + def _proxy_settings() -> "tuple[int, str]": + print() + print(DIVIDER) + print() + print(f" {bold('Proxy settings')}") + print() + port_raw = input(f" {blue('❯')} Port {grey('[4000]')}: ").strip() + port = int(port_raw) if port_raw.isdigit() else 4000 + key_raw = input(f" {blue('❯')} Master key {grey('[auto-generate]')}: ").strip() + master_key = key_raw if key_raw else f"sk-{secrets.token_urlsafe(32)}" + return port, master_key + # ── config generation ──────────────────────────────────────────────────── -# --------------------------------------------------------------------------- -# Config generation -# --------------------------------------------------------------------------- + @staticmethod + def _build_config( + providers: List[Dict], + env_vars: Dict[str, str], + port: int, + master_key: str, + ) -> str: + lines = ["model_list:"] + for p in providers: + models = p["models"] if p["models"] else (["azure/gpt-4o"] if p["id"] == "azure" else []) + for model in models: + display = model.split("/")[-1] if "/" in model else model + lines += [ + f" - model_name: {display}", + f" litellm_params:", + f" model: {model}", + ] + if p["env_key"] and p["env_key"] in env_vars: + lines.append(f" api_key: os.environ/{p['env_key']}") + if p.get("api_base"): + lines.append(f" api_base: {p['api_base']}") + elif p.get("needs_api_base"): + azure_key = f"_LITELLM_AZURE_API_BASE_{p['id'].upper()}" + if azure_key in env_vars: + lines.append(f" api_base: {env_vars.pop(azure_key)}") + if p.get("api_version"): + lines.append(f" api_version: {p['api_version']}") -def _build_config( - providers: List[Dict], - env_vars: Dict[str, str], - port: int, - master_key: str, -) -> str: - lines = ["model_list:"] + lines += ["", "general_settings:", f" master_key: {master_key}", ""] - for p in providers: - if not p["models"] and p["id"] == "azure": - # Azure — add a generic placeholder - models_to_add = ["azure/gpt-4o"] - else: - models_to_add = p["models"] + real_vars = {k: v for k, v in env_vars.items() if not k.startswith("_LITELLM_")} + if real_vars: + lines.append("environment_variables:") + for k, v in real_vars.items(): + lines.append(f' {k}: "{v}"') + lines.append("") - for model in models_to_add: - # User-facing model name (strip provider prefix for display) - display_name = model.split("/")[-1] if "/" in model else model - lines.append(f" - model_name: {display_name}") - lines.append(f" litellm_params:") - lines.append(f" model: {model}") + return "\n".join(lines) - if p["env_key"] and p["env_key"] in env_vars: - lines.append(f" api_key: os.environ/{p['env_key']}") + # ── success + launch ───────────────────────────────────────────────────── - if p.get("api_base"): - lines.append(f" api_base: {p['api_base']}") - elif p.get("needs_api_base"): - azure_base_key = f"_LITELLM_AZURE_API_BASE_{p['id'].upper()}" - if azure_base_key in env_vars: - lines.append(f" api_base: {env_vars.pop(azure_base_key)}") - if p.get("api_version"): - lines.append(f" api_version: {p['api_version']}") + @staticmethod + def _print_success(config_path: Path, port: int, master_key: str) -> None: + print() + print(DIVIDER) + print() + print(f" {green(_CHECK + ' Config saved')} → {bold(str(config_path))}") + print() + print(f" {bold('To start your proxy:')}") + print() + print(f" {grey('$')} litellm --config {config_path}") + print() + print(f" {bold('Then set your client:')}") + print() + print(f" export OPENAI_BASE_URL=http://localhost:{port}") + print(f" export OPENAI_API_KEY={master_key}") + print() + print(DIVIDER) + print() - lines.append("") - lines.append("general_settings:") - lines.append(f" master_key: {master_key}") - lines.append("") + @staticmethod + def _offer_start(config_path: Path, port: int, master_key: str) -> None: + start = input(f" {blue('❯')} Start the proxy now? {grey('(Y/n)')}: ").strip().lower() + if start not in ("", "y", "yes"): + print() + print(f" Run {bold(f'litellm --config {config_path}')} whenever you're ready.") + print() + print(grey(f" Quick test once running: curl http://localhost:{port}/health")) + print() + return - # Write env vars inline so the config is self-contained - real_env_vars = {k: v for k, v in env_vars.items() if not k.startswith("_LITELLM_")} - if real_env_vars: - lines.append("environment_variables:") - for k, v in real_env_vars.items(): - lines.append(f" {k}: \"{v}\"") - lines.append("") - - return "\n".join(lines) - - -# --------------------------------------------------------------------------- -# Proxy settings -# --------------------------------------------------------------------------- - -def _proxy_settings() -> tuple[int, str]: - print() - print(DIVIDER) - print() - print(f" {bold('Proxy settings')}") - print() - - port_raw = input(f" {blue('❯')} Port {grey('[4000]')}: ").strip() - port: int = int(port_raw) if port_raw.isdigit() else 4000 - - key_raw = input( - f" {blue('❯')} Master key {grey('[auto-generate]')}: " - ).strip() - master_key = key_raw if key_raw else f"sk-{secrets.token_urlsafe(32)}" - - return port, master_key - - -# --------------------------------------------------------------------------- -# Main wizard entrypoint -# --------------------------------------------------------------------------- - -def run_setup_wizard() -> Optional[str]: - """ - Run the interactive setup wizard. - - Returns the path to the generated config file, or None if aborted. - """ - try: - _run_wizard() - except (KeyboardInterrupt, EOFError): - print(f"\n\n {grey('Setup cancelled.')}\n") - return None - return None # caller receives path via side effect (printed to stdout) - - -def _run_wizard() -> None: - _print_welcome() - - print(f" {bold('Lets get started.')}") - print() - - # Step 1: providers - providers = _select_providers() - - # Step 2: API keys - env_vars = _collect_keys(providers) - - # Step 3: proxy settings - port, master_key = _proxy_settings() - - # Step 4: write config - config_content = _build_config(providers, env_vars, port, master_key) - - config_path = Path(os.getcwd()) / "litellm_config.yaml" - config_path.write_text(config_content) - - # Step 5: print success - print() - print(DIVIDER) - print() - print(f" {green(_CHECK + ' Config saved')} → {bold(str(config_path))}") - print() - print(f" {bold('To start your proxy:')}") - print() - print(f" {grey('$')} litellm --config {config_path}") - print() - print(f" {bold('Then set your client:')}") - print() - print(f" export OPENAI_BASE_URL=http://localhost:{port}") - print(f" export OPENAI_API_KEY={master_key}") - print() - print(DIVIDER) - print() - - # Step 6: offer to start now - start = input(f" {blue('❯')} Start the proxy now? {grey('(Y/n)')}: ").strip().lower() - if start in ("", "y", "yes"): print() print(DIVIDER) print() @@ -574,22 +516,17 @@ def _run_wizard() -> None: print() print(f" {green(_CHECK)} Starting… {grey('(Ctrl+C to stop)')}") print() - # exec replaces this process with the proxy server. - # Use the litellm console script (same bin dir as the running Python) - # rather than `python -m litellm` — litellm has no __main__.py. + scripts_dir = sysconfig.get_path("scripts") litellm_bin = os.path.join(scripts_dir, "litellm") - os.execlp( # noqa: S606 - litellm_bin, - litellm_bin, - "--config", - str(config_path), - "--port", - str(port), - ) - else: - print() - print(f" Run {bold(f'litellm --config {config_path}')} whenever you're ready.") - print() - print(grey(f" Quick test once running: curl http://localhost:{port}/health")) - print() + os.execlp(litellm_bin, litellm_bin, "--config", str(config_path), "--port", str(port)) # noqa: S606 + + +# --------------------------------------------------------------------------- +# Public entrypoint +# --------------------------------------------------------------------------- + +def run_setup_wizard() -> Optional[str]: + """Run the interactive setup wizard. Called by `litellm --setup`.""" + SetupWizard.run() + return None