diff --git a/litellm/proxy/client/cli/commands/agents.py b/litellm/proxy/client/cli/commands/agents.py index 67d2d96e9d6..8e3698985af 100644 --- a/litellm/proxy/client/cli/commands/agents.py +++ b/litellm/proxy/client/cli/commands/agents.py @@ -4,6 +4,7 @@ import re import shutil import subprocess import sys +import tempfile from collections.abc import Callable, Mapping, Sequence from dataclasses import dataclass from pathlib import Path @@ -416,7 +417,7 @@ class _CodexCatalog(BaseModel): models: tuple[_CodexModel, ...] -def codex_model_catalog(models: Sequence[ListedModel]) -> str | None: +def codex_model_catalog(models: Sequence[ListedModel], instructions: str) -> str | None: """The `model_catalog_json` body listing the proxy's chat models, or None if there are none. Codex refuses an empty catalog, hence None instead of `{"models": []}`. @@ -427,7 +428,6 @@ def codex_model_catalog(models: Sequence[ListedModel]) -> str | None: chat_models: Final = _chat_models(models) if not chat_models: return None - instructions: Final = _CODEX_BASE_INSTRUCTIONS_PATH.read_text(encoding="utf-8") catalog: Final = _CodexCatalog( models=tuple( _CodexModel( @@ -443,37 +443,49 @@ def codex_model_catalog(models: Sequence[ListedModel]) -> str | None: return catalog.model_dump_json() -def codex_model_catalog_path(env: Mapping[str, str]) -> Path: +def codex_model_catalog_path(env: Mapping[str, str], *, home: Callable[[], Path] = Path.home) -> Path: override: Final = env.get(CODEX_HOME_ENV) - root: Final = Path(override) if override else Path.home() / ".codex" + root: Final = Path(override) if override else home() / ".codex" return root / CODEX_MODEL_CATALOG_FILENAME +def _replace_file(path: Path, text: str) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + with tempfile.NamedTemporaryFile("w", encoding="utf-8", dir=path.parent, delete=False) as tmp: + _ = tmp.write(text) + os.replace(tmp.name, path) + + def codex_model_sync_args( base_env: Mapping[str, str], base_url: str, api_key: str, *, get: Callable[..., requests.Response] = requests.get, + home: Callable[[], Path] = Path.home, + instructions_path: Path = _CODEX_BASE_INSTRUCTIONS_PATH, ) -> ModelSyncArgs | ModelSyncSkipped: """`-c model_catalog_json=...` pointing Codex at the proxy's model list, or why it was skipped. Codex has no env or inline equivalent of OPENCODE_CONFIG_CONTENT: the catalog must be a file, so it is written under $CODEX_HOME (default ~/.codex) and - rewritten on every launch. The key never lands in the file. A failed fetch - or write is reported rather than raised: Codex still launches with its - built-in catalog and takes a proxy model by name via -m. + atomically replaced on every launch. The key never lands in the file. A + failed fetch, read or write is reported rather than raised: Codex still + launches with its built-in catalog and takes a proxy model by name via -m. """ listing: Final = _fetch_model_listing(base_url, api_key, get=get) if isinstance(listing, ModelSyncSkipped): return listing - catalog: Final = codex_model_catalog(listing) + try: + instructions: Final = instructions_path.read_text(encoding="utf-8") + except OSError as e: + return ModelSyncSkipped(f"could not read {instructions_path}: {e}") + catalog: Final = codex_model_catalog(listing, instructions) if catalog is None: return ModelSyncSkipped(f"{base_url.rstrip('/')}/v1/models lists no chat models") - path: Final = codex_model_catalog_path(base_env) + path: Final = codex_model_catalog_path(base_env, home=home) try: - path.parent.mkdir(parents=True, exist_ok=True) - path.write_text(catalog, encoding="utf-8") + _replace_file(path, catalog) except OSError as e: return ModelSyncSkipped(f"could not write {path}: {e}") return ModelSyncArgs(("-c", f"model_catalog_json={json.dumps(str(path))}")) diff --git a/tests/test_litellm/proxy/client/cli/test_agents.py b/tests/test_litellm/proxy/client/cli/test_agents.py index bb4b99506a2..75e42c3eaa0 100644 --- a/tests/test_litellm/proxy/client/cli/test_agents.py +++ b/tests/test_litellm/proxy/client/cli/test_agents.py @@ -2,6 +2,7 @@ import inspect import json import os import sys +from pathlib import Path from unittest.mock import patch import click @@ -90,9 +91,7 @@ class TestAgentProfile: class TestBuildAgentEnv: def test_anthropic_profile_uses_bare_root_and_bearer(self): - env = build_agent_env( - {}, "http://localhost:4000/", "sk-key", frozenset({"anthropic"}) - ) + env = build_agent_env({}, "http://localhost:4000/", "sk-key", frozenset({"anthropic"})) assert env["ANTHROPIC_BASE_URL"] == "http://localhost:4000" assert env["ANTHROPIC_AUTH_TOKEN"] == "sk-key" assert env["ENABLE_TOOL_SEARCH"] == "true" @@ -128,9 +127,7 @@ class TestBuildAgentEnv: assert "ANTHROPIC_API_KEY" not in env def test_openai_profile_appends_v1(self): - env = build_agent_env( - {}, "http://localhost:4000/", "sk-key", frozenset({"openai"}) - ) + env = build_agent_env({}, "http://localhost:4000/", "sk-key", frozenset({"openai"})) assert env["OPENAI_BASE_URL"] == "http://localhost:4000/v1" assert env["OPENAI_API_KEY"] == "sk-key" assert "ANTHROPIC_BASE_URL" not in env @@ -138,9 +135,7 @@ class TestBuildAgentEnv: assert "CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY" not in env def test_both_profiles_set_everything(self): - env = build_agent_env( - {}, "http://localhost:4000", "sk-key", frozenset({"anthropic", "openai"}) - ) + env = build_agent_env({}, "http://localhost:4000", "sk-key", frozenset({"anthropic", "openai"})) assert env["ANTHROPIC_BASE_URL"] == "http://localhost:4000" assert env["OPENAI_BASE_URL"] == "http://localhost:4000/v1" assert env["ANTHROPIC_AUTH_TOKEN"] == "sk-key" @@ -148,9 +143,7 @@ class TestBuildAgentEnv: assert env["ENABLE_TOOL_SEARCH"] == "true" def test_litellm_profile_exports_only_the_proxy_key(self): - env = build_agent_env( - {}, "http://localhost:4000/", "sk-key", frozenset({"litellm"}) - ) + env = build_agent_env({}, "http://localhost:4000/", "sk-key", frozenset({"litellm"})) assert env["LITELLM_PROXY_API_KEY"] == "sk-key" assert "ANTHROPIC_BASE_URL" not in env assert "OPENAI_BASE_URL" not in env @@ -158,9 +151,7 @@ class TestBuildAgentEnv: def test_preserves_unrelated_env_and_does_not_mutate_input(self): base = {"PATH": "/usr/bin", "ANTHROPIC_API_KEY": "real-key"} - env = build_agent_env( - base, "http://localhost:4000", "sk-key", frozenset({"anthropic"}) - ) + env = build_agent_env(base, "http://localhost:4000", "sk-key", frozenset({"anthropic"})) assert env["PATH"] == "/usr/bin" assert base == {"PATH": "/usr/bin", "ANTHROPIC_API_KEY": "real-key"} @@ -320,9 +311,7 @@ class TestOpencodeModelSync: assert "refused" in result.reason def test_non_200_is_reported(self): - result = opencode_model_sync_env( - {}, "http://localhost:4000", "sk-key", get=lambda *a, **k: _FakeResponse(500) - ) + result = opencode_model_sync_env({}, "http://localhost:4000", "sk-key", get=lambda *a, **k: _FakeResponse(500)) assert isinstance(result, ModelSyncSkipped) assert "HTTP 500" in result.reason @@ -454,13 +443,37 @@ class TestCodexModelSync: slugs = [m["slug"] for m in json.loads((tmp_path / "litellm-models.json").read_text())["models"]] assert slugs == ["new"] - def test_defaults_to_dot_codex_in_home(self, tmp_path, monkeypatch): - monkeypatch.setattr("pathlib.Path.home", classmethod(lambda cls: tmp_path)) + def test_catalog_is_replaced_whole_and_leaves_no_temp_files(self, tmp_path): + self._sync(self._listing(*(self._row(f"m{i}") for i in range(50))), tmp_path) + self._sync(self._listing(self._row("new")), tmp_path) + assert [p.name for p in tmp_path.iterdir()] == ["litellm-models.json"] + assert json.loads((tmp_path / "litellm-models.json").read_text())["models"][0]["slug"] == "new" + + def test_defaults_to_dot_codex_in_home(self, tmp_path): result = codex_model_sync_args( - {}, "http://localhost:4000", "sk-key", get=lambda *a, **k: _FakeResponse(200, self._listing(self._row("m"))) + {}, + "http://localhost:4000", + "sk-key", + get=lambda *a, **k: _FakeResponse(200, self._listing(self._row("m"))), + home=lambda: tmp_path, ) assert self._catalog_path(result) == str(tmp_path / ".codex" / "litellm-models.json") + def test_default_home_is_the_users(self): + assert _default_of(codex_model_sync_args, "home") == Path.home + + def test_missing_base_instructions_is_reported_not_raised(self, tmp_path): + result = codex_model_sync_args( + {"CODEX_HOME": str(tmp_path)}, + "http://localhost:4000", + "sk-key", + get=lambda *a, **k: _FakeResponse(200, self._listing(self._row("m"))), + instructions_path=tmp_path / "missing.md", + ) + assert isinstance(result, ModelSyncSkipped) + assert "could not read" in result.reason + assert not (tmp_path / "litellm-models.json").exists() + def test_unwritable_catalog_path_is_reported_not_raised(self, tmp_path): blocker = tmp_path / "file" blocker.write_text("") @@ -528,7 +541,9 @@ class TestRunAgent: args = calls["args"] assert args[-2:] == ("exec", "hi") assert args[args.index('model_catalog_json="/tmp/c.json"') - 1] == "-c" - assert args.index('model_provider="litellm"') < args.index('model_catalog_json="/tmp/c.json"') < args.index("exec") + assert ( + args.index('model_provider="litellm"') < args.index('model_catalog_json="/tmp/c.json"') < args.index("exec") + ) assert calls["env"]["OPENAI_API_KEY"] == "sk-key" assert "model_catalog_json" not in json.dumps(calls["env"]) @@ -1214,10 +1229,7 @@ class TestAgentCommands: assert captured["api_key"] == "sk-key" assert captured["command"] == ["claude", "--resume", "-p", "hi"] assert captured["skip_verify"] is False - assert ( - "routing Claude Code through proxy at http://localhost:4000" - in result.output - ) + assert "routing Claude Code through proxy at http://localhost:4000" in result.output def test_codex_shows_friendly_name(self): captured = {} @@ -1290,14 +1302,10 @@ class TestAgentCommands: with ( patch(f"{AGENTS_MODULE}._is_interactive", return_value=True), patch(f"{AGENTS_MODULE}.login", fake_login), - patch( - f"{AGENTS_MODULE}.get_stored_api_key", return_value="sk-after-login" - ) as mock_get, + patch(f"{AGENTS_MODULE}.get_stored_api_key", return_value="sk-after-login") as mock_get, patch( f"{AGENTS_MODULE}.run_agent", - side_effect=lambda base_url, api_key, command, **k: captured.update( - api_key=api_key - ), + side_effect=lambda base_url, api_key, command, **k: captured.update(api_key=api_key), ), ): result = self.runner.invoke(