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

1"""Remote model discovery and task classification.""" 

2 

3import logging 

4import time 

5from collections.abc import Callable 

6from functools import lru_cache 

7from threading import Lock 

8 

9import httpx 

10 

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 

32 

33log = logging.getLogger(__name__) 

34 

35_EMBEDDING_FAMILIES = frozenset({"bert", "nomic-bert", "e5", "bge"}) 

36 

37_CLASSIFY_DEFAULT_TIMEOUT_S = 5.0 

38 

39 

40@lru_cache(maxsize=1) 

41def _discovery_client() -> httpx.Client: 

42 """One shared client for the local-server discovery probes. 

43 

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() 

54 

55 

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) 

59 

60 

61def _classify_remote_task(name: str, family: str) -> ModelTask: 

62 """Classify a remote model as rerank, embedding, vision, or chat (in that order). 

63 

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 

78 

79 

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. 

87 

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) 

96 

97 

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 

107 

108 

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 [] 

120 

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 

138 

139 

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. 

144 

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 [] 

157 

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 

174 

175 

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} 

182 

183 

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 ] 

196 

197 

198def discover_api_model_groups() -> list[ApiModelGroup]: 

199 """Hosted chat models for every provider with a key set, with the key's status. 

200 

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 ] 

218 

219 

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 } 

227 

228 

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] 

232 

233 

234def gather_known_model_refs() -> set[str]: 

235 """Canonical refs from the native registry, every configured local server, and APIs. 

236 

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 

249 

250 

251class KnownModelCache: 

252 """TTL-cached union of native + remote + frontier model refs. 

253 

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 """ 

260 

261 DEFAULT_TTL_S = 30.0 

262 

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() 

269 

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 

282 

283 def resolve(self, model: str) -> str | None: 

284 """Resolve *model* to its canonical ref, or None if unknown. 

285 

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 

303 

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