core/cli: /models and remote catalog for model selection

Cache GET /models, merge with config ids for Tab and choose; free-form model ids allowed.
This commit is contained in:
2026-07-15 17:32:04 +08:00
parent e1ea43cb8f
commit 6bd80cb915
8 changed files with 348 additions and 34 deletions
+1
View File
@@ -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 |
+62
View File
@@ -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)
+19 -7
View File
@@ -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)
+94 -8
View File
@@ -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}")
+102 -15
View File
@@ -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
+52
View File
@@ -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())
+1 -1
View File
@@ -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"
+17 -3
View File
@@ -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"