import sys import tempfile import time import unittest from pathlib import Path sys.path.insert(0, str(Path(__file__).resolve().parents[1])) from server import ApiError, TrainingManager # noqa: E402 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_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_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_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()