diff --git a/README.md b/README.md index 63a9cdb..668d648 100644 --- a/README.md +++ b/README.md @@ -124,6 +124,7 @@ Type `/help` in the REPL for the live list. Common ones: | `/stream` `/verbose` `/tools` `/rounds` | Toggles and limits | | `/retry` | Re-run incomplete last user turn (after error/cancel) | | `/provider` `/model` | Switch without restarting | +| `/models` | List config + remote `GET /models` (`--refresh` bypasses cache) | | `/config` | Edit `plyngent.toml` in `$EDITOR` and reload | | `/quit` | Leave the REPL | diff --git a/src/plyngent/cli/models_source.py b/src/plyngent/cli/models_source.py new file mode 100644 index 0000000..25249b4 --- /dev/null +++ b/src/plyngent/cli/models_source.py @@ -0,0 +1,62 @@ +from __future__ import annotations + +import inspect +from typing import TYPE_CHECKING, Protocol, cast, runtime_checkable + +if TYPE_CHECKING: + from collections.abc import Iterable, Sequence + + from plyngent.config.models import Provider + +# Cache remote catalog this long (seconds) unless /models --refresh. +DEFAULT_MODELS_CACHE_TTL = 300.0 + + +@runtime_checkable +class SupportsModels(Protocol): + async def models(self) -> list[str]: ... + + +def config_model_ids(provider: Provider) -> list[str]: + """Sorted model ids declared in provider config.""" + return sorted(provider.models.keys()) + + +def merge_model_choices( + config_ids: Iterable[str], + remote_ids: Iterable[str] | None = None, +) -> list[str]: + """Union config and remote ids (sorted, unique).""" + merged: set[str] = {i for i in config_ids if i} + if remote_ids is not None: + merged.update(i for i in remote_ids if i) + return sorted(merged) + + +def client_supports_models(client: object) -> bool: + """True when *client* exposes OpenAI-compatible ``models()``.""" + return isinstance(client, SupportsModels) or callable(getattr(client, "models", None)) + + +async def fetch_remote_model_ids(client: object) -> list[str]: + """Call ``client.models()``; raise if missing or the call fails.""" + method = getattr(client, "models", None) + if not callable(method): + msg = "client does not support listing models" + raise TypeError(msg) + result = method() + if inspect.isawaitable(result): + result = await result + if not isinstance(result, list): + msg = f"models() returned unexpected type {type(result)!r}" + raise TypeError(msg) + return [str(item) for item in cast("list[object]", result) if item] + + +def model_choices_for_provider( + provider: Provider, + *, + remote_ids: Sequence[str] | None = None, +) -> list[str]: + """Config plus remote catalog for selection / Tab complete.""" + return merge_model_choices(config_model_ids(provider), remote_ids) diff --git a/src/plyngent/cli/selection.py b/src/plyngent/cli/selection.py index 4d8ac2b..6af8cc8 100644 --- a/src/plyngent/cli/selection.py +++ b/src/plyngent/cli/selection.py @@ -4,10 +4,11 @@ from typing import TYPE_CHECKING import click +from plyngent.cli.models_source import config_model_ids, model_choices_for_provider from plyngent.prompting import ChoiceOption, choose if TYPE_CHECKING: - from collections.abc import Mapping + from collections.abc import Mapping, Sequence from plyngent.config.models import Provider @@ -56,14 +57,20 @@ def select_model( *, preferred: str | None = None, interactive: bool = True, + choices: Sequence[str] | None = None, ) -> str: - """Pick a model id from provider.models or free-form prompt (readline + Tab).""" - model_names = sorted(provider.models.keys()) + """Pick a model id from config/remote choices or free-form prompt. + + *choices* overrides the default list (config keys). Explicit *preferred* + is accepted even when not in the list (API validates at chat time). + """ + model_names = list(choices) if choices is not None else config_model_ids(provider) if preferred is not None: - if model_names and preferred not in provider.models: - msg = f"unknown model {preferred!r}; available: {', '.join(model_names)}" + token = preferred.strip() + if not token: + msg = "model id must not be empty" raise click.ClickException(msg) - return preferred + return token if len(model_names) == 1: model = model_names[0] @@ -81,9 +88,14 @@ def select_model( return choose( "Select model", model_names, - allow_custom=False, + allow_custom=True, ) from plyngent.prompting import ask return ask("Model id (not listed in config)") + + +def default_model_choices(provider: Provider, remote_ids: Sequence[str] | None = None) -> list[str]: + """Config plus optional remote ids (for Tab / interactive pick).""" + return model_choices_for_provider(provider, remote_ids=remote_ids) diff --git a/src/plyngent/cli/slash.py b/src/plyngent/cli/slash.py index 5b6474d..a3879a5 100644 --- a/src/plyngent/cli/slash.py +++ b/src/plyngent/cli/slash.py @@ -8,6 +8,7 @@ import click from click.shell_completion import CompletionItem from msgspec import UNSET +from plyngent.cli.models_source import DEFAULT_MODELS_CACHE_TTL, model_choices_for_provider from plyngent.cli.retry import retry_pending_with_retries from plyngent.cli.selection import select_model, select_provider from plyngent.lmproto.openai_compatible.model import ( @@ -117,7 +118,7 @@ class ProviderNameParam(click.ParamType[str]): state = _repl_state(ctx) if state is None: return [] - return _filter_choices(incomplete, sorted(state.config.providers.keys())) + return _filter_choices(incomplete, sorted(state.config.selectable_providers().keys())) PROVIDER_NAME = ProviderNameParam() @@ -137,7 +138,7 @@ class ModelIdParam(click.ParamType[str]): state = _repl_state(ctx) if state is None: return [] - return _filter_choices(incomplete, sorted(state.provider.models.keys())) + return _filter_choices(incomplete, state.model_choice_ids()) MODEL_ID = ModelIdParam() @@ -316,6 +317,12 @@ def config_cmd(state: ReplState) -> None: click.secho(f"error: config reload failed: {exc}", fg="red") click.echo(f"config file: {path}") return + if state.config.recoverable_providers: + names = ", ".join(sorted(state.config.recoverable_providers.keys())) + click.secho( + f"recoverable providers (empty models): {names}", + fg="yellow", + ) if state.config.bad_providers: names = ", ".join(sorted(state.config.bad_providers.keys())) click.secho(f"warning: ignored bad providers: {names}", fg="yellow") @@ -554,16 +561,37 @@ def provider_cmd(state: ReplState, name: str | None) -> None: click.echo(f"provider={state.provider_name}") return try: - pname, provider = select_provider(state.config.providers, preferred=name.strip()) + from plyngent.cli.provider_recovery import ensure_provider_ready + + pname, provider = select_provider( + state.config.selectable_providers(), + preferred=name.strip(), + ) + provider = _await( + ensure_provider_ready( + state.config, + pname, + provider, + preferred_model=state.model, + interactive=True, + ) + ) prev_model = state.model state.provider_name = pname state.provider = provider - if prev_model in provider.models: + state.rebuild_client() + choices = _await(state.merged_model_choices(refresh=False)) + if prev_model and (prev_model in choices or prev_model in provider.models): state.model = prev_model else: # Current model not on the new provider — pick one (prompt when interactive). try: - state.model = select_model(provider, preferred=None, interactive=True) + state.model = select_model( + provider, + preferred=None, + interactive=True, + choices=choices, + ) except click.ClickException as exc: click.echo(f"error: switched provider but model selection failed: {exc}") return @@ -572,23 +600,81 @@ def provider_cmd(state: ReplState, name: str | None) -> None: f"model {prev_model!r} is not available on {pname}; using {state.model!r}", fg="yellow", ) - state.rebuild_client() + state.rebuild_client() _await(state.persist_llm_selection()) click.echo(f"switched provider to {pname} model={state.model}") except (click.ClickException, ProviderNotSupportedError) as exc: click.echo(f"error: {exc}") +@slash.command("models") +@click.option("--refresh", is_flag=True, help="Bypass cache and re-fetch GET /models.") +@click.pass_obj +def models_cmd(state: ReplState, *, refresh: bool) -> None: + """List models (config plus remote GET /models).""" + remote: list[str] | None = None + remote_err: str | None = None + try: + remote = _await(state.ensure_remote_models(refresh=refresh)) + except (RuntimeError, TypeError, OSError, ValueError) as exc: + remote_err = str(exc) + remote = state.cached_remote_models() + + # Promote empty-models recoverable provider after a successful remote list. + if remote and state.provider_name in state.config.recoverable_providers: + try: + state.provider = state.config.promote_provider(state.provider_name, remote) + state.rebuild_client() + click.secho( + f"recovered provider {state.provider_name!r} from remote catalog", + fg="yellow", + err=True, + ) + except (KeyError, ValueError) as exc: + click.secho(f"could not recover provider: {exc}", fg="yellow", err=True) + + config_ids = set(state.config_model_ids()) + choices = model_choices_for_provider(state.provider, remote_ids=remote) + + if not choices: + click.echo("(no models in config or remote catalog)") + else: + remote_set = set(remote or ()) + for mid in choices: + tags: list[str] = [] + if mid in config_ids: + tags.append("config") + if mid in remote_set: + tags.append("remote") + suffix = f" ({', '.join(tags)})" if tags else "" + mark = " *" if mid == state.model else "" + click.echo(f"{mid}{mark}{suffix}") + + if remote_err is not None: + click.secho(f"remote list unavailable: {remote_err}", fg="yellow", err=True) + elif remote is not None: + click.echo( + f"({len(remote)} remote, {len(config_ids)} config; " + f"cache TTL {int(DEFAULT_MODELS_CACHE_TTL)}s)", + err=True, + ) + + @slash.command("model") @click.argument("model_id", required=False, type=MODEL_ID) @click.pass_obj def model_cmd(state: ReplState, model_id: str | None) -> None: - """Show or switch model.""" + """Show or switch model (Tab: config plus cached remote).""" if not model_id: click.echo(f"model={state.model}") return try: - state.model = select_model(state.provider, preferred=model_id.strip()) + choices = _await(state.merged_model_choices(refresh=False)) + state.model = select_model( + state.provider, + preferred=model_id.strip(), + choices=choices, + ) state.rebuild_client() _await(state.persist_llm_selection()) click.echo(f"switched model to {state.model}") diff --git a/src/plyngent/cli/state.py b/src/plyngent/cli/state.py index 2413dc6..9d7dd38 100644 --- a/src/plyngent/cli/state.py +++ b/src/plyngent/cli/state.py @@ -1,11 +1,20 @@ from __future__ import annotations +import contextlib +import time from dataclasses import dataclass, field from pathlib import Path from typing import TYPE_CHECKING, cast from plyngent.agent import ChatAgent, ChatClient, ToolRegistry from plyngent.agent.loop import DEFAULT_MAX_ROUNDS +from plyngent.cli.models_source import ( + DEFAULT_MODELS_CACHE_TTL, + client_supports_models, + config_model_ids, + fetch_remote_model_ids, + model_choices_for_provider, +) from plyngent.memory.database.store import normalize_workspace from plyngent.runtime import create_client from plyngent.tools import DEFAULT_TOOLS, set_workspace_root @@ -40,6 +49,11 @@ class ReplState: client: ChatClient = field(init=False) agent: ChatAgent = field(init=False) session_id: int | None = None + # Remote GET /models cache (per provider base). + _remote_models: list[str] | None = field(default=None, init=False, repr=False) + _remote_models_fetched_at: float | None = field(default=None, init=False, repr=False) + _remote_models_key: tuple[str, str] | None = field(default=None, init=False, repr=False) + _remote_models_error: str | None = field(default=None, init=False, repr=False) def __post_init__(self) -> None: # DeepSeek client uses a compatible but distinct param type; treat as ChatClient. @@ -110,6 +124,73 @@ class ReplState: self.agent = self._make_agent() self.agent.messages = messages self.sync_display_flags() + # Drop remote catalog when provider identity/url changed (not on model-only switch). + if self._remote_models_key is not None and self._remote_models_key != self._models_cache_key(): + self.invalidate_remote_models() + + def _models_cache_key(self) -> tuple[str, str]: + url = getattr(self.provider, "url", "") or "" + return (self.provider_name, str(url)) + + def invalidate_remote_models(self) -> None: + """Drop cached remote model catalog (provider/client change).""" + self._remote_models = None + self._remote_models_fetched_at = None + self._remote_models_key = None + self._remote_models_error = None + + def cached_remote_models(self) -> list[str] | None: + """Return cached remote ids if still valid for the current provider.""" + if self._remote_models is None or self._remote_models_fetched_at is None: + return None + if self._remote_models_key != self._models_cache_key(): + return None + age = time.monotonic() - self._remote_models_fetched_at + if age > DEFAULT_MODELS_CACHE_TTL: + return None + return list(self._remote_models) + + def model_choice_ids(self, *, include_remote_cache: bool = True) -> list[str]: + """Config plus optional cached remote ids (no network).""" + remote = self.cached_remote_models() if include_remote_cache else None + return model_choices_for_provider(self.provider, remote_ids=remote) + + async def ensure_remote_models(self, *, refresh: bool = False) -> list[str]: + """Fetch ``GET /models`` (cached) and return remote ids. + + Raises RuntimeError/TypeError when the client cannot list models or + the request fails. On failure the previous cache is left unchanged. + """ + if not refresh: + cached = self.cached_remote_models() + if cached is not None: + return cached + if not client_supports_models(self.client): + msg = "client does not support listing models" + self._remote_models_error = msg + raise TypeError(msg) + try: + ids = await fetch_remote_model_ids(self.client) + except (RuntimeError, TypeError, OSError, ValueError) as exc: + self._remote_models_error = str(exc) + raise + self._remote_models = list(ids) + self._remote_models_fetched_at = time.monotonic() + self._remote_models_key = self._models_cache_key() + self._remote_models_error = None + return list(ids) + + async def merged_model_choices(self, *, refresh: bool = False) -> list[str]: + """Config plus remote catalog; remote fetch best-effort when refresh/missing.""" + remote: list[str] | None + try: + remote = await self.ensure_remote_models(refresh=refresh) + except (RuntimeError, TypeError, OSError, ValueError): + remote = self.cached_remote_models() + return model_choices_for_provider(self.provider, remote_ids=remote) + + def config_model_ids(self) -> list[str]: + return config_model_ids(self.provider) def reload_config_from_disk(self) -> None: """Re-read TOML config and re-bind provider/model when still valid.""" @@ -122,27 +203,32 @@ class ReplState: self.config.reload() set_path_denylist(self.config.agent_config.path_denylist or None) - preferred_provider = self.provider_name if self.provider_name in self.config.providers else None + selectable = self.config.selectable_providers() + preferred_provider = self.provider_name if self.provider_name in selectable else None preferred_model = self.model try: pname, provider = select_provider( - self.config.providers, + selectable, preferred=preferred_provider, interactive=False, ) except (click.ClickException, ProviderNotSupportedError) as exc: msg = f"config reloaded but provider selection failed: {exc}" raise ValueError(msg) from exc + # Empty-models providers stay recoverable until next use /models promote. + if not provider.models and preferred_model: + with contextlib.suppress(KeyError, ValueError): + provider = self.config.promote_provider(pname, [preferred_model]) try: model_id = select_model(provider, preferred=preferred_model, interactive=False) except click.ClickException: - # Previous model not on this provider; pick first listed or keep free-form. model_id = next(iter(sorted(provider.models.keys()))) if provider.models else preferred_model self.provider_name = pname self.provider = provider self.model = model_id self.rebuild_client() + self.invalidate_remote_models() def _set_workspace(self, path: Path) -> None: """Update REPL + tool workspace root.""" @@ -175,16 +261,22 @@ class ReplState: from plyngent.cli.selection import select_provider from plyngent.runtime import ProviderNotSupportedError - if pname not in self.config.providers: + if pname not in self.config.selectable_providers(): return False try: name, provider = select_provider( - self.config.providers, + self.config.selectable_providers(), preferred=pname, interactive=False, ) - except click.ClickException, ProviderNotSupportedError: + except (click.ClickException, ProviderNotSupportedError): return False + # Session resume: seed empty recoverable with remembered model if any. + if not provider.models and self.model: + try: + provider = self.config.promote_provider(name, [self.model]) + except (KeyError, ValueError): + return False if name != self.provider_name or provider is not self.provider: self.provider_name = name self.provider = provider @@ -192,16 +284,11 @@ class ReplState: return False def _try_set_model(self, model_id: str) -> bool: - import click - - from plyngent.cli.selection import select_model - - try: - resolved = select_model(self.provider, preferred=model_id, interactive=False) - except click.ClickException: + token = model_id.strip() + if not token: return False - if resolved != self.model: - self.model = resolved + if token != self.model: + self.model = token return True return False diff --git a/tests/test_cli/test_models_source.py b/tests/test_cli/test_models_source.py new file mode 100644 index 0000000..dd92fd4 --- /dev/null +++ b/tests/test_cli/test_models_source.py @@ -0,0 +1,52 @@ +from __future__ import annotations + +import pytest + +from plyngent.cli.models_source import ( + client_supports_models, + config_model_ids, + fetch_remote_model_ids, + merge_model_choices, + model_choices_for_provider, +) +from plyngent.config.models import ModelConfig, OpenAICompatibleProvider + + +def test_merge_model_choices_union() -> None: + assert merge_model_choices(["b", "a"], ["a", "c"]) == ["a", "b", "c"] + assert merge_model_choices(["a"], None) == ["a"] + assert merge_model_choices([], ["z"]) == ["z"] + + +def test_model_choices_for_provider() -> None: + provider = OpenAICompatibleProvider( + access_key_or_token="sk", + url="https://x/v1", + models={"cfg": ModelConfig()}, + ) + assert config_model_ids(provider) == ["cfg"] + assert model_choices_for_provider(provider, remote_ids=["remote", "cfg"]) == ["cfg", "remote"] + + +def test_client_supports_models() -> None: + class Ok: + async def models(self) -> list[str]: + return ["m"] + + class No: + pass + + assert client_supports_models(Ok()) + assert not client_supports_models(No()) + + +@pytest.mark.asyncio +async def test_fetch_remote_model_ids() -> None: + class Ok: + async def models(self) -> list[str]: + return ["z", "a"] + + assert await fetch_remote_model_ids(Ok()) == ["z", "a"] + + with pytest.raises(TypeError, match="does not support"): + _ = await fetch_remote_model_ids(object()) diff --git a/tests/test_cli/test_repl_commands.py b/tests/test_cli/test_repl_commands.py index ccbc108..cef9f2f 100644 --- a/tests/test_cli/test_repl_commands.py +++ b/tests/test_cli/test_repl_commands.py @@ -235,7 +235,7 @@ async def test_provider_switch_prompts_when_model_missing( # When switching to b, only-a is missing → select_model is invoked interactively. monkeypatch.setattr( "plyngent.cli.slash.select_model", - lambda provider, preferred=None, interactive=True: "only-b", + lambda provider, preferred=None, interactive=True, choices=None: "only-b", ) assert await handle_slash(state, "/provider b") is True assert state.provider_name == "b" diff --git a/tests/test_cli/test_selection.py b/tests/test_cli/test_selection.py index 7c7a181..66c38ea 100644 --- a/tests/test_cli/test_selection.py +++ b/tests/test_cli/test_selection.py @@ -64,11 +64,25 @@ def test_select_provider_interactive_choose() -> None: assert name == "b" -def test_select_model_when_preferred_missing_raises() -> None: +def test_select_model_preferred_not_in_config_allowed() -> None: + """Explicit model ids are accepted; the API validates at chat time.""" provider = OpenAICompatibleProvider( access_key_or_token="sk", url="https://x/v1", models={"m1": ModelConfig()}, ) - with pytest.raises(Exception, match="unknown model"): - _ = select_model(provider, preferred="nope") + assert select_model(provider, preferred="nope") == "nope" + assert select_model(provider, preferred=" custom ") == "custom" + + +def test_select_model_choices_override() -> None: + from plyngent.prompting import temporary_backend + from tests.test_prompting import ScriptedBackend + + provider = OpenAICompatibleProvider( + access_key_or_token="sk", + url="https://x/v1", + models={"m1": ModelConfig()}, + ) + with temporary_backend(ScriptedBackend(["remote-x"])): + assert select_model(provider, choices=["remote-x", "m1"]) == "remote-x"