Files
Mujoco_WASM/training_server/tests/test_reward_preset_tasks.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

138 lines
5.8 KiB
Python

"""Persisted preset identity -> resolver -> training validation; no training/API calls."""
import json
import sys
import tempfile
import unittest
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from server import ApiError, TrainingManager # noqa: E402
from tuning.manager import TuningManager # noqa: E402
from tuning.schema import ( # noqa: E402
FLAT_TASK,
OBSTACLE_TASK,
RewardConfigError,
base_configuration,
)
from tuning.storage import TuningStorage # noqa: E402
class RewardPresetTaskTest(unittest.TestCase):
def setUp(self):
self.temp = tempfile.TemporaryDirectory()
root = Path(self.temp.name)
(root / "scripts").mkdir()
(root / "scripts/train.py").write_text("raise AssertionError('must not train')")
self.storage = TuningStorage(root / "tuning.sqlite3")
# Real resolver with real storage; no advisor, processes or training workers are needed.
self.tuning = object.__new__(TuningManager)
self.tuning.storage = self.storage
self.training = TrainingManager(
root, sys.executable, (FLAT_TASK, OBSTACLE_TASK), check_environment=False
)
self.training.preset_resolver = self.tuning.preset_config
def tearDown(self):
self.storage.connection().close()
self.temp.cleanup()
def save(self, config, reward=None):
session = self.storage.create_session("approval", config, {}, False)
reward = reward or base_configuration(config.get("taskId", FLAT_TASK))
trial = self.storage.create_trial(session["id"], 0, 0, 1, reward, None, "trial")
return self.storage.save_preset(session["id"], session["id"], trial["id"], reward)
def payload(self, preset):
return dict(
taskId=FLAT_TASK,
rewardPresetId=preset["id"],
numEnvs=2,
maxIterations=1,
seed=42,
device="cpu",
)
def assert_rejected(self, preset):
for entry in (self.training.parse_config, self.training.start):
with self.assertRaises(ApiError) as raised:
entry(self.payload(preset))
self.assertEqual(raised.exception.status, 400)
self.assertEqual(self.training.jobs, {})
self.assertIsNone(self.training.lease.public())
def test_obstacle_list_identity_and_cross_task_rejected_before_job_creation(self):
preset = self.save({"taskId": OBSTACLE_TASK})
self.assertEqual(preset["taskId"], OBSTACLE_TASK)
self.assertEqual(self.storage.list_presets(), [preset])
self.assertEqual(self.storage.get_preset(preset["id"]), preset)
self.assertEqual(
self.tuning.preset_config(preset["id"], OBSTACLE_TASK), preset["rewardConfig"]
)
with self.assertRaisesRegex(RewardConfigError, "任务"):
self.tuning.preset_config(preset["id"])
self.assert_rejected(preset)
def test_historical_flat_without_task_id_and_explicit_flat_remain_valid(self):
for config in ({}, {"taskId": FLAT_TASK}):
with self.subTest(config=config):
preset = self.save(config)
self.assertEqual(preset["taskId"], FLAT_TASK)
self.assertEqual(self.tuning.preset_config(preset["id"]), preset["rewardConfig"])
parsed = self.training.parse_config(self.payload(preset))
self.assertEqual(parsed.reward_config, base_configuration())
self.assertIn("--reward-config-json", self.training.command_for(parsed))
self.assertEqual({p["taskId"] for p in self.storage.list_presets()}, {FLAT_TASK})
def test_malformed_preset_and_corrupt_or_missing_source_fail_closed(self):
preset = self.save({})
connection = self.storage.connection()
valid_reward = json.dumps(preset["rewardConfig"])
for bad_reward in (
'{"weights":{},"params":{}}',
"null",
"{broken",
json.dumps(base_configuration(OBSTACLE_TASK)),
):
with self.subTest(reward=bad_reward):
connection.execute(
"UPDATE presets SET reward_config_json=? WHERE id=?", (bad_reward, preset["id"])
)
with self.assertRaises(RewardConfigError):
self.storage.list_presets()
self.assert_rejected(preset)
connection.execute(
"UPDATE presets SET reward_config_json=? WHERE id=?", (valid_reward, preset["id"])
)
for source in (
"null",
"[]",
"{broken",
'{"taskId":null}',
'{"taskId":"unknown"}',
json.dumps({"taskId": OBSTACLE_TASK}),
):
with self.subTest(source=source):
connection.execute(
"UPDATE sessions SET config_json=? WHERE id=?", (source, preset["sessionId"])
)
with self.assertRaises(RewardConfigError):
self.storage.get_preset(preset["id"])
self.assert_rejected(preset)
connection.execute("DELETE FROM sessions WHERE id=?", (preset["sessionId"],))
with self.assertRaises(RewardConfigError):
self.storage.list_presets()
self.assert_rejected(preset)
def test_save_and_service_both_validate_complete_task_schema(self):
with self.assertRaises(RewardConfigError):
self.save({}, base_configuration(OBSTACLE_TASK))
self.assertEqual(self.storage.list_presets(), [])
# The service itself must reject partial/malformed configs even from a broken resolver.
self.training.preset_resolver = lambda _id, _task: {"weights": {"pose": 1}, "params": {}}
self.assert_rejected({"id": "f" * 32})
if __name__ == "__main__":
unittest.main()