diff options
| author | Danilo M. <danix@danix.xyz> | 2026-09-18 20:25:38 +0200 |
|---|---|---|
| committer | Danilo M. <danix@danix.xyz> | 2026-09-18 20:25:38 +0200 |
| commit | b724d2e9f2b02111e562d8759df326f314f66edc (patch) | |
| tree | a99ac5f846e37f0990f54d0ed4109168259262bc | |
| parent | 108031af8ec10411fdba9af4cc3ef2703c60c7de (diff) | |
| download | llamachat-b724d2e9f2b02111e562d8759df326f314f66edc.tar.gz llamachat-b724d2e9f2b02111e562d8759df326f314f66edc.zip | |
feat: detect model input modalities from the router
| -rw-r--r-- | llamachat/backend.py | 40 | ||||
| -rw-r--r-- | llamachat/ui.py | 10 | ||||
| -rwxr-xr-x | test_llamachat.py | 70 |
3 files changed, 105 insertions, 15 deletions
diff --git a/llamachat/backend.py b/llamachat/backend.py index b79040f..8b9115c 100644 --- a/llamachat/backend.py +++ b/llamachat/backend.py @@ -260,14 +260,20 @@ class Client: """Bearer auth when the provider needs it, nothing when it does not.""" return {"Authorization": f"Bearer {self.api_key}"} if self.api_key else {} - def models(self, list_timeout: int = 30) -> list[str]: - """Model ids the router currently offers. + def models( + self, list_timeout: int = 30 + ) -> tuple[list[str], dict[str, frozenset[str]]]: + """Model ids the router currently offers, plus reported input modalities. `list_timeout` is the overall deadline for the listing call, separate from the per-turn `self.timeout` used by chat: the picker has to populate fast, and a dead provider must not hang it for minutes. The connect deadline is shorter still, so an unreachable host is reported quickly while a reachable one still gets the full window. + + The second return value maps a model id to the set of inputs the + server reports it accepts. It is empty for a server that does not + report modalities. """ try: resp = httpx.get( @@ -286,7 +292,20 @@ class Client: # Some OpenAI-compatible providers return a bare array instead of # the standard {"data": [...]} envelope. Accept both. data = payload if isinstance(payload, list) else payload.get("data", []) - return [m["id"] for m in data if "id" in m] + ids: list[str] = [] + modalities: dict[str, frozenset[str]] = {} + for entry in data: + if "id" not in entry: + continue + ids.append(entry["id"]) + # The llama.cpp model router reports which inputs a model accepts. + # A plain llama-server and every cloud provider omit it, which is + # why the map is allowed to be empty rather than assumed. + arch = entry.get("architecture") + inputs = arch.get("input_modalities") if isinstance(arch, dict) else None + if isinstance(inputs, list): + modalities[entry["id"]] = frozenset(str(x) for x in inputs) + return ids, modalities def complete(self, model: str, messages: list[dict], max_tokens: int = 48) -> str: """One short non-streaming reply, for side errands like titling. @@ -608,8 +627,8 @@ class MultiClient: _, model = providers_mod.split(model_id, self.table) return model - def models(self) -> tuple[list[str], list[str]]: - """Every offered model id, plus notes about what went wrong. + def models(self) -> tuple[list[str], list[str], dict[str, frozenset[str]]]: + """Every offered model id, plus notes and reported input modalities. A provider that is unreachable or whose filter matched nothing must not stop the others being listed: local models have to stay usable @@ -621,9 +640,10 @@ class MultiClient: """ listed: list[str] = [] problems: list[str] = [] + modalities: dict[str, frozenset[str]] = {} for name, provider in self.table.items(): try: - available = self._listing_client(provider).models( + available, reported = self._listing_client(provider).models( list_timeout=self.list_timeout ) except (BackendError, providers_mod.KeyResolutionError) as exc: @@ -634,8 +654,12 @@ class MultiClient: problems.append( f"{name}: 0 of {len(available)} models matched filter" ) - listed.extend(providers_mod.qualify(name, m) for m in kept) - return listed, problems + for model in kept: + qualified = providers_mod.qualify(name, model) + listed.append(qualified) + if model in reported: + modalities[qualified] = reported[model] + return listed, problems, modalities def _listing_client(self, provider) -> Client: """Listing needs a client too, and needs the key for a private API.""" diff --git a/llamachat/ui.py b/llamachat/ui.py index db705c3..9a7aa0d 100644 --- a/llamachat/ui.py +++ b/llamachat/ui.py @@ -519,6 +519,9 @@ class ChatWindow(QMainWindow): self.skills = skills.SkillStore(cfg.skills_dir) self.loaded_skills: list[str] = [] + # Router-reported input modalities per model id, refreshed on listing. + self.modalities: dict[str, frozenset[str]] = {} + # Which system prompt this conversation uses: a name from the store, # the NONE sentinel, or CUSTOM with text held in prompt_custom. self.prompt_name = cfg.default_prompt @@ -892,7 +895,8 @@ class ChatWindow(QMainWindow): def refresh_models(self) -> None: """Repopulate the picker from every provider, keeping the selection.""" previous = self.model_box.currentText() - available, listing_problems = self.client.models() + available, listing_problems, modalities = self.client.models() + self.modalities = modalities if not available: self.show_status( "; ".join(listing_problems) or "No models available", error=True @@ -946,7 +950,9 @@ class ChatWindow(QMainWindow): updates["vision"] = preset.vision if updates: info = dataclasses.replace(info, **updates) - return info + # A reported modality is a fact and overrides the stored, provider and + # preset guesses above. + return models_mod.apply_inputs(info, self.modalities.get(model_id)) def current_info(self): return self.model_info(self.current_model()) diff --git a/test_llamachat.py b/test_llamachat.py index faa0a00..285492d 100755 --- a/test_llamachat.py +++ b/test_llamachat.py @@ -1915,6 +1915,25 @@ def test_audio_metadata_layers(): print("ok audio metadata layers") +def test_apply_inputs(): + """Router modalities overlay a ModelInfo; None means nothing reported.""" + from llamachat import models + + base = models.ModelInfo(audio=True, vision=True) + # No report -> returned unchanged, identity preserved. + assert models.apply_inputs(base, None) is base + + out = models.apply_inputs(models.ModelInfo(), frozenset({"text", "audio"})) + assert out.audio is True + assert out.vision is False + + # A reported modality overrides a stored True. + overridden = models.apply_inputs(base, frozenset({"text"})) + assert overridden.audio is False + assert overridden.vision is False + print("ok router modalities overlay") + + def test_metadata_and_cost(): """models.ini beats provider defaults beats unknown; cost sums per model.""" from llamachat import models, providers @@ -2361,7 +2380,7 @@ def test_client_auth_header(): import httpx with _patched(httpx, "get", _recorder): - assert backend.Client("http://x.example.org").models() == ["m1"] + assert backend.Client("http://x.example.org").models() == (["m1"], {}) assert "Authorization" not in sent["headers"] backend.Client( @@ -2374,7 +2393,7 @@ def test_client_auth_header(): def _bare_list(url, timeout=None, headers=None): return _HttpxResponse([{"id": "bare"}]) with _patched(httpx, "get", _bare_list): - assert backend.Client("http://x.example.org").models() == ["bare"] + assert backend.Client("http://x.example.org").models() == (["bare"], {}) # The streaming path is the one that carries every real turn, and no # other test reaches it: the search tests all stub _stream_once out. @@ -2524,18 +2543,21 @@ def test_multi_client(): def models(self, list_timeout=30): if self.base_url not in listings: raise backend.BackendError(f"cannot reach {self.base_url}") - return listings[self.base_url] + return listings[self.base_url], {} multi = backend.MultiClient( table, timeout=300, resolver=providers.KeyResolver(), client_factory=_StubClient, ) - listed, problems = multi.models() + listed, problems, modalities = multi.models() # Local models stay bare, cloud ones are prefixed, and the filter cut # the Llama model out of together's listing. assert listed == ["gemma4", "qwen3.5-9b", "together:Qwen/Qwen2.5-72B"] + # The stub reports no modalities, so the map is empty. + assert modalities == {} + # The unreachable provider is reported, and did not break the rest. assert any("down" in p for p in problems) @@ -2559,13 +2581,49 @@ def test_multi_client(): table, timeout=300, resolver=providers.KeyResolver(), client_factory=_StubClient, ) - _, notes = empty.models() + _, notes, _ = empty.models() assert any("together: 0 of 2" in n for n in notes) del os.environ["MULTI_TEST_KEY"] print("ok multi-provider client") +def test_model_inputs(): + """The router's per-model input modalities are read from the listing.""" + class _Resp: + status_code = 200 + + def __init__(self, payload): + self._payload = payload + + def raise_for_status(self): + return None + + def json(self): + return self._payload + + payload = { + "data": [ + { + "id": "gemma4", + "architecture": {"input_modalities": ["text", "image", "audio"]}, + }, + {"id": "qwen", "architecture": {"input_modalities": ["text"]}}, + {"id": "plain"}, + ] + } + import httpx + with _patched(httpx, "get", lambda *a, **k: _Resp(payload)): + ids, modalities = backend.Client("http://x.example.org").models() + + assert ids == ["gemma4", "qwen", "plain"] + assert modalities["gemma4"] == frozenset({"text", "image", "audio"}) + assert modalities["qwen"] == frozenset({"text"}) + # A server that reports nothing yields an empty map, not a guess. + assert "plain" not in modalities + print("ok model input modalities") + + def test_model_dialog_values(): """The dialog's field text converts to ModelInfo, blanks meaning unknown.""" from llamachat import models, modeldialog @@ -3985,6 +4043,7 @@ if __name__ == "__main__": test_replays_reasoning() test_models_store() test_audio_metadata_layers() + test_apply_inputs() test_metadata_and_cost() test_token_column_migration() test_usage_columns() @@ -3994,6 +4053,7 @@ if __name__ == "__main__": test_debug_log() test_stream_body_for_cloud() test_multi_client() + test_model_inputs() test_model_dialog_values() test_cost_label_text() test_search_tool_schema() |
