import importlib.util import json import os import pathlib import unittest from unittest.mock import patch os.environ.setdefault("CONTROLLER_TOKEN", "x" * 48) PATH = pathlib.Path(__file__).parents[1] / "platform/docker/profile-controller/profile_controller.py" SPEC = importlib.util.spec_from_file_location("profile_controller", PATH) controller = importlib.util.module_from_spec(SPEC) assert SPEC and SPEC.loader SPEC.loader.exec_module(controller) def item(profile, state="exited"): return { "Id": f"id-{profile}", "State": state, "Labels": {controller.LABEL_KEY: profile}, } def image_item(state="exited"): return {"Id": "id-flux", "State": state, "Labels": {controller.IMAGE_LABEL_KEY: controller.IMAGE_WORKER}} def restore_item(state="exited"): return {"Id": "id-restore", "State": state, "Labels": {controller.IMAGE_LABEL_KEY: controller.RESTORE_WORKER}} def tts_item(state="running"): return {"Id": "id-tts", "State": state, "Labels": {controller.TTS_LABEL_KEY: controller.TTS_WORKER}} def music_item(state="exited"): return {"Id": "id-music", "State": state, "Labels": {controller.MUSIC_LABEL_KEY: "acestep"}} class ProfileControllerTests(unittest.TestCase): def test_status_keeps_llm_when_optional_container_is_missing(self): payload = json.dumps([item("medium", "running")]).encode() with patch.object(controller, "MUSIC_WORKER", "missing-music"), \ patch.object(controller, "docker_request", return_value=(200, payload)) as request: result = controller.status_snapshot() self.assertEqual(result["active_profile"], "medium") self.assertEqual(result["music_worker"], "missing") self.assertIn("music", result["worker_errors"]) request.assert_called_once_with("GET", "/containers/json?all=1") def test_empty_snapshot_does_not_repeat_docker_query(self): with patch.object(controller, "docker_request", return_value=(200, b"[]")) as request: result = controller.status_snapshot() self.assertIsNone(result["active_profile"]) self.assertEqual(request.call_count, 1) def test_video_ui_failure_does_not_hide_llm(self): video = {"Id": "video", "State": "running", "Labels": {controller.VIDEO_LABEL_KEY: "ltx2"}} payload = json.dumps([item("medium", "running"), video]).encode() with patch.object(controller, "VIDEO_WORKER", "ltx2"), \ patch.object(controller, "docker_request", return_value=(200, payload)), \ patch.object(controller, "cached_video_ui_state", side_effect=RuntimeError("exec failed")): result = controller.status_snapshot() self.assertEqual(result["active_profile"], "medium") self.assertEqual(result["video_worker"], "running") self.assertIn("video_ui", result["worker_errors"]) def test_music_start_exclusively_stops_gpu_workers(self): profiles = {name: item(name) for name in controller.ALLOWED} profiles["ultra"] = item("ultra", "running") calls = [] def request(method, path): calls.append((method, path)) return 204, b"" with patch.object(controller, "MUSIC_WORKER", "acestep"), \ patch.object(controller, "containers", return_value=profiles), \ patch.object(controller, "music_container", return_value=music_item()), \ patch.object(controller, "image_containers", return_value=[image_item("running")]), \ patch.object(controller, "tts_container", return_value=tts_item()), \ patch.object(controller, "docker_request", side_effect=request): result = controller.set_music_worker(True) self.assertEqual(result, {"music_worker": "acestep", "state": "running"}) self.assertEqual(calls, [ ("POST", "/containers/id-ultra/stop?t=120"), ("POST", "/containers/id-flux/stop?t=20"), ("POST", "/containers/id-tts/stop?t=30"), ("POST", "/containers/id-music/start"), ]) def test_rejects_unknown_profile_before_docker_call(self): with patch.object(controller, "docker_request") as request: with self.assertRaises(ValueError): controller.activate("shell") request.assert_not_called() def test_switch_stops_running_profile_then_starts_target(self): profiles = {name: item(name) for name in controller.ALLOWED} profiles["fast"] = item("fast", "running") calls = [] def request(method, path): calls.append((method, path)) return 204, b"" with patch.object(controller, "containers", return_value=profiles), \ patch.object(controller, "image_containers", return_value=[image_item()]), \ patch.object(controller, "tts_container", return_value=tts_item()), \ patch.object(controller, "docker_request", side_effect=request): result = controller.activate("medium") self.assertEqual(result, {"active_profile": "medium", "changed": True}) self.assertEqual(calls, [ ("POST", "/containers/id-fast/stop?t=120"), ("POST", "/containers/id-medium/start"), ]) def test_fails_if_profile_container_is_missing(self): profiles = {name: item(name) for name in controller.ALLOWED[:-1]} with patch.object(controller, "containers", return_value=profiles), \ patch.object(controller, "image_containers", return_value=[image_item()]), \ patch.object(controller, "tts_container", return_value=tts_item()), \ patch.object(controller, "docker_request") as request: with self.assertRaisesRegex(RuntimeError, "missing"): controller.activate("fast") request.assert_not_called() def test_image_start_stops_inference_first(self): profiles = {name: item(name) for name in controller.ALLOWED} profiles["medium"] = item("medium", "running") calls = [] def request(method, path): calls.append((method, path)) return 204, b"" with patch.object(controller, "containers", return_value=profiles), \ patch.object(controller, "image_container", return_value=image_item()), \ patch.object(controller, "image_containers", return_value=[image_item(), restore_item()]), \ patch.object(controller, "tts_container", return_value=tts_item()), \ patch.object(controller, "docker_request", side_effect=request): controller.set_image_worker(True) self.assertEqual(calls, [ ("POST", "/containers/id-medium/stop?t=120"), ("POST", "/containers/id-tts/stop?t=30"), ("POST", "/containers/id-flux/start"), ]) def test_restore_start_stops_flux_and_starts_restore(self): profiles = {name: item(name) for name in controller.ALLOWED} calls = [] def request(method, path): calls.append((method, path)) return 204, b"" with patch.object(controller, "containers", return_value=profiles), \ patch.object(controller, "image_container", return_value=restore_item()), \ patch.object(controller, "image_containers", return_value=[image_item("running"), restore_item()]), \ patch.object(controller, "tts_container", return_value=tts_item()), \ patch.object(controller, "docker_request", side_effect=request): controller.set_image_worker(True, controller.RESTORE_WORKER) self.assertEqual(calls, [ ("POST", "/containers/id-tts/stop?t=30"), ("POST", "/containers/id-flux/stop?t=20"), ("POST", "/containers/id-restore/start"), ]) def test_profile_activation_stops_image_worker_first(self): profiles = {name: item(name) for name in controller.ALLOWED} calls = [] def request(method, path): calls.append((method, path)) return 204, b"" with patch.object(controller, "containers", return_value=profiles), \ patch.object(controller, "image_containers", return_value=[image_item("running"), restore_item()]), \ patch.object(controller, "tts_container", return_value=tts_item()), \ patch.object(controller, "docker_request", side_effect=request): controller.activate("fast") self.assertEqual(calls, [ ("POST", "/containers/id-flux/stop?t=120"), ("POST", "/containers/id-fast/start"), ]) if __name__ == "__main__": unittest.main()