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
« 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.
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"""
12from __future__ import annotations
14import re
15from collections.abc import Callable
16from dataclasses import dataclass, field
17from enum import StrEnum
18from typing import Any, Literal
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
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"
31class CatalogRowKind(StrEnum):
32 """Discriminator for the sealed CatalogRow union."""
34 LOCAL = "local"
35 FRONTIER = "frontier"
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"
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)
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()}
71class SourceMode(StrEnum):
72 """Per-task tab filter for which row source backs the visible cards.
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 """
80 LOCAL = "local"
81 CLOUD = "cloud"
82 BOTH = "both"
85_SOURCE_MODE_CYCLE: tuple[SourceMode, ...] = (SourceMode.LOCAL, SourceMode.CLOUD, SourceMode.BOTH)
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)]
94def task_to_tab_id(task: ModelTask | str) -> str:
95 """Return the per-task tab id for a ModelTask or its string value.
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
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
116@dataclass(frozen=True)
117class SizeVariant:
118 """One size/quant variant of a model family for the family-as-card strip.
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 """
126 label: str
127 quant: str
128 size_gb: float
129 ref: str
130 fit: FitChip | None = None
133@dataclass
134class LocalCatalogRow:
135 """A row in the catalog backed by a local GGUF (installable or installed).
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 """
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
170@dataclass
171class FrontierCatalogRow:
172 """A row in the catalog backed by a cloud provider's chat API.
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 """
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
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
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
197 match = PARAM_COUNT_RE.search(name)
198 return match.group(1).upper() if match else "--"
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)
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"
218def format_size_gb(size_gb: float) -> str:
219 """Format a browse row's size in GB, marked approximate.
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"
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))
235def family_to_size_variants(family: ModelFamily) -> list[SizeVariant]:
236 """Build the size-chip strip for a featured ModelFamily.
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 ]
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 "--"
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 )
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 )
312def remote_to_row(rm: RemoteModel) -> LocalCatalogRow:
313 """Convert a RemoteModel to a LocalCatalogRow.
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 )
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.
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 )
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}
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
376def row_delete_id(row: CatalogRow) -> str | None:
377 """Return the model_manager-compatible identifier for *row*.
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
390def matches_search(row: CatalogRow, search: str) -> bool:
391 """Return True if the row matches the search text (hyphen/underscore-insensitive).
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 )
411def _normalize_for_search(value: str) -> str:
412 return value.lower().replace("-", " ").replace("_", " ")