Coverage for src/lilbee/cli/tui/screens/catalog_utils.py: 100%

143 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-09-28 17:20 +0000

1"""Catalog data types, row builders, and formatting helpers. 

2 

3The catalog renders two distinct row shapes side by side: locally 

4installed / installable GGUFs (``LocalCatalogRow``) and cloud chat 

5models accessed through a provider's API (``FrontierCatalogRow``). 

6They share enough surface area that grouping and search reuse the 

7same helpers, but they carry different metadata and pull from 

8different sources, so they're separate types under a sealed 

9``CatalogRow`` union rather than a single optional-fields dataclass. 

10""" 

11 

12from __future__ import annotations 

13 

14import re 

15from collections.abc import Callable 

16from dataclasses import dataclass, field 

17from enum import StrEnum 

18from typing import Any, Literal 

19 

20from lilbee.catalog import PARAM_COUNT_RE, CatalogModel, ModelFamily, ModelVariant, extract_quant 

21from lilbee.catalog.types import KeyStatus, ModelCompat, ModelTask 

22from lilbee.modelhub.model_manager import RemoteModel 

23from lilbee.providers.model_ref import format_remote_ref 

24from lilbee.runtime.hardware import FitChip 

25 

26# Backend label for local GGUF models. Renderers drop this pill (the backend is 

27# implied for local models) and only show it for non-native SDK backends. 

28NATIVE_BACKEND = "native" 

29 

30 

31class CatalogRowKind(StrEnum): 

32 """Discriminator for the sealed CatalogRow union.""" 

33 

34 LOCAL = "local" 

35 FRONTIER = "frontier" 

36 

37 

38# Tab IDs for the 6-tab catalog shell. Discover is the curated landing, 

39# the four task tabs each render a single per-task grid, Library is the 

40# personal-encyclopedia view of installed local + activated cloud APIs. 

41TAB_DISCOVER = "discover" 

42TAB_CHAT = "chat" 

43TAB_EMBED = "embed" 

44TAB_VISION = "vision" 

45TAB_RERANK = "rerank" 

46TAB_LIBRARY = "library" 

47 

48# Order matters: numbered shortcuts 1-6 follow this sequence. 

49ALL_TAB_IDS: tuple[str, ...] = ( 

50 TAB_DISCOVER, 

51 TAB_CHAT, 

52 TAB_EMBED, 

53 TAB_VISION, 

54 TAB_RERANK, 

55 TAB_LIBRARY, 

56) 

57TASK_TAB_IDS: tuple[str, ...] = (TAB_CHAT, TAB_EMBED, TAB_VISION, TAB_RERANK) 

58 

59# Maps ModelTask -> the per-task tab id that renders its rows. Featured 

60# items still appear pinned at the top of their task tab; cross-task 

61# discovery happens on the Discover landing. 

62TASK_TO_TAB_ID: dict[ModelTask, str] = { 

63 ModelTask.CHAT: TAB_CHAT, 

64 ModelTask.EMBEDDING: TAB_EMBED, 

65 ModelTask.VISION: TAB_VISION, 

66 ModelTask.RERANK: TAB_RERANK, 

67} 

68TAB_ID_TO_TASK: dict[str, ModelTask] = {v: k for k, v in TASK_TO_TAB_ID.items()} 

69 

70 

71class SourceMode(StrEnum): 

72 """Per-task tab filter for which row source backs the visible cards. 

73 

74 LOCAL hides frontier rows (the default; mirrors the legacy mega-grid 

75 behavior). CLOUD shows only frontier rows for the active task. BOTH 

76 unions them so users can compare a local Llama with a cloud Llama 

77 side by side. Cycled via the ``c`` keybinding. 

78 """ 

79 

80 LOCAL = "local" 

81 CLOUD = "cloud" 

82 BOTH = "both" 

83 

84 

85_SOURCE_MODE_CYCLE: tuple[SourceMode, ...] = (SourceMode.LOCAL, SourceMode.CLOUD, SourceMode.BOTH) 

86 

87 

88def next_source_mode(current: SourceMode) -> SourceMode: 

89 """Return the next SourceMode in the LOCAL -> CLOUD -> BOTH -> LOCAL cycle.""" 

90 idx = _SOURCE_MODE_CYCLE.index(current) 

