chore(web-platform): release V0.6.1 工程质量优化
This commit is contained in:
@@ -1,20 +1,30 @@
|
||||
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 ApiError, TrainingManager # noqa: E402
|
||||
from server import ( # noqa: E402
|
||||
MAX_JOBS,
|
||||
ApiError,
|
||||
TrainingJob,
|
||||
TrainingManager,
|
||||
TrainingRequestHandler,
|
||||
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
|
||||
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)
|
||||
@@ -23,58 +33,162 @@ 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)
|
||||
encoding="utf-8",
|
||||
)
|
||||
self.manager = TrainingManager(
|
||||
self.root, sys.executable, ("Unitree-Go2-Flat",), check_environment=False
|
||||
)
|
||||
|
||||
def tearDown(self):
|
||||
self.temporary.cleanup()
|
||||
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
|
||||
@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_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_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")
|
||||
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()
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user