Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
71 changes: 57 additions & 14 deletions src/openenv/core/harness/collect.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand All @@ -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()``.
Expand Down
113 changes: 113 additions & 0 deletions tests/core/test_harness_collect.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Minor / non-blocking: this test doesn't quite verify its namesake. The fixture is written with a trailing newline ('...ep1"} '), so _append_separator takes the 1-byte fast path (handle.read(1) == b" " → early return) and never enters the chunked backward scan. Separately, the implementation reads through the open file handle (handle.read(...)), not Path.read_bytes, so patching read_bytes doesn't actually guard against a full read either. test_append_scans_a_long_final_record_across_chunks (no trailing newline) is the test that truly exercises the scan. If you want to lock in the "bounded read" guarantee here, drop the trailing newline and assert on bytes read through the handle.

):
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()."""
Expand Down Expand Up @@ -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(),
Expand Down