Coverage for src/lilbee/modelhub/role_validator.py: 100%
77 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-09-28 17:20 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-09-28 17:20 +0000
1"""Role-slot assignment validation for the four model config fields."""
3import logging
4import os
5import sys
6from typing import cast
8from lilbee.catalog import CatalogModel, find_pick
9from lilbee.catalog.query import reclassify_by_name
10from lilbee.catalog.refs import is_bare_hf_repo
11from lilbee.catalog.types import ModelTask
12from lilbee.core.config import Config, cfg
13from lilbee.modelhub.install_state import InstallState, install_state
14from lilbee.modelhub.registry import ModelRegistry
15from lilbee.providers.model_ref import PROVIDER_PREFIXES, is_native_gguf_ref
16from lilbee.providers.roles import MODEL_ROLE_FIELDS
18log = logging.getLogger(__name__)
20# Test-only bypass. Both the env var and pytest must be present so a
21# leaked env var cannot disable validation in production.
22_SKIP_MODEL_TASK_VALIDATION_ENV = "LILBEE_SKIP_MODEL_TASK_VALIDATION"
24MODEL_FIELD_TO_TASK: dict[str, str] = {
25 "chat_model": "chat",
26 "embedding_model": "embedding",
27 "vision_model": "vision",
28 "reranker_model": "rerank",
29}
32class TaskMismatchError(ValueError):
33 """A role slot was assigned a model whose catalog task does not match.
35 Carries the structured fields so each surface (HTTP, CLI, TUI, MCP)
36 can format its own user-facing message. The default ``str()`` form is
37 surface-neutral so it is safe to surface unmodified.
38 """
40 def __init__(self, ref: str, entry_task: ModelTask, expected_task: ModelTask) -> None:
41 self.ref = ref
42 self.entry_task = entry_task
43 self.expected_task = expected_task
44 super().__init__(f"Model '{ref}' is a {entry_task} model, not {expected_task}.")
47def _model_task_validation_bypassed() -> bool:
48 if not os.environ.get(_SKIP_MODEL_TASK_VALIDATION_ENV):
49 return False
50 return sys.modules.get("pytest") is not None
53def _resolve_installed_task(registry: ModelRegistry, ref: str) -> ModelTask | None:
54 """Return the manifest's ``ModelTask`` for *ref*, name-reclassified, or ``None``."""
55 manifest = registry.get_manifest(ref)
56 if manifest is None:
57 return None
58 return ModelTask(reclassify_by_name(ref, manifest.task))
61def _is_registry_ref(ref: str) -> bool:
62 """Whether *ref* names something the model registry can hold."""
63 if not ref or not ref.strip():
64 return False
65 return ref.split("/", 1)[0] not in PROVIDER_PREFIXES
68def _skips_catalog_check(ref: str, *, allow_bypass: bool) -> bool:
69 """Whether *ref* skips the catalog check."""
70 if not _is_registry_ref(ref):
71 return True
72 return allow_bypass and _model_task_validation_bypassed()
75def _canonical_pick_ref(ref: str, entry: CatalogModel, want: ModelTask) -> str:
76 """Role-check a current pick and choose the canonical ref to persist."""
77 if entry.task != want:
78 raise TaskMismatchError(ref, ModelTask(entry.task), want)
79 # Keep a full ``<repo>/<file>.gguf`` so resolve_model_path lands on the
80 # exact installed quant; fall back to the pick's own ref otherwise.
81 if is_native_gguf_ref(ref):
82 return ref
83 canonical: str = entry.ref
84 return canonical
87def _installed_ref_and_task(ref: str) -> tuple[str, str | None]:
88 """Canonical installed ref for *ref* and its manifest task, task None if absent.
90 A bare ``<org>/<repo>`` ref canonicalizes to its installed quant's full ref
91 so the persisted value always names the exact GGUF file.
92 """
93 registry = ModelRegistry(cfg.models_dir)
94 if is_bare_hf_repo(ref):
95 ref = registry.installed_ref_for_repo(ref) or ref
96 return ref, _resolve_installed_task(registry, ref)
99def _not_installed(ref: str) -> ValueError:
100 """The error for a ref that is neither a current pick nor installed."""
101 return ValueError(
102 f"Model '{ref}' is not installed. "
103 "Install it with 'lilbee model pull <ref>' "
104 "(or POST /api/models/pull) before assigning it to a role."
105 )
108def validate_model_task_assignment(field_name: str, ref: str, *, allow_bypass: bool = True) -> str:
109 """Check *ref* is assignable to *field_name*; return the canonical ref.
111 A current pick carries its own task, so it can be role-checked before it is
112 installed, which is what the catalog UI offers. Anything else is checked
113 against the installed manifest, the only other thing that can vouch for a
114 model's role. Raises ``TaskMismatchError`` on role mismatch and
115 ``ValueError`` when the model is neither a pick nor installed.
116 """
117 if _skips_catalog_check(ref, allow_bypass=allow_bypass):
118 return ref
119 want = ModelTask(MODEL_FIELD_TO_TASK[field_name])
120 # The manifest is consulted first because it answers without touching the
121 # network. This runs on the TUI main thread from /model and the model-bar
122 # picker, where resolving picks would block the UI.
123 installed_ref, installed_task = _installed_ref_and_task(ref)
124 if installed_task is not None:
125 if installed_task != want:
126 raise TaskMismatchError(installed_ref, ModelTask(installed_task), want)
127 return installed_ref
128 entry = find_pick(ref)
129 if entry is not None:
130 return _canonical_pick_ref(ref, entry, want)
131 raise _not_installed(installed_ref)
134_MISSING_ROLE_WARNING = (
135 "%s is set to '%s', which the model registry does not hold. No model listing "
136 "shows it and the HTTP chat route cannot resolve it. Install a model with "
137 "'lilbee model pull <ref>' (or POST /api/models/pull), then set the role to that ref."
138)
140_LOOSE_FILE_ROLE_WARNING = (
141 "%s is set to '%s', a GGUF file outside the model registry. The TUI and CLI load it, "
142 "but no model listing shows it and the HTTP chat route cannot resolve it. Install a "
143 "catalog ref with 'lilbee model pull <ref>' (or POST /api/models/pull) to make the "
144 "role routable over HTTP."
145)
148def configured_role_refs(config: Config) -> dict[str, str]:
149 """The model-role fields of *config*, keyed by field name.
151 The field set comes from the role registry, so a new role is reported
152 without editing this module.
153 """
154 return cast("dict[str, str]", config.model_dump(include=set(MODEL_ROLE_FIELDS)))
157def unregistered_role_refs(config: Config, registry: ModelRegistry) -> dict[str, str]:
158 """Role fields of *config* naming a ref *registry* does not hold.
160 Blank and provider-prefixed refs are excluded; neither belongs to the
161 registry. Reporting refuses nothing, so it runs on every start.
162 """
163 return {
164 field_name: ref
165 for field_name, ref in configured_role_refs(config).items()
166 if _is_registry_ref(ref) and install_state(ref, registry) is not InstallState.REGISTERED
167 }
170def warn_unregistered_role_refs(config: Config, registry: ModelRegistry) -> None:
171 """Log one warning per role field of *config* that *registry* does not hold.
173 A ref that names a GGUF file on disk gets its own message: it loads, so
174 telling it to pull that path would prescribe a command that cannot run.
175 """
176 for field_name, ref in sorted(unregistered_role_refs(config, registry).items()):
177 loose = install_state(ref, registry) is InstallState.LOOSE_FILE
178 template = _LOOSE_FILE_ROLE_WARNING if loose else _MISSING_ROLE_WARNING
179 log.warning(template, field_name, ref)