129 lines
4.9 KiB
Python
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=""))
|