138 lines
5.8 KiB
Python
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()
|