aboutsummaryrefslogtreecommitdiffstats
path: root/test_llamachat.py
diff options
context:
space:
mode:
Diffstat (limited to 'test_llamachat.py')
-rwxr-xr-xtest_llamachat.py70
1 files changed, 65 insertions, 5 deletions
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()