diff --git a/xr_rm_teleop/test/test_act_episode_recorder.py b/xr_rm_teleop/test/test_act_episode_recorder.py index ff0ae29..948ce01 100644 --- a/xr_rm_teleop/test/test_act_episode_recorder.py +++ b/xr_rm_teleop/test/test_act_episode_recorder.py @@ -744,13 +744,7 @@ def _push_recording_frames(recorder, control_ns, frame_number): ) -@requires_h5py -def test_end_to_end_fake_episode_saves_and_returns_idle(tmp_path): - recorder = _recorder_for_test(tmp_path) - recorder._handle_right_b(grip=False) - - assert recorder.state is RecordingState.ARMED - +def _finish_fake_episode(recorder): start_ns = recorder._now_ns() frame_number = 1001 for offset, seq in enumerate(range(100, 107)): @@ -768,12 +762,39 @@ def test_end_to_end_fake_episode_saves_and_returns_idle(tmp_path): _push_recording_frames(recorder, final_ns, frame_number) recorder._on_control_sample(_control_message(108, final_ns, grip=False)) + +@requires_h5py +def test_end_to_end_fake_episode_saves_and_returns_idle(tmp_path): + recorder = _recorder_for_test(tmp_path) + recorder._handle_right_b(grip=False) + + assert recorder.state is RecordingState.ARMED + + _finish_fake_episode(recorder) + assert (tmp_path / "tomato_pick" / "episode_0.hdf5").is_file() assert recorder.state is RecordingState.IDLE assert "SAVING" in recorder._status_pub.messages assert recorder._status_pub.messages[-2:] == ["SAVED", "IDLE"] +@requires_h5py +def test_final_publish_rejects_allocated_episode_number_conflict(tmp_path): + recorder = _recorder_for_test(tmp_path) + recorder._handle_right_b(grip=False) + existing = recorder._task_dir / "episode_0.hdf5" + existing.write_bytes(b"existing") + + _finish_fake_episode(recorder) + + assert existing.read_bytes() == b"existing" + assert not (recorder._task_dir / "episode_1.hdf5").exists() + rejected = list((recorder._task_dir / "rejected").glob("*.hdf5")) + assert len(rejected) == 1 + with h5py.File(rejected[0], "r") as root: + assert root.attrs["reject_reason"] == "episode_number_conflict" + + @requires_h5py def test_discard_and_interrupt_only_process_current_partial(tmp_path): recorder = _recorder_for_test(tmp_path) diff --git a/xr_rm_teleop/xr_rm_teleop/act_episode_recorder.py b/xr_rm_teleop/xr_rm_teleop/act_episode_recorder.py index aa72e56..dd76db8 100644 --- a/xr_rm_teleop/xr_rm_teleop/act_episode_recorder.py +++ b/xr_rm_teleop/xr_rm_teleop/act_episode_recorder.py @@ -1099,6 +1099,7 @@ class ActEpisodeRecorder(Node): self._writer: EpisodeWriter | None = None self._store: EpisodeStore | None = None self._partial_path: Path | None = None + self._episode_index: int | None = None self._camera_baselines: tuple[CameraStats, CameraStats] | None = None self._saving_deadline_ns: int | None = None @@ -1285,6 +1286,7 @@ class ActEpisodeRecorder(Node): self._publish_state(RecordingState.IDLE, "disk_write_error") return self._store = store + self._episode_index = episode_index self._writer = EpisodeWriter( store, queue_size=self._writer_queue_size, @@ -1533,6 +1535,7 @@ class ActEpisodeRecorder(Node): def _complete_save(self) -> None: assert self._store is not None assert self._partial_path is not None + assert self._episode_index is not None partial = self._partial_path self._store.close() report = validate_episode(partial, self._quality_limits) @@ -1546,8 +1549,9 @@ class ActEpisodeRecorder(Node): root.attrs["interrupted"] = np.bool_(False) for name, value in report.metrics.items(): root.attrs[name] = value - episode_index = next_episode_index(self._task_dir) - destination = self._task_dir / f"episode_{episode_index}.hdf5" + destination = ( + self._task_dir / f"episode_{self._episode_index}.hdf5" + ) try: publish_without_overwrite(partial, destination) except FileExistsError: @@ -1625,6 +1629,7 @@ class ActEpisodeRecorder(Node): self._writer = None self._store = None self._partial_path = None + self._episode_index = None self._camera_baselines = None self._saving_deadline_ns = None self._publish_state(RecordingState.IDLE)