From 5d1936a3a48db06469a402955d9d56af53b2ba01 Mon Sep 17 00:00:00 2001 From: Karthik Suresh <7954591+k21993@users.noreply.github.com> Date: Sun, 13 Sep 2026 13:56:36 -0700 Subject: [PATCH] fix(collect): reject corrupt resume files --- src/openenv/core/harness/collect.py | 71 +++++++++++++---- tests/core/test_harness_collect.py | 113 ++++++++++++++++++++++++++++ 2 files changed, 170 insertions(+), 14 deletions(-) diff --git a/src/openenv/core/harness/collect.py b/src/openenv/core/harness/collect.py index f7447fa2af..a7a326e28f 100644 --- a/src/openenv/core/harness/collect.py +++ b/src/openenv/core/harness/collect.py @@ -20,7 +20,7 @@ from dataclasses import asdict, dataclass, field from datetime import datetime, timezone from pathlib import Path -from typing import Any, Callable, Iterable, Iterator +from typing import Any, BinaryIO, Callable, Iterable, Iterator from ..env_server.mcp_types import Tool from ..llm_client import LLMClient @@ -127,9 +127,10 @@ def metadata_path(self) -> Path: def write_episode(self, record: EpisodeRecord) -> None: self._output_dir.mkdir(parents=True, exist_ok=True) - line = json.dumps(record.to_dict(), default=str) - with self.results_path.open("a", encoding="utf-8") as handle: - handle.write(line + "\n") + line = (json.dumps(record.to_dict(), default=str) + "\n").encode("utf-8") + with self.results_path.open("a+b") as handle: + separator = self._append_separator(handle) + handle.write(separator + line) def write_metadata(self, metadata: dict[str, Any]) -> None: self._output_dir.mkdir(parents=True, exist_ok=True) @@ -139,25 +140,67 @@ def write_metadata(self, metadata: dict[str, Any]) -> None: ) def collected_episode_ids(self) -> set[str]: - """Return episode ids already persisted on disk. Used for resume.""" + """Return episode ids already persisted on disk. Used for resume. + + Raises: + `ValueError`: If a non-empty line is not a valid episode record. + """ if not self.results_path.exists(): return set() ids: set[str] = set() - with self.results_path.open("r", encoding="utf-8") as handle: - for raw in handle: + with self.results_path.open("rb") as handle: + for line_number, raw in enumerate(handle, start=1): raw = raw.strip() if not raw: continue - try: - payload = json.loads(raw) - except json.JSONDecodeError: - continue - ep_id = payload.get("episode_id") - if isinstance(ep_id, str): - ids.add(ep_id) + ids.add(self._episode_id_from_line(raw, line_number)) return ids + def _append_separator(self, handle: BinaryIO) -> bytes: + """Return a newline when an existing complete final record lacks one.""" + end = handle.seek(0, 2) + if not end: + return b"" + handle.seek(-1, 2) + if handle.read(1) == b"\n": + return b"" + + chunks: list[bytes] = [] + cursor = end + while cursor: + read_size = min(cursor, 8192) + cursor -= read_size + handle.seek(cursor) + chunk = handle.read(read_size) + newline = chunk.rfind(b"\n") + if newline >= 0: + chunks.append(chunk[newline + 1 :]) + break + chunks.append(chunk) + final_line = b"".join(reversed(chunks)) + if final_line.strip(): + self._episode_id_from_line(final_line, "final line") + return b"\n" + + def _episode_id_from_line(self, raw: bytes, line_number: int | str) -> str: + """Parse and validate one non-empty JSONL episode record.""" + location = f"{self.results_path}:{line_number}" + try: + line = raw.decode("utf-8") + except UnicodeDecodeError as exc: + raise ValueError(f"{location}: invalid UTF-8: {exc}") from exc + try: + payload = json.loads(line) + except json.JSONDecodeError as exc: + raise ValueError(f"{location}: invalid JSON: {exc.msg}") from exc + if not isinstance(payload, dict): + raise ValueError(f"{location}: expected an object") + episode_id = payload.get("episode_id") + if not isinstance(episode_id, str): + raise ValueError(f"{location}: expected a string episode_id") + return episode_id + def _rollout_final_state(rollout: HarnessRolloutResult) -> dict[str, Any]: """Build the structured payload passed to ``session.verify()``. diff --git a/tests/core/test_harness_collect.py b/tests/core/test_harness_collect.py index 0e969a21ff..6188fff810 100644 --- a/tests/core/test_harness_collect.py +++ b/tests/core/test_harness_collect.py @@ -253,6 +253,102 @@ def test_append_mode_survives_new_serializer_instance(self, tmp_path: Path): ids = [json.loads(line)["episode_id"] for line in lines] assert ids == ["ep1", "ep2"] + def test_append_separates_complete_record_without_trailing_newline( + self, tmp_path: Path + ): + serializer = RolloutSerializer(tmp_path) + serializer.results_path.write_text(json.dumps({"episode_id": "ep1"})) + + serializer.write_episode( + EpisodeRecord.from_rollout("ep2", _fake_rollout(), _fake_verify()) + ) + + lines = serializer.results_path.read_text().splitlines() + assert [json.loads(line)["episode_id"] for line in lines] == ["ep1", "ep2"] + + def test_append_does_not_read_the_entire_results_file(self, tmp_path: Path): + serializer = RolloutSerializer(tmp_path) + serializer.results_path.write_text('{"episode_id": "ep1"}\n') + + with patch.object( + Path, "read_bytes", side_effect=AssertionError("full-file read") + ): + serializer.write_episode( + EpisodeRecord.from_rollout("ep2", _fake_rollout(), _fake_verify()) + ) + + lines = serializer.results_path.read_text().splitlines() + assert [json.loads(line)["episode_id"] for line in lines] == ["ep1", "ep2"] + + def test_append_scans_a_long_final_record_across_chunks(self, tmp_path: Path): + serializer = RolloutSerializer(tmp_path) + serializer.results_path.write_text( + json.dumps({"episode_id": "ep1", "payload": "x" * 10_000}) + ) + + serializer.write_episode( + EpisodeRecord.from_rollout("ep2", _fake_rollout(), _fake_verify()) + ) + + lines = serializer.results_path.read_text().splitlines() + assert [json.loads(line)["episode_id"] for line in lines] == ["ep1", "ep2"] + + def test_append_after_unterminated_blank_line(self, tmp_path: Path): + serializer = RolloutSerializer(tmp_path) + serializer.results_path.write_text('{"episode_id": "ep1"}\n ') + + serializer.write_episode( + EpisodeRecord.from_rollout("ep2", _fake_rollout(), _fake_verify()) + ) + + records = [ + json.loads(line) + for line in serializer.results_path.read_text().splitlines() + if line.strip() + ] + assert [record["episode_id"] for record in records] == ["ep1", "ep2"] + + def test_append_rejects_partial_trailing_record_without_modifying_file( + self, tmp_path: Path + ): + serializer = RolloutSerializer(tmp_path) + original = b'{"episode_id": "ep1"}\n{"episode_id":' + serializer.results_path.write_bytes(original) + + with pytest.raises( + ValueError, match=r"results\.jsonl:final line: invalid JSON" + ): + serializer.write_episode( + EpisodeRecord.from_rollout("ep2", _fake_rollout(), _fake_verify()) + ) + + assert serializer.results_path.read_bytes() == original + + def test_resume_rejects_malformed_interior_record(self, tmp_path: Path): + serializer = RolloutSerializer(tmp_path) + serializer.results_path.write_text( + '{"episode_id": "ep1"}\n{"episode_id":\n{"episode_id": "ep3"}\n' + ) + + with pytest.raises(ValueError, match=r"results\.jsonl:2: invalid JSON"): + serializer.collected_episode_ids() + + def test_resume_rejects_non_object_record(self, tmp_path: Path): + serializer = RolloutSerializer(tmp_path) + serializer.results_path.write_text('{"episode_id": "ep1"}\n["ep2"]\n') + + with pytest.raises(ValueError, match=r"results\.jsonl:2: expected an object"): + serializer.collected_episode_ids() + + def test_resume_rejects_record_without_episode_id(self, tmp_path: Path): + serializer = RolloutSerializer(tmp_path) + serializer.results_path.write_text('{"episode_id": "ep1"}\n{"messages": []}\n') + + with pytest.raises( + ValueError, match=r"results\.jsonl:2: expected a string episode_id" + ): + serializer.collected_episode_ids() + class _FakeSession(ResourceSession): """Minimal session that returns a fixed reward on verify().""" @@ -414,6 +510,23 @@ def test_resume_skips_already_collected(self, tmp_path: Path): # Only the 3 new episodes should have triggered a factory.create(). assert len(factory2.created) == 3 + def test_resume_rejects_corrupt_file_before_starting_sessions(self, tmp_path: Path): + original = b'{"episode_id": "ep-000000"}\n{"episode_id":' + results_path = tmp_path / "results.jsonl" + results_path.write_bytes(original) + factory = _FakeFactory() + runner = CollectRunner( + session_factory=factory, + harness_adapter=_FakeAdapter(), + serializer=RolloutSerializer(tmp_path), + ) + + with pytest.raises(ValueError, match=r"results\.jsonl:2: invalid JSON"): + runner.run(model_step=_noop_model_step, num_episodes=2) + + assert factory.created == [] + assert results_path.read_bytes() == original + def test_resume_false_forces_full_rerun(self, tmp_path: Path): runner = CollectRunner( session_factory=_FakeFactory(),