agent-harness / src /agent_harness /repository.py
cuber12's picture
Publish agent harness research code and paper artifacts
d61821a verified
Raw
History Blame Contribute Delete
4.5 kB
"""Read immutable source snapshots directly from a local Git object database."""
from __future__ import annotations
from dataclasses import dataclass
from hashlib import sha256
from pathlib import Path
import subprocess
from typing import Iterator
class GitSnapshotError(RuntimeError):
"""Raised when a repository or pinned commit cannot be read safely."""
@dataclass(frozen=True, slots=True)
class SourceFile:
path: str
text: str
@dataclass(frozen=True, slots=True)
class SourceChunk:
chunk_id: str
path: str
line_start: int
line_end: int
text: str
class GitSnapshot:
def __init__(self, repository: Path):
self.repository = repository.resolve()
if not (self.repository / ".git").exists():
raise GitSnapshotError(f"Not a Git checkout: {self.repository}")
def _git(self, arguments: list[str], timeout_seconds: int = 60) -> bytes:
result = subprocess.run(
["git", *arguments],
cwd=self.repository,
check=False,
capture_output=True,
timeout=timeout_seconds,
)
if result.returncode != 0:
detail = result.stderr.decode("utf-8", errors="replace").strip()
raise GitSnapshotError(f"git {' '.join(arguments)} failed: {detail}")
return result.stdout
def verify_commit(self, commit: str) -> None:
self._git(["cat-file", "-e", f"{commit}^{{commit}}"])
def tracked_paths(self, commit: str, suffixes: tuple[str, ...] = (".go",)) -> tuple[str, ...]:
self.verify_commit(commit)
raw = self._git(["ls-tree", "-r", "--name-only", "-z", commit])
paths = tuple(
path
for path in raw.decode("utf-8", errors="surrogateescape").split("\0")
if path and path.endswith(suffixes)
)
return tuple(sorted(paths))
def read_file(self, commit: str, path: str) -> SourceFile:
if Path(path).is_absolute() or ".." in Path(path).parts:
raise GitSnapshotError(f"Unsafe repository path: {path}")
raw = self._git(["show", f"{commit}:{path}"])
return SourceFile(path=path, text=raw.decode("utf-8", errors="replace"))
def iter_files(
self,
commit: str,
suffixes: tuple[str, ...] = (".go",),
) -> Iterator[SourceFile]:
for path in self.tracked_paths(commit, suffixes):
yield self.read_file(commit, path)
def chunk_file(
source: SourceFile,
chunk_lines: int,
overlap_lines: int,
char_limit: int,
) -> tuple[SourceChunk, ...]:
if chunk_lines <= 0 or char_limit <= 0 or not 0 <= overlap_lines < chunk_lines:
raise ValueError("invalid line chunking policy")
lines = source.text.splitlines(keepends=True)
if not lines:
return ()
stride = chunk_lines - overlap_lines
chunks: list[SourceChunk] = []
for start in range(0, len(lines), stride):
selected = lines[start : start + chunk_lines]
if not selected:
break
block = "".join(selected)
block_line_start = start + 1
char_stride = max(char_limit - 512, 1)
for char_start in range(0, len(block), char_stride):
text = block[char_start : char_start + char_limit]
if not text:
break
line_start = block_line_start + block[:char_start].count("\n")
line_end = line_start + text.count("\n")
if text and not text.endswith("\n"):
line_end += 1
identity = sha256(
f"{source.path}\0{line_start}\0{line_end}\0{text}".encode("utf-8")
).hexdigest()
chunks.append(
SourceChunk(
chunk_id=identity,
path=source.path,
line_start=line_start,
line_end=line_end,
text=text,
)
)
if char_start + char_limit >= len(block):
break
if line_end == len(lines):
break
return tuple(chunks)
def chunk_snapshot(
snapshot: GitSnapshot,
commit: str,
chunk_lines: int,
overlap_lines: int,
char_limit: int,
suffixes: tuple[str, ...] = (".go",),
) -> tuple[SourceChunk, ...]:
chunks: list[SourceChunk] = []
for source in snapshot.iter_files(commit, suffixes):
chunks.extend(chunk_file(source, chunk_lines, overlap_lines, char_limit))
return tuple(chunks)