first commit
This commit is contained in:
@@ -0,0 +1,128 @@
|
||||
"""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=""))
|
||||
Reference in New Issue
Block a user