91 return _SOURCE_MODE_CYCLE[(idx + 1) % len(_SOURCE_MODE_CYCLE)] 

92 

93 

94def task_to_tab_id(task: ModelTask | str) -> str: 

95 """Return the per-task tab id for a ModelTask or its string value. 

96 

97 Accepts the string form because catalog rows carry ``task`` as a 

98 raw string (matching how HF API and the row builders return it), 

99 while the routing tables are keyed on the enum. 

100 """ 

101 if isinstance(task, ModelTask): 

102 return TASK_TO_TAB_ID[task] 

103 try: 

104 return TASK_TO_TAB_ID[ModelTask(task)] 

105 except (KeyError, ValueError) as exc: 

106 raise KeyError(f"unknown task: {task!r}") from exc 

107 

108 

109# SI thresholds for short download counts ("12.3M" / "456K") and binary 

110# thresholds for sizes ("4.2 GB" / "768 MB"). 

111_DOWNLOADS_PER_M = 1_000_000 

112_DOWNLOADS_PER_K = 1_000 

113_MB_PER_GB = 1024 

114 

115 

116@dataclass(frozen=True) 

117class SizeVariant: 

118 """One size/quant variant of a model family for the family-as-card strip. 

119 

120 ``label`` renders inline on the card (e.g. "8B Q4_K_M"). ``ref`` is 

121 the canonical pull target for this specific variant. ``fit`` is the 

122 fit chip computed against the host's available memory; ``None`` 

123 when the hardware probe has not yet run. 

124 """ 

125 

126 label: str 

127 quant: str 

128 size_gb: float 

129 ref: str 

130 fit: FitChip | None = None 

131 

132 

133@dataclass 

134class LocalCatalogRow: 

135 """A row in the catalog backed by a local GGUF (installable or installed). 

136 

137 ``name`` is the human-readable display label (e.g. "Qwen3 0.6B"). 

138 ``ref`` is the canonical identifier used for config persistence: 

139 ``hf_repo`` for catalog rows, ``hf_repo/filename`` for installed 

140 native models, and the provider's ref shape for remote/API rows. 

141 ``size_variants`` carries every quant for a family-aggregated row, 

142 so the card can render an inline chip strip and the detail drawer 

143 can list all sizes. ``fit`` is the chip for the row's primary 

144 variant. 

145 """ 

146 

147 name: str 

148 task: str 

149 params: str 

150 size: str 

151 quant: str 

152 downloads: str 

153 featured: bool 

154 installed: bool 

155 sort_downloads: int 

156 sort_size: float 

157 ref: str = "" 

158 backend: str = "" 

159 variant: ModelVariant | None = None 

160 family: ModelFamily | None = None 

161 catalog_model: CatalogModel | None = None 

162 remote_model: RemoteModel | None = None 

163 size_variants: list[SizeVariant] = field(default_factory=list) 

164 fit: FitChip | None = None 

165 compat: ModelCompat = ModelCompat.UNKNOWN 

166 safety_stripped: bool = False 

167 kind: Literal[CatalogRowKind.LOCAL] = CatalogRowKind.LOCAL 

168 

169 

170@dataclass 

171class FrontierCatalogRow: 

172 """A row in the catalog backed by a cloud provider's chat API. 

173 

174 Frontier rows skip the local-model fields (size on disk, quant, 

175 GGUF filename) because they don't apply: the model lives on the 

176 provider's infrastructure. 

177 """ 

178 

179 name: str 

180 ref: str 

181 task: str 

182 provider: str # Display label, e.g. "Gemini" / "OpenAI" / "Anthropic". 

183 provider_id: str # Canonical id used for the API key field, e.g. "gemini". 

184 key_status: KeyStatus 

185 kind: Literal[CatalogRowKind.FRONTIER] = CatalogRowKind.FRONTIER 

186 

187 

188# Sealed union discriminated on .kind. Pattern-match (or compare) on row.kind 

189# to dispatch instead of isinstance, so adding a new row type is one place. 

190CatalogRow = LocalCatalogRow | FrontierCatalogRow 

191 

192 

193def parse_param_label(name: str) -> str: 

194 """Extract parameter count label from model name (e.g. '8B', '0.6B').""" 

195 from lilbee.catalog import PARAM_COUNT_RE 

196 

197 match = PARAM_COUNT_RE.search(name) 

