Files
AI-Profile-Router/dev/test_profile_controller.py
T

234 lines
10 KiB
Python

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 qwen_image_test_item(state="exited"):
return {"Id": "id-qwen-image-test", "State": state,
"Labels": {controller.IMAGE_LABEL_KEY: controller.QWEN_IMAGE_TEST_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_qwen_image_test_is_allowlisted_and_exclusive(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=qwen_image_test_item()), \
patch.object(controller, "image_containers", return_value=[
image_item("running"), restore_item(), qwen_image_test_item()]), \
patch.object(controller, "tts_container", return_value=tts_item()), \
patch.object(controller, "docker_request", side_effect=request):
result = controller.set_image_worker(
True, controller.QWEN_IMAGE_TEST_WORKER)
self.assertEqual(result, {
"image_worker": "qwen-image-2.1-test", "state": "running"})
self.assertEqual(calls, [
("POST", "/containers/id-medium/stop?t=120"),
("POST", "/containers/id-tts/stop?t=30"),
("POST", "/containers/id-flux/stop?t=20"),
("POST", "/containers/id-qwen-image-test/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()