diff --git a/dev/test_dashboard_modes.py b/dev/test_dashboard_modes.py new file mode 100644 index 0000000..ccadbfb --- /dev/null +++ b/dev/test_dashboard_modes.py @@ -0,0 +1,65 @@ +import importlib.util +import json +import os +import tempfile +import unittest +from pathlib import Path +from unittest.mock import patch + + +ROOT = Path(__file__).resolve().parents[1] +MODULE_PATH = ROOT / "platform/llama-dashboard/app.py" + + +class _Response: + status = 202 + + def __enter__(self): + return self + + def __exit__(self, *_args): + return False + + def read(self): + return b'{"status":"accepted"}' + + +class DashboardModeTests(unittest.TestCase): + @classmethod + def setUpClass(cls): + cls.tempdir = tempfile.TemporaryDirectory() + with patch.dict(os.environ, { + "DASHBOARD_HISTORY_DB": str(Path(cls.tempdir.name) / "history.sqlite3"), + "ROUTER_URL": "http://router.test:8081", + "ROUTER_API_KEY": "test-key", + }): + spec = importlib.util.spec_from_file_location("dashboard_app_test", MODULE_PATH) + cls.dashboard = importlib.util.module_from_spec(spec) + assert spec.loader is not None + spec.loader.exec_module(cls.dashboard) + + @classmethod + def tearDownClass(cls): + cls.tempdir.cleanup() + + def test_separation_mode_is_forwarded_to_router(self): + with patch.object(self.dashboard.urllib.request, "urlopen", return_value=_Response()) as urlopen: + status, body = self.dashboard.change_mode("separation") + + self.assertEqual(status, 202) + self.assertEqual(body, {"status": "accepted"}) + request = urlopen.call_args.args[0] + self.assertEqual(json.loads(request.data), {"mode": "separation"}) + self.assertEqual(request.get_header("Authorization"), "Bearer test-key") + + def test_unknown_mode_is_rejected_without_router_request(self): + with patch.object(self.dashboard.urllib.request, "urlopen") as urlopen: + status, body = self.dashboard.change_mode("unknown") + + self.assertEqual(status, 400) + self.assertEqual(body, {"error": "invalid mode"}) + urlopen.assert_not_called() + + +if __name__ == "__main__": + unittest.main() diff --git a/platform/llama-dashboard/app.py b/platform/llama-dashboard/app.py index f1b9d32..eb5f61e 100644 --- a/platform/llama-dashboard/app.py +++ b/platform/llama-dashboard/app.py @@ -277,7 +277,7 @@ def router_status() -> tuple[dict[str, Any], str | None]: def change_mode(mode: str) -> tuple[int, dict[str, Any]]: - if mode not in {"llm", "music"}: + if mode not in {"llm", "music", "separation"}: return 400, {"error": "invalid mode"} headers = {"Accept": "application/json", "Content-Type": "application/json"} if ROUTER_API_KEY: