"""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=""))