import json import subprocess import sys import tempfile import threading import time import unittest from pathlib import Path from unittest.mock import patch sys.path.insert(0, str(Path(__file__).resolve().parents[1])) from server import ( # noqa: E402 DEFAULT_TASKS, MAX_JOBS, ApiError, TrainingJob, TrainingManager, TrainingRequestHandler, default_trainer_root, termination_signal_handler, ) class TrainingManagerTest(unittest.TestCase): def setUp(self): self.temporary = tempfile.TemporaryDirectory() self.root = Path(self.temporary.name) (self.root / "scripts").mkdir() (self.root / "scripts" / "train.py").write_text( """import os, pathlib, time print('WANDB_MODE=' + os.environ.get('WANDB_MODE', ''), flush=True) print('Learning iteration 1 / 2', flush=True) time.sleep(0.02) print('Learning iteration 2 / 2', flush=True) out=pathlib.Path('logs/rsl_rl/test/run/policy.onnx') out.parent.mkdir(parents=True, exist_ok=True) out.write_bytes(b'onnx') """, encoding="utf-8", ) self.manager = TrainingManager( self.root, sys.executable, ("Unitree-Go2-Flat",), check_environment=False ) def tearDown(self): self.temporary.cleanup() @staticmethod def payload(**patch): value = { "taskId": "Unitree-Go2-Flat", "numEnvs": 16, "maxIterations": 2, "seed": 42, "runName": "browser-test", "device": "cpu", "gpuIds": [], "wandbMode": "offline", } value.update(patch) return value def test_default_trainer_is_bundled_with_go2_assets(self): trainer_root = default_trainer_root() self.assertEqual(trainer_root, Path(__file__).resolve().parents[1] / "rl") self.assertTrue((trainer_root / "scripts" / "train.py").is_file()) self.assertTrue( ( trainer_root / "src" / "assets" / "robots" / "unitree_go2" / "xmls" / "go2.xml" ).is_file() ) def test_validates_allowlist_and_limits(self): with self.assertRaises(ApiError): self.manager.parse_config(self.payload(taskId="shell injection")) with self.assertRaises(ApiError): self.manager.parse_config(self.payload(numEnvs=0)) with self.assertRaises(ApiError): self.manager.parse_config(self.payload(runName="bad name")) with self.assertRaises(ApiError): self.manager.parse_config(self.payload(wandbMode="login")) def test_custom_task_metadata_and_validated_deployment(self): self.manager.tasks = DEFAULT_TASKS health = self.manager.health() task = next( item for item in health["taskMetadata"] if item["id"] == "Unitree-Go2-ObstacleAvoidance" ) self.assertEqual(task["sensorTypes"], ["raycast"]) config = self.manager.parse_config( self.payload( taskId=task["id"], terrainPreset="discrete_obstacles", terrainParams={"obstacle_count": 8, "friction": 0.9}, sensorCfg={"type": "raycast", "fov": 100, "maxDistance": 5}, ) ) self.assertEqual(config.deployment["observationSize"], 81) self.assertEqual(config.deployment["terrain"]["actualObstacleCount"], 8) self.assertEqual(config.deployment["sensorCfg"]["rayCount"], 32) self.assertEqual( TrainingJob(id="a" * 32, config=config).public()["deployment"], config.deployment ) rough = self.manager.parse_config(self.payload(taskId="Unitree-Go2-Rough")) self.assertFalse(rough.deployment["browserCompatible"]) self.assertEqual(rough.deployment["observationSize"], 234) def test_custom_config_rejects_unknown_nonfinite_and_incompatible_fields(self): self.manager.tasks = DEFAULT_TASKS for fields in ( {"terrainPreset": "plane;touch /tmp/injected"}, {"terrainPreset": ["plane"]}, {"terrainParams": {"friction": float("nan")}}, {"terrainParams": {"size": float("inf")}}, {"terrainParams": {"size": 10**400}}, {"terrainParams": {"obstacle_count": True}}, {"terrainParams": {"obstacle_count": 1.5}}, {"terrainParams": {"mjcf": ""}}, {"terrainParams": {"obstacle_height_min": 1, "obstacle_height_max": 0.1}}, {"sensorCfg": {"rayCount": 1_000_000}}, {"sensorCfg": {"maxDistance": 1, "safetyDistance": 1}}, {"sensorCfg": {"fov": float("nan")}}, {"sensorType": "camera_depth"}, {"rewardPresetId": "a" * 32}, ): with self.subTest(fields=fields), self.assertRaises(ApiError): self.manager.parse_config( self.payload(taskId="Unitree-Go2-ObstacleAvoidance", **fields) ) with self.assertRaises(ApiError): self.manager.parse_config(self.payload(sensorCfg={"fov": 90})) def test_custom_job_passes_server_owned_json_file_without_shell(self): self.manager.tasks = DEFAULT_TASKS config = self.manager.parse_config(self.payload(taskId="Unitree-Go2-ObstacleAvoidance")) job = TrainingJob(id="b" * 32, config=config) self.manager.jobs[job.id] = job with patch("server.subprocess.Popen", wraps=subprocess.Popen) as popen: self.manager._run(job) args = popen.call_args.args[0] self.assertNotIn("shell", popen.call_args.kwargs) path = Path(args[args.index("--task-config") + 1]) self.assertTrue(path.is_relative_to(self.root)) self.assertEqual(json.loads(path.read_text()), config.task_config) self.assertEqual(job.state, "succeeded") def test_builds_argument_array_without_shell(self): config = self.manager.parse_config(self.payload(device="gpu", gpuIds=[0, 2])) command = self.manager.command_for(config) self.assertEqual( command[:4], [sys.executable, "-u", "scripts/train.py", "Unitree-Go2-Flat"] ) self.assertEqual(command[-2:], ["--gpu-ids", "[0,2]"]) def test_resolves_reward_preset_to_inline_validated_trainer_argument(self): preset_id = "f" * 32 from tuning.schema import base_configuration reward_config = base_configuration() reward_config["weights"]["pose"] = 1.2 self.manager.preset_resolver = lambda value, task: ( reward_config if value == preset_id and task == "Unitree-Go2-Flat" else None ) config = self.manager.parse_config(self.payload(rewardPresetId=preset_id)) command = self.manager.command_for(config) index = command.index("--reward-config-json") self.assertEqual(json.loads(command[index + 1]), reward_config) def test_requires_local_host_origin_and_bearer_token(self): handler = object.__new__(TrainingRequestHandler) handler.access_token = "secret-token-1234" handler.allowed_origins = () handler.headers = { "Host": "127.0.0.1:8765", "Origin": "http://localhost:5173", "Authorization": "Bearer secret-token-1234", } handler._ensure_request() handler.headers["Authorization"] = "Bearer wrong-token" with self.assertRaises(ApiError) as error: handler._ensure_request() self.assertEqual(error.exception.status, 401) handler.headers["Authorization"] = "Bearer secret-token-1234" handler.headers["Host"] = "attacker.example" with self.assertRaises(ApiError) as error: handler._ensure_request() self.assertEqual(error.exception.status, 403) def test_caps_completed_job_history(self): config = self.manager.parse_config(self.payload()) for index in range(MAX_JOBS): job_id = f"{index:032x}" self.manager.jobs[job_id] = TrainingJob( id=job_id, config=config, state="succeeded", ) with patch("server.threading.Thread") as thread: created = self.manager.start(self.payload()) self.assertEqual(len(self.manager.jobs), MAX_JOBS) self.assertNotIn(f"{0:032x}", self.manager.jobs) self.assertIn(created["id"], self.manager.jobs) thread.return_value.start.assert_called_once() def test_cancel_waits_until_starting_process_is_registered(self): entered_popen = threading.Event() release_popen = threading.Event() terminated = threading.Event() class FakeStdout: def __iter__(self): terminated.wait(2) return iter(()) def close(self): pass class FakeProcess: pid = 1234 stdout = FakeStdout() @staticmethod def poll(): return -15 if terminated.is_set() else None @staticmethod def wait(timeout=None): if not terminated.wait(timeout): raise subprocess.TimeoutExpired("fake-training", timeout) return -15 def create_process(*_args, **_kwargs): entered_popen.set() self.assertTrue(release_popen.wait(2)) return FakeProcess() config = self.manager.parse_config(self.payload()) job = TrainingJob(id="a" * 32, config=config) self.manager.jobs[job.id] = job runner = threading.Thread(target=self.manager._run, args=(job,)) cancel_done = threading.Event() def cancel(): self.manager.cancel(job.id) cancel_done.set() with ( patch("server.subprocess.Popen", side_effect=create_process), patch("server.os.killpg", side_effect=lambda *_args: terminated.set()) as killpg, ): runner.start() self.assertTrue(entered_popen.wait(2)) canceller = threading.Thread(target=cancel) canceller.start() self.assertFalse(cancel_done.wait(0.05)) release_popen.set() canceller.join(2) runner.join(2) self.assertFalse(runner.is_alive()) self.assertFalse(canceller.is_alive()) killpg.assert_called_once_with(FakeProcess.pid, 15) self.assertEqual(self.manager.get(job.id)["state"], "cancelled") def test_sigterm_enters_controlled_shutdown(self): with self.assertRaises(KeyboardInterrupt): termination_signal_handler(15, None) def test_runs_job_and_exposes_new_onnx_artifact(self): job = self.manager.start(self.payload()) deadline = time.monotonic() + 5 while time.monotonic() < deadline: job = self.manager.get(job["id"]) if job["state"] not in ("queued", "running"): break time.sleep(0.02) self.assertEqual(job["state"], "succeeded") self.assertEqual(job["iteration"], 2) self.assertIn("WANDB_MODE=offline", job["logs"]) self.assertTrue(job["artifactReady"]) self.assertEqual(self.manager.artifact(job["id"]).read_bytes(), b"onnx") if __name__ == "__main__": unittest.main()