308 lines
14 KiB
Python
308 lines
14 KiB
Python
import importlib.util
|
|
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"}}
|
|
|
|
|
|
def yue2_item(state="exited"):
|
|
return {"Id": "id-yue2", "State": state,
|
|
"Labels": {controller.MUSIC_LABEL_KEY: "yue2"}}
|
|
|
|
|
|
def separator_item(state="exited"):
|
|
return {"Id": "id-separator", "State": state,
|
|
"Labels": {controller.SEPARATOR_LABEL_KEY: "bs-roformer"}}
|
|
|
|
|
|
def voice_item(state="exited"):
|
|
return {"Id": "id-voice", "State": state,
|
|
"Labels": {controller.VOICE_LABEL_KEY: "vevo2"}}
|
|
|
|
|
|
def voice_change_item(state="exited"):
|
|
return {"Id": "id-xvc", "State": state,
|
|
"Labels": {controller.VOICE_CHANGE_LABEL_KEY: "xvc"}}
|
|
|
|
|
|
class ProfileControllerTests(unittest.TestCase):
|
|
def test_yue2_start_exclusively_stops_llm_and_ace_step(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, "YUE2_WORKER", "yue2"), \
|
|
patch.object(controller, "MUSIC_WORKER", "acestep"), \
|
|
patch.object(controller, "containers", return_value=profiles), \
|
|
patch.object(controller, "yue2_container", return_value=yue2_item()), \
|
|
patch.object(controller, "music_container", return_value=music_item("running")), \
|
|
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_yue2_worker(True)
|
|
|
|
self.assertEqual(result, {"yue2_worker": "yue2", "state": "running"})
|
|
self.assertIn(("POST", "/containers/id-ultra/stop?t=120"), calls)
|
|
self.assertIn(("POST", "/containers/id-music/stop?t=30"), calls)
|
|
self.assertEqual(calls[-1], ("POST", "/containers/id-yue2/start"))
|
|
|
|
def test_voice_change_start_exclusively_stops_gpu_workers(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, "VOICE_CHANGE_WORKER", "xvc"), \
|
|
patch.object(controller, "VOICE_WORKER", "vevo2"), \
|
|
patch.object(controller, "MUSIC_WORKER", "acestep"), \
|
|
patch.object(controller, "SEPARATOR_WORKER", "bs-roformer"), \
|
|
patch.object(controller, "containers", return_value=profiles), \
|
|
patch.object(controller, "voice_change_container", return_value=voice_change_item()), \
|
|
patch.object(controller, "voice_container", return_value=voice_item("running")), \
|
|
patch.object(controller, "music_container", return_value=music_item("running")), \
|
|
patch.object(controller, "separator_container", return_value=separator_item("running")), \
|
|
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_voice_change_worker(True)
|
|
|
|
self.assertEqual(result, {"voice_change_worker": "xvc", "state": "running"})
|
|
self.assertEqual(calls, [
|
|
("POST", "/containers/id-medium/stop?t=120"),
|
|
("POST", "/containers/id-flux/stop?t=20"),
|
|
("POST", "/containers/id-tts/stop?t=30"),
|
|
("POST", "/containers/id-music/stop?t=30"),
|
|
("POST", "/containers/id-separator/stop?t=30"),
|
|
("POST", "/containers/id-voice/stop?t=30"),
|
|
("POST", "/containers/id-xvc/start"),
|
|
])
|
|
|
|
def test_voice_start_exclusively_stops_gpu_workers(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, "VOICE_WORKER", "vevo2"), \
|
|
patch.object(controller, "MUSIC_WORKER", "acestep"), \
|
|
patch.object(controller, "SEPARATOR_WORKER", "bs-roformer"), \
|
|
patch.object(controller, "containers", return_value=profiles), \
|
|
patch.object(controller, "voice_container", return_value=voice_item()), \
|
|
patch.object(controller, "music_container", return_value=music_item("running")), \
|
|
patch.object(controller, "separator_container", return_value=separator_item("running")), \
|
|
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_voice_worker(True)
|
|
|
|
self.assertEqual(result, {"voice_worker": "vevo2", "state": "running"})
|
|
self.assertEqual(calls, [
|
|
("POST", "/containers/id-medium/stop?t=120"),
|
|
("POST", "/containers/id-flux/stop?t=20"),
|
|
("POST", "/containers/id-tts/stop?t=30"),
|
|
("POST", "/containers/id-music/stop?t=30"),
|
|
("POST", "/containers/id-separator/stop?t=30"),
|
|
("POST", "/containers/id-voice/start"),
|
|
])
|
|
|
|
def test_separator_start_exclusively_stops_gpu_workers(self):
|
|
profiles = {name: item(name) for name in controller.ALLOWED}
|
|
profiles["large"] = item("large", "running")
|
|
calls = []
|
|
|
|
def request(method, path):
|
|
calls.append((method, path))
|
|
return 204, b""
|
|
|
|
with patch.object(controller, "SEPARATOR_WORKER", "bs-roformer"), \
|
|
patch.object(controller, "MUSIC_WORKER", "acestep"), \
|
|
patch.object(controller, "containers", return_value=profiles), \
|
|
patch.object(controller, "separator_container", return_value=separator_item()), \
|
|
patch.object(controller, "music_container", return_value=music_item("running")), \
|
|
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_separator_worker(True)
|
|
|
|
self.assertEqual(result, {"separator_worker": "bs-roformer", "state": "running"})
|
|
self.assertEqual(calls, [
|
|
("POST", "/containers/id-large/stop?t=120"),
|
|
("POST", "/containers/id-flux/stop?t=20"),
|
|
("POST", "/containers/id-tts/stop?t=30"),
|
|
("POST", "/containers/id-music/stop?t=30"),
|
|
("POST", "/containers/id-separator/start"),
|
|
])
|
|
|
|
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()):
|
|
with self.assertRaisesRegex(RuntimeError, "missing"):
|
|
controller.activate("fast")
|
|
|
|
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()
|