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:
+33
-5
@@ -46,7 +46,7 @@ def _load_config(config_path: Path | None) -> ConfigStore:
|
||||
|
||||
|
||||
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:
|
||||
return
|
||||
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)
|
||||
|
||||
|
||||
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:
|
||||
raw = dict(store.database)
|
||||
# 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
|
||||
|
||||
|
||||
async def _run_chat( # noqa: C901 — chat orchestration
|
||||
async def _run_chat( # noqa: C901, PLR0912 — chat orchestration
|
||||
*,
|
||||
config_path: Path | None,
|
||||
provider_name: str | None,
|
||||
@@ -187,14 +200,19 @@ async def _run_chat( # noqa: C901 — chat orchestration
|
||||
configure_prompting(backend=NonInteractiveBackend())
|
||||
|
||||
store = _load_config(config_path)
|
||||
if store.bad_providers and not quiet:
|
||||
if not quiet:
|
||||
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)
|
||||
confirm_destructive: bool | None = False if yes else None
|
||||
|
||||
memory = await MemoryStore.open(_database_config(store, quiet=quiet or oneshot))
|
||||
try:
|
||||
from plyngent.cli.provider_recovery import ensure_provider_ready
|
||||
|
||||
# Prefer session-remembered LLM when resuming (unless flags override).
|
||||
preferred_provider = provider_name
|
||||
preferred_model = model
|
||||
@@ -211,10 +229,17 @@ async def _run_chat( # noqa: C901 — chat orchestration
|
||||
|
||||
try:
|
||||
pname, provider = select_provider(
|
||||
store.providers,
|
||||
store.selectable_providers(),
|
||||
preferred=preferred_provider,
|
||||
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)
|
||||
_ = create_client(provider)
|
||||
except ProviderNotSupportedError as exc:
|
||||
@@ -414,12 +439,15 @@ def chat_cmd(
|
||||
def providers_cmd(config_path: Path | None) -> None:
|
||||
"""List configured providers."""
|
||||
store = _load_config(config_path)
|
||||
if not store.providers:
|
||||
if not store.providers and not store.recoverable_providers:
|
||||
click.echo("(no providers)")
|
||||
for name, provider in sorted(store.providers.items()):
|
||||
tag = type(provider).__struct_config__.tag
|
||||
models = ", ".join(sorted(provider.models.keys())) or "(none listed)"
|
||||
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:
|
||||
_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 tomlkit
|
||||
|
||||
from .models import AgentConfig, DatabaseConfig, Provider
|
||||
from .models import AgentConfig, DatabaseConfig, ModelConfig, Provider
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Mapping, MutableMapping
|
||||
from collections.abc import Mapping, MutableMapping, Sequence
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
@@ -35,19 +35,25 @@ def _parse_agent(raw: dict[str, object]) -> AgentConfig:
|
||||
|
||||
def _parse_providers(
|
||||
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.
|
||||
|
||||
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] = {}
|
||||
bad_providers: dict[str, object] = {}
|
||||
recoverable: dict[str, Provider] = {}
|
||||
|
||||
raw: dict[str, object] = document.unwrap()
|
||||
providers_raw: object = raw.get("providers", {})
|
||||
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():
|
||||
if not isinstance(raw_entry, dict):
|
||||
@@ -70,15 +76,15 @@ def _parse_providers(
|
||||
bad_providers[name] = raw_entry
|
||||
continue
|
||||
|
||||
# Usable providers must list at least one model (DeepSeek seeds defaults
|
||||
# when models is omitted; an explicit empty models={} is invalid).
|
||||
# Empty models: keep as recoverable (DeepSeek seeds defaults when
|
||||
# models is omitted, so only explicit models={} lands here).
|
||||
if not provider.models:
|
||||
bad_providers[name] = {**cast("dict[str, object]", raw_entry), "_reason": "no models"}
|
||||
recoverable[name] = provider
|
||||
continue
|
||||
|
||||
providers[name] = provider
|
||||
|
||||
return providers, bad_providers
|
||||
return providers, bad_providers, recoverable
|
||||
|
||||
|
||||
class ConfigStore:
|
||||
@@ -90,6 +96,7 @@ class ConfigStore:
|
||||
_agent: AgentConfig
|
||||
_providers: dict[str, Provider]
|
||||
_bad_providers: dict[str, object]
|
||||
_recoverable_providers: dict[str, Provider]
|
||||
|
||||
def __init__(self, path: Path, document: tomlkit.TOMLDocument) -> None:
|
||||
self._path = path
|
||||
@@ -97,7 +104,7 @@ class ConfigStore:
|
||||
raw: dict[str, object] = document.unwrap()
|
||||
self._database = _parse_database(cast("dict[str, object]", raw.get("database", {})))
|
||||
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
|
||||
def path(self) -> Path:
|
||||
@@ -135,9 +142,10 @@ class ConfigStore:
|
||||
|
||||
@providers.setter
|
||||
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._bad_providers = {}
|
||||
self._recoverable_providers = {}
|
||||
|
||||
# -- bad_providers (read-only) --
|
||||
|
||||
@@ -146,6 +154,46 @@ class ConfigStore:
|
||||
"""Read-only view of unrecognised / malformed provider entries."""
|
||||
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 --
|
||||
|
||||
def write(self) -> None:
|
||||
@@ -162,7 +210,7 @@ class ConfigStore:
|
||||
raw: dict[str, object] = self._document.unwrap()
|
||||
self._database = _parse_database(cast("dict[str, object]", raw.get("database", {})))
|
||||
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 --
|
||||
|
||||
@@ -194,14 +242,15 @@ class ConfigStore:
|
||||
self._sync_section("agent", self._agent)
|
||||
|
||||
def _sync_providers_section(self) -> None:
|
||||
"""Sync ``[providers]`` to the document."""
|
||||
"""Sync ``[providers]`` to the document (ready + recoverable)."""
|
||||
section = self._toml_table("providers")
|
||||
keep = set(self._providers) | set(self._recoverable_providers)
|
||||
|
||||
for name in list(section.keys()):
|
||||
if name not in self._providers:
|
||||
if name not in keep:
|
||||
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)
|
||||
if name in section:
|
||||
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)
|
||||
|
||||
|
||||
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.write_text(
|
||||
"""
|
||||
@@ -95,7 +95,29 @@ models = {}
|
||||
)
|
||||
config = plyngent.config.load(path)
|
||||
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:
|
||||
|
||||
Reference in New Issue
Block a user