mirror of
https://github.com/NCBM/plyngent.git
synced 2026-07-23 05:55:16 +08:00
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:
@@ -124,6 +124,7 @@ Type `/help` in the REPL for the live list. Common ones:
|
|||||||
| `/stream` `/verbose` `/tools` `/rounds` | Toggles and limits |
|
| `/stream` `/verbose` `/tools` `/rounds` | Toggles and limits |
|
||||||
| `/retry` | Re-run incomplete last user turn (after error/cancel) |
|
| `/retry` | Re-run incomplete last user turn (after error/cancel) |
|
||||||
| `/provider` `/model` | Switch without restarting |
|
| `/provider` `/model` | Switch without restarting |
|
||||||
|
| `/models` | List config + remote `GET /models` (`--refresh` bypasses cache) |
|
||||||
| `/config` | Edit `plyngent.toml` in `$EDITOR` and reload |
|
| `/config` | Edit `plyngent.toml` in `$EDITOR` and reload |
|
||||||
| `/quit` | Leave the REPL |
|
| `/quit` | Leave the REPL |
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
@@ -4,10 +4,11 @@ from typing import TYPE_CHECKING
|
|||||||
|
|
||||||
import click
|
import click
|
||||||
|
|
||||||
|
from plyngent.cli.models_source import config_model_ids, model_choices_for_provider
|
||||||
from plyngent.prompting import ChoiceOption, choose
|
from plyngent.prompting import ChoiceOption, choose
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from collections.abc import Mapping
|
from collections.abc import Mapping, Sequence
|
||||||
|
|
||||||
from plyngent.config.models import Provider
|
from plyngent.config.models import Provider
|
||||||
|
|
||||||
@@ -56,14 +57,20 @@ def select_model(
|
|||||||
*,
|
*,
|
||||||
preferred: str | None = None,
|
preferred: str | None = None,
|
||||||
interactive: bool = True,
|
interactive: bool = True,
|
||||||
|
choices: Sequence[str] | None = None,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Pick a model id from provider.models or free-form prompt (readline + Tab)."""
|
"""Pick a model id from config/remote choices or free-form prompt.
|
||||||
model_names = sorted(provider.models.keys())
|
|
||||||
|
*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 preferred is not None:
|
||||||
if model_names and preferred not in provider.models:
|
token = preferred.strip()
|
||||||
msg = f"unknown model {preferred!r}; available: {', '.join(model_names)}"
|
if not token:
|
||||||
|
msg = "model id must not be empty"
|
||||||
raise click.ClickException(msg)
|
raise click.ClickException(msg)
|
||||||
return preferred
|
return token
|
||||||
|
|
||||||
if len(model_names) == 1:
|
if len(model_names) == 1:
|
||||||
model = model_names[0]
|
model = model_names[0]
|
||||||
@@ -81,9 +88,14 @@ def select_model(
|
|||||||
return choose(
|
return choose(
|
||||||
"Select model",
|
"Select model",
|
||||||
model_names,
|
model_names,
|
||||||
allow_custom=False,
|
allow_custom=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
from plyngent.prompting import ask
|
from plyngent.prompting import ask
|
||||||
|
|
||||||
return ask("Model id (not listed in config)")
|
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)
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ import click
|
|||||||
from click.shell_completion import CompletionItem
|
from click.shell_completion import CompletionItem
|
||||||
from msgspec import UNSET
|
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.retry import retry_pending_with_retries
|
||||||
from plyngent.cli.selection import select_model, select_provider
|
from plyngent.cli.selection import select_model, select_provider
|
||||||
from plyngent.lmproto.openai_compatible.model import (
|
from plyngent.lmproto.openai_compatible.model import (
|
||||||
@@ -117,7 +118,7 @@ class ProviderNameParam(click.ParamType[str]):
|
|||||||
state = _repl_state(ctx)
|
state = _repl_state(ctx)
|
||||||
if state is None:
|
if state is None:
|
||||||
return []
|
return []
|
||||||
return _filter_choices(incomplete, sorted(state.config.providers.keys()))
|
return _filter_choices(incomplete, sorted(state.config.selectable_providers().keys()))
|
||||||
|
|
||||||
|
|
||||||
PROVIDER_NAME = ProviderNameParam()
|
PROVIDER_NAME = ProviderNameParam()
|
||||||
@@ -137,7 +138,7 @@ class ModelIdParam(click.ParamType[str]):
|
|||||||
state = _repl_state(ctx)
|
state = _repl_state(ctx)
|
||||||
if state is None:
|
if state is None:
|
||||||
return []
|
return []
|
||||||
return _filter_choices(incomplete, sorted(state.provider.models.keys()))
|
return _filter_choices(incomplete, state.model_choice_ids())
|
||||||
|
|
||||||
|
|
||||||
MODEL_ID = ModelIdParam()
|
MODEL_ID = ModelIdParam()
|
||||||
@@ -316,6 +317,12 @@ def config_cmd(state: ReplState) -> None:
|
|||||||
click.secho(f"error: config reload failed: {exc}", fg="red")
|
click.secho(f"error: config reload failed: {exc}", fg="red")
|
||||||
click.echo(f"config file: {path}")
|
click.echo(f"config file: {path}")
|
||||||
return
|
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:
|
if state.config.bad_providers:
|
||||||
names = ", ".join(sorted(state.config.bad_providers.keys()))
|
names = ", ".join(sorted(state.config.bad_providers.keys()))
|
||||||
click.secho(f"warning: ignored bad providers: {names}", fg="yellow")
|
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}")
|
click.echo(f"provider={state.provider_name}")
|
||||||
return
|
return
|
||||||
try:
|
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
|
prev_model = state.model
|
||||||
state.provider_name = pname
|
state.provider_name = pname
|
||||||
state.provider = provider
|
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
|
state.model = prev_model
|
||||||
else:
|
else:
|
||||||
# Current model not on the new provider — pick one (prompt when interactive).
|
# Current model not on the new provider — pick one (prompt when interactive).
|
||||||
try:
|
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:
|
except click.ClickException as exc:
|
||||||
click.echo(f"error: switched provider but model selection failed: {exc}")
|
click.echo(f"error: switched provider but model selection failed: {exc}")
|
||||||
return
|
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}",
|
f"model {prev_model!r} is not available on {pname}; using {state.model!r}",
|
||||||
fg="yellow",
|
fg="yellow",
|
||||||
)
|
)
|
||||||
state.rebuild_client()
|
state.rebuild_client()
|
||||||
_await(state.persist_llm_selection())
|
_await(state.persist_llm_selection())
|
||||||
click.echo(f"switched provider to {pname} model={state.model}")
|
click.echo(f"switched provider to {pname} model={state.model}")
|
||||||
except (click.ClickException, ProviderNotSupportedError) as exc:
|
except (click.ClickException, ProviderNotSupportedError) as exc:
|
||||||
click.echo(f"error: {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")
|
@slash.command("model")
|
||||||
@click.argument("model_id", required=False, type=MODEL_ID)
|
@click.argument("model_id", required=False, type=MODEL_ID)
|
||||||
@click.pass_obj
|
@click.pass_obj
|
||||||
def model_cmd(state: ReplState, model_id: str | None) -> None:
|
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:
|
if not model_id:
|
||||||
click.echo(f"model={state.model}")
|
click.echo(f"model={state.model}")
|
||||||
return
|
return
|
||||||
try:
|
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()
|
state.rebuild_client()
|
||||||
_await(state.persist_llm_selection())
|
_await(state.persist_llm_selection())
|
||||||
click.echo(f"switched model to {state.model}")
|
click.echo(f"switched model to {state.model}")
|
||||||
|
|||||||
+102
-15
@@ -1,11 +1,20 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import contextlib
|
||||||
|
import time
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING, cast
|
from typing import TYPE_CHECKING, cast
|
||||||
|
|
||||||
from plyngent.agent import ChatAgent, ChatClient, ToolRegistry
|
from plyngent.agent import ChatAgent, ChatClient, ToolRegistry
|
||||||
from plyngent.agent.loop import DEFAULT_MAX_ROUNDS
|
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.memory.database.store import normalize_workspace
|
||||||
from plyngent.runtime import create_client
|
from plyngent.runtime import create_client
|
||||||
from plyngent.tools import DEFAULT_TOOLS, set_workspace_root
|
from plyngent.tools import DEFAULT_TOOLS, set_workspace_root
|
||||||
@@ -40,6 +49,11 @@ class ReplState:
|
|||||||
client: ChatClient = field(init=False)
|
client: ChatClient = field(init=False)
|
||||||
agent: ChatAgent = field(init=False)
|
agent: ChatAgent = field(init=False)
|
||||||
session_id: int | None = None
|
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:
|
def __post_init__(self) -> None:
|
||||||
# DeepSeek client uses a compatible but distinct param type; treat as ChatClient.
|
# 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 = self._make_agent()
|
||||||
self.agent.messages = messages
|
self.agent.messages = messages
|
||||||
self.sync_display_flags()
|
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:
|
def reload_config_from_disk(self) -> None:
|
||||||
"""Re-read TOML config and re-bind provider/model when still valid."""
|
"""Re-read TOML config and re-bind provider/model when still valid."""
|
||||||
@@ -122,27 +203,32 @@ class ReplState:
|
|||||||
self.config.reload()
|
self.config.reload()
|
||||||
set_path_denylist(self.config.agent_config.path_denylist or None)
|
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
|
preferred_model = self.model
|
||||||
try:
|
try:
|
||||||
pname, provider = select_provider(
|
pname, provider = select_provider(
|
||||||
self.config.providers,
|
selectable,
|
||||||
preferred=preferred_provider,
|
preferred=preferred_provider,
|
||||||
interactive=False,
|
interactive=False,
|
||||||
)
|
)
|
||||||
except (click.ClickException, ProviderNotSupportedError) as exc:
|
except (click.ClickException, ProviderNotSupportedError) as exc:
|
||||||
msg = f"config reloaded but provider selection failed: {exc}"
|
msg = f"config reloaded but provider selection failed: {exc}"
|
||||||
raise ValueError(msg) from 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:
|
try:
|
||||||
model_id = select_model(provider, preferred=preferred_model, interactive=False)
|
model_id = select_model(provider, preferred=preferred_model, interactive=False)
|
||||||
except click.ClickException:
|
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
|
model_id = next(iter(sorted(provider.models.keys()))) if provider.models else preferred_model
|
||||||
|
|
||||||
self.provider_name = pname
|
self.provider_name = pname
|
||||||
self.provider = provider
|
self.provider = provider
|
||||||
self.model = model_id
|
self.model = model_id
|
||||||
self.rebuild_client()
|
self.rebuild_client()
|
||||||
|
self.invalidate_remote_models()
|
||||||
|
|
||||||
def _set_workspace(self, path: Path) -> None:
|
def _set_workspace(self, path: Path) -> None:
|
||||||
"""Update REPL + tool workspace root."""
|
"""Update REPL + tool workspace root."""
|
||||||
@@ -175,16 +261,22 @@ class ReplState:
|
|||||||
from plyngent.cli.selection import select_provider
|
from plyngent.cli.selection import select_provider
|
||||||
from plyngent.runtime import ProviderNotSupportedError
|
from plyngent.runtime import ProviderNotSupportedError
|
||||||
|
|
||||||
if pname not in self.config.providers:
|
if pname not in self.config.selectable_providers():
|
||||||
return False
|
return False
|
||||||
try:
|
try:
|
||||||
name, provider = select_provider(
|
name, provider = select_provider(
|
||||||
self.config.providers,
|
self.config.selectable_providers(),
|
||||||
preferred=pname,
|
preferred=pname,
|
||||||
interactive=False,
|
interactive=False,
|
||||||
)
|
)
|
||||||
except click.ClickException, ProviderNotSupportedError:
|
except (click.ClickException, ProviderNotSupportedError):
|
||||||
return False
|
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:
|
if name != self.provider_name or provider is not self.provider:
|
||||||
self.provider_name = name
|
self.provider_name = name
|
||||||
self.provider = provider
|
self.provider = provider
|
||||||
@@ -192,16 +284,11 @@ class ReplState:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
def _try_set_model(self, model_id: str) -> bool:
|
def _try_set_model(self, model_id: str) -> bool:
|
||||||
import click
|
token = model_id.strip()
|
||||||
|
if not token:
|
||||||
from plyngent.cli.selection import select_model
|
|
||||||
|
|
||||||
try:
|
|
||||||
resolved = select_model(self.provider, preferred=model_id, interactive=False)
|
|
||||||
except click.ClickException:
|
|
||||||
return False
|
return False
|
||||||
if resolved != self.model:
|
if token != self.model:
|
||||||
self.model = resolved
|
self.model = token
|
||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|||||||
@@ -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())
|
||||||
@@ -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.
|
# When switching to b, only-a is missing → select_model is invoked interactively.
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"plyngent.cli.slash.select_model",
|
"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 await handle_slash(state, "/provider b") is True
|
||||||
assert state.provider_name == "b"
|
assert state.provider_name == "b"
|
||||||
|
|||||||
@@ -64,11 +64,25 @@ def test_select_provider_interactive_choose() -> None:
|
|||||||
assert name == "b"
|
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(
|
provider = OpenAICompatibleProvider(
|
||||||
access_key_or_token="sk",
|
access_key_or_token="sk",
|
||||||
url="https://x/v1",
|
url="https://x/v1",
|
||||||
models={"m1": ModelConfig()},
|
models={"m1": ModelConfig()},
|
||||||
)
|
)
|
||||||
with pytest.raises(Exception, match="unknown model"):
|
assert select_model(provider, preferred="nope") == "nope"
|
||||||
_ = select_model(provider, preferred="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"
|
||||||
|
|||||||
Reference in New Issue
Block a user