"""Registered source authority, immutable snapshots, and same-trial continuation.""" import hashlib import json import os import sys import tempfile import threading import unittest from pathlib import Path from unittest.mock import patch sys.path.insert(0, str(Path(__file__).resolve().parents[1])) from pretrained_sources import PretrainedSources, SourceError, regular_bytes from server import ApiError, TrainingManager from test_tuning_manager import BASE_METRICS, FakeAdvisor from tuning.manager import TuningError, TuningManager from tuning.process import GpuLease from tuning.schema import BASE_REWARD_CONFIGURATION, RewardConfigError, validate_proposal def fake_validate(_self, directory, checkpoint, *_args): artifacts = {} for key, relative in { "checkpoint": checkpoint, "onnx": "policy.onnx", "env": "params/env.yaml", "agent": "params/agent.yaml", }.items(): data = (directory / relative).read_bytes() artifacts[key] = { "name": Path(relative).name, "sha256": hashlib.sha256(data).hexdigest(), "bytes": len(data), } return { "source_id": hashlib.sha256(json.dumps(artifacts, sort_keys=True).encode()).hexdigest(), "artifacts": artifacts, "source_iteration": 10000, } class RegisteredSourcesTest(unittest.TestCase): def setUp(self): self.temp = tempfile.TemporaryDirectory() self.root = Path(self.temp.name) self.original = self.root / "original" (self.original / "params").mkdir(parents=True) for relative in ("model_10000.pt", "policy.onnx", "params/env.yaml", "params/agent.yaml"): (self.original / relative).write_bytes(relative.encode()) self.config = self.root / "sources.json" self.config.write_text( json.dumps( { "allowedRoots": [str(self.original)], "sources": [ { "id": "base", "label": "基础行走", "checkpoint": str(self.original / "model_10000.pt"), "onnx": str(self.original / "policy.onnx"), } ], } ) ) self.validator = patch.object(PretrainedSources, "_validate", fake_validate) self.validator.start() self.registry = PretrainedSources( self.config, self.root / "snapshots", sys.executable, Path(__file__).resolve().parents[1] / "rl", ) def tearDown(self): self.validator.stop() self.temp.cleanup() def test_registered_snapshot_never_rereads_mutated_original_and_checks_sha(self): bound = self.registry.bind(self.registry.catalog()[0]["id"], "Unitree-Go2-Flat") original_sha = bound["manifest"]["artifacts"]["checkpoint"]["sha256"] (self.original / "model_10000.pt").write_bytes(b"changed after registration") self.assertEqual( self.registry.bind(self.registry.catalog()[0]["id"], "Unitree-Go2-Flat"), bound ) args = self.registry.arguments(bound) checkpoint = Path(args[args.index("--pretrained-checkpoint") + 1]) self.assertEqual(hashlib.sha256(checkpoint.read_bytes()).hexdigest(), original_sha) restored = PretrainedSources(None, self.root / "snapshots", sys.executable, self.root) self.assertEqual(restored.arguments(bound), args) re_registered = PretrainedSources( self.config, self.root / "snapshots", sys.executable, self.root ) with self.assertRaisesRegex(SourceError, "已变化"): re_registered.bind(bound["sourceId"], "Unitree-Go2-Flat") self.assertEqual(restored.arguments(bound), args) checkpoint.chmod(0o644) checkpoint.write_bytes(b"tamper") with self.assertRaisesRegex(SourceError, "SHA"): restored.arguments(bound) checkpoint.parent.chmod(0o755) checkpoint.unlink() with self.assertRaises(SourceError): restored.verify(bound) def test_rejects_path_symlink_fifo_size_and_invalid_registry(self): with self.assertRaises(SourceError): self.registry.bind(str(self.original / "model_10000.pt"), "Unitree-Go2-Flat") with self.assertRaisesRegex(SourceError, "Rough"): self.registry.bind(self.registry.catalog()[0]["id"], "Unitree-Go2-Rough") (self.original / "link.pt").symlink_to(self.original / "model_10000.pt") with self.assertRaises(SourceError): regular_bytes(self.original / "link.pt", self.original, 256) with self.assertRaises(SourceError): regular_bytes(self.original / ".." / "sources.json", self.original, 256) with self.assertRaises(SourceError): regular_bytes(self.original / "policy.onnx", self.original, 1) os.mkfifo(self.original / "pipe") with self.assertRaises(SourceError): regular_bytes(self.original / "pipe", self.original, 256) self.config.write_text('{"allowedRoots": [], "sources": [], "python": "evil"}') with self.assertRaises(SourceError): PretrainedSources(self.config, self.root / "other", sys.executable, self.root) def test_missing_checkpoint_is_visible_failure_not_random_fallback(self): (self.original / "model_10000.pt").unlink() registry = PretrainedSources(self.config, self.root / "other", sys.executable, self.root) self.assertFalse(registry.catalog()[0]["ready"]) with self.assertRaisesRegex(SourceError, "pt"): registry.bind("base", "Unitree-Go2-Flat") def test_training_request_only_id_and_public_identity(self): manager = TrainingManager( Path(__file__).resolve().parents[1] / "rl", sys.executable, ("Unitree-Go2-Flat", "Unitree-Go2-Rough"), check_environment=False, sources=self.registry, ) payload = { "taskId": "Unitree-Go2-Flat", "numEnvs": 4, "maxIterations": 1, "seed": 42, "device": "cpu", "pretrainedSourceId": self.registry.catalog()[0]["id"], } config = manager.parse_config(payload) self.assertEqual( config.pretrained, self.registry.bind(self.registry.catalog()[0]["id"], "Unitree-Go2-Flat"), ) command = manager.command_for(config) self.assertIn("--pretrained-source-id", command) self.assertNotIn("--resume-checkpoint", command) self.assertNotIn(str(self.original), json.dumps(manager.health()["pretrainedSources"])) for key in ("pretrainedCheckpoint", "allowedRoots", "pretrained", "resumeCheckpoint"): with self.assertRaises(ApiError): manager.parse_config({**payload, key: "/etc/passwd"}) with self.assertRaises(ApiError): manager.parse_config({**payload, "pretrainedSourceId": None}) with self.assertRaises(ApiError): manager.parse_config({**payload, "taskId": "Unitree-Go2-Rough"}) old = manager.parse_config({k: v for k, v in payload.items() if k != "pretrainedSourceId"}) self.assertIsNone(old.pretrained) self.assertNotIn("--pretrained-checkpoint", manager.command_for(old)) def tuning(self): return TuningManager( Path(__file__).resolve().parents[1] / "rl", sys.executable, self.root / "tuning", GpuLease(), advisor=FakeAdvisor(), sources=self.registry, ) def session(self, manager, mode="approval"): _, config, objective, fallback = manager.parse_create( { "mode": mode, "pretrainedSourceId": self.registry.catalog()[0]["id"], "trialCount": 2, "numEnvs": 4, "initialIterations": 1, "middleIterations": 2, "finalIterations": 3, } ) return manager.storage.create_session(mode, config, objective, fallback) def test_baseline_new_trial_warmstart_and_rung_resumes_only_own_checkpoint(self): manager = self.tuning() session = self.session(manager) manager.cancel_events[session["id"]] = threading.Event() commands = [] def run(_session, command, _cwd, _env, _log): commands.append(command) if "--output-dir" in command: directory = Path(command[command.index("--output-dir") + 1]) (directory / "model_0.pt").write_bytes(b"trial checkpoint") (directory / "policy.onnx").write_bytes(b"trial policy") (directory / "initialization.json").write_text( json.dumps(session["config"]["pretrained"]["manifest"]) ) else: Path(command[command.index("--output") + 1]).write_text( json.dumps({"metrics": BASE_METRICS}) ) return 0 with patch.object(manager, "_run_command", run): parent = None for number, rung in ((0, 0), (1, 0), (1, 1)): trial = manager.storage.create_trial( session["id"], number, rung, rung + 1, BASE_REWARD_CONFIGURATION, None, f"trial-{number}-{rung}", ) result = manager._execute_trial(session, trial, parent if rung else None) parent = manager._session_root(session["id"]) / result["checkpointPath"] train = [c for c in commands if "--output-dir" in c] self.assertEqual( train[0][train[0].index("--pretrained-source-id") + 1], train[1][train[1].index("--pretrained-source-id") + 1], ) self.assertIn("--resume-checkpoint", train[2]) self.assertNotIn("--pretrained-checkpoint", train[2]) self.assertIn("trial-1-0/model_0.pt", train[2][-1]) self.assertEqual(len([c for c in commands if "--steps-per-seed=1000" in c]), 3) with self.assertRaises(RewardConfigError): validate_proposal({"pretrainedSourceId": "other"}, BASE_REWARD_CONFIGURATION) def test_restart_states_source_constraints_and_explicit_recovery(self): manager = self.tuning() sessions = [] for state in ("queued", "paused", "awaiting_approval", "succeeded"): session = self.session(manager) trial = manager.storage.create_trial( session["id"], 0, 0, 1, BASE_REWARD_CONFIGURATION, None, "baseline" ) manager.storage.update_trial( trial["id"], state="completed", evaluation={"metrics": BASE_METRICS}, eligible=True, score=0, ) proposal = manager.storage.create_proposal( session["id"], trial["id"], {"weights": {"pose": 1.1}}, "proposal", {}, 0.8 ) manager.storage.replace_constraints( session["id"], 0, {"weights.pose": {"kind": "range", "min": 0, "max": 2}} ) manager.storage.update_session(session["id"], state=state) sessions.append((session, proposal, state)) restored = self.tuning() # Actual SQLite reopening, no worker/API calls on construction. self.assertFalse(restored.workers) for session, _proposal, previous in sessions: value = restored.detail(session["id"]) self.assertEqual( value["state"], "succeeded" if previous == "succeeded" else "interrupted" ) self.assertEqual(value["config"]["pretrained"], session["config"]["pretrained"]) self.assertEqual(value["control"]["constraintsRevision"], 1) self.assertEqual(value["mode"], "approval") self.assertEqual(value["proposals"][0]["state"], "pending") session, proposal, _ = sessions[0] # Stop before dispatch: recovery must invalidate pending first, never execute it. with ( patch.object(restored, "_wait_for_dispatch", side_effect=RuntimeError("test boundary")), patch.object(restored.advisor, "propose") as advisor, ): restored._run_session(session["id"], True, threading.Event()) advisor.assert_not_called() self.assertEqual(restored.storage.get_proposal(proposal["id"])["state"], "rejected") self.assertIn( "recovery_invalidated", restored.storage.get_proposal(proposal["id"])["feedback"] ) self.assertEqual(len(restored.storage.list_trials(session["id"])), 1) session = sessions[1][0] with patch.object(restored, "_start_worker") as start: restored.resume(session["id"]) with self.assertRaises(TuningError): restored.resume(session["id"]) start.assert_called_once() session = sessions[2][0] snapshot = self.registry.verify(session["config"]["pretrained"]) file = snapshot / "policy.onnx" file.chmod(0o644) file.write_bytes(b"corrupt") with patch.object(restored, "_start_worker") as start: with self.assertRaises(SourceError): restored.resume(session["id"]) start.assert_not_called() self.assertEqual(restored.storage.get_session(session["id"])["state"], "interrupted") if __name__ == "__main__": unittest.main()