aboutsummaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorDanilo M. <danix@danix.xyz>2026-09-18 20:25:38 +0200
committerDanilo M. <danix@danix.xyz>2026-09-18 20:25:38 +0200
commitb724d2e9f2b02111e562d8759df326f314f66edc (patch)
treea99ac5f846e37f0990f54d0ed4109168259262bc
parent108031af8ec10411fdba9af4cc3ef2703c60c7de (diff)
downloadllamachat-b724d2e9f2b02111e562d8759df326f314f66edc.tar.gz
llamachat-b724d2e9f2b02111e562d8759df326f314f66edc.zip
feat: detect model input modalities from the router
-rw-r--r--llamachat/backend.py40
-rw-r--r--llamachat/ui.py10
-rwxr-xr-xtest_llamachat.py70
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()