diff options
Diffstat (limited to 'test_llamachat.py')
| -rwxr-xr-x | test_llamachat.py | 70 |
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() |
