diff --git a/myagents/backends.py b/myagents/backends.py new file mode 100644 index 0000000..683d9fa --- /dev/null +++ b/myagents/backends.py @@ -0,0 +1,141 @@ +"""Provider metadata for Claude Code backend switching.""" + +from __future__ import annotations + +from dataclasses import dataclass + + +@dataclass(frozen=True) +class Provider: + """A Claude Code-compatible model provider.""" + + id: str + display_name: str + base_url: str | None + default_model: str | None + key_name: str + auth_env_var: str = "ANTHROPIC_AUTH_TOKEN" + + +#: Built-in providers. Keys are normalized provider ids (hyphens). +PROVIDERS: dict[str, Provider] = { + "deepseek": Provider( + id="deepseek", + display_name="DeepSeek", + base_url="https://api.deepseek.com/anthropic", + default_model="deepseek-v4-pro", + key_name="deepseek", + ), + "kimi": Provider( + id="kimi", + display_name="Kimi", + base_url="https://api.moonshot.ai/anthropic", + default_model="kimi-k2.6", + key_name="kimi", + ), + "kimi-code": Provider( + id="kimi-code", + display_name="Kimi Code", + base_url="https://api.kimi.com/coding", + default_model="kimi-for-coding", + key_name="kimi_code", + ), + "claude": Provider( + id="claude", + display_name="Claude (Anthropic)", + base_url=None, + default_model=None, + key_name="claude", + auth_env_var="ANTHROPIC_API_KEY", + ), +} + +#: Env vars that control Claude Code model selection / endpoint. +AUTH_ENV_VARS: tuple[str, ...] = ( + "ANTHROPIC_API_KEY", + "ANTHROPIC_AUTH_TOKEN", + "ANTHROPIC_BASE_URL", + "ANTHROPIC_MODEL", +) + +#: Env vars that map Claude model tiers to a single provider model. +MODEL_TIER_ENV_VARS: tuple[str, ...] = ( + "ANTHROPIC_DEFAULT_FABLE_MODEL", + "ANTHROPIC_DEFAULT_FABLE_MODEL_NAME", + "ANTHROPIC_DEFAULT_OPUS_MODEL", + "ANTHROPIC_DEFAULT_OPUS_MODEL_NAME", + "ANTHROPIC_DEFAULT_SONNET_MODEL", + "ANTHROPIC_DEFAULT_SONNET_MODEL_NAME", + "ANTHROPIC_DEFAULT_HAIKU_MODEL", + "ANTHROPIC_DEFAULT_HAIKU_MODEL_NAME", + "CLAUDE_CODE_SUBAGENT_MODEL", +) + +#: All env vars managed by the backend switcher. +MANAGED_ENV_VARS: tuple[str, ...] = AUTH_ENV_VARS + MODEL_TIER_ENV_VARS + +#: Provider ids that route to a third-party endpoint. +THIRD_PARTY_PROVIDER_IDS: tuple[str, ...] = ("deepseek", "kimi", "kimi-code") + + +def normalize_provider_id(name: str) -> str: + """Normalize user input: underscores and case-insensitive to hyphens.""" + return name.strip().lower().replace("_", "-") + + +def get_provider(name: str) -> Provider: + """Return a provider by id, accepting aliases like ``kimi_code``.""" + normalized = normalize_provider_id(name) + if normalized not in PROVIDERS: + raise KeyError(normalized) + return PROVIDERS[normalized] + + +def list_providers() -> list[Provider]: + """Return built-in providers in a stable order.""" + return list(PROVIDERS.values()) + + +def is_third_party(provider: Provider) -> bool: + """Whether the provider is a non-Anthropic endpoint.""" + return provider.id in THIRD_PARTY_PROVIDER_IDS + + +def build_provider_env(provider: Provider, key: str, model: str | None = None) -> dict[str, str]: + """Return the env dict to apply for a provider. + + For third-party providers this sets the auth token, base url, and model + tier env vars. For the official ``claude`` provider the dict is empty; + the caller should remove managed env vars instead. + """ + if not is_third_party(provider): + return {} + + env: dict[str, str] = { + provider.auth_env_var: key, + } + if provider.base_url: + env["ANTHROPIC_BASE_URL"] = provider.base_url + + active_model = model or provider.default_model or "" + if active_model: + env["ANTHROPIC_MODEL"] = active_model + for tier_var in MODEL_TIER_ENV_VARS: + env[tier_var] = active_model + + return env + + +def detect_provider(env: dict[str, str]) -> Provider | None: + """Best-effort detect the active provider from a Claude Code env dict.""" + base_url = env.get("ANTHROPIC_BASE_URL", "") + if not base_url: + # No base_url means official Claude or unset. + if env.get("ANTHROPIC_API_KEY") and not env.get("ANTHROPIC_AUTH_TOKEN"): + return PROVIDERS["claude"] + return None + + for provider in PROVIDERS.values(): + if provider.base_url and provider.base_url.rstrip("/") == base_url.rstrip("/"): + return provider + return None diff --git a/myagents/claude_settings.py b/myagents/claude_settings.py new file mode 100644 index 0000000..1b8f259 --- /dev/null +++ b/myagents/claude_settings.py @@ -0,0 +1,125 @@ +"""Read/write project-level Claude Code settings.json for backend switching. + +Claude Code configuration precedence is: + + local settings > project settings > user settings + +This module targets the project-level file at +``workspace/.agents/settings.json`` (symlinked from ``workspace/.claude/settings.json``), +so the selected backend only affects the current workspace without polluting the +user's global ``~/.claude/settings.json``. +""" + +from __future__ import annotations + +import json +from pathlib import Path +from typing import Any + +from myagents.backends import ( + MANAGED_ENV_VARS, + PROVIDERS, + Provider, + build_provider_env, + detect_provider, +) +from myagents.project_root import get_workspace_root + + +def _settings_path() -> Path: + """Return the project-level Claude Code settings path.""" + return get_workspace_root() / ".agents" / "settings.json" + + +def _load_json(path: Path) -> dict[str, Any]: + """Load JSON; return empty dict on missing/corrupt.""" + try: + return json.loads(path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError): + return {} + + +def _save_json(path: Path, data: dict[str, Any]) -> None: + """Persist settings atomically with safe permissions.""" + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(json.dumps(data, indent=2, ensure_ascii=False) + "\n", encoding="utf-8") + + +def load_project_settings() -> dict[str, Any]: + """Load the project-level Claude Code settings dict.""" + return _load_json(_settings_path()) + + +def save_project_settings(data: dict[str, Any]) -> None: + """Persist the project-level Claude Code settings dict.""" + _save_json(_settings_path(), data) + + +def get_project_env() -> dict[str, str]: + """Return the current ``env`` block from project settings.""" + env = load_project_settings().get("env", {}) + return {k: v for k, v in env.items() if isinstance(v, str)} + + +def get_active_provider() -> Provider | None: + """Detect the active provider from project-level env.""" + return detect_provider(get_project_env()) + + +def apply_provider( + provider: Provider, + key: str | None = None, + model: str | None = None, + base_url: str | None = None, +) -> dict[str, Any]: + """Return updated project settings with the given provider applied. + + This does not write to disk; callers should pass the result to + ``save_project_settings``. + """ + settings = load_project_settings() + env: dict[str, str] = { + k: v for k, v in settings.get("env", {}).items() if isinstance(v, str) + } + + # Remove stale managed env vars first. + for var in MANAGED_ENV_VARS: + env.pop(var, None) + + if provider.id == "claude": + # Official Claude: no third-party env vars needed. Any ANTHROPIC_API_KEY + # the user manages separately is left untouched. + pass + else: + if not key: + raise ValueError(f"API key required for provider '{provider.id}'") + provider_env = build_provider_env(provider, key, model=model) + if base_url: + provider_env["ANTHROPIC_BASE_URL"] = base_url + env.update(provider_env) + + if env: + settings["env"] = env + else: + settings.pop("env", None) + + return settings + + +def reset_backend() -> dict[str, Any]: + """Convenience: return settings with all managed backend env vars removed.""" + return apply_provider(PROVIDERS["claude"]) + + +def describe_active_backend() -> str: + """Human-readable description of the active project-level backend.""" + provider = get_active_provider() + env = get_project_env() + if provider is None: + if env.get("ANTHROPIC_API_KEY") and not env.get("ANTHROPIC_AUTH_TOKEN"): + return "Claude (official API key)" + return "Claude (default / not configured)" + if provider.id == "claude": + return "Claude (official)" + model = env.get("ANTHROPIC_MODEL") or provider.default_model or "unknown" + return f"{provider.display_name} ({model})" diff --git a/myagents/commands/backend.py b/myagents/commands/backend.py new file mode 100644 index 0000000..94265cb --- /dev/null +++ b/myagents/commands/backend.py @@ -0,0 +1,212 @@ +"""``xiaohe backend`` — switch Claude Code's model backend.""" + +from __future__ import annotations + +import click +from rich.console import Console + +from myagents.backends import ( + Provider, + get_provider, + is_third_party, + list_providers, +) +from myagents.claude_settings import ( + apply_provider, + describe_active_backend, + get_active_provider, + save_project_settings, +) +from myagents.secrets import get_key, has_key, remove_key, set_key +from myagents.settings import set_setting + +stderr_console = Console(stderr=True) +console = Console() + + +def _secure_prompt_key(provider: Provider) -> str: + """Prompt securely for an API key, falling back to /dev/tty when needed.""" + prompt_text = f"Enter {provider.display_name} API key" + try: + return click.prompt(prompt_text, hide_input=True, err=True) + except click.UsageError: + # click.prompt raises when stdin is not a TTY. Fall back to /dev/tty. + try: + with open("/dev/tty", encoding="utf-8") as tty: # noqa: PTH123 + return tty.readline().rstrip("\n") + except OSError as exc: + raise click.ClickException( + f"Cannot read API key interactively: {exc}" + ) from exc + + +@click.group("backend") +def backend_cmd() -> None: + """Switch Claude Code's model backend (DeepSeek, Kimi, Kimi Code, Claude).""" + + +@backend_cmd.command("list") +def backend_list() -> None: + """List built-in backends and the active one.""" + active = get_active_provider() + console.print("[bold]Built-in backends:[/bold]") + for provider in list_providers(): + marker = "" + if active and active.id == provider.id: + marker = " [green](current)[/green]" + key_status = "[dim]no key[/dim]" + if has_key(provider.key_name): + key_status = "[cyan]key stored[/cyan]" + if provider.id == "claude": + key_status = "[dim]official[/dim]" + console.print( + f" {provider.id:12} {provider.display_name:18} " + f"{provider.default_model or '-':20} {key_status}{marker}" + ) + + +@backend_cmd.command("current") +def backend_current() -> None: + """Show the active backend for this workspace.""" + console.print(f"[bold]Current backend:[/bold] {describe_active_backend()}") + + +@backend_cmd.command("use") +@click.argument("provider_id") +@click.option("--key", help="API key (scripting only; appears in shell history)") +@click.option("--model", help="Override the default model for this provider") +@click.option("--base-url", help="Override the provider base URL") +@click.option("--yes", "-y", is_flag=True, help="Skip confirmation prompt") +def backend_use( + provider_id: str, + key: str | None, + model: str | None, + base_url: str | None, + yes: bool, +) -> None: + """Switch Claude Code to the given backend.""" + try: + provider = get_provider(provider_id) + except KeyError: + known = ", ".join(p.id for p in list_providers()) + raise click.ClickException( + f"Unknown backend '{provider_id}'. Choose from: {known}" + ) + + if is_third_party(provider): + resolved_key = key or get_key(provider.key_name) + if not resolved_key: + resolved_key = _secure_prompt_key(provider) + if not resolved_key: + raise click.ClickException("API key is required for this backend.") + else: + resolved_key = None + + summary = provider.display_name + if model: + summary += f" (model: {model})" + if base_url: + summary += f" (base-url: {base_url})" + + if not yes and not click.confirm( + f"Set backend to {summary} for this workspace?", + default=True, + err=True, + ): + console.print("Cancelled.") + return + + settings = apply_provider( + provider, + key=resolved_key, + model=model, + base_url=base_url, + ) + save_project_settings(settings) + set_setting("backend_provider", provider.id) + + console.print(f"[bold green]Backend switched to {summary}.[/bold green]") + if provider.id == "claude": + console.print( + "[dim]Cleared third-party backend env vars from workspace " + ".agents/settings.json.[/dim]" + ) + + +@backend_cmd.command("reset") +@click.option("--yes", "-y", is_flag=True, help="Skip confirmation prompt") +def backend_reset(yes: bool) -> None: + """Reset to the official Claude backend.""" + ctx = click.get_current_context() + ctx.invoke(backend_use, provider_id="claude", yes=yes) + + +@backend_cmd.group("key") +def backend_key() -> None: + """Manage stored API keys.""" + + +@backend_key.command("set") +@click.argument("provider_id") +@click.option("--key", help="API key (scripting only; appears in shell history)") +def backend_key_set(provider_id: str, key: str | None) -> None: + """Store an API key for a backend without switching to it.""" + try: + provider = get_provider(provider_id) + except KeyError: + known = ", ".join(p.id for p in list_providers()) + raise click.ClickException( + f"Unknown backend '{provider_id}'. Choose from: {known}" + ) + + if provider.id == "claude": + raise click.ClickException( + "Use the ANTHROPIC_API_KEY environment variable for official Claude auth." + ) + + resolved_key = key + if not resolved_key: + resolved_key = _secure_prompt_key(provider) + if not resolved_key: + raise click.ClickException("API key cannot be empty.") + + set_key(provider.key_name, resolved_key) + console.print( + f"[green]Stored {provider.display_name} key in ~/.xiaohe/agent/config.json.[/green]" + ) + + +@backend_key.command("rm") +@click.argument("provider_id") +@click.option( + "--yes", + "-y", + is_flag=True, + help="Skip the confirmation prompt", +) +def backend_key_rm(provider_id: str, yes: bool) -> None: + """Remove the stored API key for a backend.""" + try: + provider = get_provider(provider_id) + except KeyError: + known = ", ".join(p.id for p in list_providers()) + raise click.ClickException( + f"Unknown backend '{provider_id}'. Choose from: {known}" + ) + + if not yes and not click.confirm( + f"Remove stored key for {provider.display_name}?", + default=False, + err=True, + ): + console.print("Cancelled.") + return + + if remove_key(provider.key_name): + console.print(f"[green]Removed {provider.display_name} key.[/green]") + else: + console.print(f"[yellow]No stored key for {provider.display_name}.[/yellow]") + + +if __name__ == "__main__": + backend_cmd() diff --git a/myagents/entrypoints.py b/myagents/entrypoints.py index 7a43679..786f96f 100644 --- a/myagents/entrypoints.py +++ b/myagents/entrypoints.py @@ -27,6 +27,7 @@ def default_agent() -> str: def build_xiaohe_cli(): """The ``xiaohe`` command group: agent forwarding + init/sync/upgrade/switch/uninstall/version.""" + from myagents.commands.backend import backend_cmd from myagents.commands.switch import switch_cmd from myagents.commands.sync_workspace import init_cmd, sync_cmd from myagents.commands.uninstall import uninstall_cmd @@ -34,6 +35,7 @@ def build_xiaohe_cli(): from myagents.commands.version import version_cmd xiaohe_cli = build_cli(default_agent(), prog_name="xiaohe", offer_install=True) + xiaohe_cli.add_command(backend_cmd) xiaohe_cli.add_command(init_cmd) xiaohe_cli.add_command(sync_cmd) xiaohe_cli.add_command(upgrade_cmd) diff --git a/myagents/secrets.py b/myagents/secrets.py new file mode 100644 index 0000000..3a22260 --- /dev/null +++ b/myagents/secrets.py @@ -0,0 +1,126 @@ +"""Secure API key storage for backend switching. + +Reads keys from two sources, in priority order: + +1. ``~/.xiaohe/agent/config.json`` under ``keys.*`` — the primary store for + myagents/xiaohe secrets. +2. ``~/.mytoolkit/config.json`` under ``keys.*`` — for backward compatibility + with keys already managed by mytoolkit. + +Writes always go to ``~/.xiaohe/agent/config.json`` so that xiaohe-managed +keys shadow mytoolkit keys without modifying them. +""" + +from __future__ import annotations + +import json +import os +import stat +import tempfile +from pathlib import Path +from typing import Any + + +XIAOHE_CONFIG_DIR = Path.home() / ".xiaohe" / "agent" +XIAOHE_CONFIG_PATH = XIAOHE_CONFIG_DIR / "config.json" + +MYTOOLKIT_CONFIG_PATH = Path.home() / ".mytoolkit" / "config.json" + + +def _load_json(path: Path) -> dict[str, Any]: + """Load JSON from path; return empty dict on missing/corrupt.""" + try: + text = path.read_text(encoding="utf-8") + except (OSError, UnicodeDecodeError): + return {} + try: + data = json.loads(text) + except json.JSONDecodeError: + return {} + return data if isinstance(data, dict) else {} + + +def _atomic_write(path: Path, data: dict[str, Any]) -> None: + """Write JSON atomically and restrict permissions on Unix.""" + path.parent.mkdir(parents=True, exist_ok=True) + serialized = json.dumps(data, indent=2, ensure_ascii=False) + "\n" + + fd, tmp = tempfile.mkstemp(dir=path.parent, prefix=f".{path.name}.") + try: + with os.fdopen(fd, "w", encoding="utf-8") as fh: + fh.write(serialized) + if os.name != "nt": + os.chmod(tmp, stat.S_IRUSR | stat.S_IWUSR) + os.replace(tmp, path) + except Exception: + try: + os.unlink(tmp) + except OSError: + pass + raise + + +def _read_key_from_config(path: Path, key_name: str) -> str | None: + """Read a single key from a config file's ``keys`` section.""" + data = _load_json(path) + value = data.get("keys", {}).get(key_name) + return value if isinstance(value, str) and value else None + + +def load_xiaohe_config() -> dict[str, Any]: + """Load the full ``~/.xiaohe/agent/config.json``.""" + return _load_json(XIAOHE_CONFIG_PATH) + + +def save_xiaohe_config(data: dict[str, Any]) -> None: + """Persist the full ``~/.xiaohe/agent/config.json``.""" + _atomic_write(XIAOHE_CONFIG_PATH, data) + + +def get_key(key_name: str) -> str | None: + """Return a key, preferring the xiaohe store over mytoolkit. + + Returns ``None`` when no key is stored or the stored value is empty. + """ + value = _read_key_from_config(XIAOHE_CONFIG_PATH, key_name) + if value: + return value + return _read_key_from_config(MYTOOLKIT_CONFIG_PATH, key_name) + + +def set_key(key_name: str, value: str) -> None: + """Store a key in ``~/.xiaohe/agent/config.json``. + + Empty values are rejected. + """ + if not isinstance(value, str) or not value.strip(): + raise ValueError("API key cannot be empty") + + data = load_xiaohe_config() + if "keys" not in data or not isinstance(data["keys"], dict): + data["keys"] = {} + data["keys"][key_name] = value + save_xiaohe_config(data) + + +def remove_key(key_name: str) -> bool: + """Remove a key from ``~/.xiaohe/agent/config.json``. + + Returns ``True`` if the key existed and was removed. + """ + data = load_xiaohe_config() + keys = data.get("keys", {}) + if not isinstance(keys, dict): + return False + if key_name not in keys: + return False + del keys[key_name] + if not keys: + data.pop("keys", None) + save_xiaohe_config(data) + return True + + +def has_key(key_name: str) -> bool: + """Return whether a non-empty key exists in either store.""" + return get_key(key_name) is not None diff --git a/tests/test_backends.py b/tests/test_backends.py new file mode 100644 index 0000000..b031989 --- /dev/null +++ b/tests/test_backends.py @@ -0,0 +1,96 @@ +"""Tests for myagents.backends.""" + +import pytest + +from myagents import backends as backends_mod +from myagents.backends import ( + Provider, + build_provider_env, + detect_provider, + get_provider, + is_third_party, + list_providers, + normalize_provider_id, +) + + +class TestProviderRegistry: + def test_list_providers_includes_all(self) -> None: + ids = {p.id for p in list_providers()} + assert ids == {"deepseek", "kimi", "kimi-code", "claude"} + + def test_get_provider_by_id(self) -> None: + provider = get_provider("kimi") + assert provider.id == "kimi" + assert provider.base_url == "https://api.moonshot.ai/anthropic" + + def test_get_provider_normalizes_name(self) -> None: + assert get_provider("Kimi").id == "kimi" + assert get_provider("kimi_code").id == "kimi-code" + + def test_unknown_provider_raises(self) -> None: + with pytest.raises(KeyError): + get_provider("openai") + + +class TestNormalizeProviderId: + def test_hyphenates_underscores(self) -> None: + assert normalize_provider_id("kimi_code") == "kimi-code" + + def test_lowercases(self) -> None: + assert normalize_provider_id("DeepSeek") == "deepseek" + + +class TestBuildProviderEnv: + def test_deepseek_env(self) -> None: + provider = get_provider("deepseek") + env = build_provider_env(provider, "sk-test") + assert env["ANTHROPIC_AUTH_TOKEN"] == "sk-test" + assert env["ANTHROPIC_BASE_URL"] == "https://api.deepseek.com/anthropic" + assert env["ANTHROPIC_MODEL"] == "deepseek-v4-pro" + assert env["ANTHROPIC_DEFAULT_FABLE_MODEL"] == "deepseek-v4-pro" + + def test_kimi_code_env(self) -> None: + provider = get_provider("kimi-code") + env = build_provider_env(provider, "sk-test") + assert env["ANTHROPIC_BASE_URL"] == "https://api.kimi.com/coding" + assert env["ANTHROPIC_MODEL"] == "kimi-for-coding" + + def test_model_override(self) -> None: + provider = get_provider("kimi") + env = build_provider_env(provider, "sk-test", model="kimi-k2.5") + assert env["ANTHROPIC_MODEL"] == "kimi-k2.5" + + def test_official_claude_returns_empty(self) -> None: + provider = get_provider("claude") + assert build_provider_env(provider, "sk-test") == {} + + +class TestDetectProvider: + def test_detects_deepseek(self) -> None: + env = {"ANTHROPIC_BASE_URL": "https://api.deepseek.com/anthropic"} + assert detect_provider(env) == get_provider("deepseek") + + def test_detects_kimi_code(self) -> None: + env = {"ANTHROPIC_BASE_URL": "https://api.kimi.com/coding"} + assert detect_provider(env) == get_provider("kimi-code") + + def test_no_base_url_with_api_key_is_claude(self) -> None: + env = {"ANTHROPIC_API_KEY": "sk-test"} + assert detect_provider(env) == get_provider("claude") + + def test_empty_env_returns_none(self) -> None: + assert detect_provider({}) is None + + def test_unknown_base_url_returns_none(self) -> None: + assert detect_provider({"ANTHROPIC_BASE_URL": "https://example.com"}) is None + + +class TestIsThirdParty: + def test_third_party_providers(self) -> None: + assert is_third_party(get_provider("deepseek")) + assert is_third_party(get_provider("kimi")) + assert is_third_party(get_provider("kimi-code")) + + def test_claude_is_not_third_party(self) -> None: + assert not is_third_party(get_provider("claude"))