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 |
|
||||
| `/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 |
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
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)
|
||||
|
||||
@@ -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
|
||||
@@ -579,16 +607,74 @@ def provider_cmd(state: ReplState, name: str | None) -> None:
|
||||
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
@@ -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,15 +261,21 @@ 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
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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.
|
||||
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"
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user