Files
Mujoco_WASM/training_server/tests/test_mobile_training.py
T
chenlin f3a8a38acd
web-platform-ci / Standalone decision service (no cloud credentials) (push) Has been cancelled
web-platform-ci / TypeScript, lint, unit, build (push) Has been cancelled
web-platform-ci / Playwright E2E (push) Has been cancelled
lekiwi-compatibility / cpu-compatibility (push) Has been cancelled
web-platform-ci / Standalone decision service (no cloud credentials) (pull_request) Has been cancelled
web-platform-ci / TypeScript, lint, unit, build (pull_request) Has been cancelled
web-platform-ci / Playwright E2E (pull_request) Has been cancelled
lekiwi-compatibility / cpu-compatibility (pull_request) Has been cancelled
feat: release v1.0.1 CADWorld 网站与 LeKiwi 智能抓放
集成同源 BYOK 会话隔离、精简模型设置、官方订阅入口和 HTTPS 发布运维;保留本地训练/调参与控制能力。同步 npm 版本及 CHANGELOG,记录公网真实 API 验收仍待用户凭据。
2026-09-24 09:57:41 +08:00

274 lines
12 KiB
Python

"""One-click API/runner tests. No CUDA, mjlab or large robot assets required."""
import hashlib
import io
import json
import sys
import tempfile
import threading
import time
import unittest
import zipfile
from http.server import ThreadingHTTPServer
from pathlib import Path
from unittest.mock import patch
from urllib.error import HTTPError
from urllib.request import Request, urlopen
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from mobile_manipulator.config import MOBILE_TASKS, ROBOTS, TASK
from server import DEFAULT_TASKS, ApiError, TrainingJob, TrainingManager, TrainingRequestHandler
TASK_ID = next(iter(MOBILE_TASKS))
def archive(robot_id="lekiwi-v1", extra=None):
files = {
"robot.json": json.dumps(ROBOTS[robot_id], separators=(",", ":")),
"task.json": json.dumps(TASK),
"environment.json": json.dumps(
dict(scene="scene.xml", robotId=robot_id, taskId=TASK["id"], mujoco="3.11.0")
),
"scene.xml": "<mujoco><worldbody/></mujoco>",
}
files.update(extra or {})
stream = io.BytesIO()
with zipfile.ZipFile(stream, "w") as z:
for name, value in files.items():
z.writestr(zipfile.ZipInfo(name, date_time=(1980, 1, 1, 0, 0, 0)), value)
return stream.getvalue()
class MobileTrainingTests(unittest.TestCase):
def setUp(self):
self.temp = tempfile.TemporaryDirectory()
self.root = Path(self.temp.name)
self.manager = TrainingManager(
self.root, sys.executable, DEFAULT_TASKS, check_environment=False
)
raw = archive()
self.package = self.manager.mobile_packages.receive(io.BytesIO(raw), len(raw))
def tearDown(self):
self.manager.shutdown()
self.temp.cleanup()
def payload(self, **changes):
result = dict(
taskId=TASK_ID,
numEnvs=2,
maxIterations=2,
seed=123,
runName="mobile",
device="cpu",
gpuIds=[],
mobilePackageId=self.package["id"],
mobileParams=dict(rolloutSteps=8, goalPosition=[0.5, 0.4, 0.05]),
)
result.update(changes)
return result
def test_registry_readiness_is_per_family(self):
health = self.manager.health()
self.assertTrue(health["ready"])
mobile = [m for m in health["taskMetadata"] if m.get("family") == "mobile-manipulator"]
self.assertEqual(len(mobile), 2)
self.assertTrue(all(m["ready"] and not m["terrainPresets"] for m in mobile))
self.assertFalse(health["taskMetadata"][0]["ready"])
def test_validation_variant_and_parameter_bounds(self):
config = self.manager.parse_config(self.payload())
self.assertEqual(config.seed, 123)
self.assertEqual(config.mobile_params["rolloutSteps"], 8)
self.assertEqual(config.mobile_params["goalPosition"], [0.5, 0.4, 0.05])
self.assertEqual(config.deployment["sceneSha256"], self.package["sceneSha256"])
for fields in [
dict(taskId="MobileManipulator-LeKiwi-Bundle"),
dict(numEnvs=65),
dict(mobilePackageId="../../etc"),
dict(terrainPreset="plane"),
dict(mobileParams={"rolloutSteps": True}),
dict(mobileParams={"rolloutSteps": 7}),
dict(mobileParams={"stage": "fly"}),
dict(mobileParams={"stage": "reach"}),
dict(mobileParams={"sourceJobId": "../../untrusted.zip"}),
dict(mobileParams={"positionJitter": float("nan")}),
dict(mobileParams={"evaluationEpisodes": True}),
dict(mobileParams={"goalPosition": [0, float("nan"), 1]}),
dict(mobileParams={"goalPosition": [0, 0, -1]}),
dict(device="gpu", gpuIds=[0, 1]),
dict(pretrainedSourceId="abc"),
]:
with self.subTest(fields=fields), self.assertRaises(ApiError):
self.manager.parse_config(self.payload(**fields))
bundle = archive("lekiwi-bundle")
uploaded = self.manager.mobile_packages.receive(io.BytesIO(bundle), len(bundle))
self.assertEqual(
self.manager.parse_config(
self.payload(
taskId="MobileManipulator-LeKiwi-Bundle", mobilePackageId=uploaded["id"]
)
).deployment["robotId"],
"lekiwi-bundle",
)
def test_runner_arguments_use_server_owned_paths_and_mobile_interpreter(self):
self.manager.mobile_python = "/isolated/mobile/python"
args = self.manager.command_for(
self.manager.parse_config(self.payload(device="gpu", gpuIds=[2]))
)
self.assertEqual(args[0], self.manager.mobile_python)
self.assertIn("training_server.mobile_manipulator.train", args)
self.assertIn("cuda:2", args)
self.assertEqual(args[args.index("--seed") + 1], "123")
self.assertEqual(args[args.index("--num-envs") + 1], "2")
self.assertTrue(Path(args[args.index("--package") + 1]).is_relative_to(self.root))
def test_upload_rejects_traversal_xml_external_paths_and_wrong_contract(self):
for extra in [
{"../escape": "x"},
{"scene.xml": '<mujoco><include file="a.xml"/></mujoco>'},
{"scene.xml": '<mujoco><asset><mesh file="/etc/passwd"/></asset></mujoco>'},
{"scene.xml": '<mujoco><extension><plugin plugin="bad"/></extension></mujoco>'},
{"robot.json": "{}"},
{"robot.json": "[]"},
{"environment.json": "[]"},
{"task.json": "{}"},
{"./scene.xml": "<mujoco/>"},
{"scene.xml": '<mujoco><asset><texture fileleft="/etc/passwd"/></asset></mujoco>'},
]:
data = archive(extra=extra)
with self.subTest(extra=extra), self.assertRaises(ValueError):
self.manager.mobile_packages.receive(io.BytesIO(data), len(data))
with self.assertRaises(ValueError):
self.manager.mobile_packages.receive(io.BytesIO(b"bad"), 3)
with patch("mobile_manipulator.packages.MAX_EXPANDED", 1), self.assertRaises(ValueError):
data = archive(extra={"extra.txt": "other"})
self.manager.mobile_packages.receive(io.BytesIO(data), len(data))
def test_staged_resume_requires_matching_successful_evaluated_job(self):
source = TrainingJob(
id="d" * 32, config=self.manager.parse_config(self.payload()), state="succeeded"
)
source.artifact = self.root / "policy.onnx"
source.artifact.write_bytes(b"onnx")
source.artifact.with_suffix(".ppo.zip").write_bytes(b"trusted-checkpoint")
self.manager.jobs[source.id] = source
params = source.config.mobile_params | {"stage": "reach", "sourceJobId": source.id}
with self.assertRaisesRegex(ApiError, "尚未达标"):
self.manager.parse_config(self.payload(mobileParams=params))
source.config.deployment["evaluation"] = {
"episodes": 10,
"successRate": 0.8,
"safetyStops": 0,
}
config = self.manager.parse_config(self.payload(mobileParams=params))
command = self.manager.command_for(config)
self.assertEqual(
command[command.index("--resume") + 1], str(source.artifact.with_suffix(".ppo.zip"))
)
with self.assertRaisesRegex(ApiError, "依次推进"):
self.manager.parse_config(self.payload(mobileParams=params | {"stage": "pick-place"}))
source.config.deployment["evaluation"]["safetyStops"] = 1
with self.assertRaisesRegex(ApiError, "尚未达标"):
self.manager.parse_config(self.payload(mobileParams=params))
# Poor quality still permits same-stage continuation, never a stage promotion.
self.manager.parse_config(self.payload(mobileParams=params | {"stage": "navigate"}))
changed = archive(extra={"mesh.txt": "different bytes"})
other_package = self.manager.mobile_packages.receive(io.BytesIO(changed), len(changed))
with self.assertRaisesRegex(ApiError, "资产快照不匹配"):
self.manager.parse_config(
self.payload(
mobilePackageId=other_package["id"], mobileParams=params | {"stage": "navigate"}
)
)
source.config.deployment["sceneSha256"] = "other"
with self.assertRaisesRegex(ApiError, "不匹配"):
self.manager.parse_config(self.payload(mobileParams=params | {"stage": "navigate"}))
def fake_command(self, config, _path=None):
# Real subprocess and lifecycle, deterministic stand-in only for expensive PPO/export.
script = self.root / "fake.py"
metadata = config.deployment | {"modelSha256": hashlib.sha256(b"onnx").hexdigest()}
script.write_text(
"import pathlib,sys,json,time\n"
"out=pathlib.Path(sys.argv[sys.argv.index('--output')+1])\n"
"print('Learning iteration 1 / 2',flush=True)\n"
"print('Mean surrogate loss: -0.25',flush=True)\n"
"time.sleep(.05)\n"
"out.write_bytes(b'onnx')\n"
f"out.with_name('deployment.json').write_text({json.dumps(json.dumps(metadata))})\n"
)
return [sys.executable, "-u", str(script)]
def test_http_create_poll_download_and_failures(self):
manager = self.manager
class Handler(TrainingRequestHandler):
access_token = "test-token"
def log_message(self, *_args):
pass
Handler.manager = manager
http = ThreadingHTTPServer(("127.0.0.1", 0), Handler)
thread = threading.Thread(target=http.serve_forever, daemon=True)
thread.start()
base = f"http://127.0.0.1:{http.server_port}"
def request(path, body=None, content_type="application/json", token="test-token"):
req = Request(
base + path,
data=body,
headers={"Authorization": f"Bearer {token}", "Content-Type": content_type},
)
return urlopen(req, timeout=5)
try:
with self.assertRaises(HTTPError) as error:
request(
"/api/training/mobile-packages", archive(), "application/zip", token="wrong"
)
self.assertEqual(error.exception.code, 401)
error.exception.close()
with request("/api/training/mobile-packages", archive(), "application/zip") as response:
self.assertEqual(json.load(response)["id"], self.package["id"])
with patch.object(manager, "command_for", side_effect=self.fake_command):
with request("/api/training/jobs", json.dumps(self.payload()).encode()) as response:
self.assertEqual(response.status, 202)
job = json.load(response)
for _ in range(100):
with request("/api/training/jobs/" + job["id"]) as response:
job = json.load(response)
if job["state"] not in ("queued", "running"):
break
time.sleep(0.02)
self.assertEqual(job["state"], "succeeded", job)
self.assertEqual(job["progress"], 1)
self.assertIn("Mean surrogate loss: -0.25", job["logs"])
for filename in ("policy.onnx", "deployment.json"):
with request(f"/api/training/jobs/{job['id']}/artifacts/{filename}") as response:
self.assertTrue(response.read())
self.assertIsNone(manager.lease.public())
failed = TrainingJob(id="f" * 32, config=manager.parse_config(self.payload()))
with patch.object(
manager, "command_for", return_value=[sys.executable, "-c", "raise SystemExit(7)"]
):
manager._run(failed)
self.assertEqual(failed.state, "failed")
self.assertIsNone(failed.artifact)
finally:
http.shutdown()
http.server_close()
thread.join()
def test_cancel_before_launch_and_progress_parser(self):
job = TrainingJob(id="c" * 32, config=self.manager.parse_config(self.payload()))
self.manager._update_from_log(job, "\x1b[32mLearning iteration 1 / 2\x1b[0m")
self.assertEqual(job.public()["progress"], 0.5)
job.cancel_requested = True
with patch("server.subprocess.Popen") as popen:
self.manager._run(job)
popen.assert_not_called()
self.assertEqual(job.state, "cancelled")