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

129 lines
4.9 KiB
Python

"""Dataset builder: apply mutation rules to samples and persist Ground Truth."""
import random
from collections import Counter
from pathlib import Path
from typing import List, Optional
from sqlalchemy.orm import Session
from app.dataset.diff_extractor import _detect_language, extract_diff_bundle, bundle_to_sample_dict
from app.dataset.git_parser import select_candidates
from app.dataset.rules.registry import get_registry
from app.models import Sample, Defect
def build_dataset_from_repo(
db: Session,
repo_path: str | Path,
count: int = 12,
languages: Optional[List[str]] = None,
) -> List[Sample]:
"""Build a dataset from a Git repository with injected defects.
Only candidates that actually receive a mutation are kept: every sample
in the dataset must carry Ground Truth, so candidates whose changed files
match no rule are skipped and the scan continues until ``count`` mutated
samples have been collected.
"""
languages = languages or ["python", "java", "javascript"]
registry = get_registry()
candidates = select_candidates(repo_path, count=max(count * 5, 30), languages=languages)
samples = []
type_counts: Counter = Counter()
for candidate in candidates:
if len(samples) >= count:
break
bundle = extract_diff_bundle(repo_path, candidate.sha)
primary_lang = _infer_primary_language(bundle.after_files or bundle.before_files, languages)
sample_dict = bundle_to_sample_dict(bundle, primary_language=primary_lang)
sample = Sample(**sample_dict)
# Collect every (file, rule) mutation opportunity for this commit.
# Rules parse and mutate a single source file, so each changed file
# of the primary language is tried individually; line numbers are
# shifted into the coordinate space of the joined source stored on
# the sample. Among all opportunities, prefer the defect type that is
# currently least represented in this build, so frequent patterns
# (e.g. `&&` swaps) do not dominate the dataset.
after_source = _join_source(bundle.after_files)
matches = []
if primary_lang in registry.all_rules():
rules = list(registry.rules_for(primary_lang))
random.Random(candidate.sha).shuffle(rules)
line_offset = 0
for path, content in (bundle.after_files or {}).items():
if _detect_language(path) != primary_lang:
line_offset += len(content.splitlines()) + 2
continue
for rule in rules:
m = rule.detect_and_mutate(content, filename=path)
if m:
m.line_start += line_offset
m.line_end += line_offset
matches.append((path, m))
line_offset += len(content.splitlines()) + 2
if not matches:
# Skip candidates whose files match no rule: samples without
# Ground Truth are useless for the experiment.
continue
rng = random.Random(candidate.sha)
min_count = min(type_counts[m.defect_type] for _, m in matches)
best = [(p, m) for p, m in matches if type_counts[m.defect_type] == min_count]
path, mutation = rng.choice(best)
type_counts[mutation.defect_type] += 1
files = dict(bundle.after_files)
files[path] = mutation.mutated_source
mutated_joined = _join_source(files)
mutation.description = f"{mutation.description} (file: {path})"
sample.diff = _compute_diff_from_mutated(after_source, mutated_joined)
sample.after_context = {"mutated": mutated_joined}
defect = Defect(
sample=sample,
defect_type=mutation.defect_type,
language=mutation.language,
line_start=mutation.line_start,
line_end=mutation.line_end,
description=mutation.description,
reference_fix=mutation.reference_fix,
)
sample.defects.append(defect)
db.add(sample)
samples.append(sample)
db.commit()
for sample in samples:
db.refresh(sample)
return samples
def _infer_primary_language(files: dict, languages: List[str]) -> str:
from app.dataset.diff_extractor import _detect_language
counts = {}
for path in files:
lang = _detect_language(path)
if lang in languages:
counts[lang] = counts.get(lang, 0) + 1
if counts:
return max(counts, key=counts.get)
return languages[0]
def _join_source(files: dict) -> str:
return "\n\n".join(files.values())
def _compute_diff_from_mutated(original: str, mutated: str) -> str:
"""Produce a simple unified-diff-like string from original and mutated."""
import difflib
orig_lines = original.splitlines(keepends=True)
mut_lines = mutated.splitlines(keepends=True)
return "".join(difflib.unified_diff(orig_lines, mut_lines, lineterm=""))