mirror of
https://github.com/NCBM/plyngent.git
synced 2026-07-23 05:55:16 +08:00
core/cli: remote-first model lists; always fetch GET /models
Prefer provider catalog over config for selection/Tab; fetch at startup, /provider, /model, and /models. Config-only ids remain as fallback.
This commit is contained in:
+25
-3
@@ -180,7 +180,7 @@ async def _run_oneshot(state: ReplState, prompt_text: str) -> int:
|
|||||||
return EXIT_OK if ok else EXIT_TURN_FAILED
|
return EXIT_OK if ok else EXIT_TURN_FAILED
|
||||||
|
|
||||||
|
|
||||||
async def _run_chat( # noqa: C901, PLR0912 — chat orchestration
|
async def _run_chat( # noqa: C901, PLR0912, PLR0915 — chat orchestration
|
||||||
*,
|
*,
|
||||||
config_path: Path | None,
|
config_path: Path | None,
|
||||||
provider_name: str | None,
|
provider_name: str | None,
|
||||||
@@ -245,8 +245,27 @@ async def _run_chat( # noqa: C901, PLR0912 — chat orchestration
|
|||||||
preferred_model=preferred_model,
|
preferred_model=preferred_model,
|
||||||
interactive=interactive,
|
interactive=interactive,
|
||||||
)
|
)
|
||||||
model_id = select_model(provider, preferred=preferred_model, interactive=interactive)
|
# Build client early so we can always try GET /models for remote-first lists.
|
||||||
_ = create_client(provider)
|
from plyngent.cli.models_source import (
|
||||||
|
client_supports_models,
|
||||||
|
fetch_remote_model_ids,
|
||||||
|
model_choices_for_provider,
|
||||||
|
)
|
||||||
|
|
||||||
|
client = create_client(provider)
|
||||||
|
remote_ids: list[str] | None = None
|
||||||
|
try:
|
||||||
|
if client_supports_models(client):
|
||||||
|
remote_ids = await fetch_remote_model_ids(client)
|
||||||
|
except (RuntimeError, TypeError, OSError, ValueError):
|
||||||
|
remote_ids = None
|
||||||
|
choices = model_choices_for_provider(provider, remote_ids=remote_ids)
|
||||||
|
model_id = select_model(
|
||||||
|
provider,
|
||||||
|
preferred=preferred_model,
|
||||||
|
interactive=interactive,
|
||||||
|
choices=choices,
|
||||||
|
)
|
||||||
except ProviderNotSupportedError as exc:
|
except ProviderNotSupportedError as exc:
|
||||||
raise click.ClickException(str(exc)) from exc
|
raise click.ClickException(str(exc)) from exc
|
||||||
|
|
||||||
@@ -263,6 +282,9 @@ async def _run_chat( # noqa: C901, PLR0912 — chat orchestration
|
|||||||
interactive_limits=interactive,
|
interactive_limits=interactive,
|
||||||
yolo=yolo,
|
yolo=yolo,
|
||||||
)
|
)
|
||||||
|
# Seed remote model cache from startup fetch so Tab/complete stays warm.
|
||||||
|
if remote_ids is not None:
|
||||||
|
state.seed_remote_models(remote_ids)
|
||||||
if not quiet and not oneshot:
|
if not quiet and not oneshot:
|
||||||
click.secho(f"workspace: {state.workspace}", fg="bright_black", err=True)
|
click.secho(f"workspace: {state.workspace}", fg="bright_black", err=True)
|
||||||
|
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import inspect
|
import inspect
|
||||||
from typing import TYPE_CHECKING, Protocol, cast, runtime_checkable
|
from typing import TYPE_CHECKING, Literal, Protocol, cast, runtime_checkable
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from collections.abc import Iterable, Sequence
|
from collections.abc import Iterable, Sequence
|
||||||
@@ -11,6 +11,8 @@ if TYPE_CHECKING:
|
|||||||
# Cache remote catalog this long (seconds) unless /models --refresh.
|
# Cache remote catalog this long (seconds) unless /models --refresh.
|
||||||
DEFAULT_MODELS_CACHE_TTL = 300.0
|
DEFAULT_MODELS_CACHE_TTL = 300.0
|
||||||
|
|
||||||
|
type ModelListPrefer = Literal["remote", "union", "config"]
|
||||||
|
|
||||||
|
|
||||||
@runtime_checkable
|
@runtime_checkable
|
||||||
class SupportsModels(Protocol):
|
class SupportsModels(Protocol):
|
||||||
@@ -25,12 +27,30 @@ def config_model_ids(provider: Provider) -> list[str]:
|
|||||||
def merge_model_choices(
|
def merge_model_choices(
|
||||||
config_ids: Iterable[str],
|
config_ids: Iterable[str],
|
||||||
remote_ids: Iterable[str] | None = None,
|
remote_ids: Iterable[str] | None = None,
|
||||||
|
*,
|
||||||
|
prefer: ModelListPrefer = "remote",
|
||||||
) -> list[str]:
|
) -> list[str]:
|
||||||
"""Union config and remote ids (sorted, unique)."""
|
"""Merge config and remote model ids.
|
||||||
merged: set[str] = {i for i in config_ids if i}
|
|
||||||
if remote_ids is not None:
|
*prefer*:
|
||||||
merged.update(i for i in remote_ids if i)
|
- ``remote`` (default): remote catalog first (sorted), then config-only ids
|
||||||
return sorted(merged)
|
- ``union``: sorted unique union
|
||||||
|
- ``config``: config first, then remote-only ids
|
||||||
|
"""
|
||||||
|
config_list = [i for i in config_ids if i]
|
||||||
|
remote_list = [i for i in (remote_ids or ()) if i]
|
||||||
|
if not remote_list:
|
||||||
|
return sorted(set(config_list))
|
||||||
|
if prefer == "union":
|
||||||
|
return sorted(set(config_list) | set(remote_list))
|
||||||
|
remote_sorted = sorted(set(remote_list))
|
||||||
|
config_only = sorted(set(config_list) - set(remote_sorted))
|
||||||
|
if prefer == "remote":
|
||||||
|
return [*remote_sorted, *config_only]
|
||||||
|
# config first
|
||||||
|
config_sorted = sorted(set(config_list))
|
||||||
|
remote_only = sorted(set(remote_list) - set(config_sorted))
|
||||||
|
return [*config_sorted, *remote_only]
|
||||||
|
|
||||||
|
|
||||||
def client_supports_models(client: object) -> bool:
|
def client_supports_models(client: object) -> bool:
|
||||||
@@ -57,6 +77,7 @@ def model_choices_for_provider(
|
|||||||
provider: Provider,
|
provider: Provider,
|
||||||
*,
|
*,
|
||||||
remote_ids: Sequence[str] | None = None,
|
remote_ids: Sequence[str] | None = None,
|
||||||
|
prefer: ModelListPrefer = "remote",
|
||||||
) -> list[str]:
|
) -> list[str]:
|
||||||
"""Config plus remote catalog for selection / Tab complete."""
|
"""Config plus remote catalog for selection / Tab complete (remote-first)."""
|
||||||
return merge_model_choices(config_model_ids(provider), remote_ids)
|
return merge_model_choices(config_model_ids(provider), remote_ids, prefer=prefer)
|
||||||
|
|||||||
@@ -607,7 +607,8 @@ def provider_cmd(state: ReplState, name: str | None) -> None:
|
|||||||
state.provider_name = pname
|
state.provider_name = pname
|
||||||
state.provider = provider
|
state.provider = provider
|
||||||
state.rebuild_client()
|
state.rebuild_client()
|
||||||
choices = _await(state.merged_model_choices(refresh=False))
|
# Always request remote catalog for the new provider (bypass stale cache).
|
||||||
|
choices = _await(state.merged_model_choices(refresh=True))
|
||||||
if prev_model and (prev_model in choices or prev_model in provider.models):
|
if prev_model and (prev_model in choices or prev_model in provider.models):
|
||||||
state.model = prev_model
|
state.model = prev_model
|
||||||
else:
|
else:
|
||||||
@@ -638,11 +639,12 @@ def provider_cmd(state: ReplState, name: str | None) -> None:
|
|||||||
@click.option("--refresh", is_flag=True, help="Bypass cache and re-fetch GET /models.")
|
@click.option("--refresh", is_flag=True, help="Bypass cache and re-fetch GET /models.")
|
||||||
@click.pass_obj
|
@click.pass_obj
|
||||||
def models_cmd(state: ReplState, *, refresh: bool) -> None:
|
def models_cmd(state: ReplState, *, refresh: bool) -> None:
|
||||||
"""List models (config plus remote GET /models)."""
|
"""List models (remote-first, plus config-only ids). Always tries GET /models."""
|
||||||
|
del refresh # always re-fetch; flag kept for CLI compatibility / docs
|
||||||
remote: list[str] | None = None
|
remote: list[str] | None = None
|
||||||
remote_err: str | None = None
|
remote_err: str | None = None
|
||||||
try:
|
try:
|
||||||
remote = _await(state.ensure_remote_models(refresh=refresh))
|
remote = _await(state.ensure_remote_models(refresh=True))
|
||||||
except (RuntimeError, TypeError, OSError, ValueError) as exc:
|
except (RuntimeError, TypeError, OSError, ValueError) as exc:
|
||||||
remote_err = str(exc)
|
remote_err = str(exc)
|
||||||
remote = state.cached_remote_models()
|
remote = state.cached_remote_models()
|
||||||
@@ -669,10 +671,10 @@ def models_cmd(state: ReplState, *, refresh: bool) -> None:
|
|||||||
remote_set = set(remote or ())
|
remote_set = set(remote or ())
|
||||||
for mid in choices:
|
for mid in choices:
|
||||||
tags: list[str] = []
|
tags: list[str] = []
|
||||||
if mid in config_ids:
|
|
||||||
tags.append("config")
|
|
||||||
if mid in remote_set:
|
if mid in remote_set:
|
||||||
tags.append("remote")
|
tags.append("remote")
|
||||||
|
if mid in config_ids:
|
||||||
|
tags.append("config")
|
||||||
suffix = f" ({', '.join(tags)})" if tags else ""
|
suffix = f" ({', '.join(tags)})" if tags else ""
|
||||||
mark = " *" if mid == state.model else ""
|
mark = " *" if mid == state.model else ""
|
||||||
click.echo(f"{mid}{mark}{suffix}")
|
click.echo(f"{mid}{mark}{suffix}")
|
||||||
@@ -681,7 +683,8 @@ def models_cmd(state: ReplState, *, refresh: bool) -> None:
|
|||||||
click.secho(f"remote list unavailable: {remote_err}", fg="yellow", err=True)
|
click.secho(f"remote list unavailable: {remote_err}", fg="yellow", err=True)
|
||||||
elif remote is not None:
|
elif remote is not None:
|
||||||
click.echo(
|
click.echo(
|
||||||
f"({len(remote)} remote, {len(config_ids)} config; cache TTL {int(DEFAULT_MODELS_CACHE_TTL)}s)",
|
f"(remote-first: {len(remote)} remote, {len(config_ids)} config; "
|
||||||
|
f"cache TTL {int(DEFAULT_MODELS_CACHE_TTL)}s)",
|
||||||
err=True,
|
err=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -690,12 +693,12 @@ def models_cmd(state: ReplState, *, refresh: bool) -> None:
|
|||||||
@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 (Tab: config plus cached remote)."""
|
"""Show or switch model (Tab: remote-first plus config; live fetch on pick)."""
|
||||||
if not model_id:
|
if not model_id:
|
||||||
click.echo(f"model={state.model}")
|
click.echo(f"model={state.model}")
|
||||||
return
|
return
|
||||||
try:
|
try:
|
||||||
choices = _await(state.merged_model_choices(refresh=False))
|
choices = _await(state.merged_model_choices(refresh=True))
|
||||||
state.model = select_model(
|
state.model = select_model(
|
||||||
state.provider,
|
state.provider,
|
||||||
preferred=model_id.strip(),
|
preferred=model_id.strip(),
|
||||||
|
|||||||
@@ -167,6 +167,13 @@ class ReplState:
|
|||||||
self._remote_models_key = None
|
self._remote_models_key = None
|
||||||
self._remote_models_error = None
|
self._remote_models_error = None
|
||||||
|
|
||||||
|
def seed_remote_models(self, ids: list[str]) -> None:
|
||||||
|
"""Install a freshly fetched remote catalog into the session cache."""
|
||||||
|
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
|
||||||
|
|
||||||
def cached_remote_models(self) -> list[str] | None:
|
def cached_remote_models(self) -> list[str] | None:
|
||||||
"""Return cached remote ids if still valid for the current provider."""
|
"""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:
|
if self._remote_models is None or self._remote_models_fetched_at is None:
|
||||||
|
|||||||
@@ -12,10 +12,14 @@ from plyngent.cli.models_source import (
|
|||||||
from plyngent.config.models import ModelConfig, OpenAICompatibleProvider
|
from plyngent.config.models import ModelConfig, OpenAICompatibleProvider
|
||||||
|
|
||||||
|
|
||||||
def test_merge_model_choices_union() -> None:
|
def test_merge_model_choices_remote_first() -> None:
|
||||||
assert merge_model_choices(["b", "a"], ["a", "c"]) == ["a", "b", "c"]
|
# remote-first: remote sorted, then config-only
|
||||||
|
assert merge_model_choices(["b", "a"], ["a", "c"]) == ["a", "c", "b"]
|
||||||
assert merge_model_choices(["a"], None) == ["a"]
|
assert merge_model_choices(["a"], None) == ["a"]
|
||||||
assert merge_model_choices([], ["z"]) == ["z"]
|
assert merge_model_choices([], ["z"]) == ["z"]
|
||||||
|
assert merge_model_choices(["cfg"], ["remote", "cfg"], prefer="remote") == ["cfg", "remote"]
|
||||||
|
assert merge_model_choices(["b", "a"], ["a", "c"], prefer="union") == ["a", "b", "c"]
|
||||||
|
assert merge_model_choices(["b", "a"], ["a", "c"], prefer="config") == ["a", "b", "c"]
|
||||||
|
|
||||||
|
|
||||||
def test_model_choices_for_provider() -> None:
|
def test_model_choices_for_provider() -> None:
|
||||||
@@ -26,6 +30,7 @@ def test_model_choices_for_provider() -> None:
|
|||||||
)
|
)
|
||||||
assert config_model_ids(provider) == ["cfg"]
|
assert config_model_ids(provider) == ["cfg"]
|
||||||
assert model_choices_for_provider(provider, remote_ids=["remote", "cfg"]) == ["cfg", "remote"]
|
assert model_choices_for_provider(provider, remote_ids=["remote", "cfg"]) == ["cfg", "remote"]
|
||||||
|
assert model_choices_for_provider(provider, remote_ids=["remote"]) == ["remote", "cfg"]
|
||||||
|
|
||||||
|
|
||||||
def test_client_supports_models() -> None:
|
def test_client_supports_models() -> None:
|
||||||
|
|||||||
Reference in New Issue
Block a user