Coverage for src/lilbee/modelhub/model_manager/discovery.py: 100%
135 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"""Remote model discovery and task classification."""
3import logging
4import time
5from collections.abc import Callable
6from functools import lru_cache
7from threading import Lock
9import httpx
11from lilbee.app.services import get_services
12from lilbee.catalog.formatting import agent_model_id
13from lilbee.catalog.query import (
14 EMBEDDING_NAME_PATTERNS,
15 RERANKER_NAME_PATTERNS,
16 VISION_NAME_PATTERNS,
17)
18from lilbee.catalog.types import KeyStatus, ModelTask
19from lilbee.modelhub.model_manager.types import ApiModelGroup, RemoteModel
20from lilbee.providers.backend_names import BackendName
21from lilbee.providers.base import LLMProvider
22from lilbee.providers.key_check import provider_key_set, provider_key_statuses
23from lilbee.providers.local_servers import (
24 LM_STUDIO,
25 OLLAMA,
26 LocalServerSpec,
27 openai_models_url,
28)
29from lilbee.providers.local_servers.config_urls import configured_local_servers
30from lilbee.providers.model_ref import format_remote_ref
31from lilbee.providers.sdk_backend import PROVIDER_KEYS
33log = logging.getLogger(__name__)
35_EMBEDDING_FAMILIES = frozenset({"bert", "nomic-bert", "e5", "bge"})
37_CLASSIFY_DEFAULT_TIMEOUT_S = 5.0
40@lru_cache(maxsize=1)
41def _discovery_client() -> httpx.Client:
42 """One shared client for the local-server discovery probes.
44 ``httpx.get`` builds a fresh ``Client`` per call, and every ``Client``
45 construction creates an SSL context that loads the system CA bundle. The
46 catalog re-runs discovery on tab activations and refreshes, which made
47 ``ssl.create_default_context`` 16% of a real-terminal TUI py-spy session.
48 One client builds it once and reuses connections between probes (same
49 pattern as the engine probes in ``fleet.swap_manager``). Unlike those
50 loopback-only probes, ``trust_env`` stays on: these base URLs are
51 user-configurable and a LAN server may sit behind a proxy.
52 """
53 return httpx.Client()
56def _http_get(url: str, *, timeout: float) -> httpx.Response:
57 """GET via the shared discovery client (module seam; tests stub here)."""
58 return _discovery_client().get(url, timeout=timeout)
61def _classify_remote_task(name: str, family: str) -> ModelTask:
62 """Classify a remote model as rerank, embedding, vision, or chat (in that order).
64 Embedding matches by family tag or name pattern; the name path covers
65 servers like LM Studio that report no family.
66 """
67 name_lower = name.lower()
68 if any(rp in name_lower for rp in RERANKER_NAME_PATTERNS):
69 return ModelTask.RERANK
70 family_lower = family.lower()
71 if any(ef in family_lower for ef in _EMBEDDING_FAMILIES) or any(
72 ep in name_lower for ep in EMBEDDING_NAME_PATTERNS
73 ):
74 return ModelTask.EMBEDDING
75 if any(vp in name_lower for vp in VISION_NAME_PATTERNS):
76 return ModelTask.VISION
77 return ModelTask.CHAT
80def classify_remote_models(
81 base_url: str,
82 spec: LocalServerSpec,
83 *,
84 timeout: float = _CLASSIFY_DEFAULT_TIMEOUT_S,
85) -> list[RemoteModel]:
86 """Discover and classify all models from one local server by task.
88 The strategy and provider label come from *spec* (Ollama ``/api/tags`` vs
89 LM Studio ``/v1/models``), so a server reached at a non-default host is
90 classified correctly. A transport failure or a non-JSON body yields ``[]``
91 so read-only callers stay responsive when the backend is down. A listing
92 the strategy cannot walk raises: the parse runs outside the request guard.
93 """
94 discover = _DISCOVERY_BY_KEY[spec.key]
95 return discover(base_url, spec.display_name, timeout)
98def classify_all_remote_models(
99 *,
100 timeout: float = _CLASSIFY_DEFAULT_TIMEOUT_S,
101) -> list[RemoteModel]:
102 """Classify models across every configured local server, source-labeled."""
103 result: list[RemoteModel] = []
104 for spec, base_url in configured_local_servers():
105 result.extend(classify_remote_models(base_url, spec, timeout=timeout))
106 return result
109def _discover_via_ollama_tags(
110 base_url: str, provider: BackendName, timeout: float
111) -> list[RemoteModel]:
112 """Classify models from Ollama's ``/api/tags`` using family metadata."""
113 try:
114 resp = _http_get(f"{base_url}/api/tags", timeout=timeout)
115 resp.raise_for_status()
116 raw_models = resp.json().get("models", [])
117 except Exception:
118 log.debug("Failed to classify remote models", exc_info=True)
119 return []
121 result: list[RemoteModel] = []
122 for model in raw_models:
123 name = model.get("name", "")
124 details = model.get("details", {})
125 family = details.get("family", "")
126 param_size = details.get("parameter_size", "")
127 task = _classify_remote_task(name, family)
128 result.append(
129 RemoteModel(
130 name=name,
131 task=task,
132 family=family,
133 parameter_size=param_size,
134 provider=provider,
135 )
136 )
137 return result
140def _discover_via_openai_models(
141 base_url: str, provider: BackendName, timeout: float
142) -> list[RemoteModel]:
143 """Classify models from an OpenAI-compatible ``/v1/models`` endpoint.
145 These servers report only ids (no family), so task detection runs off the
146 name patterns, which LM Studio ids usually carry. Every id is surfaced: LM
147 Studio presents LM Link remote/cloud models here as if local, so the list
148 is intentionally not filtered to locally-downloaded models.
149 """
150 try:
151 resp = _http_get(openai_models_url(base_url), timeout=timeout)
152 resp.raise_for_status()
153 raw_models = resp.json().get("data", [])
154 except Exception:
155 log.debug("Failed to classify remote models", exc_info=True)
156 return []
158 result: list[RemoteModel] = []
159 for model in raw_models:
160 name = model.get("id", "")
161 if not name:
162 continue
163 task = _classify_remote_task(name, "")
164 result.append(
165 RemoteModel(
166 name=name,
167 task=task,
168 family="",
169 parameter_size="",
170 provider=provider,
171 )
172 )
173 return result
176# Listing strategy per local-server routing key. Module-level so it stays a
177# single source of truth as servers are added to the registry.
178_DISCOVERY_BY_KEY: dict[str, Callable[[str, BackendName, float], list[RemoteModel]]] = {
179 OLLAMA.key: _discover_via_ollama_tags,
180 LM_STUDIO.key: _discover_via_openai_models,
181}
184def _chat_models_for(provider: LLMProvider, prov: str, display_name: str) -> list[RemoteModel]:
185 """The backend's chat models for hosted provider *prov*, labeled *display_name*."""
186 return [
187 RemoteModel(
188 name=model_name,
189 task=ModelTask.CHAT,
190 family="",
191 parameter_size="",
192 provider=display_name,
193 )
194 for model_name in provider.list_chat_models(prov)
195 ]
198def discover_api_model_groups() -> list[ApiModelGroup]:
199 """Hosted chat models for every provider with a key set, with the key's status.
201 Short-circuits before touching the SDK when no keys are present. Only
202 providers that list models have their key checked.
203 """
204 keyed = [(prov, label) for prov, _cfg, _env, label in PROVIDER_KEYS if provider_key_set(prov)]
205 if not keyed:
206 return []
207 provider = get_services().provider
208 listed = [
209 (prov, label, models)
210 for prov, label in keyed
211 if (models := _chat_models_for(provider, prov, label))
212 ]
213 statuses = provider_key_statuses([prov for prov, _label, _models in listed])
214 return [
215 ApiModelGroup(provider=prov, display_name=label, key_status=statuses[prov], models=models)
216 for prov, label, models in listed
217 ]
220def discover_api_models() -> dict[str, list[RemoteModel]]:
221 """Hosted chat models grouped by provider display name, for accepted keys only."""
222 return {
223 group.display_name: group.models
224 for group in discover_api_model_groups()
225 if group.key_status is KeyStatus.READY
226 }
229def detect_remote_embedding_models() -> list[str]:
230 """Return embedding-model names across every configured local server."""
231 return [m.name for m in classify_all_remote_models() if m.task == ModelTask.EMBEDDING]
234def gather_known_model_refs() -> set[str]:
235 """Canonical refs from the native registry, every configured local server, and APIs.
237 A local server or API that is down contributes an empty subset. An
238 unreadable native registry raises ``OSError`` instead: answering from the
239 remote sources alone would route a request for an installed local model
240 to a hosted provider.
241 """
242 refs = {m.ref for m in get_services().registry.list_installed()}
243 for rm in classify_all_remote_models():
244 refs.add(format_remote_ref(rm.name, rm.provider))
245 for models in discover_api_models().values():
246 for rm in models:
247 refs.add(format_remote_ref(rm.name, rm.provider))
248 return refs
251class KnownModelCache:
252 """TTL-cached union of native + remote + frontier model refs.
254 Not a ``cachetools.TTLCache``: the generation counter is the point. A
255 fan-out already in flight when :meth:`invalidate` runs would otherwise
256 install its pre-pull answer with a full TTL, hiding a freshly pulled model
257 for 30s; bumping the generation makes that late writer publish without
258 renewing the expiry. The fan-out runs off the lock because it hits network.
259 """
261 DEFAULT_TTL_S = 30.0
263 def __init__(self, ttl_s: float = DEFAULT_TTL_S) -> None:
264 self._ttl_s = ttl_s
265 self._refs: frozenset[str] = frozenset()
266 self._expires_at: float = 0.0
267 self._generation: int = 0
268 self._lock = Lock()
270 def refs(self) -> frozenset[str]:
271 """Cached canonical-ref set, refreshing past the TTL (fan-out runs off the lock)."""
272 with self._lock:
273 if time.monotonic() < self._expires_at:
274 return self._refs
275 captured_generation = self._generation
276 fresh = frozenset(gather_known_model_refs())
277 with self._lock:
278 self._refs = fresh
279 if self._generation == captured_generation:
280 self._expires_at = time.monotonic() + self._ttl_s
281 return self._refs
283 def resolve(self, model: str) -> str | None:
284 """Resolve *model* to its canonical ref, or None if unknown.
286 Accepts the canonical ref, an Ollama ``name:tag`` shorthand, and the clean
287 agent-facing id (:func:`agent_model_id`) an agent config pins in place of
288 the full GGUF path. The clean id resolves only when exactly one known ref
289 produces it, so two same-labelled quants stay unresolved rather than
290 routing to the wrong file.
291 """
292 refs = self.refs()
293 if model in refs:
294 return model
295 if "/" not in model and ":" in model:
296 prefixed = OLLAMA.qualify(model)
297 if prefixed in refs:
298 return prefixed
299 aliased = [ref for ref in refs if agent_model_id(ref) == model]
300 if len(aliased) == 1:
301 return aliased[0]
302 return None
304 def invalidate(self) -> None:
305 """Force the next ``refs()`` call to re-probe."""
306 with self._lock:
307 self._expires_at = 0.0
308 self._generation += 1