diff --git a/src/plyngent/cli/app.py b/src/plyngent/cli/app.py index 911cf61..e3e318b 100644 --- a/src/plyngent/cli/app.py +++ b/src/plyngent/cli/app.py @@ -46,7 +46,7 @@ def _load_config(config_path: Path | None) -> ConfigStore: def _warn_bad_providers(bad: Mapping[str, object]) -> None: - """Surface ignored provider entries (parse errors, empty models, …).""" + """Surface ignored provider entries (parse errors, unknown fields, …).""" if not bad: return names = ", ".join(sorted(bad.keys())) @@ -68,6 +68,19 @@ def _warn_bad_providers(bad: Mapping[str, object]) -> None: click.secho(f" - {name}: {reason}", fg="yellow", err=True) +def _warn_recoverable_providers(recoverable: Mapping[str, object]) -> None: + """Surface providers with empty models (recoverable via GET /models).""" + if not recoverable: + return + names = ", ".join(sorted(recoverable.keys())) + click.secho( + f"warning: providers with empty models ({len(recoverable)}): {names} " + f"(will try GET /models or --model on use)", + fg="yellow", + err=True, + ) + + def _database_config(store: ConfigStore, *, quiet: bool = False) -> DatabaseConfig: raw = dict(store.database) # Prefer a durable file DB so sessions survive CLI restarts. @@ -165,7 +178,7 @@ async def _run_oneshot(state: ReplState, prompt_text: str) -> int: return EXIT_OK if ok else EXIT_TURN_FAILED -async def _run_chat( # noqa: C901 — chat orchestration +async def _run_chat( # noqa: C901, PLR0912 — chat orchestration *, config_path: Path | None, provider_name: str | None, @@ -187,14 +200,19 @@ async def _run_chat( # noqa: C901 — chat orchestration configure_prompting(backend=NonInteractiveBackend()) store = _load_config(config_path) - if store.bad_providers and not quiet: - _warn_bad_providers(store.bad_providers) + if not quiet: + if store.bad_providers: + _warn_bad_providers(store.bad_providers) + if store.recoverable_providers: + _warn_recoverable_providers(store.recoverable_providers) _setup_workspace_and_hooks(store, workspace, interactive=interactive) confirm_destructive: bool | None = False if yes else None memory = await MemoryStore.open(_database_config(store, quiet=quiet or oneshot)) try: + from plyngent.cli.provider_recovery import ensure_provider_ready + # Prefer session-remembered LLM when resuming (unless flags override). preferred_provider = provider_name preferred_model = model @@ -211,10 +229,17 @@ async def _run_chat( # noqa: C901 — chat orchestration try: pname, provider = select_provider( - store.providers, + store.selectable_providers(), preferred=preferred_provider, interactive=interactive, ) + provider = await ensure_provider_ready( + store, + pname, + provider, + preferred_model=preferred_model, + interactive=interactive, + ) model_id = select_model(provider, preferred=preferred_model, interactive=interactive) _ = create_client(provider) except ProviderNotSupportedError as exc: @@ -414,12 +439,15 @@ def chat_cmd( def providers_cmd(config_path: Path | None) -> None: """List configured providers.""" store = _load_config(config_path) - if not store.providers: + if not store.providers and not store.recoverable_providers: click.echo("(no providers)") for name, provider in sorted(store.providers.items()): tag = type(provider).__struct_config__.tag models = ", ".join(sorted(provider.models.keys())) or "(none listed)" click.echo(f"{name}\tpreset={tag}\tmodels={models}") + for name, provider in sorted(store.recoverable_providers.items()): + tag = type(provider).__struct_config__.tag + click.echo(f"{name}\tpreset={tag}\tmodels=(empty; recoverable)") if store.bad_providers: _warn_bad_providers(store.bad_providers) diff --git a/src/plyngent/cli/provider_recovery.py b/src/plyngent/cli/provider_recovery.py new file mode 100644 index 0000000..fac47a8 --- /dev/null +++ b/src/plyngent/cli/provider_recovery.py @@ -0,0 +1,117 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +import click + +from plyngent.cli.models_source import fetch_remote_model_ids +from plyngent.runtime import ProviderNotSupportedError, create_client + +if TYPE_CHECKING: + from collections.abc import Sequence + + from plyngent.config.models import Provider + from plyngent.config.store import ConfigStore + + +async def discover_model_ids(provider: Provider) -> list[str]: + """``GET /models`` for *provider*; empty list if unsupported or empty catalog.""" + try: + client = create_client(provider) + except ProviderNotSupportedError: + return [] + try: + return await fetch_remote_model_ids(client) + except (RuntimeError, TypeError, OSError, ValueError): + return [] + + +async def try_promote_provider( + store: ConfigStore, + name: str, + *, + seed_model_ids: Sequence[str] | None = None, +) -> Provider | None: + """Promote recoverable *name* into ready providers. + + Prefers *seed_model_ids* when non-empty; otherwise remote ``models()``. + Returns the promoted provider, or ``None`` if recovery failed. + """ + if name in store.providers: + return store.providers[name] + if name not in store.recoverable_providers: + return None + + provider = store.recoverable_providers[name] + ids: list[str] = [] + if seed_model_ids: + ids = [mid.strip() for mid in seed_model_ids if mid and str(mid).strip()] + if not ids: + ids = await discover_model_ids(provider) + if not ids: + return None + return store.promote_provider(name, ids) + + +async def ensure_provider_ready( + store: ConfigStore, + name: str, + provider: Provider, + *, + preferred_model: str | None = None, + interactive: bool = True, +) -> Provider: + """Return a ready provider; recover empty-models entries when possible. + + Raises: + click.ClickException: When recovery is required but fails. + """ + if provider.models: + return provider + if name not in store.recoverable_providers and name not in store.providers: + msg = f"provider {name!r} has no models and is not recoverable" + raise click.ClickException(msg) + + seeds: list[str] = [] + if preferred_model and preferred_model.strip(): + seeds = [preferred_model.strip()] + + promoted = await try_promote_provider(store, name, seed_model_ids=seeds or None) + if promoted is not None: + click.secho( + f"recovered provider {name!r} with {len(promoted.models)} model(s) " + f"(was empty models in config)", + fg="yellow", + err=True, + ) + return promoted + + # Remote failed and no preferred seed: interactive free-form model id. + if interactive and not seeds: + from plyngent.prompting import ask + + mid = ask(f"Model id for provider {name!r} (empty models; remote list failed)") + if mid.strip(): + promoted = store.promote_provider(name, [mid.strip()]) + click.secho( + f"recovered provider {name!r} with model {mid.strip()!r}", + fg="yellow", + err=True, + ) + return promoted + + if seeds: + # Prefer explicit model even when remote list failed. + promoted = store.promote_provider(name, seeds) + click.secho( + f"recovered provider {name!r} with model {seeds[0]!r}", + fg="yellow", + err=True, + ) + return promoted + + msg = ( + f"provider {name!r} has empty models and could not be recovered " + f"(pass --model or fix GET /models)" + ) + raise click.ClickException(msg) diff --git a/src/plyngent/config/store.py b/src/plyngent/config/store.py index d8d1564..3c071b8 100644 --- a/src/plyngent/config/store.py +++ b/src/plyngent/config/store.py @@ -6,10 +6,10 @@ from typing import TYPE_CHECKING, cast import msgspec import tomlkit -from .models import AgentConfig, DatabaseConfig, Provider +from .models import AgentConfig, DatabaseConfig, ModelConfig, Provider if TYPE_CHECKING: - from collections.abc import Mapping, MutableMapping + from collections.abc import Mapping, MutableMapping, Sequence from pathlib import Path @@ -35,19 +35,25 @@ def _parse_agent(raw: dict[str, object]) -> AgentConfig: def _parse_providers( document: tomlkit.TOMLDocument, -) -> tuple[dict[str, Provider], dict[str, object]]: +) -> tuple[dict[str, Provider], dict[str, object], dict[str, Provider]]: """Parse provider entries from the document. Returns: - (providers, bad_providers) — valid and invalid entries respectively. + (providers, bad_providers, recoverable_providers) + + *providers* — ready to use (non-empty ``models``). + *bad_providers* — unparseable / unknown fields. + *recoverable_providers* — parsed OK but ``models`` empty; may be + promoted after a successful remote ``GET /models`` (or explicit model id). """ providers: dict[str, Provider] = {} bad_providers: dict[str, object] = {} + recoverable: dict[str, Provider] = {} raw: dict[str, object] = document.unwrap() providers_raw: object = raw.get("providers", {}) if not isinstance(providers_raw, dict): - return providers, bad_providers + return providers, bad_providers, recoverable for name, raw_entry in cast("dict[str, object]", providers_raw).items(): if not isinstance(raw_entry, dict): @@ -70,15 +76,15 @@ def _parse_providers( bad_providers[name] = raw_entry continue - # Usable providers must list at least one model (DeepSeek seeds defaults - # when models is omitted; an explicit empty models={} is invalid). + # Empty models: keep as recoverable (DeepSeek seeds defaults when + # models is omitted, so only explicit models={} lands here). if not provider.models: - bad_providers[name] = {**cast("dict[str, object]", raw_entry), "_reason": "no models"} + recoverable[name] = provider continue providers[name] = provider - return providers, bad_providers + return providers, bad_providers, recoverable class ConfigStore: @@ -90,6 +96,7 @@ class ConfigStore: _agent: AgentConfig _providers: dict[str, Provider] _bad_providers: dict[str, object] + _recoverable_providers: dict[str, Provider] def __init__(self, path: Path, document: tomlkit.TOMLDocument) -> None: self._path = path @@ -97,7 +104,7 @@ class ConfigStore: raw: dict[str, object] = document.unwrap() self._database = _parse_database(cast("dict[str, object]", raw.get("database", {}))) self._agent = _parse_agent(cast("dict[str, object]", raw.get("agent", {}))) - self._providers, self._bad_providers = _parse_providers(document) + self._providers, self._bad_providers, self._recoverable_providers = _parse_providers(document) @property def path(self) -> Path: @@ -135,9 +142,10 @@ class ConfigStore: @providers.setter def providers(self, value: Mapping[str, Provider]) -> None: # pyright: ignore[reportPropertyTypeMismatch] - """Replace all providers.""" + """Replace all ready providers (clears recoverable/bad).""" self._providers = dict(value) self._bad_providers = {} + self._recoverable_providers = {} # -- bad_providers (read-only) -- @@ -146,6 +154,46 @@ class ConfigStore: """Read-only view of unrecognised / malformed provider entries.""" return MappingProxyType(self._bad_providers) + @property + def recoverable_providers(self) -> MappingProxyType[str, Provider]: + """Parsed providers with empty ``models`` (recoverable via remote list).""" + return MappingProxyType(self._recoverable_providers) + + def selectable_providers(self) -> dict[str, Provider]: + """Ready providers plus recoverable ones (ready wins on name clash).""" + return {**self._recoverable_providers, **self._providers} + + def get_provider(self, name: str) -> Provider | None: + """Look up a ready or recoverable provider by config name.""" + if name in self._providers: + return self._providers[name] + return self._recoverable_providers.get(name) + + def promote_provider(self, name: str, model_ids: Sequence[str]) -> Provider: + """Seed ``models`` from *model_ids* and move recoverable → ready. + + Also re-seeds an already-ready provider that somehow has empty models. + Does not write the TOML file (in-memory session only unless ``write()``). + """ + ids = [mid.strip() for mid in model_ids if mid and mid.strip()] + if not ids: + msg = f"cannot promote provider {name!r}: no model ids" + raise ValueError(msg) + + if name in self._providers: + provider = self._providers[name] + elif name in self._recoverable_providers: + provider = self._recoverable_providers.pop(name) + else: + msg = f"unknown provider {name!r}" + raise KeyError(msg) + + models = {mid: ModelConfig() for mid in ids} + promoted = msgspec.structs.replace(provider, models=models) + self._providers[name] = promoted + _ = self._bad_providers.pop(name, None) + return promoted + # -- persistence -- def write(self) -> None: @@ -162,7 +210,7 @@ class ConfigStore: raw: dict[str, object] = self._document.unwrap() self._database = _parse_database(cast("dict[str, object]", raw.get("database", {}))) self._agent = _parse_agent(cast("dict[str, object]", raw.get("agent", {}))) - self._providers, self._bad_providers = _parse_providers(self._document) + self._providers, self._bad_providers, self._recoverable_providers = _parse_providers(self._document) # -- internal sync helpers -- @@ -194,14 +242,15 @@ class ConfigStore: self._sync_section("agent", self._agent) def _sync_providers_section(self) -> None: - """Sync ``[providers]`` to the document.""" + """Sync ``[providers]`` to the document (ready + recoverable).""" section = self._toml_table("providers") + keep = set(self._providers) | set(self._recoverable_providers) for name in list(section.keys()): - if name not in self._providers: + if name not in keep: del section[name] - for name, provider in self._providers.items(): + for name, provider in {**self._recoverable_providers, **self._providers}.items(): raw: dict[str, object] = msgspec.to_builtins(provider) if name in section: entry = cast("MutableMapping[str, object]", section[name]) diff --git a/tests/test_cli/test_provider_recovery.py b/tests/test_cli/test_provider_recovery.py new file mode 100644 index 0000000..90049d9 --- /dev/null +++ b/tests/test_cli/test_provider_recovery.py @@ -0,0 +1,75 @@ +from __future__ import annotations + +from pathlib import Path + +import pytest + +import plyngent.config +from plyngent.cli.provider_recovery import ensure_provider_ready, try_promote_provider +from plyngent.config.models import ModelConfig, OpenAICompatibleProvider + + +def _hollow_store(tmp_path: Path): + path = tmp_path / "cfg.toml" + _ = path.write_text( + """ +[providers.hollow] +preset = "openai-compatible" +url = "https://example.com/v1" +access_key_or_token = "sk-test" +models = {} +""", + encoding="utf-8", + ) + return plyngent.config.load(path) + + +@pytest.mark.asyncio +async def test_try_promote_with_seed(tmp_path: Path) -> None: + store = _hollow_store(tmp_path) + promoted = await try_promote_provider(store, "hollow", seed_model_ids=["gpt-x"]) + assert promoted is not None + assert "gpt-x" in promoted.models + assert "hollow" in store.providers + assert "hollow" not in store.recoverable_providers + + +@pytest.mark.asyncio +async def test_try_promote_via_remote(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + store = _hollow_store(tmp_path) + + async def fake_discover(provider: object) -> list[str]: + del provider + return ["remote-a", "remote-b"] + + monkeypatch.setattr("plyngent.cli.provider_recovery.discover_model_ids", fake_discover) + promoted = await try_promote_provider(store, "hollow") + assert promoted is not None + assert set(promoted.models) == {"remote-a", "remote-b"} + + +@pytest.mark.asyncio +async def test_ensure_provider_ready_already_ready(tmp_path: Path) -> None: + store = _hollow_store(tmp_path) + ready = OpenAICompatibleProvider( + access_key_or_token="sk", + url="https://x/v1", + models={"m": ModelConfig()}, + ) + store.providers = {"ready": ready} + out = await ensure_provider_ready(store, "ready", ready, interactive=False) + assert out is ready + + +@pytest.mark.asyncio +async def test_ensure_provider_ready_seed_model(tmp_path: Path) -> None: + store = _hollow_store(tmp_path) + provider = store.recoverable_providers["hollow"] + out = await ensure_provider_ready( + store, + "hollow", + provider, + preferred_model="explicit", + interactive=False, + ) + assert "explicit" in out.models diff --git a/tests/test_config/test_config.py b/tests/test_config/test_config.py index 15c8b1f..9944e69 100644 --- a/tests/test_config/test_config.py +++ b/tests/test_config/test_config.py @@ -81,7 +81,7 @@ def test_read_bad_config() -> None: assert isinstance(config.bad_providers, Mapping) -def test_provider_with_empty_models_is_bad(tmp_path: Path) -> None: +def test_provider_with_empty_models_is_recoverable(tmp_path: Path) -> None: path = tmp_path / "empty-models.toml" _ = path.write_text( """ @@ -95,7 +95,29 @@ models = {} ) config = plyngent.config.load(path) assert "hollow" not in config.providers - assert "hollow" in config.bad_providers + assert "hollow" not in config.bad_providers + assert "hollow" in config.recoverable_providers + promoted = config.promote_provider("hollow", ["m1", "m2"]) + assert "hollow" in config.providers + assert "hollow" not in config.recoverable_providers + assert set(promoted.models) == {"m1", "m2"} + + +def test_promote_provider_requires_ids(tmp_path: Path) -> None: + path = tmp_path / "empty-models.toml" + _ = path.write_text( + """ +[providers.hollow] +preset = "openai-compatible" +url = "https://example.com/v1" +access_key_or_token = "sk-test" +models = {} +""", + encoding="utf-8", + ) + config = plyngent.config.load(path) + with pytest.raises(ValueError, match="no model ids"): + _ = config.promote_provider("hollow", []) def test_read_invalid_config() -> None: