Files
2026-09-19 12:54:45 +08:00

95 lines
2.8 KiB
Python

"""Unidiff extraction with before/after context."""
from dataclasses import dataclass
from pathlib import Path
from typing import Optional
import git
from git import Repo
@dataclass
class DiffBundle:
repo: str
commit_sha: str
diff: str
before_files: dict
after_files: dict
def _detect_language(filename: str) -> str:
ext = Path(filename).suffix.lower()
mapping = {
".py": "python",
".java": "java",
".js": "javascript",
".ts": "javascript",
".jsx": "javascript",
".tsx": "javascript",
}
return mapping.get(ext, "unknown")
def extract_diff_bundle(
repo_path: str,
commit_sha: str,
parent_index: int = 0,
) -> DiffBundle:
"""Extract unified diff and before/after file snapshots for a commit."""
repo = Repo(str(repo_path))
commit = repo.commit(commit_sha)
parents = commit.parents
if parents:
base = parents[parent_index]
diff = base.diff(commit, create_patch=True, unified=3)
else:
# Initial commit: diff against empty tree
diff = commit.diff(git.Git(repo).hash_object("/dev/null", t=None), create_patch=True, unified=3)
diff_text = "\n".join(d.diff.decode("utf-8", errors="replace") for d in diff if d.diff)
before_files = {}
after_files = {}
for d in diff:
a_path = d.a_path or d.b_path
b_path = d.b_path or d.a_path
if a_path and d.a_blob:
try:
before_files[a_path] = d.a_blob.data_stream.read().decode("utf-8", errors="replace")
except Exception:
before_files[a_path] = ""
if b_path and d.b_blob:
try:
after_files[b_path] = d.b_blob.data_stream.read().decode("utf-8", errors="replace")
except Exception:
after_files[b_path] = ""
return DiffBundle(
repo=Path(repo_path).name,
commit_sha=commit_sha,
diff=diff_text,
before_files=before_files,
after_files=after_files,
)
def bundle_to_sample_dict(bundle: DiffBundle, primary_language: Optional[str] = None) -> dict:
"""Convert a DiffBundle into a dict matching the Sample schema."""
if not primary_language:
# Infer from changed files
exts = {Path(p).suffix.lower() for p in bundle.after_files or bundle.before_files}
for ext, lang in {".py": "python", ".java": "java", ".js": "javascript"}.items():
if ext in exts:
primary_language = lang
break
primary_language = primary_language or "unknown"
return {
"repo": bundle.repo,
"commit_sha": bundle.commit_sha,
"language": primary_language,
"diff": bundle.diff,
"before_context": bundle.before_files,
"after_context": bundle.after_files,
}