198 return match.group(1).upper() if match else "--" 

199 

200 

201def _format_downloads(n: int) -> str: 

202 if n >= _DOWNLOADS_PER_M: 

203 return f"{n / _DOWNLOADS_PER_M:.1f}M" 

204 if n >= _DOWNLOADS_PER_K: 

205 return f"{n / _DOWNLOADS_PER_K:.0f}K" 

206 return str(n) 

207 

208 

209def _format_size_mb(size_mb: int) -> str: 

210 """Format size in MB to a human-readable string.""" 

211 if size_mb == 0: 

212 return "--" 

213 if size_mb >= _MB_PER_GB: 

214 return f"{size_mb / _MB_PER_GB:.1f} GB" 

215 return f"{size_mb} MB" 

216 

217 

218def format_size_gb(size_gb: float) -> str: 

219 """Format a browse row's size in GB, marked approximate. 

220 

221 A listing row carries no per-file byte count, so its size comes from the 

222 parameter count and the quant's ggml type. The tilde says so: the exact 

223 figure lands when a pull names one file. 

224 """ 

225 if size_gb <= 0: 

226 return "--" 

227 return f"~{size_gb:.1f} GB" 

228 

229 

230def _is_param_count(label: str) -> bool: 

231 """True when label looks like a parameter count (e.g. '8B', '0.6B').""" 

232 return bool(PARAM_COUNT_RE.fullmatch(label)) 

233 

234 

235def family_to_size_variants(family: ModelFamily) -> list[SizeVariant]: 

236 """Build the size-chip strip for a featured ModelFamily. 

237 

238 Variants are returned in increasing size order so the chip strip 

239 reads compact-to-large left-to-right. ``fit`` is left ``None``; 

240 the catalog screen fills it in once the hardware probe has run. 

241 """ 

242 variants = sorted(family.variants, key=lambda v: v.size_mb) 

243 return [ 

244 SizeVariant( 

245 label=_size_variant_label(v), 

246 quant=v.quant or "--", 

247 size_gb=v.size_mb / 1024, 

248 ref=v.hf_repo, 

249 fit=None, 

250 ) 

251 for v in variants 

252 ] 

253 

254 

255def _size_variant_label(v: ModelVariant) -> str: 

256 """Render a compact label for a ModelVariant chip (e.g. '8B Q4_K_M').""" 

257 pieces = [p for p in (v.param_count, v.quant) if p] 

258 return " ".join(pieces) if pieces else "--" 

259 

260 

261def variant_to_row(v: ModelVariant, f: ModelFamily, installed: bool) -> LocalCatalogRow: 

262 """Convert a ModelVariant + family to a LocalCatalogRow.""" 

263 # Avoid duplicating the param count when the family name already ends with it. 

264 if v.param_count and not f.name.endswith(v.param_count): 

265 label = f"{f.name} {v.param_count}" 

266 else: 

267 label = f.name 

268 params = v.param_count if _is_param_count(v.param_count) else "--" 

269 return LocalCatalogRow( 

270 name=label, 

271 task=f.task, 

272 params=params, 

273 size=_format_size_mb(v.size_mb), 

274 quant=v.quant or "--", 

275 downloads="--", 

276 featured=True, 

277 installed=installed, 

278 sort_downloads=0, 

279 sort_size=v.size_mb / 1024, 

280 ref=v.hf_repo, 

281 backend=NATIVE_BACKEND, 

282 variant=v, 

283 family=f, 

284 compat=v.compat, 

285 safety_stripped=v.safety_stripped, 

286 ) 

287 

288 

289def catalog_to_row(m: CatalogModel, installed: bool) -> LocalCatalogRow: 

290 """Convert a CatalogModel to a LocalCatalogRow.""" 

291 quant = extract_quant(m.gguf_filename) 

292 return LocalCatalogRow( 

293 name=m.display_name, 

294 task=m.task, 

295 params=parse_param_label(m.display_name), 

296 size=format_size_gb(m.size_gb), 

297 quant=quant or "--", 

298 downloads=_format_downloads(m.downloads) if m.downloads > 0 else "--", 

299 featured=m.featured, 

300 installed=installed, 

301 sort_downloads=m.downloads, 

302 sort_size=m.size_gb, 

303 ref=m.ref, 

304 backend=NATIVE_BACKEND, 

305 catalog_model=m, 

306 # An installed model demonstrably runs, whatever the catalog probe said. 

307 compat=ModelCompat.SUPPORTED if installed else m.compat, 

308 safety_stripped=m.safety_stripped, 

309 ) 

