Files
Mujoco_WASM/training_server/tests/test_server.py
T
chenlin 438e56bcc8
web-platform-ci / TypeScript, lint, unit, build (push) Has been cancelled
web-platform-ci / Playwright E2E (push) Has been cancelled
feat(training): release V0.9.1 避障训练与基础策略迁移
2026-09-08 10:50:13 +08:00

286 lines
11 KiB
Python

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": "<include/>"}},
{"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()