mirror of
https://github.com/NCBM/plyngent.git
synced 2026-07-23 05:55:16 +08:00
core/config+cli: recover empty-models providers via models()
Keep empty models as recoverable; promote to ready after remote list or --model.
This commit is contained in:
+34
-6
@@ -46,7 +46,7 @@ def _load_config(config_path: Path | None) -> ConfigStore:
|
|||||||
|
|
||||||
|
|
||||||
def _warn_bad_providers(bad: Mapping[str, object]) -> None:
|
def _warn_bad_providers(bad: Mapping[str, object]) -> None:
|
||||||
"""Surface ignored provider entries (parse errors, empty models, …)."""
|
"""Surface ignored provider entries (parse errors, unknown fields, …)."""
|
||||||
if not bad:
|
if not bad:
|
||||||
return
|
return
|
||||||
names = ", ".join(sorted(bad.keys()))
|
names = ", ".join(sorted(bad.keys()))
|
||||||
@@ -68,6 +68,19 @@ def _warn_bad_providers(bad: Mapping[str, object]) -> None:
|
|||||||
click.secho(f" - {name}: {reason}", fg="yellow", err=True)
|
click.secho(f" - {name}: {reason}", fg="yellow", err=True)
|
||||||
|
|
||||||
|
|
||||||
|
def _warn_recoverable_providers(recoverable: Mapping[str, object]) -> None:
|
||||||
|
"""Surface providers with empty models (recoverable via GET /models)."""
|
||||||
|
if not recoverable:
|
||||||
|
return
|
||||||
|
names = ", ".join(sorted(recoverable.keys()))
|
||||||
|
click.secho(
|
||||||
|
f"warning: providers with empty models ({len(recoverable)}): {names} "
|
||||||
|
f"(will try GET /models or --model on use)",
|
||||||
|
fg="yellow",
|
||||||
|
err=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _database_config(store: ConfigStore, *, quiet: bool = False) -> DatabaseConfig:
|
def _database_config(store: ConfigStore, *, quiet: bool = False) -> DatabaseConfig:
|
||||||
raw = dict(store.database)
|
raw = dict(store.database)
|
||||||
# Prefer a durable file DB so sessions survive CLI restarts.
|
# Prefer a durable file DB so sessions survive CLI restarts.
|
||||||
@@ -165,7 +178,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 — chat orchestration
|
async def _run_chat( # noqa: C901, PLR0912 — chat orchestration
|
||||||
*,
|
*,
|
||||||
config_path: Path | None,
|
config_path: Path | None,
|
||||||
provider_name: str | None,
|
provider_name: str | None,
|
||||||
@@ -187,14 +200,19 @@ async def _run_chat( # noqa: C901 — chat orchestration
|
|||||||
configure_prompting(backend=NonInteractiveBackend())
|
configure_prompting(backend=NonInteractiveBackend())
|
||||||
|
|
||||||
store = _load_config(config_path)
|
store = _load_config(config_path)
|
||||||
if store.bad_providers and not quiet:
|
if not quiet:
|
||||||
_warn_bad_providers(store.bad_providers)
|
if store.bad_providers:
|
||||||
|
_warn_bad_providers(store.bad_providers)
|
||||||
|
if store.recoverable_providers:
|
||||||
|
_warn_recoverable_providers(store.recoverable_providers)
|
||||||
|
|
||||||
_setup_workspace_and_hooks(store, workspace, interactive=interactive)
|
_setup_workspace_and_hooks(store, workspace, interactive=interactive)
|
||||||
confirm_destructive: bool | None = False if yes else None
|
confirm_destructive: bool | None = False if yes else None
|
||||||
|
|
||||||
memory = await MemoryStore.open(_database_config(store, quiet=quiet or oneshot))
|
memory = await MemoryStore.open(_database_config(store, quiet=quiet or oneshot))
|
||||||
try:
|
try:
|
||||||
|
from plyngent.cli.provider_recovery import ensure_provider_ready
|
||||||
|
|
||||||
# Prefer session-remembered LLM when resuming (unless flags override).
|
# Prefer session-remembered LLM when resuming (unless flags override).
|
||||||
preferred_provider = provider_name
|
preferred_provider = provider_name
|
||||||
preferred_model = model
|
preferred_model = model
|
||||||
@@ -211,10 +229,17 @@ async def _run_chat( # noqa: C901 — chat orchestration
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
pname, provider = select_provider(
|
pname, provider = select_provider(
|
||||||
store.providers,
|
store.selectable_providers(),
|
||||||
preferred=preferred_provider,
|
preferred=preferred_provider,
|
||||||
interactive=interactive,
|
interactive=interactive,
|
||||||
)
|
)
|
||||||
|
provider = await ensure_provider_ready(
|
||||||
|
store,
|
||||||
|
pname,
|
||||||
|
provider,
|
||||||
|
preferred_model=preferred_model,
|
||||||
|
interactive=interactive,
|
||||||
|
)
|
||||||
model_id = select_model(provider, preferred=preferred_model, interactive=interactive)
|
model_id = select_model(provider, preferred=preferred_model, interactive=interactive)
|
||||||
_ = create_client(provider)
|
_ = create_client(provider)
|
||||||
except ProviderNotSupportedError as exc:
|
except ProviderNotSupportedError as exc:
|
||||||
@@ -414,12 +439,15 @@ def chat_cmd(
|
|||||||
def providers_cmd(config_path: Path | None) -> None:
|
def providers_cmd(config_path: Path | None) -> None:
|
||||||
"""List configured providers."""
|
"""List configured providers."""
|
||||||
store = _load_config(config_path)
|
store = _load_config(config_path)
|
||||||
if not store.providers:
|
if not store.providers and not store.recoverable_providers:
|
||||||
click.echo("(no providers)")
|
click.echo("(no providers)")
|
||||||
for name, provider in sorted(store.providers.items()):
|
for name, provider in sorted(store.providers.items()):
|
||||||
tag = type(provider).__struct_config__.tag
|
tag = type(provider).__struct_config__.tag
|
||||||
models = ", ".join(sorted(provider.models.keys())) or "(none listed)"
|
models = ", ".join(sorted(provider.models.keys())) or "(none listed)"
|
||||||
click.echo(f"{name}\tpreset={tag}\tmodels={models}")
|
click.echo(f"{name}\tpreset={tag}\tmodels={models}")
|
||||||
|
for name, provider in sorted(store.recoverable_providers.items()):
|
||||||
|
tag = type(provider).__struct_config__.tag
|
||||||
|
click.echo(f"{name}\tpreset={tag}\tmodels=(empty; recoverable)")
|
||||||
if store.bad_providers:
|
if store.bad_providers:
|
||||||
_warn_bad_providers(store.bad_providers)
|
_warn_bad_providers(store.bad_providers)
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,117 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import click
|
||||||
|
|
||||||
|
from plyngent.cli.models_source import fetch_remote_model_ids
|
||||||
|
from plyngent.runtime import ProviderNotSupportedError, create_client
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
from plyngent.config.models import Provider
|
||||||
|
from plyngent.config.store import ConfigStore
|
||||||
|
|
||||||
|
|
||||||
|
async def discover_model_ids(provider: Provider) -> list[str]:
|
||||||
|
"""``GET /models`` for *provider*; empty list if unsupported or empty catalog."""
|
||||||
|
try:
|
||||||
|
client = create_client(provider)
|
||||||
|
except ProviderNotSupportedError:
|
||||||
|
return []
|
||||||
|
try:
|
||||||
|
return await fetch_remote_model_ids(client)
|
||||||
|
except (RuntimeError, TypeError, OSError, ValueError):
|
||||||
|
return []
|
||||||
|
|
||||||
|
|
||||||
|
async def try_promote_provider(
|
||||||
|
store: ConfigStore,
|
||||||
|
name: str,
|
||||||
|
*,
|
||||||
|
seed_model_ids: Sequence[str] | None = None,
|
||||||
|
) -> Provider | None:
|
||||||
|
"""Promote recoverable *name* into ready providers.
|
||||||
|
|
||||||
|
Prefers *seed_model_ids* when non-empty; otherwise remote ``models()``.
|
||||||
|
Returns the promoted provider, or ``None`` if recovery failed.
|
||||||
|
"""
|
||||||
|
if name in store.providers:
|
||||||
|
return store.providers[name]
|
||||||
|
if name not in store.recoverable_providers:
|
||||||
|
return None
|
||||||
|
|
||||||
|
provider = store.recoverable_providers[name]
|
||||||
|
ids: list[str] = []
|
||||||
|
if seed_model_ids:
|
||||||
|
ids = [mid.strip() for mid in seed_model_ids if mid and str(mid).strip()]
|
||||||
|
if not ids:
|
||||||
|
ids = await discover_model_ids(provider)
|
||||||
|
if not ids:
|
||||||
|
return None
|
||||||
|
return store.promote_provider(name, ids)
|
||||||
|
|
||||||
|
|
||||||
|
async def ensure_provider_ready(
|
||||||
|
store: ConfigStore,
|
||||||
|
name: str,
|
||||||
|
provider: Provider,
|
||||||
|
*,
|
||||||
|
preferred_model: str | None = None,
|
||||||
|
interactive: bool = True,
|
||||||
|
) -> Provider:
|
||||||
|
"""Return a ready provider; recover empty-models entries when possible.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
click.ClickException: When recovery is required but fails.
|
||||||
|
"""
|
||||||
|
if provider.models:
|
||||||
|
return provider
|
||||||
|
if name not in store.recoverable_providers and name not in store.providers:
|
||||||
|
msg = f"provider {name!r} has no models and is not recoverable"
|
||||||
|
raise click.ClickException(msg)
|
||||||
|
|
||||||
|
seeds: list[str] = []
|
||||||
|
if preferred_model and preferred_model.strip():
|
||||||
|
seeds = [preferred_model.strip()]
|
||||||
|
|
||||||
|
promoted = await try_promote_provider(store, name, seed_model_ids=seeds or None)
|
||||||
|
if promoted is not None:
|
||||||
|
click.secho(
|
||||||
|
f"recovered provider {name!r} with {len(promoted.models)} model(s) "
|
||||||
|
f"(was empty models in config)",
|
||||||
|
fg="yellow",
|
||||||
|
err=True,
|
||||||
|
)
|
||||||
|
return promoted
|
||||||
|
|
||||||
|
# Remote failed and no preferred seed: interactive free-form model id.
|
||||||
|
if interactive and not seeds:
|
||||||
|
from plyngent.prompting import ask
|
||||||
|
|
||||||
|
mid = ask(f"Model id for provider {name!r} (empty models; remote list failed)")
|
||||||
|
if mid.strip():
|
||||||
|
promoted = store.promote_provider(name, [mid.strip()])
|
||||||
|
click.secho(
|
||||||
|
f"recovered provider {name!r} with model {mid.strip()!r}",
|
||||||
|
fg="yellow",
|
||||||
|
err=True,
|
||||||
|
)
|
||||||
|
return promoted
|
||||||
|
|
||||||
|
if seeds:
|
||||||
|
# Prefer explicit model even when remote list failed.
|
||||||
|
promoted = store.promote_provider(name, seeds)
|
||||||
|
click.secho(
|
||||||
|
f"recovered provider {name!r} with model {seeds[0]!r}",
|
||||||
|
fg="yellow",
|
||||||
|
err=True,
|
||||||
|
)
|
||||||
|
return promoted
|
||||||
|
|
||||||
|
msg = (
|
||||||
|
f"provider {name!r} has empty models and could not be recovered "
|
||||||
|
f"(pass --model or fix GET /models)"
|
||||||
|
)
|
||||||
|
raise click.ClickException(msg)
|
||||||
@@ -6,10 +6,10 @@ from typing import TYPE_CHECKING, cast
|
|||||||
import msgspec
|
import msgspec
|
||||||
import tomlkit
|
import tomlkit
|
||||||
|
|
||||||
from .models import AgentConfig, DatabaseConfig, Provider
|
from .models import AgentConfig, DatabaseConfig, ModelConfig, Provider
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from collections.abc import Mapping, MutableMapping
|
from collections.abc import Mapping, MutableMapping, Sequence
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
@@ -35,19 +35,25 @@ def _parse_agent(raw: dict[str, object]) -> AgentConfig:
|
|||||||
|
|
||||||
def _parse_providers(
|
def _parse_providers(
|
||||||
document: tomlkit.TOMLDocument,
|
document: tomlkit.TOMLDocument,
|
||||||
) -> tuple[dict[str, Provider], dict[str, object]]:
|
) -> tuple[dict[str, Provider], dict[str, object], dict[str, Provider]]:
|
||||||
"""Parse provider entries from the document.
|
"""Parse provider entries from the document.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(providers, bad_providers) — valid and invalid entries respectively.
|
(providers, bad_providers, recoverable_providers)
|
||||||
|
|
||||||
|
*providers* — ready to use (non-empty ``models``).
|
||||||
|
*bad_providers* — unparseable / unknown fields.
|
||||||
|
*recoverable_providers* — parsed OK but ``models`` empty; may be
|
||||||
|
promoted after a successful remote ``GET /models`` (or explicit model id).
|
||||||
"""
|
"""
|
||||||
providers: dict[str, Provider] = {}
|
providers: dict[str, Provider] = {}
|
||||||
bad_providers: dict[str, object] = {}
|
bad_providers: dict[str, object] = {}
|
||||||
|
recoverable: dict[str, Provider] = {}
|
||||||
|
|
||||||
raw: dict[str, object] = document.unwrap()
|
raw: dict[str, object] = document.unwrap()
|
||||||
providers_raw: object = raw.get("providers", {})
|
providers_raw: object = raw.get("providers", {})
|
||||||
if not isinstance(providers_raw, dict):
|
if not isinstance(providers_raw, dict):
|
||||||
return providers, bad_providers
|
return providers, bad_providers, recoverable
|
||||||
|
|
||||||
for name, raw_entry in cast("dict[str, object]", providers_raw).items():
|
for name, raw_entry in cast("dict[str, object]", providers_raw).items():
|
||||||
if not isinstance(raw_entry, dict):
|
if not isinstance(raw_entry, dict):
|
||||||
@@ -70,15 +76,15 @@ def _parse_providers(
|
|||||||
bad_providers[name] = raw_entry
|
bad_providers[name] = raw_entry
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# Usable providers must list at least one model (DeepSeek seeds defaults
|
# Empty models: keep as recoverable (DeepSeek seeds defaults when
|
||||||
# when models is omitted; an explicit empty models={} is invalid).
|
# models is omitted, so only explicit models={} lands here).
|
||||||
if not provider.models:
|
if not provider.models:
|
||||||
bad_providers[name] = {**cast("dict[str, object]", raw_entry), "_reason": "no models"}
|
recoverable[name] = provider
|
||||||
continue
|
continue
|
||||||
|
|
||||||
providers[name] = provider
|
providers[name] = provider
|
||||||
|
|
||||||
return providers, bad_providers
|
return providers, bad_providers, recoverable
|
||||||
|
|
||||||
|
|
||||||
class ConfigStore:
|
class ConfigStore:
|
||||||
@@ -90,6 +96,7 @@ class ConfigStore:
|
|||||||
_agent: AgentConfig
|
_agent: AgentConfig
|
||||||
_providers: dict[str, Provider]
|
_providers: dict[str, Provider]
|
||||||
_bad_providers: dict[str, object]
|
_bad_providers: dict[str, object]
|
||||||
|
_recoverable_providers: dict[str, Provider]
|
||||||
|
|
||||||
def __init__(self, path: Path, document: tomlkit.TOMLDocument) -> None:
|
def __init__(self, path: Path, document: tomlkit.TOMLDocument) -> None:
|
||||||
self._path = path
|
self._path = path
|
||||||
@@ -97,7 +104,7 @@ class ConfigStore:
|
|||||||
raw: dict[str, object] = document.unwrap()
|
raw: dict[str, object] = document.unwrap()
|
||||||
self._database = _parse_database(cast("dict[str, object]", raw.get("database", {})))
|
self._database = _parse_database(cast("dict[str, object]", raw.get("database", {})))
|
||||||
self._agent = _parse_agent(cast("dict[str, object]", raw.get("agent", {})))
|
self._agent = _parse_agent(cast("dict[str, object]", raw.get("agent", {})))
|
||||||
self._providers, self._bad_providers = _parse_providers(document)
|
self._providers, self._bad_providers, self._recoverable_providers = _parse_providers(document)
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def path(self) -> Path:
|
def path(self) -> Path:
|
||||||
@@ -135,9 +142,10 @@ class ConfigStore:
|
|||||||
|
|
||||||
@providers.setter
|
@providers.setter
|
||||||
def providers(self, value: Mapping[str, Provider]) -> None: # pyright: ignore[reportPropertyTypeMismatch]
|
def providers(self, value: Mapping[str, Provider]) -> None: # pyright: ignore[reportPropertyTypeMismatch]
|
||||||
"""Replace all providers."""
|
"""Replace all ready providers (clears recoverable/bad)."""
|
||||||
self._providers = dict(value)
|
self._providers = dict(value)
|
||||||
self._bad_providers = {}
|
self._bad_providers = {}
|
||||||
|
self._recoverable_providers = {}
|
||||||
|
|
||||||
# -- bad_providers (read-only) --
|
# -- bad_providers (read-only) --
|
||||||
|
|
||||||
@@ -146,6 +154,46 @@ class ConfigStore:
|
|||||||
"""Read-only view of unrecognised / malformed provider entries."""
|
"""Read-only view of unrecognised / malformed provider entries."""
|
||||||
return MappingProxyType(self._bad_providers)
|
return MappingProxyType(self._bad_providers)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def recoverable_providers(self) -> MappingProxyType[str, Provider]:
|
||||||
|
"""Parsed providers with empty ``models`` (recoverable via remote list)."""
|
||||||
|
return MappingProxyType(self._recoverable_providers)
|
||||||
|
|
||||||
|
def selectable_providers(self) -> dict[str, Provider]:
|
||||||
|
"""Ready providers plus recoverable ones (ready wins on name clash)."""
|
||||||
|
return {**self._recoverable_providers, **self._providers}
|
||||||
|
|
||||||
|
def get_provider(self, name: str) -> Provider | None:
|
||||||
|
"""Look up a ready or recoverable provider by config name."""
|
||||||
|
if name in self._providers:
|
||||||
|
return self._providers[name]
|
||||||
|
return self._recoverable_providers.get(name)
|
||||||
|
|
||||||
|
def promote_provider(self, name: str, model_ids: Sequence[str]) -> Provider:
|
||||||
|
"""Seed ``models`` from *model_ids* and move recoverable → ready.
|
||||||
|
|
||||||
|
Also re-seeds an already-ready provider that somehow has empty models.
|
||||||
|
Does not write the TOML file (in-memory session only unless ``write()``).
|
||||||
|
"""
|
||||||
|
ids = [mid.strip() for mid in model_ids if mid and mid.strip()]
|
||||||
|
if not ids:
|
||||||
|
msg = f"cannot promote provider {name!r}: no model ids"
|
||||||
|
raise ValueError(msg)
|
||||||
|
|
||||||
|
if name in self._providers:
|
||||||
|
provider = self._providers[name]
|
||||||
|
elif name in self._recoverable_providers:
|
||||||
|
provider = self._recoverable_providers.pop(name)
|
||||||
|
else:
|
||||||
|
msg = f"unknown provider {name!r}"
|
||||||
|
raise KeyError(msg)
|
||||||
|
|
||||||
|
models = {mid: ModelConfig() for mid in ids}
|
||||||
|
promoted = msgspec.structs.replace(provider, models=models)
|
||||||
|
self._providers[name] = promoted
|
||||||
|
_ = self._bad_providers.pop(name, None)
|
||||||
|
return promoted
|
||||||
|
|
||||||
# -- persistence --
|
# -- persistence --
|
||||||
|
|
||||||
def write(self) -> None:
|
def write(self) -> None:
|
||||||
@@ -162,7 +210,7 @@ class ConfigStore:
|
|||||||
raw: dict[str, object] = self._document.unwrap()
|
raw: dict[str, object] = self._document.unwrap()
|
||||||
self._database = _parse_database(cast("dict[str, object]", raw.get("database", {})))
|
self._database = _parse_database(cast("dict[str, object]", raw.get("database", {})))
|
||||||
self._agent = _parse_agent(cast("dict[str, object]", raw.get("agent", {})))
|
self._agent = _parse_agent(cast("dict[str, object]", raw.get("agent", {})))
|
||||||
self._providers, self._bad_providers = _parse_providers(self._document)
|
self._providers, self._bad_providers, self._recoverable_providers = _parse_providers(self._document)
|
||||||
|
|
||||||
# -- internal sync helpers --
|
# -- internal sync helpers --
|
||||||
|
|
||||||
@@ -194,14 +242,15 @@ class ConfigStore:
|
|||||||
self._sync_section("agent", self._agent)
|
self._sync_section("agent", self._agent)
|
||||||
|
|
||||||
def _sync_providers_section(self) -> None:
|
def _sync_providers_section(self) -> None:
|
||||||
"""Sync ``[providers]`` to the document."""
|
"""Sync ``[providers]`` to the document (ready + recoverable)."""
|
||||||
section = self._toml_table("providers")
|
section = self._toml_table("providers")
|
||||||
|
keep = set(self._providers) | set(self._recoverable_providers)
|
||||||
|
|
||||||
for name in list(section.keys()):
|
for name in list(section.keys()):
|
||||||
if name not in self._providers:
|
if name not in keep:
|
||||||
del section[name]
|
del section[name]
|
||||||
|
|
||||||
for name, provider in self._providers.items():
|
for name, provider in {**self._recoverable_providers, **self._providers}.items():
|
||||||
raw: dict[str, object] = msgspec.to_builtins(provider)
|
raw: dict[str, object] = msgspec.to_builtins(provider)
|
||||||
if name in section:
|
if name in section:
|
||||||
entry = cast("MutableMapping[str, object]", section[name])
|
entry = cast("MutableMapping[str, object]", section[name])
|
||||||
|
|||||||
@@ -0,0 +1,75 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
import plyngent.config
|
||||||
|
from plyngent.cli.provider_recovery import ensure_provider_ready, try_promote_provider
|
||||||
|
from plyngent.config.models import ModelConfig, OpenAICompatibleProvider
|
||||||
|
|
||||||
|
|
||||||
|
def _hollow_store(tmp_path: Path):
|
||||||
|
path = tmp_path / "cfg.toml"
|
||||||
|
_ = path.write_text(
|
||||||
|
"""
|
||||||
|
[providers.hollow]
|
||||||
|
preset = "openai-compatible"
|
||||||
|
url = "https://example.com/v1"
|
||||||
|
access_key_or_token = "sk-test"
|
||||||
|
models = {}
|
||||||
|
""",
|
||||||
|
encoding="utf-8",
|
||||||
|
)
|
||||||
|
return plyngent.config.load(path)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_try_promote_with_seed(tmp_path: Path) -> None:
|
||||||
|
store = _hollow_store(tmp_path)
|
||||||
|
promoted = await try_promote_provider(store, "hollow", seed_model_ids=["gpt-x"])
|
||||||
|
assert promoted is not None
|
||||||
|
assert "gpt-x" in promoted.models
|
||||||
|
assert "hollow" in store.providers
|
||||||
|
assert "hollow" not in store.recoverable_providers
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_try_promote_via_remote(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
|
store = _hollow_store(tmp_path)
|
||||||
|
|
||||||
|
async def fake_discover(provider: object) -> list[str]:
|
||||||
|
del provider
|
||||||
|
return ["remote-a", "remote-b"]
|
||||||
|
|
||||||
|
monkeypatch.setattr("plyngent.cli.provider_recovery.discover_model_ids", fake_discover)
|
||||||
|
promoted = await try_promote_provider(store, "hollow")
|
||||||
|
assert promoted is not None
|
||||||
|
assert set(promoted.models) == {"remote-a", "remote-b"}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_ensure_provider_ready_already_ready(tmp_path: Path) -> None:
|
||||||
|
store = _hollow_store(tmp_path)
|
||||||
|
ready = OpenAICompatibleProvider(
|
||||||
|
access_key_or_token="sk",
|
||||||
|
url="https://x/v1",
|
||||||
|
models={"m": ModelConfig()},
|
||||||
|
)
|
||||||
|
store.providers = {"ready": ready}
|
||||||
|
out = await ensure_provider_ready(store, "ready", ready, interactive=False)
|
||||||
|
assert out is ready
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_ensure_provider_ready_seed_model(tmp_path: Path) -> None:
|
||||||
|
store = _hollow_store(tmp_path)
|
||||||
|
provider = store.recoverable_providers["hollow"]
|
||||||
|
out = await ensure_provider_ready(
|
||||||
|
store,
|
||||||
|
"hollow",
|
||||||
|
provider,
|
||||||
|
preferred_model="explicit",
|
||||||
|
interactive=False,
|
||||||
|
)
|
||||||
|
assert "explicit" in out.models
|
||||||
@@ -81,7 +81,7 @@ def test_read_bad_config() -> None:
|
|||||||
assert isinstance(config.bad_providers, Mapping)
|
assert isinstance(config.bad_providers, Mapping)
|
||||||
|
|
||||||
|
|
||||||
def test_provider_with_empty_models_is_bad(tmp_path: Path) -> None:
|
def test_provider_with_empty_models_is_recoverable(tmp_path: Path) -> None:
|
||||||
path = tmp_path / "empty-models.toml"
|
path = tmp_path / "empty-models.toml"
|
||||||
_ = path.write_text(
|
_ = path.write_text(
|
||||||
"""
|
"""
|
||||||
@@ -95,7 +95,29 @@ models = {}
|
|||||||
)
|
)
|
||||||
config = plyngent.config.load(path)
|
config = plyngent.config.load(path)
|
||||||
assert "hollow" not in config.providers
|
assert "hollow" not in config.providers
|
||||||
assert "hollow" in config.bad_providers
|
assert "hollow" not in config.bad_providers
|
||||||
|
assert "hollow" in config.recoverable_providers
|
||||||
|
promoted = config.promote_provider("hollow", ["m1", "m2"])
|
||||||
|
assert "hollow" in config.providers
|
||||||
|
assert "hollow" not in config.recoverable_providers
|
||||||
|
assert set(promoted.models) == {"m1", "m2"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_promote_provider_requires_ids(tmp_path: Path) -> None:
|
||||||
|
path = tmp_path / "empty-models.toml"
|
||||||
|
_ = path.write_text(
|
||||||
|
"""
|
||||||
|
[providers.hollow]
|
||||||
|
preset = "openai-compatible"
|
||||||
|
url = "https://example.com/v1"
|
||||||
|
access_key_or_token = "sk-test"
|
||||||
|
models = {}
|
||||||
|
""",
|
||||||
|
encoding="utf-8",
|
||||||
|
)
|
||||||
|
config = plyngent.config.load(path)
|
||||||
|
with pytest.raises(ValueError, match="no model ids"):
|
||||||
|
_ = config.promote_provider("hollow", [])
|
||||||
|
|
||||||
|
|
||||||
def test_read_invalid_config() -> None:
|
def test_read_invalid_config() -> None:
|
||||||
|
|||||||
Reference in New Issue
Block a user