310 

311 

312def remote_to_row(rm: RemoteModel) -> LocalCatalogRow: 

313 """Convert a RemoteModel to a LocalCatalogRow. 

314 

315 ``ref`` is the canonical ``provider/name`` form so it round-trips 

316 through ``Config.chat_model``'s validator without a per-call-site 

317 fixup. 

318 """ 

319 return LocalCatalogRow( 

320 name=rm.name, 

321 task=rm.task, 

322 params=rm.parameter_size or "--", 

323 size="--", 

324 quant="--", 

325 downloads="--", 

326 featured=False, 

327 installed=True, 

328 sort_downloads=0, 

329 sort_size=0.0, 

330 ref=format_remote_ref(rm.name, rm.provider), 

331 backend=rm.provider.lower(), 

332 remote_model=rm, 

333 # The model is live on the reporting server, so it demonstrably runs. 

334 compat=ModelCompat.SUPPORTED, 

335 ) 

336 

337 

338def frontier_row_from_remote( 

339 rm: RemoteModel, *, provider_id: str, key_status: KeyStatus 

340) -> FrontierCatalogRow: 

341 """Convert a discovered cloud chat model to a FrontierCatalogRow. 

342 

343 ``ref`` is the canonical ``provider/name`` form so callers pass it 

344 straight to ``Config.chat_model`` without re-prefixing. 

345 """ 

346 return FrontierCatalogRow( 

347 name=rm.name, 

348 ref=format_remote_ref(rm.name, rm.provider), 

349 task=rm.task, 

350 provider=rm.provider, 

351 provider_id=provider_id, 

352 key_status=key_status, 

353 ) 

354 

355 

356# Column sort key extractors. Local-only because every column except 

357# Name reads a field FrontierCatalogRow doesn't carry, and the catalog 

358# screen sorts local and frontier rows independently before concat. 

359SORT_KEYS: dict[str, Callable[[LocalCatalogRow], Any]] = { 

360 "Name": lambda r: r.name.lower(), 

361 "Task": lambda r: r.task, 

362 "Backend": lambda r: r.backend.lower(), 

363 "Params": lambda r: _param_sort_value(r.params), 

364 "Size": lambda r: r.sort_size, 

365 "Quant": lambda r: r.quant, 

366 "Downloads": lambda r: r.sort_downloads, 

367} 

368 

369 

370def _param_sort_value(params: str) -> float: 

371 """Convert param label to sortable float (e.g. '8B' -> 8.0).""" 

372 match = re.search(r"(\d+\.?\d*)", params) 

373 return float(match.group(1)) if match else 0.0 

374 

375 

376def row_delete_id(row: CatalogRow) -> str | None: 

377 """Return the model_manager-compatible identifier for *row*. 

378 

379 Remote rows hand back the bare ``RemoteModel.name`` because the 

380 Ollama HTTP API keys models by bare name, while ``ref`` carries the 

381 canonical ``ollama/<name>`` chat_model form. 

382 """ 

383 if row.kind == CatalogRowKind.FRONTIER: 

384 return row.ref or None 

385 if row.remote_model is not None: 

386 return row.remote_model.name or None 

387 return row.ref or None 

388 

389 

390def matches_search(row: CatalogRow, search: str) -> bool: 

391 """Return True if the row matches the search text (hyphen/underscore-insensitive). 

392 

393 Local rows match against name/task/params/quant/backend; frontier 

394 rows match against name + provider so users can type "gemini" and 

395 see every Gemini model regardless of suffix. 

396 """ 

397 if not search: 

398 return True 

399 needle = _normalize_for_search(search) 

400 if row.kind == CatalogRowKind.FRONTIER: 

401 return any( 

402 needle in _normalize_for_search(field) 

403 for field in (row.name, row.provider, row.provider_id) 

404 ) 

405 return any( 

406 needle in _normalize_for_search(field) 

407 for field in (row.name, row.task, row.params, row.quant, row.backend) 

408 ) 

409 

410 

411def _normalize_for_search(value: str) -> str: 

412 return value.lower().replace("-", " ").replace("_", " ")