first commit
This commit is contained in:
@@ -0,0 +1,80 @@
|
||||
"""Matplotlib-based chart generation for paper figures."""
|
||||
|
||||
import base64
|
||||
from io import BytesIO
|
||||
from typing import Dict, List
|
||||
|
||||
import matplotlib
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
|
||||
matplotlib.rcParams["font.sans-serif"] = ["DejaVu Sans"]
|
||||
matplotlib.rcParams["axes.unicode_minus"] = False
|
||||
|
||||
|
||||
def _to_base64(fig: matplotlib.figure.Figure) -> str:
|
||||
buf = BytesIO()
|
||||
fig.savefig(buf, format="png", dpi=150, bbox_inches="tight")
|
||||
buf.seek(0)
|
||||
return base64.b64encode(buf.read()).decode("utf-8")
|
||||
|
||||
|
||||
def heatmap(data: Dict[str, Dict[str, float]], title: str = "Heatmap") -> str:
|
||||
"""Generate a heatmap from a nested dict (rows × columns)."""
|
||||
rows = list(data.keys())
|
||||
cols = sorted({c for row in data.values() for c in row.keys()})
|
||||
matrix = np.array([[data[row].get(col, 0.0) for col in cols] for row in rows])
|
||||
|
||||
fig, ax = plt.subplots(figsize=(8, 6))
|
||||
im = ax.imshow(matrix, cmap="YlOrRd", aspect="auto")
|
||||
ax.set_xticks(np.arange(len(cols)))
|
||||
ax.set_yticks(np.arange(len(rows)))
|
||||
ax.set_xticklabels(cols)
|
||||
ax.set_yticklabels(rows)
|
||||
ax.set_title(title)
|
||||
|
||||
for i in range(len(rows)):
|
||||
for j in range(len(cols)):
|
||||
text = ax.text(j, i, f"{matrix[i, j]:.2f}", ha="center", va="center", color="black")
|
||||
|
||||
fig.colorbar(im, ax=ax)
|
||||
encoded = _to_base64(fig)
|
||||
plt.close(fig)
|
||||
return encoded
|
||||
|
||||
|
||||
def boxplot(groups: Dict[str, List[float]], title: str = "Boxplot") -> str:
|
||||
fig, ax = plt.subplots(figsize=(8, 6))
|
||||
labels = list(groups.keys())
|
||||
values = [groups[label] for label in labels]
|
||||
ax.boxplot(values)
|
||||
ax.set_xticklabels(labels)
|
||||
ax.set_title(title)
|
||||
ax.set_ylabel("Score")
|
||||
encoded = _to_base64(fig)
|
||||
plt.close(fig)
|
||||
return encoded
|
||||
|
||||
|
||||
def grouped_bar(
|
||||
data: Dict[str, Dict[str, float]],
|
||||
title: str = "Grouped Bar Chart",
|
||||
) -> str:
|
||||
fig, ax = plt.subplots(figsize=(10, 6))
|
||||
categories = list(data.keys())
|
||||
subcategories = sorted({sc for row in data.values() for sc in row.keys()})
|
||||
x = np.arange(len(categories))
|
||||
width = 0.8 / len(subcategories)
|
||||
|
||||
for idx, subcat in enumerate(subcategories):
|
||||
values = [data[cat].get(subcat, 0.0) for cat in categories]
|
||||
ax.bar(x + idx * width, values, width, label=subcat)
|
||||
|
||||
ax.set_xticks(x + width * (len(subcategories) - 1) / 2)
|
||||
ax.set_xticklabels(categories)
|
||||
ax.set_ylabel("Score")
|
||||
ax.set_title(title)
|
||||
ax.legend()
|
||||
encoded = _to_base64(fig)
|
||||
plt.close(fig)
|
||||
return encoded
|
||||
@@ -0,0 +1,111 @@
|
||||
"""LLM-as-judge evaluation of review outputs against Ground Truth.
|
||||
|
||||
Rule-based matching (position hit + type match) cannot score L1 outputs,
|
||||
which deliberately contain no line numbers or defect-type labels. Following
|
||||
the cross-evaluation methodology of Liang et al., a fixed judge model
|
||||
(temperature 0, reasoning disabled) decides whether a review output
|
||||
semantically identifies each injected defect, and counts false alarms.
|
||||
|
||||
The judge verdict is stored alongside rule-based metrics so both remain
|
||||
auditable.
|
||||
"""
|
||||
|
||||
import json
|
||||
import re
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
JUDGE_PROMPT = """You are the judge in an automated code-review experiment. Exactly one defect was deliberately injected into the code under review.
|
||||
|
||||
## Injected defect (Ground Truth)
|
||||
- Type: {defect_type}
|
||||
- Location: lines {line_start}-{line_end} of the presented diff
|
||||
- Description: {description}
|
||||
- Reference fix: {reference_fix}
|
||||
|
||||
## Review output under evaluation
|
||||
\"\"\"
|
||||
{raw_output}
|
||||
\"\"\"
|
||||
|
||||
Tasks:
|
||||
1. Decide whether the review output correctly identifies the injected defect — i.e. it points out essentially the same problem, even if the wording, defect label, or line numbers differ or are absent.
|
||||
2. Count how many DISTINCT additional problems the review claims that are clearly NOT the injected defect (false alarms). Ignore stylistic nits that are part of describing the injected defect.
|
||||
3. If the review output mentions any line number(s) for the defect it identifies, list them; otherwise use null.
|
||||
|
||||
Answer with JSON only, no other text:
|
||||
{{"detected": true or false, "false_alarms": <integer>, "lines_reported": [<integer>, ...] or null, "reason": "<one short sentence>"}}"""
|
||||
|
||||
|
||||
def build_judge_prompt(raw_output: str, gt: Dict[str, Any]) -> str:
|
||||
return JUDGE_PROMPT.format(
|
||||
defect_type=gt["defect_type"],
|
||||
line_start=gt.get("line_start"),
|
||||
line_end=gt.get("line_end"),
|
||||
description=gt.get("description") or "",
|
||||
reference_fix=gt.get("reference_fix") or "",
|
||||
raw_output=(raw_output or "").strip()[:12000],
|
||||
)
|
||||
|
||||
|
||||
def parse_verdict(text: str) -> Optional[Dict[str, Any]]:
|
||||
"""Extract the JSON verdict from the judge's reply."""
|
||||
if not text:
|
||||
return None
|
||||
match = re.search(r"\{.*\}", text, re.DOTALL)
|
||||
if not match:
|
||||
return None
|
||||
try:
|
||||
data = json.loads(match.group(0))
|
||||
except json.JSONDecodeError:
|
||||
return None
|
||||
if "detected" not in data:
|
||||
return None
|
||||
lines = data.get("lines_reported")
|
||||
if isinstance(lines, list):
|
||||
lines = [int(x) for x in lines if isinstance(x, (int, float))]
|
||||
else:
|
||||
lines = None
|
||||
return {
|
||||
"detected": bool(data["detected"]),
|
||||
"false_alarms": int(data.get("false_alarms") or 0),
|
||||
"lines_reported": lines or None,
|
||||
"reason": str(data.get("reason") or ""),
|
||||
}
|
||||
|
||||
|
||||
def verdict_to_metrics(verdict: Dict[str, Any], n_ground_truth: int = 1) -> Dict[str, float]:
|
||||
"""Convert a judge verdict into detection/false-positive rates.
|
||||
|
||||
detection_rate: share of injected defects identified (0 or 1 per run).
|
||||
false_positive_rate: false alarms / (false alarms + hits), matching the
|
||||
thesis definition "proportion of reported problems that hit no injected
|
||||
defect".
|
||||
"""
|
||||
hits = 1 if verdict["detected"] else 0
|
||||
fa = max(0, verdict["false_alarms"])
|
||||
detection_rate = hits / n_ground_truth if n_ground_truth else 0.0
|
||||
total_reported = hits + fa
|
||||
fpr = fa / total_reported if total_reported else 0.0
|
||||
return {
|
||||
"detection_rate": round(detection_rate, 4),
|
||||
"false_positive_rate": round(fpr, 4),
|
||||
}
|
||||
|
||||
|
||||
def coverage_from_verdict(verdict: Dict[str, Any], ground_truth: List[Dict[str, Any]], tolerance: int = 3) -> Optional[float]:
|
||||
"""Line coverage from judge-extracted line numbers.
|
||||
|
||||
Returns None when the review reported no line numbers (expected for L1),
|
||||
so coverage stays NULL instead of polluting aggregates with zeros.
|
||||
"""
|
||||
lines = verdict.get("lines_reported")
|
||||
if not lines:
|
||||
return None
|
||||
targets = [(gt["line_start"], gt.get("line_end") or gt["line_start"]) for gt in ground_truth if gt.get("line_start") is not None]
|
||||
if not targets:
|
||||
return None
|
||||
hits = 0
|
||||
for gt_start, gt_end in targets:
|
||||
if any(gt_start - tolerance <= ln <= gt_end + tolerance for ln in lines):
|
||||
hits += 1
|
||||
return round(hits / len(targets), 4)
|
||||
@@ -0,0 +1,54 @@
|
||||
"""Likert scale scoring storage and aggregation."""
|
||||
|
||||
from typing import Dict, List
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.models import Result
|
||||
|
||||
|
||||
def save_likert_score(db: Session, run_id: str, score: int) -> Result:
|
||||
if not 1 <= score <= 5:
|
||||
raise ValueError("Likert score must be between 1 and 5")
|
||||
result = db.query(Result).filter_by(run_id=run_id).first()
|
||||
if not result:
|
||||
raise ValueError(f"Result not found for run {run_id}")
|
||||
result.likert_score = score
|
||||
db.commit()
|
||||
db.refresh(result)
|
||||
return result
|
||||
|
||||
|
||||
def aggregate_likert_by_model_and_level(db: Session) -> Dict[str, Dict[str, Dict[str, float]]]:
|
||||
"""Aggregate Likert scores by model_id and level.
|
||||
|
||||
Returns mean and frequency distribution per (model, level).
|
||||
"""
|
||||
rows = (
|
||||
db.query(Result, ExperimentRun)
|
||||
.join(ExperimentRun, Result.run_id == ExperimentRun.id)
|
||||
.all()
|
||||
)
|
||||
|
||||
grouped: Dict[str, Dict[str, List[int]]] = {}
|
||||
for result, run in rows:
|
||||
if result.likert_score is None:
|
||||
continue
|
||||
key_model = run.model_id
|
||||
key_level = run.template_version.template.level
|
||||
grouped.setdefault(key_model, {}).setdefault(key_level, []).append(result.likert_score)
|
||||
|
||||
output = {}
|
||||
for model, levels in grouped.items():
|
||||
output[model] = {}
|
||||
for level, scores in levels.items():
|
||||
total = len(scores)
|
||||
output[model][level] = {
|
||||
"mean": round(sum(scores) / total, 2) if total else 0.0,
|
||||
"count": total,
|
||||
"distribution": {i: scores.count(i) for i in range(1, 6)},
|
||||
}
|
||||
return output
|
||||
|
||||
|
||||
from app.models import ExperimentRun # noqa: E402
|
||||
@@ -0,0 +1,129 @@
|
||||
"""Parse model outputs and compare against Ground Truth."""
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
|
||||
@dataclass
|
||||
class Finding:
|
||||
defect_type: str
|
||||
line_start: Optional[int]
|
||||
line_end: Optional[int]
|
||||
description: str
|
||||
|
||||
|
||||
def parse_output(output: str, level: str) -> List[Finding]:
|
||||
"""Parse model output into structured findings.
|
||||
|
||||
Does not guess: if output is empty or unparseable, returns empty list.
|
||||
"""
|
||||
if not output or not output.strip():
|
||||
return []
|
||||
|
||||
findings = []
|
||||
if level == "L1":
|
||||
# Expect bullet list of defect types or short descriptions
|
||||
for line in output.splitlines():
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
if line.startswith(("-", "*", "•", "1.", "2.", "3.")):
|
||||
item = re.sub(r"^[-*•0-9.\s]+", "", line)
|
||||
findings.append(
|
||||
Finding(
|
||||
defect_type=item.split(":", 1)[0].strip(),
|
||||
line_start=None,
|
||||
line_end=None,
|
||||
description=item,
|
||||
)
|
||||
)
|
||||
else:
|
||||
# Parse L2/L3 structured blocks
|
||||
current: Dict[str, str] = {}
|
||||
for raw in output.splitlines():
|
||||
line = raw.strip()
|
||||
if line.startswith(("-", "*", "•")):
|
||||
if current:
|
||||
findings.append(_build_finding(current))
|
||||
current = {}
|
||||
key, _, value = line.lstrip("-*• ").partition(":")
|
||||
current[key.strip().lower()] = value.strip()
|
||||
elif line and current:
|
||||
key, _, value = line.partition(":")
|
||||
current[key.strip().lower()] = value.strip()
|
||||
if current:
|
||||
findings.append(_build_finding(current))
|
||||
|
||||
return findings
|
||||
|
||||
|
||||
def _build_finding(fields: Dict[str, str]) -> Finding:
|
||||
defect_type = fields.get("type", "unknown")
|
||||
lines = fields.get("lines", "")
|
||||
line_start, line_end = None, None
|
||||
if lines:
|
||||
parts = re.split(r"[-,\s]+", lines)
|
||||
try:
|
||||
line_start = int(parts[0])
|
||||
line_end = int(parts[-1]) if len(parts) > 1 else line_start
|
||||
except ValueError:
|
||||
pass
|
||||
description = fields.get("explanation", fields.get("fix", ""))
|
||||
return Finding(defect_type, line_start, line_end, description)
|
||||
|
||||
|
||||
def compare_findings(
|
||||
findings: List[Finding],
|
||||
ground_truth: List[dict],
|
||||
line_tolerance: int = 3,
|
||||
) -> Dict[str, float]:
|
||||
"""Compare parsed findings to Ground Truth defects.
|
||||
|
||||
Returns detection_rate, false_positive_rate, coverage_rate.
|
||||
"""
|
||||
if not ground_truth:
|
||||
return {"detection_rate": 0.0, "false_positive_rate": 0.0, "coverage_rate": 0.0}
|
||||
|
||||
detected = set()
|
||||
false_positives = 0
|
||||
|
||||
for finding in findings:
|
||||
matched = False
|
||||
for gt in ground_truth:
|
||||
type_match = finding.defect_type.lower() in gt["defect_type"].lower() or gt[
|
||||
"defect_type"
|
||||
].lower() in finding.defect_type.lower()
|
||||
line_match = False
|
||||
if finding.line_start is not None and gt.get("line_start") is not None:
|
||||
gt_start = gt["line_start"]
|
||||
gt_end = gt.get("line_end", gt_start)
|
||||
if (
|
||||
min(finding.line_start, finding.line_end or finding.line_start) - line_tolerance
|
||||
<= gt_end
|
||||
and max(finding.line_start, finding.line_end or finding.line_start)
|
||||
+ line_tolerance
|
||||
>= gt_start
|
||||
):
|
||||
line_match = True
|
||||
if type_match or line_match:
|
||||
matched = True
|
||||
detected.add(gt.get("id", id(gt)))
|
||||
break
|
||||
if not matched:
|
||||
false_positives += 1
|
||||
|
||||
tp = len(detected)
|
||||
fp = false_positives
|
||||
fn = len(ground_truth) - tp
|
||||
|
||||
detection_rate = tp / len(ground_truth)
|
||||
precision = tp / (tp + fp) if (tp + fp) > 0 else 0.0
|
||||
false_positive_rate = 1.0 - precision
|
||||
coverage_rate = tp / len(ground_truth)
|
||||
|
||||
return {
|
||||
"detection_rate": round(detection_rate, 4),
|
||||
"false_positive_rate": round(false_positive_rate, 4),
|
||||
"coverage_rate": round(coverage_rate, 4),
|
||||
}
|
||||
@@ -0,0 +1,185 @@
|
||||
"""Analysis service: compute metrics and produce charts/JSON for frontend."""
|
||||
|
||||
from typing import Any, Dict, List
|
||||
from uuid import UUID
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.analysis.charts import boxplot, grouped_bar, heatmap
|
||||
from app.analysis.likert import aggregate_likert_by_model_and_level
|
||||
from app.analysis.parser import compare_findings, parse_output
|
||||
from app.analysis.statistics import anova, descriptive_stats, paired_t_test
|
||||
from app.models import Experiment, ExperimentRun, Result
|
||||
|
||||
|
||||
class AnalysisService:
|
||||
def __init__(self, db: Session):
|
||||
self.db = db
|
||||
|
||||
def get_experiment_results(self, experiment_id: str) -> List[Dict[str, Any]]:
|
||||
runs = (
|
||||
self.db.query(ExperimentRun)
|
||||
.filter_by(experiment_id=UUID(experiment_id))
|
||||
.all()
|
||||
)
|
||||
return [
|
||||
{
|
||||
"run_id": str(run.id),
|
||||
"model_id": run.model_id,
|
||||
"level": run.template_version.template.level,
|
||||
"sample_id": str(run.sample_id),
|
||||
"repeat_index": run.repeat_index,
|
||||
"status": run.status,
|
||||
"result": self._serialize_result(run.result) if run.result else None,
|
||||
}
|
||||
for run in runs
|
||||
]
|
||||
|
||||
def _serialize_result(self, result: Result) -> Dict[str, Any]:
|
||||
return {
|
||||
"raw_output": result.raw_output,
|
||||
"token_usage": result.token_usage,
|
||||
"latency_ms": result.latency_ms,
|
||||
"parsed_findings": result.parsed_findings,
|
||||
"detection_rate": result.detection_rate,
|
||||
"false_positive_rate": result.false_positive_rate,
|
||||
"coverage_rate": result.coverage_rate,
|
||||
"stability_score": result.stability_score,
|
||||
"likert_score": result.likert_score,
|
||||
}
|
||||
|
||||
def compute_metrics_for_run(self, run_id: str) -> Dict[str, Any]:
|
||||
run = self.db.query(ExperimentRun).filter_by(id=UUID(run_id)).first()
|
||||
if not run or not run.result:
|
||||
return {"error": "Run or result not found"}
|
||||
|
||||
level = run.template_version.template.level
|
||||
findings = parse_output(run.result.raw_output or "", level)
|
||||
gt = [
|
||||
{
|
||||
"id": str(d.id),
|
||||
"defect_type": d.defect_type,
|
||||
"line_start": d.line_start,
|
||||
"line_end": d.line_end,
|
||||
}
|
||||
for d in run.sample.defects
|
||||
]
|
||||
metrics = compare_findings(findings, gt)
|
||||
|
||||
run.result.parsed_findings = [self._finding_to_dict(f) for f in findings]
|
||||
run.result.detection_rate = metrics["detection_rate"]
|
||||
run.result.false_positive_rate = metrics["false_positive_rate"]
|
||||
run.result.coverage_rate = metrics["coverage_rate"]
|
||||
self.db.commit()
|
||||
|
||||
return {
|
||||
"run_id": run_id,
|
||||
"findings": [self._finding_to_dict(f) for f in findings],
|
||||
**metrics,
|
||||
}
|
||||
|
||||
def _finding_to_dict(self, finding) -> Dict[str, Any]:
|
||||
return {
|
||||
"defect_type": finding.defect_type,
|
||||
"line_start": finding.line_start,
|
||||
"line_end": finding.line_end,
|
||||
"description": finding.description,
|
||||
}
|
||||
|
||||
def aggregate_metrics(self, experiment_id: str) -> Dict[str, Any]:
|
||||
runs = (
|
||||
self.db.query(ExperimentRun)
|
||||
.filter_by(experiment_id=UUID(experiment_id), status="done")
|
||||
.all()
|
||||
)
|
||||
detection_rates = []
|
||||
fp_rates = []
|
||||
coverage_rates = []
|
||||
for run in runs:
|
||||
if not run.result:
|
||||
continue
|
||||
if run.result.detection_rate is not None:
|
||||
detection_rates.append(run.result.detection_rate)
|
||||
if run.result.false_positive_rate is not None:
|
||||
fp_rates.append(run.result.false_positive_rate)
|
||||
if run.result.coverage_rate is not None:
|
||||
coverage_rates.append(run.result.coverage_rate)
|
||||
|
||||
return {
|
||||
"detection_rate": descriptive_stats(detection_rates),
|
||||
"false_positive_rate": descriptive_stats(fp_rates),
|
||||
"coverage_rate": descriptive_stats(coverage_rates),
|
||||
}
|
||||
|
||||
def likert_aggregation(self) -> Dict[str, Any]:
|
||||
return aggregate_likert_by_model_and_level(self.db)
|
||||
|
||||
def generate_charts(self, experiment_id: str) -> Dict[str, str]:
|
||||
runs = (
|
||||
self.db.query(ExperimentRun)
|
||||
.filter_by(experiment_id=UUID(experiment_id), status="done")
|
||||
.all()
|
||||
)
|
||||
heatmap_data: Dict[str, Dict[str, float]] = {}
|
||||
box_groups: Dict[str, List[float]] = {}
|
||||
bar_data: Dict[str, Dict[str, float]] = {}
|
||||
|
||||
for run in runs:
|
||||
if not run.result:
|
||||
continue
|
||||
model = run.model_id
|
||||
level = run.template_version.template.level
|
||||
dr = run.result.detection_rate or 0.0
|
||||
heatmap_data.setdefault(model, {})
|
||||
bar_data.setdefault(model, {})
|
||||
heatmap_data[model][level] = heatmap_data[model].get(level, 0.0) + dr
|
||||
box_groups.setdefault(f"{model}-{level}", []).append(dr)
|
||||
bar_data[model][level] = bar_data[model].get(level, 0.0) + dr
|
||||
|
||||
# average heatmap and bar values
|
||||
counts: Dict[str, Dict[str, int]] = {}
|
||||
for run in runs:
|
||||
if not run.result:
|
||||
continue
|
||||
model = run.model_id
|
||||
level = run.template_version.template.level
|
||||
counts.setdefault(model, {}).setdefault(level, 0)
|
||||
counts[model][level] += 1
|
||||
for model in heatmap_data:
|
||||
for level in heatmap_data[model]:
|
||||
heatmap_data[model][level] /= counts[model][level]
|
||||
bar_data[model][level] /= counts[model][level]
|
||||
|
||||
return {
|
||||
"heatmap": heatmap(heatmap_data, title="Detection Rate Heatmap"),
|
||||
"boxplot": boxplot(box_groups, title="Detection Rate Distribution"),
|
||||
"grouped_bar": grouped_bar(bar_data, title="Detection Rate by Model and Level"),
|
||||
}
|
||||
|
||||
def run_anova(self, experiment_id: str) -> Dict[str, Any]:
|
||||
runs = (
|
||||
self.db.query(ExperimentRun)
|
||||
.filter_by(experiment_id=UUID(experiment_id), status="done")
|
||||
.all()
|
||||
)
|
||||
groups: Dict[str, List[float]] = {}
|
||||
for run in runs:
|
||||
if not run.result or run.result.detection_rate is None:
|
||||
continue
|
||||
key = f"{run.model_id}-{run.template_version.template.level}"
|
||||
groups.setdefault(key, []).append(run.result.detection_rate)
|
||||
return anova(list(groups.values()))
|
||||
|
||||
def run_paired_t_test(self, group_a_key: str, group_b_key: str, experiment_id: str) -> Dict[str, Any]:
|
||||
runs = (
|
||||
self.db.query(ExperimentRun)
|
||||
.filter_by(experiment_id=UUID(experiment_id), status="done")
|
||||
.all()
|
||||
)
|
||||
groups: Dict[str, List[float]] = {}
|
||||
for run in runs:
|
||||
if not run.result or run.result.detection_rate is None:
|
||||
continue
|
||||
key = f"{run.model_id}-{run.template_version.template.level}"
|
||||
groups.setdefault(key, []).append(run.result.detection_rate)
|
||||
return paired_t_test(groups.get(group_a_key, []), groups.get(group_b_key, []))
|
||||
@@ -0,0 +1,58 @@
|
||||
"""Output stability metric using Jaccard similarity."""
|
||||
|
||||
from typing import Dict, List, Set
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.analysis.parser import Finding, parse_output
|
||||
from app.models import ExperimentRun
|
||||
|
||||
|
||||
def _finding_key(finding: Finding) -> str:
|
||||
parts = [finding.defect_type.lower()]
|
||||
if finding.line_start is not None:
|
||||
parts.append(str(finding.line_start))
|
||||
if finding.line_end is not None:
|
||||
parts.append(str(finding.line_end))
|
||||
return "|".join(parts)
|
||||
|
||||
|
||||
def compute_stability_score(db: Session, experiment_id: str, model_id: str, level: str) -> float:
|
||||
"""Compute average pairwise Jaccard across three repeats for each sample."""
|
||||
from uuid import UUID
|
||||
|
||||
runs = (
|
||||
db.query(ExperimentRun)
|
||||
.filter_by(experiment_id=UUID(experiment_id), model_id=model_id)
|
||||
.all()
|
||||
)
|
||||
|
||||
by_sample: Dict[str, List[Set[str]]] = {}
|
||||
for run in runs:
|
||||
if run.template_version.template.level != level:
|
||||
continue
|
||||
if not run.result or not run.result.raw_output:
|
||||
continue
|
||||
sample_id = str(run.sample_id)
|
||||
findings = set(_finding_key(f) for f in parse_output(run.result.raw_output, level))
|
||||
by_sample.setdefault(sample_id, []).append(findings)
|
||||
|
||||
scores = []
|
||||
for sample_id, repeats in by_sample.items():
|
||||
if len(repeats) < 2:
|
||||
continue
|
||||
# pairwise Jaccard for up to 3 repeats
|
||||
pairs = [(0, 1), (0, 2), (1, 2)]
|
||||
pair_scores = []
|
||||
for i, j in pairs:
|
||||
if i < len(repeats) and j < len(repeats):
|
||||
a, b = repeats[i], repeats[j]
|
||||
union = a | b
|
||||
if not union:
|
||||
pair_scores.append(1.0)
|
||||
else:
|
||||
pair_scores.append(len(a & b) / len(union))
|
||||
if pair_scores:
|
||||
scores.append(sum(pair_scores) / len(pair_scores))
|
||||
|
||||
return round(sum(scores) / len(scores), 4) if scores else 0.0
|
||||
@@ -0,0 +1,36 @@
|
||||
"""Statistical analysis helpers."""
|
||||
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
from scipy import stats
|
||||
|
||||
|
||||
def descriptive_stats(values: List[float]) -> Dict[str, float]:
|
||||
if not values:
|
||||
return {"mean": 0.0, "std": 0.0, "min": 0.0, "max": 0.0, "median": 0.0}
|
||||
arr = np.array(values, dtype=float)
|
||||
return {
|
||||
"mean": round(float(np.mean(arr)), 4),
|
||||
"std": round(float(np.std(arr, ddof=1)), 4),
|
||||
"min": round(float(np.min(arr)), 4),
|
||||
"max": round(float(np.max(arr)), 4),
|
||||
"median": round(float(np.median(arr)), 4),
|
||||
}
|
||||
|
||||
|
||||
def anova(groups: List[List[float]]) -> Dict[str, Optional[float]]:
|
||||
"""One-way ANOVA across groups."""
|
||||
if len(groups) < 2 or any(len(g) < 2 for g in groups):
|
||||
return {"f_statistic": None, "p_value": None}
|
||||
f_stat, p_value = stats.f_oneway(*groups)
|
||||
return {"f_statistic": round(float(f_stat), 4), "p_value": round(float(p_value), 6)}
|
||||
|
||||
|
||||
def paired_t_test(a: List[float], b: List[float]) -> Dict[str, Optional[float]]:
|
||||
"""Paired t-test between two samples."""
|
||||
if len(a) != len(b) or len(a) < 2:
|
||||
return {"t_statistic": None, "p_value": None}
|
||||
t_stat, p_value = stats.ttest_rel(a, b)
|
||||
return {"t_statistic": round(float(t_stat), 4), "p_value": round(float(p_value), 6)}
|
||||
@@ -0,0 +1,24 @@
|
||||
"""FastAPI application."""
|
||||
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
from fastapi import Depends, FastAPI, HTTPException
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.api import router
|
||||
from app.db import get_db, init_db
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
init_db()
|
||||
yield
|
||||
|
||||
|
||||
app = FastAPI(title="PromptCR-Lab API", lifespan=lifespan)
|
||||
app.include_router(router.api_router, prefix="/api")
|
||||
|
||||
|
||||
@app.get("/health")
|
||||
def health_check():
|
||||
return {"status": "ok"}
|
||||
@@ -0,0 +1,280 @@
|
||||
"""FastAPI RESTful routers."""
|
||||
|
||||
from typing import Any, Dict, List, Optional
|
||||
from uuid import UUID
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.analysis.service import AnalysisService
|
||||
from app.db import get_db
|
||||
from app.dataset.builder import build_dataset_from_repo
|
||||
from app.dataset.git_parser import select_candidates
|
||||
from app.experiments.matrix import generate_full_factorial_matrix
|
||||
from app.experiments.runner import ExperimentRunner
|
||||
from app.model_adapters.factory import create_adapter, list_models
|
||||
from app.models import Experiment, ExperimentRun, PromptTemplate, PromptTemplateVersion, Sample
|
||||
from app.prompts.defaults import seed_default_templates
|
||||
from app.prompts.service import PromptService
|
||||
|
||||
api_router = APIRouter()
|
||||
|
||||
|
||||
# ---- Dataset ----
|
||||
|
||||
|
||||
class DatasetBuildRequest(BaseModel):
|
||||
repo_path: str
|
||||
count: int = 12
|
||||
languages: List[str] = None
|
||||
|
||||
|
||||
@api_router.post("/datasets/build")
|
||||
def build_dataset(req: DatasetBuildRequest, db: Session = Depends(get_db)):
|
||||
languages = req.languages or ["python", "java", "javascript"]
|
||||
samples = build_dataset_from_repo(db, req.repo_path, req.count, languages)
|
||||
return {
|
||||
"count": len(samples),
|
||||
"samples": [
|
||||
{"id": str(s.id), "repo": s.repo, "commit_sha": s.commit_sha, "language": s.language}
|
||||
for s in samples
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
@api_router.get("/datasets/samples")
|
||||
def list_samples(db: Session = Depends(get_db)):
|
||||
samples = db.query(Sample).all()
|
||||
return [
|
||||
{
|
||||
"id": str(s.id),
|
||||
"repo": s.repo,
|
||||
"commit_sha": s.commit_sha,
|
||||
"language": s.language,
|
||||
"defect_count": len(s.defects),
|
||||
}
|
||||
for s in samples
|
||||
]
|
||||
|
||||
|
||||
@api_router.get("/datasets/samples/{sample_id}")
|
||||
def get_sample(sample_id: str, db: Session = Depends(get_db)):
|
||||
try:
|
||||
sample = db.query(Sample).filter_by(id=UUID(sample_id)).first()
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=400, detail="Invalid UUID")
|
||||
if not sample:
|
||||
raise HTTPException(status_code=404, detail="Sample not found")
|
||||
return {
|
||||
"id": str(sample.id),
|
||||
"repo": sample.repo,
|
||||
"commit_sha": sample.commit_sha,
|
||||
"language": sample.language,
|
||||
"diff": sample.diff,
|
||||
"defects": [
|
||||
{
|
||||
"id": str(d.id),
|
||||
"defect_type": d.defect_type,
|
||||
"line_start": d.line_start,
|
||||
"line_end": d.line_end,
|
||||
"reference_fix": d.reference_fix,
|
||||
}
|
||||
for d in sample.defects
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
# ---- Prompts ----
|
||||
|
||||
|
||||
class PromptVersionCreate(BaseModel):
|
||||
strategy_id: str
|
||||
level: str
|
||||
body: str
|
||||
variables_schema: Optional[Dict[str, Any]] = None
|
||||
|
||||
|
||||
@api_router.get("/prompts")
|
||||
def list_prompts(db: Session = Depends(get_db)):
|
||||
templates = db.query(PromptTemplate).all()
|
||||
return [
|
||||
{
|
||||
"id": str(t.id),
|
||||
"strategy_id": t.strategy_id,
|
||||
"level": t.level,
|
||||
"version_count": len(t.versions),
|
||||
}
|
||||
for t in templates
|
||||
]
|
||||
|
||||
|
||||
@api_router.get("/prompts/{strategy_id}/{level}/versions")
|
||||
def list_prompt_versions(strategy_id: str, level: str, db: Session = Depends(get_db)):
|
||||
service = PromptService(db)
|
||||
return [
|
||||
{
|
||||
"id": str(v.id),
|
||||
"version_number": v.version_number,
|
||||
"body": v.body,
|
||||
"variables_schema": v.variables_schema,
|
||||
"created_at": v.created_at.isoformat() if v.created_at else None,
|
||||
}
|
||||
for v in service.list_versions(strategy_id, level)
|
||||
]
|
||||
|
||||
|
||||
@api_router.post("/prompts/versions")
|
||||
def create_prompt_version(req: PromptVersionCreate, db: Session = Depends(get_db)):
|
||||
service = PromptService(db)
|
||||
version = service.create_version(
|
||||
req.strategy_id, req.level, req.body, req.variables_schema
|
||||
)
|
||||
return {
|
||||
"id": str(version.id),
|
||||
"version_number": version.version_number,
|
||||
"template_id": str(version.template_id),
|
||||
}
|
||||
|
||||
|
||||
# ---- Experiments ----
|
||||
|
||||
|
||||
class ExperimentCreate(BaseModel):
|
||||
name: str
|
||||
models: List[str]
|
||||
levels: List[str]
|
||||
sample_ids: List[str]
|
||||
repeats: int = 3
|
||||
sampling_params: Optional[Dict[str, Any]] = None
|
||||
|
||||
|
||||
@api_router.post("/experiments")
|
||||
def create_experiment(req: ExperimentCreate, db: Session = Depends(get_db)):
|
||||
seed_default_templates(db)
|
||||
experiment = generate_full_factorial_matrix(
|
||||
db,
|
||||
name=req.name,
|
||||
models=req.models,
|
||||
levels=req.levels,
|
||||
sample_ids=req.sample_ids,
|
||||
repeats=req.repeats,
|
||||
sampling_params=req.sampling_params,
|
||||
)
|
||||
return {
|
||||
"id": str(experiment.id),
|
||||
"name": experiment.name,
|
||||
"status": experiment.status,
|
||||
"run_count": len(experiment.runs),
|
||||
}
|
||||
|
||||
|
||||
@api_router.get("/experiments")
|
||||
def list_experiments(db: Session = Depends(get_db)):
|
||||
experiments = db.query(Experiment).all()
|
||||
return [
|
||||
{
|
||||
"id": str(e.id),
|
||||
"name": e.name,
|
||||
"status": e.status,
|
||||
"run_count": len(e.runs),
|
||||
}
|
||||
for e in experiments
|
||||
]
|
||||
|
||||
|
||||
@api_router.post("/experiments/{experiment_id}/run")
|
||||
async def run_experiment(experiment_id: str, db: Session = Depends(get_db)):
|
||||
runner = ExperimentRunner(db)
|
||||
summary = await runner.run_experiment(experiment_id=experiment_id)
|
||||
return summary
|
||||
|
||||
|
||||
@api_router.get("/experiments/{experiment_id}/runs")
|
||||
def get_experiment_runs(experiment_id: str, db: Session = Depends(get_db)):
|
||||
try:
|
||||
runs = db.query(ExperimentRun).filter_by(experiment_id=UUID(experiment_id)).all()
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=400, detail="Invalid UUID")
|
||||
return [
|
||||
{
|
||||
"id": str(r.id),
|
||||
"run_id": r.run_id,
|
||||
"model_id": r.model_id,
|
||||
"level": r.template_version.template.level,
|
||||
"sample_id": str(r.sample_id),
|
||||
"repeat_index": r.repeat_index,
|
||||
"status": r.status,
|
||||
"retry_count": r.retry_count,
|
||||
}
|
||||
for r in runs
|
||||
]
|
||||
|
||||
|
||||
@api_router.get("/experiments/{experiment_id}/runs/{run_id}")
|
||||
def get_run(run_id: str, db: Session = Depends(get_db)):
|
||||
try:
|
||||
run = db.query(ExperimentRun).filter_by(id=UUID(run_id)).first()
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=400, detail="Invalid UUID")
|
||||
if not run:
|
||||
raise HTTPException(status_code=404, detail="Run not found")
|
||||
return {
|
||||
"id": str(run.id),
|
||||
"run_id": run.run_id,
|
||||
"model_id": run.model_id,
|
||||
"level": run.template_version.template.level,
|
||||
"sample_id": str(run.sample_id),
|
||||
"status": run.status,
|
||||
"raw_output": run.result.raw_output if run.result else None,
|
||||
"latency_ms": run.result.latency_ms if run.result else None,
|
||||
}
|
||||
|
||||
|
||||
# ---- Analysis ----
|
||||
|
||||
|
||||
@api_router.post("/analysis/{run_id}/metrics")
|
||||
def compute_run_metrics(run_id: str, db: Session = Depends(get_db)):
|
||||
service = AnalysisService(db)
|
||||
return service.compute_metrics_for_run(run_id)
|
||||
|
||||
|
||||
@api_router.get("/analysis/{experiment_id}/aggregate")
|
||||
def aggregate_metrics(experiment_id: str, db: Session = Depends(get_db)):
|
||||
service = AnalysisService(db)
|
||||
return service.aggregate_metrics(experiment_id)
|
||||
|
||||
|
||||
@api_router.get("/analysis/{experiment_id}/charts")
|
||||
def get_charts(experiment_id: str, db: Session = Depends(get_db)):
|
||||
service = AnalysisService(db)
|
||||
return service.generate_charts(experiment_id)
|
||||
|
||||
|
||||
@api_router.get("/analysis/likert")
|
||||
def likert_aggregation(db: Session = Depends(get_db)):
|
||||
service = AnalysisService(db)
|
||||
return service.likert_aggregation()
|
||||
|
||||
|
||||
@api_router.post("/analysis/likert/{run_id}")
|
||||
def save_likert(run_id: str, score: int, db: Session = Depends(get_db)):
|
||||
from app.analysis.likert import save_likert_score
|
||||
|
||||
result = save_likert_score(db, run_id, score)
|
||||
return {"run_id": run_id, "likert_score": result.likert_score}
|
||||
|
||||
|
||||
@api_router.get("/analysis/{experiment_id}/anova")
|
||||
def experiment_anova(experiment_id: str, db: Session = Depends(get_db)):
|
||||
service = AnalysisService(db)
|
||||
return service.run_anova(experiment_id)
|
||||
|
||||
|
||||
# ---- Models ----
|
||||
|
||||
|
||||
@api_router.get("/models")
|
||||
def list_available_models():
|
||||
return {"models": list_models()}
|
||||
@@ -0,0 +1,135 @@
|
||||
"""Typer CLI equivalent to the REST API."""
|
||||
|
||||
import uuid
|
||||
from typing import List, Optional
|
||||
|
||||
import typer
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.api.router import (
|
||||
AnalysisService,
|
||||
DatasetBuildRequest,
|
||||
ExperimentCreate,
|
||||
PromptService,
|
||||
create_adapter,
|
||||
generate_full_factorial_matrix,
|
||||
seed_default_templates,
|
||||
)
|
||||
from app.db import SessionLocal, init_db
|
||||
from app.experiments.runner import ExperimentRunner
|
||||
from app.model_adapters.factory import list_models
|
||||
from app.models import Sample
|
||||
|
||||
app = typer.Typer(help="PromptCR-Lab command-line interface")
|
||||
|
||||
|
||||
def get_db() -> Session:
|
||||
init_db()
|
||||
return SessionLocal()
|
||||
|
||||
|
||||
@app.command()
|
||||
def build_dataset(
|
||||
repo_path: str,
|
||||
count: int = typer.Option(12, "--count", "-c"),
|
||||
languages: Optional[List[str]] = typer.Option(None, "--language", "-l"),
|
||||
):
|
||||
"""Build dataset from a Git repository."""
|
||||
from app.dataset.builder import build_dataset_from_repo
|
||||
|
||||
db = get_db()
|
||||
languages = languages or ["python", "java", "javascript"]
|
||||
samples = build_dataset_from_repo(db, repo_path, count, languages)
|
||||
typer.echo(f"Created {len(samples)} samples")
|
||||
|
||||
|
||||
@app.command()
|
||||
def list_samples():
|
||||
"""List all samples."""
|
||||
db = get_db()
|
||||
samples = db.query(Sample).all()
|
||||
for s in samples:
|
||||
typer.echo(f"{s.id} {s.repo} {s.commit_sha} {s.language}")
|
||||
|
||||
|
||||
@app.command()
|
||||
def create_experiment(
|
||||
name: str,
|
||||
models: List[str] = typer.Option(..., "--model", "-m"),
|
||||
levels: List[str] = typer.Option(..., "--level", "-l"),
|
||||
sample_ids: List[str] = typer.Option(..., "--sample", "-s"),
|
||||
repeats: int = typer.Option(3, "--repeats", "-r"),
|
||||
):
|
||||
"""Create a full-factorial experiment."""
|
||||
db = get_db()
|
||||
seed_default_templates(db)
|
||||
experiment = generate_full_factorial_matrix(
|
||||
db,
|
||||
name=name,
|
||||
models=models,
|
||||
levels=levels,
|
||||
sample_ids=sample_ids,
|
||||
repeats=repeats,
|
||||
)
|
||||
typer.echo(f"Created experiment {experiment.id} with {len(experiment.runs)} runs")
|
||||
|
||||
|
||||
@app.command()
|
||||
def run_experiment(
|
||||
experiment_id: str,
|
||||
concurrency: int = typer.Option(5, "--concurrency", "-c"),
|
||||
):
|
||||
"""Run pending experiment units."""
|
||||
import asyncio
|
||||
|
||||
db = get_db()
|
||||
runner = ExperimentRunner(db, concurrency=concurrency)
|
||||
summary = asyncio.run(runner.run_experiment(experiment_id=uuid.UUID(experiment_id)))
|
||||
typer.echo(f"Total: {summary['total']}, Completed: {summary['completed']}")
|
||||
|
||||
|
||||
@app.command()
|
||||
def smoke(
|
||||
model: str = typer.Option("deepseek", "--model", "-m"),
|
||||
level: str = typer.Option("L1", "--level", "-l"),
|
||||
sample_id: str = typer.Option(..., "--sample", "-s"),
|
||||
):
|
||||
"""Run a 1×1×1×1 smoke test against a real model API."""
|
||||
import asyncio
|
||||
|
||||
db = get_db()
|
||||
seed_default_templates(db)
|
||||
sample = db.query(Sample).filter_by(id=uuid.UUID(sample_id)).first()
|
||||
if not sample:
|
||||
typer.echo("Sample not found", err=True)
|
||||
raise typer.Exit(1)
|
||||
|
||||
prompt_service = PromptService(db)
|
||||
prompt = prompt_service.render("code_review", None, {"language": sample.language, "diff": sample.diff})
|
||||
adapter = create_adapter(model)
|
||||
|
||||
async def call():
|
||||
response = await adapter.chat(prompt)
|
||||
typer.echo(response.text)
|
||||
|
||||
asyncio.run(call())
|
||||
|
||||
|
||||
@app.command()
|
||||
def list_models_cmd():
|
||||
"""List supported model IDs."""
|
||||
for model in list_models():
|
||||
typer.echo(model)
|
||||
|
||||
|
||||
@app.command()
|
||||
def aggregate(experiment_id: str):
|
||||
"""Aggregate metrics for an experiment."""
|
||||
db = get_db()
|
||||
service = AnalysisService(db)
|
||||
# AnalysisService converts the id internally and expects a string.
|
||||
typer.echo(service.aggregate_metrics(experiment_id))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
app()
|
||||
@@ -0,0 +1,41 @@
|
||||
from functools import lru_cache
|
||||
from typing import Optional
|
||||
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
model_config = SettingsConfigDict(
|
||||
env_file=".env",
|
||||
env_file_encoding="utf-8",
|
||||
extra="ignore",
|
||||
)
|
||||
|
||||
# Application
|
||||
app_name: str = "PromptCR-Lab"
|
||||
debug: bool = False
|
||||
|
||||
# Database (psycopg2 preferred; fall back to psycopg3 if not built)
|
||||
database_url: str = "postgresql+psycopg://postgres:postgres@localhost:5432/promptcr"
|
||||
|
||||
# Model API keys (loaded from environment, never hardcoded)
|
||||
deepseek_api_key: Optional[str] = None
|
||||
deepseek_base_url: str = "https://api.deepseek.com/v1"
|
||||
deepseek_model: str = "deepseek-chat"
|
||||
|
||||
kimi_api_key: Optional[str] = None
|
||||
kimi_base_url: str = "https://api.moonshot.cn/v1"
|
||||
kimi_model: str = "moonshot-v1-8k"
|
||||
|
||||
qwen_api_key: Optional[str] = None
|
||||
qwen_base_url: str = "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||
qwen_model: str = "qwen-turbo"
|
||||
|
||||
# Defaults
|
||||
default_model_concurrency: int = 5
|
||||
default_max_retries: int = 3
|
||||
|
||||
|
||||
@lru_cache
|
||||
def get_settings() -> Settings:
|
||||
return Settings()
|
||||
@@ -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=""))
|
||||
@@ -0,0 +1,94 @@
|
||||
"""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,
|
||||
}
|
||||
@@ -0,0 +1,121 @@
|
||||
"""Git repository parsing and commit candidate selection."""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import List, Optional
|
||||
|
||||
from git import Repo
|
||||
|
||||
|
||||
@dataclass
|
||||
class CommitCandidate:
|
||||
repo: str
|
||||
sha: str
|
||||
message: str
|
||||
author: str
|
||||
date: str
|
||||
stats: dict
|
||||
files: List[dict]
|
||||
|
||||
|
||||
def list_commits(
|
||||
repo_path: str | Path,
|
||||
max_count: Optional[int] = None,
|
||||
reverse: bool = True,
|
||||
) -> List[CommitCandidate]:
|
||||
"""List commits from a Git repository."""
|
||||
repo = Repo(str(repo_path))
|
||||
repo_name = Path(repo_path).name
|
||||
commits = []
|
||||
iterator = list(repo.iter_commits())
|
||||
if reverse:
|
||||
iterator = reversed(iterator)
|
||||
for commit in iterator:
|
||||
if max_count and len(commits) >= max_count:
|
||||
break
|
||||
stats = commit.stats.total
|
||||
files = []
|
||||
for item in commit.stats.files.items():
|
||||
filename, file_stats = item
|
||||
files.append(
|
||||
{
|
||||
"path": filename,
|
||||
"insertions": file_stats["insertions"],
|
||||
"deletions": file_stats["deletions"],
|
||||
"lines": file_stats["lines"],
|
||||
}
|
||||
)
|
||||
commits.append(
|
||||
CommitCandidate(
|
||||
repo=repo_name,
|
||||
sha=commit.hexsha,
|
||||
message=commit.message.strip(),
|
||||
author=str(commit.author),
|
||||
date=commit.committed_datetime.isoformat(),
|
||||
stats=stats,
|
||||
files=files,
|
||||
)
|
||||
)
|
||||
return commits
|
||||
|
||||
|
||||
def score_commit(commit: CommitCandidate) -> float:
|
||||
"""Score a commit by size, message quality, and language diversity."""
|
||||
total_lines = commit.stats.get("lines", 0)
|
||||
# Prefer moderate size: ~50-200 lines ideal
|
||||
size_score = 1.0 - abs(total_lines - 125) / 200.0
|
||||
size_score = max(0.0, min(1.0, size_score))
|
||||
|
||||
# Message quality: length and presence of verb/noun clues
|
||||
msg = commit.message.lower()
|
||||
msg_score = min(1.0, len(commit.message) / 40.0)
|
||||
if any(k in msg for k in ("fix", "bug", "refactor", "feature", "add", "update")):
|
||||
msg_score = min(1.0, msg_score + 0.2)
|
||||
|
||||
# Language diversity bonus based on file extensions
|
||||
exts = {Path(f["path"]).suffix.lower() for f in commit.files if Path(f["path"]).suffix}
|
||||
diversity_score = min(1.0, len(exts) / 3.0)
|
||||
|
||||
return size_score * 0.5 + msg_score * 0.3 + diversity_score * 0.2
|
||||
|
||||
|
||||
def select_candidates(
|
||||
repo_path: str | Path,
|
||||
count: int = 12,
|
||||
languages: Optional[List[str]] = None,
|
||||
scan_limit: int = 300,
|
||||
) -> List[CommitCandidate]:
|
||||
"""Select top-scoring commits, optionally balanced by language.
|
||||
|
||||
Only the most recent ``scan_limit`` commits are scanned: computing
|
||||
per-commit stats spawns a git subprocess each time, so scanning the full
|
||||
history of a large repository is prohibitively slow.
|
||||
"""
|
||||
languages = languages or ["python", "java", "javascript"]
|
||||
commits = list_commits(repo_path, max_count=scan_limit, reverse=False)
|
||||
scored = [(c, score_commit(c)) for c in commits]
|
||||
scored.sort(key=lambda x: x[1], reverse=True)
|
||||
|
||||
# Simple balancing: prefer at least one commit per target language when detectable
|
||||
by_lang = {lang: [] for lang in languages}
|
||||
others = []
|
||||
for commit, score in scored:
|
||||
ext_set = {Path(f["path"]).suffix.lower() for f in commit.files}
|
||||
placed = False
|
||||
for lang in languages:
|
||||
hint = ".py" if lang == "python" else ".java" if lang == "java" else ".js"
|
||||
if hint in ext_set:
|
||||
by_lang[lang].append((commit, score))
|
||||
placed = True
|
||||
break
|
||||
if not placed:
|
||||
others.append((commit, score))
|
||||
|
||||
result = []
|
||||
per_lang = max(1, count // len(languages))
|
||||
for lang in languages:
|
||||
result.extend(by_lang[lang][:per_lang])
|
||||
result.extend(others)
|
||||
result = result[:count]
|
||||
result.sort(key=lambda x: x[1], reverse=True)
|
||||
return [c for c, _ in result]
|
||||
@@ -0,0 +1,30 @@
|
||||
"""Base class for pluggable defect mutation rules."""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional
|
||||
|
||||
|
||||
@dataclass
|
||||
class Mutation:
|
||||
defect_type: str
|
||||
language: str
|
||||
line_start: int
|
||||
line_end: int
|
||||
mutated_source: str
|
||||
reference_fix: str
|
||||
description: str
|
||||
|
||||
|
||||
class MutationRule(ABC):
|
||||
name: str = ""
|
||||
language: str = ""
|
||||
defect_type: str = ""
|
||||
|
||||
@abstractmethod
|
||||
def detect_and_mutate(self, source: str, filename: str = "") -> Optional[Mutation]:
|
||||
"""Return a mutation if the rule can introduce a defect, else None."""
|
||||
...
|
||||
|
||||
def _line_for_position(self, source: str, position: int) -> int:
|
||||
return source[:position].count("\n") + 1
|
||||
@@ -0,0 +1,57 @@
|
||||
"""AST-level boundary error injection for Java using javalang."""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import javalang
|
||||
|
||||
from app.dataset.rules.base import Mutation, MutationRule
|
||||
|
||||
|
||||
class JavaBoundaryErrorRule(MutationRule):
|
||||
"""Mutate array length boundary from `<` to `<=`.
|
||||
|
||||
Uses `javalang` to locate a binary comparison against `.length` and flips
|
||||
the operator to introduce an off-by-one access.
|
||||
"""
|
||||
|
||||
name = "java_boundary_error"
|
||||
language = "java"
|
||||
defect_type = "boundary_condition_error"
|
||||
|
||||
def detect_and_mutate(self, source: str, filename: str = "") -> Optional[Mutation]:
|
||||
try:
|
||||
tree = javalang.parse.parse(source)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
for path, node in tree:
|
||||
if not isinstance(node, javalang.tree.BinaryOperation):
|
||||
continue
|
||||
if node.operator != "<":
|
||||
continue
|
||||
right = node.operandr
|
||||
if isinstance(right, javalang.tree.MemberReference) and right.member == "length":
|
||||
# BinaryOperation itself may lack position; use enclosing statement
|
||||
if_statement = next((n for n in path if isinstance(n, javalang.tree.IfStatement)), None)
|
||||
pos = if_statement.position if if_statement else node.position
|
||||
if not pos:
|
||||
continue
|
||||
lines = source.splitlines(keepends=True)
|
||||
line_no = pos.line
|
||||
line = lines[line_no - 1]
|
||||
mutated_line = line.replace("< ", "<= ", 1)
|
||||
if mutated_line == line:
|
||||
mutated_line = line.replace("<", "<=", 1)
|
||||
if mutated_line == line:
|
||||
continue
|
||||
mutated = "".join(lines[:line_no - 1] + [mutated_line] + lines[line_no:])
|
||||
return Mutation(
|
||||
defect_type=self.defect_type,
|
||||
language=self.language,
|
||||
line_start=line_no,
|
||||
line_end=line_no,
|
||||
mutated_source=mutated,
|
||||
reference_fix="Use strict `< length` to avoid ArrayIndexOutOfBoundsException.",
|
||||
description="Changed array boundary check to off-by-one (<= length).",
|
||||
)
|
||||
return None
|
||||
@@ -0,0 +1,69 @@
|
||||
"""AST-level concurrency issue injection for Java using javalang."""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import javalang
|
||||
|
||||
from app.dataset.rules.base import Mutation, MutationRule
|
||||
|
||||
|
||||
class JavaConcurrencyRule(MutationRule):
|
||||
"""Remove a synchronized block to expose a race condition.
|
||||
|
||||
Uses `javalang` to locate `synchronized (lock) { ... }` and replaces it
|
||||
with the bare block body.
|
||||
"""
|
||||
|
||||
name = "java_concurrency"
|
||||
language = "java"
|
||||
defect_type = "concurrency_issue"
|
||||
|
||||
def detect_and_mutate(self, source: str, filename: str = "") -> Optional[Mutation]:
|
||||
try:
|
||||
tree = javalang.parse.parse(source)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
for path, node in tree:
|
||||
if not isinstance(node, javalang.tree.SynchronizedStatement):
|
||||
continue
|
||||
pos = node.position
|
||||
if not pos:
|
||||
continue
|
||||
lines = source.splitlines(keepends=True)
|
||||
start = pos.line
|
||||
# Estimate end by braces of the synchronized block
|
||||
end = start
|
||||
depth = 0
|
||||
for idx in range(start - 1, len(lines)):
|
||||
depth += lines[idx].count("{") - lines[idx].count("}")
|
||||
if depth > 0:
|
||||
end = idx + 1
|
||||
if depth <= 0 and idx > start - 1:
|
||||
end = idx + 1
|
||||
break
|
||||
body_lines = lines[start - 1:end]
|
||||
# drop header line and closing brace line, keep body; body is at same
|
||||
# indentation as the synchronized header minus one level
|
||||
inner = body_lines[1:-1] if len(body_lines) > 2 else []
|
||||
dedented = []
|
||||
for line in inner:
|
||||
if line.startswith(" "):
|
||||
dedented.append(" " + line[12:])
|
||||
elif line.startswith(" "):
|
||||
dedented.append(" " + line[8:])
|
||||
elif line.startswith(" "):
|
||||
dedented.append(line[4:])
|
||||
else:
|
||||
dedented.append(line)
|
||||
mutated = "".join(lines[: start - 1] + dedented + lines[end:])
|
||||
return Mutation(
|
||||
defect_type=self.defect_type,
|
||||
language=self.language,
|
||||
line_start=start,
|
||||
line_end=end,
|
||||
mutated_source=mutated,
|
||||
reference_fix="Restore `synchronized (lock)` to protect the critical section.",
|
||||
description="Removed synchronized block, exposing a race condition.",
|
||||
)
|
||||
return None
|
||||
@@ -0,0 +1,52 @@
|
||||
"""AST-level logical operator misuse injection for Java using javalang."""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import javalang
|
||||
|
||||
from app.dataset.rules.base import Mutation, MutationRule
|
||||
|
||||
|
||||
class JavaLogicOperatorRule(MutationRule):
|
||||
"""Swap `&&` with `||` in a boolean expression.
|
||||
|
||||
Uses `javalang` to locate a binary operation with `&&` and replaces the
|
||||
operator with `||`.
|
||||
"""
|
||||
|
||||
name = "java_logic_operator"
|
||||
language = "java"
|
||||
defect_type = "logic_operator_misuse"
|
||||
|
||||
def detect_and_mutate(self, source: str, filename: str = "") -> Optional[Mutation]:
|
||||
try:
|
||||
tree = javalang.parse.parse(source)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
for path, node in tree:
|
||||
if not isinstance(node, javalang.tree.BinaryOperation):
|
||||
continue
|
||||
if node.operator != "&&":
|
||||
continue
|
||||
return_statement = next((n for n in path if isinstance(n, javalang.tree.ReturnStatement)), None)
|
||||
pos = return_statement.position if return_statement else node.position
|
||||
if not pos:
|
||||
continue
|
||||
lines = source.splitlines(keepends=True)
|
||||
line_no = pos.line
|
||||
line = lines[line_no - 1]
|
||||
mutated_line = line.replace("&&", "||", 1)
|
||||
if mutated_line == line:
|
||||
continue
|
||||
mutated = "".join(lines[:line_no - 1] + [mutated_line] + lines[line_no:])
|
||||
return Mutation(
|
||||
defect_type=self.defect_type,
|
||||
language=self.language,
|
||||
line_start=line_no,
|
||||
line_end=line_no,
|
||||
mutated_source=mutated,
|
||||
reference_fix="Restore `&&` for correct conjunction semantics.",
|
||||
description="Replaced boolean `&&` with `||`.",
|
||||
)
|
||||
return None
|
||||
@@ -0,0 +1,75 @@
|
||||
"""AST-level null-pointer injection for Java using javalang."""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import javalang
|
||||
|
||||
from app.dataset.rules.base import Mutation, MutationRule
|
||||
|
||||
|
||||
class JavaNoneReferenceRule(MutationRule):
|
||||
"""Remove a null-check guard in Java source.
|
||||
|
||||
Uses the pure-Python `javalang` parser to locate an `if (x != null)` guard
|
||||
and remove it, leaving the dereference unprotected. This keeps mutation
|
||||
semantics precise without regex/text replacement.
|
||||
"""
|
||||
|
||||
name = "java_none_reference"
|
||||
language = "java"
|
||||
defect_type = "null_pointer"
|
||||
|
||||
def detect_and_mutate(self, source: str, filename: str = "") -> Optional[Mutation]:
|
||||
try:
|
||||
tree = javalang.parse.parse(source)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
for path, node in tree:
|
||||
if not isinstance(node, javalang.tree.IfStatement):
|
||||
continue
|
||||
cond = node.condition
|
||||
# Match: x != null
|
||||
if (
|
||||
isinstance(cond, javalang.tree.BinaryOperation)
|
||||
and cond.operator == "!="
|
||||
and isinstance(cond.operandr, javalang.tree.Literal)
|
||||
and cond.operandr.value == "null"
|
||||
):
|
||||
var_name = getattr(cond.operandl, "member", str(cond.operandl))
|
||||
lines = source.splitlines(keepends=True)
|
||||
start = node.position.line if node.position else 1
|
||||
# Estimate end line by finding matching brace (simplistic)
|
||||
end = start
|
||||
depth = 0
|
||||
for idx in range(start - 1, len(lines)):
|
||||
depth += lines[idx].count("{") - lines[idx].count("}")
|
||||
if depth > 0:
|
||||
end = idx + 1
|
||||
if depth <= 0 and idx > start - 1:
|
||||
end = idx + 1
|
||||
break
|
||||
body_lines = lines[start - 1:end]
|
||||
# keep body lines between header and closing brace, dedent one level
|
||||
inner = body_lines[1:-1] if len(body_lines) > 2 else []
|
||||
dedented = []
|
||||
for line in inner:
|
||||
if line.startswith(" "):
|
||||
dedented.append(" " + line[12:])
|
||||
elif line.startswith(" "):
|
||||
dedented.append(" " + line[8:])
|
||||
elif line.startswith(" "):
|
||||
dedented.append(line[4:])
|
||||
else:
|
||||
dedented.append(line)
|
||||
mutated = "".join(lines[: start - 1] + dedented + lines[end:])
|
||||
return Mutation(
|
||||
defect_type=self.defect_type,
|
||||
language=self.language,
|
||||
line_start=start,
|
||||
line_end=end,
|
||||
mutated_source=mutated,
|
||||
reference_fix=f"Add `if ({var_name} != null)` guard before dereferencing.",
|
||||
description=f"Removed null-check guard for '{var_name}'.",
|
||||
)
|
||||
return None
|
||||
@@ -0,0 +1,50 @@
|
||||
"""AST-level resource leak injection for Java using javalang."""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import javalang
|
||||
|
||||
from app.dataset.rules.base import Mutation, MutationRule
|
||||
|
||||
|
||||
class JavaResourceLeakRule(MutationRule):
|
||||
"""Replace try-with-resources with a plain try block, leaking the resource.
|
||||
|
||||
Uses `javalang` to locate a try-with-resources statement and removes the
|
||||
resource specification, leaving the stream unclosed.
|
||||
"""
|
||||
|
||||
name = "java_resource_leak"
|
||||
language = "java"
|
||||
defect_type = "resource_not_closed"
|
||||
|
||||
def detect_and_mutate(self, source: str, filename: str = "") -> Optional[Mutation]:
|
||||
try:
|
||||
tree = javalang.parse.parse(source)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
for path, node in tree:
|
||||
if not isinstance(node, javalang.tree.TryStatement):
|
||||
continue
|
||||
if not node.resources:
|
||||
continue
|
||||
pos = node.position
|
||||
if not pos:
|
||||
continue
|
||||
lines = source.splitlines(keepends=True)
|
||||
start = pos.line
|
||||
# Find the resource clause line e.g. try (BufferedReader br = ...)
|
||||
resource_line = lines[start - 1]
|
||||
new_header = resource_line.split("(", 1)[0].rstrip() + " {\n"
|
||||
mutated = "".join(lines[: start - 1] + [new_header] + lines[start:])
|
||||
return Mutation(
|
||||
defect_type=self.defect_type,
|
||||
language=self.language,
|
||||
line_start=start,
|
||||
line_end=start,
|
||||
mutated_source=mutated,
|
||||
reference_fix="Use try-with-resources or explicitly close the stream in finally.",
|
||||
description="Removed try-with-resources, leaking the acquired resource.",
|
||||
)
|
||||
return None
|
||||
@@ -0,0 +1,63 @@
|
||||
"""AST-level boundary error injection for JavaScript using esprima."""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import esprima
|
||||
|
||||
from app.dataset.rules.base import Mutation, MutationRule
|
||||
|
||||
|
||||
class JSBoundaryErrorRule(MutationRule):
|
||||
"""Mutate array length boundary from `<` to `<=`.
|
||||
|
||||
Uses `esprima` to locate a binary expression comparing against `.length`
|
||||
and flips the operator.
|
||||
"""
|
||||
|
||||
name = "js_boundary_error"
|
||||
language = "javascript"
|
||||
defect_type = "boundary_condition_error"
|
||||
|
||||
def detect_and_mutate(self, source: str, filename: str = "") -> Optional[Mutation]:
|
||||
# Modern JS is usually ESM: try module grammar first, then script.
|
||||
try:
|
||||
tree = esprima.parseModule(source, loc=True)
|
||||
except Exception:
|
||||
try:
|
||||
tree = esprima.parseScript(source, loc=True)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
for node in self._walk(tree):
|
||||
if node.type != "BinaryExpression" or node.operator != "<":
|
||||
continue
|
||||
right = node.right
|
||||
if right.type == "MemberExpression" and getattr(right.property, "name", None) == "length":
|
||||
line_no = node.loc.start.line
|
||||
lines = source.splitlines(keepends=True)
|
||||
line = lines[line_no - 1]
|
||||
mutated_line = line.replace("<", "<=", 1)
|
||||
if mutated_line == line:
|
||||
continue
|
||||
mutated = "".join(lines[:line_no - 1] + [mutated_line] + lines[line_no:])
|
||||
return Mutation(
|
||||
defect_type=self.defect_type,
|
||||
language=self.language,
|
||||
line_start=line_no,
|
||||
line_end=line_no,
|
||||
mutated_source=mutated,
|
||||
reference_fix="Use strict `< length` to avoid out-of-bounds access.",
|
||||
description="Changed array boundary check to off-by-one (<= length).",
|
||||
)
|
||||
return None
|
||||
|
||||
def _walk(self, node):
|
||||
yield node
|
||||
for key in getattr(node, "__dict__", {}):
|
||||
child = getattr(node, key)
|
||||
if isinstance(child, list):
|
||||
for item in child:
|
||||
if hasattr(item, "type"):
|
||||
yield from self._walk(item)
|
||||
elif hasattr(child, "type"):
|
||||
yield from self._walk(child)
|
||||
@@ -0,0 +1,60 @@
|
||||
"""AST-level concurrency issue injection for JavaScript using esprima."""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import esprima
|
||||
|
||||
from app.dataset.rules.base import Mutation, MutationRule
|
||||
|
||||
|
||||
class JSConcurrencyRule(MutationRule):
|
||||
"""Remove an `await mutex.acquire()` / `mutex.release()` pair.
|
||||
|
||||
Uses `esprima` to locate a try block followed by a finally that releases a
|
||||
mutex and removes the finally/release, exposing a race.
|
||||
"""
|
||||
|
||||
name = "js_concurrency"
|
||||
language = "javascript"
|
||||
defect_type = "concurrency_issue"
|
||||
|
||||
def detect_and_mutate(self, source: str, filename: str = "") -> Optional[Mutation]:
|
||||
# Modern JS is usually ESM: try module grammar first, then script.
|
||||
try:
|
||||
tree = esprima.parseModule(source, loc=True)
|
||||
except Exception:
|
||||
try:
|
||||
tree = esprima.parseScript(source, loc=True)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def walk(node):
|
||||
yield node
|
||||
for key in node.__dict__:
|
||||
child = getattr(node, key)
|
||||
if isinstance(child, list):
|
||||
for item in child:
|
||||
if hasattr(item, "type"):
|
||||
yield from walk(item)
|
||||
elif hasattr(child, "type"):
|
||||
yield from walk(child)
|
||||
|
||||
for node in walk(tree):
|
||||
if node.type != "TryStatement" or not node.finalizer:
|
||||
continue
|
||||
start = node.loc.start.line
|
||||
end = node.finalizer.loc.end.line
|
||||
lines = source.splitlines(keepends=True)
|
||||
# Drop the entire finally block; the try block ends on the same line as finally starts
|
||||
finally_start = node.finalizer.loc.start.line - 1
|
||||
mutated = "".join(lines[:finally_start] + [" }\n"] + lines[end:])
|
||||
return Mutation(
|
||||
defect_type=self.defect_type,
|
||||
language=self.language,
|
||||
line_start=start,
|
||||
line_end=end,
|
||||
mutated_source=mutated,
|
||||
reference_fix="Restore mutex release in finally to protect the critical section.",
|
||||
description="Removed mutex release, exposing a race condition.",
|
||||
)
|
||||
return None
|
||||
@@ -0,0 +1,61 @@
|
||||
"""AST-level logical operator misuse injection for JavaScript using esprima."""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import esprima
|
||||
|
||||
from app.dataset.rules.base import Mutation, MutationRule
|
||||
|
||||
|
||||
class JSLogicOperatorRule(MutationRule):
|
||||
"""Swap `&&` with `||` in a boolean expression.
|
||||
|
||||
Uses `esprima` to locate a LogicalExpression using `&&` and replaces the
|
||||
operator with `||`.
|
||||
"""
|
||||
|
||||
name = "js_logic_operator"
|
||||
language = "javascript"
|
||||
defect_type = "logic_operator_misuse"
|
||||
|
||||
def detect_and_mutate(self, source: str, filename: str = "") -> Optional[Mutation]:
|
||||
# Modern JS is usually ESM: try module grammar first, then script.
|
||||
try:
|
||||
tree = esprima.parseModule(source, loc=True)
|
||||
except Exception:
|
||||
try:
|
||||
tree = esprima.parseScript(source, loc=True)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
for node in self._walk(tree):
|
||||
if node.type != "LogicalExpression" or node.operator != "&&":
|
||||
continue
|
||||
line_no = node.loc.start.line
|
||||
lines = source.splitlines(keepends=True)
|
||||
line = lines[line_no - 1]
|
||||
mutated_line = line.replace("&&", "||", 1)
|
||||
if mutated_line == line:
|
||||
continue
|
||||
mutated = "".join(lines[:line_no - 1] + [mutated_line] + lines[line_no:])
|
||||
return Mutation(
|
||||
defect_type=self.defect_type,
|
||||
language=self.language,
|
||||
line_start=line_no,
|
||||
line_end=line_no,
|
||||
mutated_source=mutated,
|
||||
reference_fix="Restore `&&` for correct short-circuit logic.",
|
||||
description="Replaced boolean `&&` with `||`.",
|
||||
)
|
||||
return None
|
||||
|
||||
def _walk(self, node):
|
||||
yield node
|
||||
for key in getattr(node, "__dict__", {}):
|
||||
child = getattr(node, key)
|
||||
if isinstance(child, list):
|
||||
for item in child:
|
||||
if hasattr(item, "type"):
|
||||
yield from self._walk(item)
|
||||
elif hasattr(child, "type"):
|
||||
yield from self._walk(child)
|
||||
@@ -0,0 +1,76 @@
|
||||
"""AST-level null-pointer injection for JavaScript using esprima."""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import esprima
|
||||
|
||||
from app.dataset.rules.base import Mutation, MutationRule
|
||||
|
||||
|
||||
class JSNoneReferenceRule(MutationRule):
|
||||
"""Remove a `if (x !== null)` guard in JavaScript.
|
||||
|
||||
Uses the Python port of `esprima` to locate the guard statement and
|
||||
replaces it with the body, leaving a potential null dereference.
|
||||
"""
|
||||
|
||||
name = "js_none_reference"
|
||||
language = "javascript"
|
||||
defect_type = "null_pointer"
|
||||
|
||||
def detect_and_mutate(self, source: str, filename: str = "") -> Optional[Mutation]:
|
||||
# Modern JS is usually ESM: try module grammar first, then script.
|
||||
try:
|
||||
tree = esprima.parseModule(source, loc=True)
|
||||
except Exception:
|
||||
try:
|
||||
tree = esprima.parseScript(source, loc=True)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def walk(node):
|
||||
yield node
|
||||
for key in node.__dict__:
|
||||
child = getattr(node, key)
|
||||
if isinstance(child, list):
|
||||
for item in child:
|
||||
if hasattr(item, "type"):
|
||||
yield from walk(item)
|
||||
elif hasattr(child, "type"):
|
||||
yield from walk(child)
|
||||
|
||||
for node in walk(tree):
|
||||
if node.type != "IfStatement":
|
||||
continue
|
||||
cond = node.test
|
||||
if (
|
||||
cond.type == "BinaryExpression"
|
||||
and cond.operator == "!=="
|
||||
and cond.right.type == "Literal"
|
||||
and cond.right.value is None
|
||||
):
|
||||
var_name = getattr(cond.left, "name", str(cond.left))
|
||||
start = node.loc.start.line
|
||||
end = node.consequent.loc.end.line
|
||||
lines = source.splitlines(keepends=True)
|
||||
# Drop guard header and closing brace, keep body (1-based -> 0-based)
|
||||
body_start = node.consequent.loc.start.line
|
||||
body_end = node.consequent.loc.end.line - 1
|
||||
body_lines = lines[body_start:body_end]
|
||||
dedented = []
|
||||
for line in body_lines:
|
||||
if line.startswith(" "):
|
||||
dedented.append(line[4:])
|
||||
else:
|
||||
dedented.append(line)
|
||||
mutated = "".join(lines[: start - 1] + dedented + lines[end:])
|
||||
return Mutation(
|
||||
defect_type=self.defect_type,
|
||||
language=self.language,
|
||||
line_start=start,
|
||||
line_end=end,
|
||||
mutated_source=mutated,
|
||||
reference_fix=f"Add `if ({var_name} !== null)` guard before dereferencing.",
|
||||
description=f"Removed null-check guard for '{var_name}'.",
|
||||
)
|
||||
return None
|
||||
@@ -0,0 +1,60 @@
|
||||
"""AST-level resource leak injection for JavaScript using esprima."""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import esprima
|
||||
|
||||
from app.dataset.rules.base import Mutation, MutationRule
|
||||
|
||||
|
||||
class JSResourceLeakRule(MutationRule):
|
||||
"""Remove a fetch Response body close/usage, leaking the reader.
|
||||
|
||||
Uses `esprima` to locate a `try/finally` that closes a reader and removes
|
||||
the finally block.
|
||||
"""
|
||||
|
||||
name = "js_resource_leak"
|
||||
language = "javascript"
|
||||
defect_type = "resource_not_closed"
|
||||
|
||||
def detect_and_mutate(self, source: str, filename: str = "") -> Optional[Mutation]:
|
||||
# Modern JS is usually ESM: try module grammar first, then script.
|
||||
try:
|
||||
tree = esprima.parseModule(source, loc=True)
|
||||
except Exception:
|
||||
try:
|
||||
tree = esprima.parseScript(source, loc=True)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def walk(node):
|
||||
yield node
|
||||
for key in node.__dict__:
|
||||
child = getattr(node, key)
|
||||
if isinstance(child, list):
|
||||
for item in child:
|
||||
if hasattr(item, "type"):
|
||||
yield from walk(item)
|
||||
elif hasattr(child, "type"):
|
||||
yield from walk(child)
|
||||
|
||||
for node in walk(tree):
|
||||
if node.type != "TryStatement" or not node.finalizer:
|
||||
continue
|
||||
start = node.loc.start.line
|
||||
end = node.finalizer.loc.end.line
|
||||
lines = source.splitlines(keepends=True)
|
||||
# Drop the finally block entirely, close the try block
|
||||
finally_start = node.finalizer.loc.start.line - 1
|
||||
mutated = "".join(lines[:finally_start] + [" }\n"] + lines[end:])
|
||||
return Mutation(
|
||||
defect_type=self.defect_type,
|
||||
language=self.language,
|
||||
line_start=start,
|
||||
line_end=end,
|
||||
mutated_source=mutated,
|
||||
reference_fix="Restore finally block to close/release resources.",
|
||||
description="Removed finally block, leaving resource unreleased.",
|
||||
)
|
||||
return None
|
||||
@@ -0,0 +1,47 @@
|
||||
"""Custom rule demonstrating pluggable extensibility (A1b)."""
|
||||
|
||||
import ast
|
||||
from typing import Optional
|
||||
|
||||
from app.dataset.rules.base import Mutation, MutationRule
|
||||
|
||||
|
||||
class UnusedVariableRule(MutationRule):
|
||||
"""A custom rule: replace a variable read with an undefined name.
|
||||
|
||||
This is intentionally simple and demonstrates that adding a new file under
|
||||
app/dataset/rules/<language>/ is enough to register a rule.
|
||||
"""
|
||||
|
||||
name = "unused_variable_demo"
|
||||
language = "python"
|
||||
defect_type = "custom_demo"
|
||||
|
||||
def detect_and_mutate(self, source: str, filename: str = "") -> Optional[Mutation]:
|
||||
try:
|
||||
tree = ast.parse(source)
|
||||
except SyntaxError:
|
||||
return None
|
||||
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.FunctionDef) and node.body:
|
||||
first = node.body[0]
|
||||
if isinstance(first, ast.Assign) and isinstance(first.targets[0], ast.Name):
|
||||
var_name = first.targets[0].id
|
||||
line_no = first.lineno
|
||||
lines = source.splitlines(keepends=True)
|
||||
line = lines[line_no - 1]
|
||||
mutated_line = line.replace(var_name, "__undefined_" + var_name, 1)
|
||||
if mutated_line == line:
|
||||
continue
|
||||
mutated = "".join(lines[:line_no - 1] + [mutated_line] + lines[line_no:])
|
||||
return Mutation(
|
||||
defect_type=self.defect_type,
|
||||
language=self.language,
|
||||
line_start=line_no,
|
||||
line_end=line_no,
|
||||
mutated_source=mutated,
|
||||
reference_fix=f"Use the original variable name '{var_name}'.",
|
||||
description=f"Custom rule: replaced '{var_name}' with an undefined placeholder.",
|
||||
)
|
||||
return None
|
||||
@@ -0,0 +1,50 @@
|
||||
"""AST-level boundary condition injection for Python."""
|
||||
|
||||
import ast
|
||||
from typing import Optional
|
||||
|
||||
from app.dataset.rules.base import Mutation, MutationRule
|
||||
|
||||
|
||||
class BoundaryErrorRule(MutationRule):
|
||||
"""Mutate a list-index boundary check from `< len(seq)` to `<= len(seq)`.
|
||||
|
||||
Uses `ast` to find comparisons guarding index access and flips the operator
|
||||
so the boundary becomes off-by-one.
|
||||
"""
|
||||
|
||||
name = "boundary_error"
|
||||
language = "python"
|
||||
defect_type = "boundary_condition_error"
|
||||
|
||||
def detect_and_mutate(self, source: str, filename: str = "") -> Optional[Mutation]:
|
||||
try:
|
||||
tree = ast.parse(source)
|
||||
except SyntaxError:
|
||||
return None
|
||||
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, ast.If):
|
||||
continue
|
||||
test = node.test
|
||||
if isinstance(test, ast.Compare) and isinstance(test.left, ast.Name):
|
||||
if len(test.ops) == 1 and isinstance(test.ops[0], ast.Lt):
|
||||
# i < len(x) -> i <= len(x)
|
||||
comparator = test.comparators[0]
|
||||
if isinstance(comparator, ast.Call) and isinstance(comparator.func, ast.Name) and comparator.func.id == "len":
|
||||
lines = source.splitlines(keepends=True)
|
||||
line = lines[node.lineno - 1]
|
||||
mutated_line = line.replace("< len(", "<= len(", 1)
|
||||
if mutated_line == line:
|
||||
continue
|
||||
mutated = "".join(lines[: node.lineno - 1] + [mutated_line] + lines[node.lineno:])
|
||||
return Mutation(
|
||||
defect_type=self.defect_type,
|
||||
language=self.language,
|
||||
line_start=node.lineno,
|
||||
line_end=getattr(node, "end_lineno", node.lineno) or node.lineno,
|
||||
mutated_source=mutated,
|
||||
reference_fix="Use strict `< len(seq)` to avoid index-out-of-range.",
|
||||
description="Changed index boundary check to off-by-one (<= len).",
|
||||
)
|
||||
return None
|
||||
@@ -0,0 +1,52 @@
|
||||
"""AST-level concurrency safety injection for Python."""
|
||||
|
||||
import ast
|
||||
from typing import Optional
|
||||
|
||||
from app.dataset.rules.base import Mutation, MutationRule
|
||||
|
||||
|
||||
class ConcurrencyRule(MutationRule):
|
||||
"""Remove a threading.Lock.acquire/release pair to introduce race condition.
|
||||
|
||||
Uses `ast` to find a with-statement using a lock and replaces it with the
|
||||
bare body, removing synchronization.
|
||||
"""
|
||||
|
||||
name = "concurrency"
|
||||
language = "python"
|
||||
defect_type = "concurrency_issue"
|
||||
|
||||
def detect_and_mutate(self, source: str, filename: str = "") -> Optional[Mutation]:
|
||||
try:
|
||||
tree = ast.parse(source)
|
||||
except SyntaxError:
|
||||
return None
|
||||
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, ast.With):
|
||||
continue
|
||||
first_item = node.items[0]
|
||||
ctx = first_item.context_expr
|
||||
if isinstance(ctx, ast.Call) and isinstance(ctx.func, ast.Attribute) and ctx.func.attr == "acquire":
|
||||
start = node.lineno
|
||||
end = getattr(node, "end_lineno", node.lineno) or node.lineno
|
||||
lines = source.splitlines(keepends=True)
|
||||
body = lines[start:end]
|
||||
dedented = []
|
||||
for line in body[1:]:
|
||||
if line.startswith(" "):
|
||||
dedented.append(line[4:])
|
||||
else:
|
||||
dedented.append(line)
|
||||
mutated = "".join(lines[: start - 1] + dedented + lines[end:])
|
||||
return Mutation(
|
||||
defect_type=self.defect_type,
|
||||
language=self.language,
|
||||
line_start=start,
|
||||
line_end=end,
|
||||
mutated_source=mutated,
|
||||
reference_fix="Restore `with lock:` to protect the critical section.",
|
||||
description="Removed lock acquisition, exposing a race condition.",
|
||||
)
|
||||
return None
|
||||
@@ -0,0 +1,44 @@
|
||||
"""AST-level logical operator misuse injection for Python."""
|
||||
|
||||
import ast
|
||||
from typing import Optional
|
||||
|
||||
from app.dataset.rules.base import Mutation, MutationRule
|
||||
|
||||
|
||||
class LogicOperatorRule(MutationRule):
|
||||
"""Swap `and` with `or` in a boolean expression.
|
||||
|
||||
Uses `ast` to locate a BoolOp using `And` and replaces it with `Or`,
|
||||
preserving exact source position via line replacement.
|
||||
"""
|
||||
|
||||
name = "logic_operator"
|
||||
language = "python"
|
||||
defect_type = "logic_operator_misuse"
|
||||
|
||||
def detect_and_mutate(self, source: str, filename: str = "") -> Optional[Mutation]:
|
||||
try:
|
||||
tree = ast.parse(source)
|
||||
except SyntaxError:
|
||||
return None
|
||||
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.BoolOp) and isinstance(node.op, ast.And):
|
||||
line_no = node.lineno
|
||||
lines = source.splitlines(keepends=True)
|
||||
line = lines[line_no - 1]
|
||||
mutated_line = line.replace(" and ", " or ", 1)
|
||||
if mutated_line == line:
|
||||
continue
|
||||
mutated = "".join(lines[:line_no - 1] + [mutated_line] + lines[line_no:])
|
||||
return Mutation(
|
||||
defect_type=self.defect_type,
|
||||
language=self.language,
|
||||
line_start=line_no,
|
||||
line_end=getattr(node, "end_lineno", line_no) or line_no,
|
||||
mutated_source=mutated,
|
||||
reference_fix="Restore the original `and` operator for correct short-circuit logic.",
|
||||
description="Replaced boolean `and` with `or`.",
|
||||
)
|
||||
return None
|
||||
@@ -0,0 +1,65 @@
|
||||
"""AST-level null-pointer / None-reference injection for Python."""
|
||||
|
||||
import ast
|
||||
from typing import Optional
|
||||
|
||||
from app.dataset.rules.base import Mutation, MutationRule
|
||||
|
||||
|
||||
class NoneReferenceRule(MutationRule):
|
||||
"""Replace a checked variable access with an unchecked None dereference.
|
||||
|
||||
This rule uses the standard library `ast` module to precisely locate a
|
||||
variable that is used after an `if x is not None:` guard, then removes the
|
||||
guard. The mutation position is derived from AST line numbers so it is
|
||||
exact and reproducible.
|
||||
"""
|
||||
|
||||
name = "none_reference"
|
||||
language = "python"
|
||||
defect_type = "null_pointer"
|
||||
|
||||
def detect_and_mutate(self, source: str, filename: str = "") -> Optional[Mutation]:
|
||||
try:
|
||||
tree = ast.parse(source)
|
||||
except SyntaxError:
|
||||
return None
|
||||
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, ast.If):
|
||||
continue
|
||||
test = node.test
|
||||
# Match: if x is not None:
|
||||
if (
|
||||
isinstance(test, ast.Compare)
|
||||
and isinstance(test.left, ast.Name)
|
||||
and len(test.ops) == 1
|
||||
and isinstance(test.ops[0], ast.IsNot)
|
||||
and len(test.comparators) == 1
|
||||
and isinstance(test.comparators[0], ast.Constant)
|
||||
and test.comparators[0].value is None
|
||||
):
|
||||
var_name = test.left.id
|
||||
lines = source.splitlines(keepends=True)
|
||||
start = node.lineno
|
||||
end = getattr(node, "end_lineno", node.lineno) or node.lineno
|
||||
body_lines = lines[start:end]
|
||||
dedented = []
|
||||
for line in body_lines[1:]:
|
||||
if line.startswith(" "):
|
||||
dedented.append(line[4:])
|
||||
elif line.startswith("\t"):
|
||||
dedented.append(line[1:])
|
||||
else:
|
||||
dedented.append(line)
|
||||
mutated = "".join(lines[: start - 1] + dedented + lines[end:])
|
||||
return Mutation(
|
||||
defect_type=self.defect_type,
|
||||
language=self.language,
|
||||
line_start=start,
|
||||
line_end=end,
|
||||
mutated_source=mutated,
|
||||
reference_fix=f"Add `if {var_name} is not None:` guard before use.",
|
||||
description=f"Removed None-check guard for variable '{var_name}'.",
|
||||
)
|
||||
return None
|
||||
@@ -0,0 +1,57 @@
|
||||
"""AST-level resource leak injection for Python."""
|
||||
|
||||
import ast
|
||||
from typing import Optional
|
||||
|
||||
from app.dataset.rules.base import Mutation, MutationRule
|
||||
|
||||
|
||||
class ResourceLeakRule(MutationRule):
|
||||
"""Convert a `with open(...)` block into an unclosed `open(...).read()`.
|
||||
|
||||
Uses `ast` to locate a with-statement managing a file resource and replaces
|
||||
it with a direct call chain that leaks the file handle.
|
||||
"""
|
||||
|
||||
name = "resource_leak"
|
||||
language = "python"
|
||||
defect_type = "resource_not_closed"
|
||||
|
||||
def detect_and_mutate(self, source: str, filename: str = "") -> Optional[Mutation]:
|
||||
try:
|
||||
tree = ast.parse(source)
|
||||
except SyntaxError:
|
||||
return None
|
||||
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, ast.With):
|
||||
continue
|
||||
first_item = node.items[0]
|
||||
ctx = first_item.context_expr
|
||||
if isinstance(ctx, ast.Call) and isinstance(ctx.func, ast.Name) and ctx.func.id == "open":
|
||||
start = node.lineno
|
||||
end = getattr(node, "end_lineno", node.lineno) or node.lineno
|
||||
lines = source.splitlines(keepends=True)
|
||||
# Keep the with header expression but replace 'with open(...)' by 'f = open(...)'
|
||||
header = lines[start - 1]
|
||||
header_expr = header.split("with ", 1)[1].split(" as ", 1)[0].strip().rstrip(":\n")
|
||||
var = header.split(" as ", 1)[1].strip().rstrip(":\n") if " as " in header else "f"
|
||||
body = lines[start:end]
|
||||
dedented = []
|
||||
for line in body[1:]:
|
||||
if line.startswith(" "):
|
||||
dedented.append(line[4:])
|
||||
else:
|
||||
dedented.append(line)
|
||||
replacement = [f"{var} = {header_expr}\n"] + dedented
|
||||
mutated = "".join(lines[: start - 1] + replacement + lines[end:])
|
||||
return Mutation(
|
||||
defect_type=self.defect_type,
|
||||
language=self.language,
|
||||
line_start=start,
|
||||
line_end=end,
|
||||
mutated_source=mutated,
|
||||
reference_fix="Use `with open(...) as f:` to ensure the file is closed.",
|
||||
description="Replaced context-managed open() with an unclosed file handle.",
|
||||
)
|
||||
return None
|
||||
@@ -0,0 +1,49 @@
|
||||
"""Registry that auto-discovers mutation rules from the rules directory."""
|
||||
|
||||
import importlib
|
||||
import inspect
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, Type
|
||||
|
||||
from app.dataset.rules.base import MutationRule
|
||||
|
||||
|
||||
class RuleRegistry:
|
||||
def __init__(self, rules_dir: Path):
|
||||
self.rules_dir = rules_dir
|
||||
self._rules: Dict[str, List[MutationRule]] = {}
|
||||
|
||||
def discover(self) -> None:
|
||||
"""Scan rules directory and register all MutationRule subclasses."""
|
||||
self._rules.clear()
|
||||
for lang_dir in self.rules_dir.iterdir():
|
||||
if not lang_dir.is_dir():
|
||||
continue
|
||||
for py_file in lang_dir.glob("*.py"):
|
||||
if py_file.name.startswith("_"):
|
||||
continue
|
||||
module_name = f"app.dataset.rules.{lang_dir.name}.{py_file.stem}"
|
||||
try:
|
||||
module = importlib.import_module(module_name)
|
||||
except Exception:
|
||||
continue
|
||||
for _, obj in inspect.getmembers(module, inspect.isclass):
|
||||
if (
|
||||
issubclass(obj, MutationRule)
|
||||
and obj is not MutationRule
|
||||
and not getattr(obj, "__abstractmethods__", False)
|
||||
):
|
||||
rule = obj()
|
||||
self._rules.setdefault(rule.language, []).append(rule)
|
||||
|
||||
def rules_for(self, language: str) -> List[MutationRule]:
|
||||
return self._rules.get(language, [])
|
||||
|
||||
def all_rules(self) -> Dict[str, List[MutationRule]]:
|
||||
return self._rules.copy()
|
||||
|
||||
|
||||
def get_registry() -> RuleRegistry:
|
||||
registry = RuleRegistry(Path(__file__).parent)
|
||||
registry.discover()
|
||||
return registry
|
||||
@@ -0,0 +1,24 @@
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker, declarative_base
|
||||
|
||||
from app.config import get_settings
|
||||
|
||||
settings = get_settings()
|
||||
|
||||
engine = create_engine(settings.database_url, pool_pre_ping=True)
|
||||
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)
|
||||
Base = declarative_base()
|
||||
|
||||
|
||||
def get_db():
|
||||
db = SessionLocal()
|
||||
try:
|
||||
yield db
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def init_db():
|
||||
# Import models so Base.metadata is populated
|
||||
from app.models import database # noqa: F401
|
||||
Base.metadata.create_all(bind=engine)
|
||||
@@ -0,0 +1,4 @@
|
||||
from app.experiments.matrix import generate_full_factorial_matrix
|
||||
from app.experiments.runner import ExperimentRunner
|
||||
|
||||
__all__ = ["generate_full_factorial_matrix", "ExperimentRunner"]
|
||||
@@ -0,0 +1,81 @@
|
||||
"""Full-factorial experiment matrix generation."""
|
||||
|
||||
from itertools import product
|
||||
from typing import Any, Dict, List
|
||||
from uuid import uuid4
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.models import Experiment, ExperimentRun, PromptTemplateVersion, Sample
|
||||
from app.prompts.service import PromptService
|
||||
|
||||
|
||||
def generate_full_factorial_matrix(
|
||||
db: Session,
|
||||
name: str,
|
||||
models: List[str],
|
||||
levels: List[str],
|
||||
sample_ids: List[str],
|
||||
repeats: int = 3,
|
||||
sampling_params: Dict[str, Any] = None,
|
||||
strategy_id: str = "code_review",
|
||||
) -> Experiment:
|
||||
"""Generate experiment runs for a full-factorial matrix.
|
||||
|
||||
Models × Levels × Samples × Repeats.
|
||||
"""
|
||||
sampling_params = sampling_params or {"temperature": 0.7, "max_tokens": 2048}
|
||||
experiment = Experiment(
|
||||
name=name,
|
||||
models=models,
|
||||
levels=levels,
|
||||
sample_ids=sample_ids,
|
||||
repeats=repeats,
|
||||
sampling_params=sampling_params,
|
||||
status="pending",
|
||||
)
|
||||
db.add(experiment)
|
||||
db.flush()
|
||||
|
||||
prompt_service = PromptService(db)
|
||||
runs = []
|
||||
for model_id, level in product(models, levels):
|
||||
# Resolve latest template version for this level
|
||||
versions = prompt_service.list_versions(strategy_id, level)
|
||||
if not versions:
|
||||
raise ValueError(f"No prompt template found for {strategy_id}/{level}")
|
||||
template_version = versions[-1]
|
||||
|
||||
for sample_id in sample_ids:
|
||||
from uuid import UUID
|
||||
|
||||
sample_uuid = UUID(sample_id) if isinstance(sample_id, str) else sample_id
|
||||
sample = db.query(Sample).filter_by(id=sample_uuid).first()
|
||||
if not sample:
|
||||
raise ValueError(f"Sample not found: {sample_id}")
|
||||
|
||||
for repeat_index in range(1, repeats + 1):
|
||||
run = ExperimentRun(
|
||||
run_id=str(uuid4()),
|
||||
experiment_id=experiment.id,
|
||||
sample_id=sample.id,
|
||||
model_id=model_id,
|
||||
template_version_id=template_version.id,
|
||||
repeat_index=repeat_index,
|
||||
status="pending",
|
||||
sampling_params=sampling_params,
|
||||
)
|
||||
runs.append(run)
|
||||
|
||||
db.add_all(runs)
|
||||
db.commit()
|
||||
db.refresh(experiment)
|
||||
return experiment
|
||||
|
||||
|
||||
def count_pending_runs(db: Session, experiment_id: str) -> int:
|
||||
return (
|
||||
db.query(ExperimentRun)
|
||||
.filter_by(experiment_id=experiment_id, status="pending")
|
||||
.count()
|
||||
)
|
||||
@@ -0,0 +1,112 @@
|
||||
"""Async experiment runner with retry, isolation, and resume support."""
|
||||
|
||||
import asyncio
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.dataset.diff_extractor import extract_diff_bundle
|
||||
from app.model_adapters.factory import create_adapter
|
||||
from app.models import ExperimentRun, Result
|
||||
from app.prompts.service import PromptService
|
||||
|
||||
|
||||
class ExperimentRunner:
|
||||
def __init__(
|
||||
self,
|
||||
db: Session,
|
||||
concurrency: int = 5,
|
||||
max_retries: int = 3,
|
||||
strategy_id: str = "code_review",
|
||||
):
|
||||
self.db = db
|
||||
self.concurrency = concurrency
|
||||
self.max_retries = max_retries
|
||||
self.strategy_id = strategy_id
|
||||
self.prompt_service = PromptService(db)
|
||||
|
||||
def _get_pending_runs(self, experiment_id: Optional[str] = None) -> List[ExperimentRun]:
|
||||
query = self.db.query(ExperimentRun)
|
||||
if experiment_id:
|
||||
query = query.filter_by(experiment_id=experiment_id)
|
||||
return query.filter(ExperimentRun.status.in_(["pending", "failed"])).all()
|
||||
|
||||
def _reset_stale_running(self) -> None:
|
||||
"""Mark runs stuck in running without result back to pending."""
|
||||
stale = (
|
||||
self.db.query(ExperimentRun)
|
||||
.filter_by(status="running")
|
||||
.filter(~ExperimentRun.result.has())
|
||||
.all()
|
||||
)
|
||||
for run in stale:
|
||||
run.status = "pending"
|
||||
self.db.commit()
|
||||
|
||||
async def run_experiment(
|
||||
self,
|
||||
experiment_id: Optional[str] = None,
|
||||
progress_callback=None,
|
||||
) -> Dict[str, Any]:
|
||||
self._reset_stale_running()
|
||||
pending = self._get_pending_runs(experiment_id)
|
||||
semaphore = asyncio.Semaphore(self.concurrency)
|
||||
|
||||
async def execute(run: ExperimentRun):
|
||||
async with semaphore:
|
||||
return await self._execute_run(run)
|
||||
|
||||
tasks = [asyncio.create_task(execute(run)) for run in pending]
|
||||
results = []
|
||||
for coro in asyncio.as_completed(tasks):
|
||||
result = await coro
|
||||
results.append(result)
|
||||
if progress_callback:
|
||||
progress_callback(result)
|
||||
|
||||
return {"total": len(pending), "completed": len(results)}
|
||||
|
||||
async def _execute_run(self, run: ExperimentRun) -> Dict[str, Any]:
|
||||
run.status = "running"
|
||||
run.started_at = datetime.now(timezone.utc)
|
||||
self.db.commit()
|
||||
|
||||
try:
|
||||
sample = run.sample
|
||||
prompt = self.prompt_service.render(
|
||||
self.strategy_id,
|
||||
run.template_version.version_number,
|
||||
{"language": sample.language, "diff": sample.diff},
|
||||
)
|
||||
adapter = create_adapter(run.model_id)
|
||||
response = await adapter.chat(prompt, run.sampling_params)
|
||||
|
||||
run.status = "done"
|
||||
run.completed_at = datetime.now(timezone.utc)
|
||||
result = Result(
|
||||
run_id=run.id,
|
||||
raw_output=response.text,
|
||||
token_usage=response.token_usage,
|
||||
latency_ms=response.latency_ms,
|
||||
)
|
||||
self.db.add(result)
|
||||
self.db.commit()
|
||||
return {
|
||||
"run_id": str(run.id),
|
||||
"status": "done",
|
||||
"model_id": run.model_id,
|
||||
}
|
||||
except Exception as e:
|
||||
run.retry_count += 1
|
||||
if run.retry_count > self.max_retries:
|
||||
run.status = "failed"
|
||||
else:
|
||||
run.status = "pending" # will be retried on next resume
|
||||
run.completed_at = datetime.now(timezone.utc)
|
||||
self.db.commit()
|
||||
return {
|
||||
"run_id": str(run.id),
|
||||
"status": run.status,
|
||||
"error": str(e),
|
||||
}
|
||||
@@ -0,0 +1,93 @@
|
||||
"""Abstract base class for model adapters."""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
import asyncio
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChatResponse:
|
||||
text: str
|
||||
token_usage: Dict[str, int]
|
||||
latency_ms: float
|
||||
|
||||
|
||||
class ModelAdapter(ABC):
|
||||
"""Unified interface for LLM vendors.
|
||||
|
||||
Subclasses only need to provide base_url, api_key, model name and any
|
||||
vendor-specific headers. Concurrency and retry logic are inherited.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str,
|
||||
model: str,
|
||||
base_url: str,
|
||||
concurrency: int = 5,
|
||||
max_retries: int = 3,
|
||||
timeout: float = 120.0,
|
||||
):
|
||||
self.api_key = api_key
|
||||
self.model = model
|
||||
self.base_url = base_url.rstrip("/")
|
||||
self.semaphore = asyncio.Semaphore(concurrency)
|
||||
self.max_retries = max_retries
|
||||
self.timeout = timeout
|
||||
|
||||
@abstractmethod
|
||||
def _build_payload(self, prompt: str, params: Dict[str, Any]) -> Dict[str, Any]:
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
def _extract_text(self, data: Dict[str, Any]) -> str:
|
||||
...
|
||||
|
||||
def _extract_token_usage(self, data: Dict[str, Any]) -> Dict[str, int]:
|
||||
usage = data.get("usage", {})
|
||||
return {
|
||||
"prompt_tokens": usage.get("prompt_tokens", 0),
|
||||
"completion_tokens": usage.get("completion_tokens", 0),
|
||||
"total_tokens": usage.get("total_tokens", 0),
|
||||
}
|
||||
|
||||
def _headers(self) -> Dict[str, str]:
|
||||
return {
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
async def chat(self, prompt: str, params: Optional[Dict[str, Any]] = None) -> ChatResponse:
|
||||
params = params or {}
|
||||
payload = self._build_payload(prompt, params)
|
||||
async with self.semaphore:
|
||||
last_exception: Optional[Exception] = None
|
||||
for attempt in range(self.max_retries + 1):
|
||||
start = time.perf_counter()
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||
response = await client.post(
|
||||
f"{self.base_url}/chat/completions",
|
||||
headers=self._headers(),
|
||||
json=payload,
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
latency_ms = (time.perf_counter() - start) * 1000
|
||||
return ChatResponse(
|
||||
text=self._extract_text(data),
|
||||
token_usage=self._extract_token_usage(data),
|
||||
latency_ms=latency_ms,
|
||||
)
|
||||
except Exception as e:
|
||||
last_exception = e
|
||||
if attempt < self.max_retries:
|
||||
wait = 2**attempt
|
||||
await asyncio.sleep(wait)
|
||||
raise RuntimeError(
|
||||
f"Model {self.model} failed after {self.max_retries} retries: {last_exception}"
|
||||
)
|
||||
@@ -0,0 +1,48 @@
|
||||
"""Factory for creating model adapters from configuration."""
|
||||
|
||||
from app.config import get_settings
|
||||
from app.model_adapters.base import ModelAdapter
|
||||
from app.model_adapters.providers import DeepSeekAdapter, KimiAdapter, QwenAdapter
|
||||
|
||||
|
||||
_ADAPTER_MAP = {
|
||||
"deepseek": DeepSeekAdapter,
|
||||
"kimi": KimiAdapter,
|
||||
"qwen": QwenAdapter,
|
||||
}
|
||||
|
||||
|
||||
def create_adapter(model_id: str, concurrency: int = 5, max_retries: int = 3) -> ModelAdapter:
|
||||
settings = get_settings()
|
||||
model_id = model_id.lower()
|
||||
adapter_cls = _ADAPTER_MAP.get(model_id)
|
||||
if not adapter_cls:
|
||||
raise ValueError(f"Unknown model_id: {model_id}. Available: {list(_ADAPTER_MAP.keys())}")
|
||||
|
||||
if model_id == "deepseek":
|
||||
return adapter_cls(
|
||||
api_key=settings.deepseek_api_key or "",
|
||||
model=settings.deepseek_model,
|
||||
base_url=settings.deepseek_base_url,
|
||||
concurrency=concurrency,
|
||||
max_retries=max_retries,
|
||||
)
|
||||
if model_id == "kimi":
|
||||
return adapter_cls(
|
||||
api_key=settings.kimi_api_key or "",
|
||||
model=settings.kimi_model,
|
||||
base_url=settings.kimi_base_url,
|
||||
concurrency=concurrency,
|
||||
max_retries=max_retries,
|
||||
)
|
||||
return adapter_cls(
|
||||
api_key=settings.qwen_api_key or "",
|
||||
model=settings.qwen_model,
|
||||
base_url=settings.qwen_base_url,
|
||||
concurrency=concurrency,
|
||||
max_retries=max_retries,
|
||||
)
|
||||
|
||||
|
||||
def list_models() -> list[str]:
|
||||
return list(_ADAPTER_MAP.keys())
|
||||
@@ -0,0 +1,57 @@
|
||||
"""Concrete model adapters for DeepSeek, Kimi, and Qwen."""
|
||||
|
||||
from typing import Any, Dict
|
||||
|
||||
from app.model_adapters.base import ModelAdapter
|
||||
|
||||
|
||||
class DeepSeekAdapter(ModelAdapter):
|
||||
def _build_payload(self, prompt: str, params: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return {
|
||||
"model": self.model,
|
||||
"messages": [{"role": "user", "content": prompt}],
|
||||
"temperature": params.get("temperature", 0.7),
|
||||
"max_tokens": params.get("max_tokens", 8192),
|
||||
# Disable vendor reasoning mode: thinking tokens would otherwise
|
||||
# exhaust max_tokens and leave `content` empty, and reasoning
|
||||
# behavior is an uncontrolled variable in the prompt-strategy
|
||||
# experiment.
|
||||
"thinking": {"type": "disabled"},
|
||||
}
|
||||
|
||||
def _extract_text(self, data: Dict[str, Any]) -> str:
|
||||
return data["choices"][0]["message"]["content"]
|
||||
|
||||
|
||||
class KimiAdapter(ModelAdapter):
|
||||
def _build_payload(self, prompt: str, params: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return {
|
||||
"model": self.model,
|
||||
"messages": [{"role": "user", "content": prompt}],
|
||||
# kimi-k2.x rejects any temperature other than 0.6.
|
||||
"temperature": params.get("temperature", 0.6),
|
||||
"max_tokens": params.get("max_tokens", 8192),
|
||||
# No `thinking` switch: kimi-k2.x rejects it, and its built-in
|
||||
# reasoning is short enough to leave room for the answer.
|
||||
}
|
||||
|
||||
def _extract_text(self, data: Dict[str, Any]) -> str:
|
||||
return data["choices"][0]["message"]["content"]
|
||||
|
||||
|
||||
class QwenAdapter(ModelAdapter):
|
||||
def _build_payload(self, prompt: str, params: Dict[str, Any]) -> Dict[str, Any]:
|
||||
return {
|
||||
"model": self.model,
|
||||
"messages": [{"role": "user", "content": prompt}],
|
||||
"temperature": params.get("temperature", 0.7),
|
||||
"max_tokens": params.get("max_tokens", 8192),
|
||||
# Disable vendor reasoning mode: thinking tokens would otherwise
|
||||
# exhaust max_tokens and leave `content` empty, and reasoning
|
||||
# behavior is an uncontrolled variable in the prompt-strategy
|
||||
# experiment.
|
||||
"thinking": {"type": "disabled"},
|
||||
}
|
||||
|
||||
def _extract_text(self, data: Dict[str, Any]) -> str:
|
||||
return data["choices"][0]["message"]["content"]
|
||||
@@ -0,0 +1,21 @@
|
||||
from app.models.database import (
|
||||
Base,
|
||||
Defect,
|
||||
Experiment,
|
||||
ExperimentRun,
|
||||
PromptTemplate,
|
||||
PromptTemplateVersion,
|
||||
Result,
|
||||
Sample,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"Base",
|
||||
"Sample",
|
||||
"Defect",
|
||||
"PromptTemplate",
|
||||
"PromptTemplateVersion",
|
||||
"Experiment",
|
||||
"ExperimentRun",
|
||||
"Result",
|
||||
]
|
||||
@@ -0,0 +1,146 @@
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy import (
|
||||
JSON,
|
||||
Column,
|
||||
DateTime,
|
||||
Float,
|
||||
ForeignKey,
|
||||
Integer,
|
||||
String,
|
||||
Text,
|
||||
UniqueConstraint,
|
||||
)
|
||||
from sqlalchemy.dialects.postgresql import UUID
|
||||
from sqlalchemy.orm import relationship
|
||||
|
||||
from app.db import Base
|
||||
|
||||
|
||||
def now_utc() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
class Sample(Base):
|
||||
__tablename__ = "samples"
|
||||
|
||||
id = Column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
|
||||
repo = Column(String(255), nullable=False)
|
||||
commit_sha = Column(String(40), nullable=False)
|
||||
language = Column(String(50), nullable=False)
|
||||
diff = Column(Text, nullable=False)
|
||||
before_context = Column(JSON, nullable=True)
|
||||
after_context = Column(JSON, nullable=True)
|
||||
created_at = Column(DateTime, default=now_utc, nullable=False)
|
||||
|
||||
defects = relationship("Defect", back_populates="sample", cascade="all, delete-orphan")
|
||||
runs = relationship("ExperimentRun", back_populates="sample")
|
||||
|
||||
|
||||
class Defect(Base):
|
||||
__tablename__ = "defects"
|
||||
|
||||
id = Column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
|
||||
sample_id = Column(UUID(as_uuid=True), ForeignKey("samples.id"), nullable=False)
|
||||
defect_type = Column(String(100), nullable=False)
|
||||
language = Column(String(50), nullable=False)
|
||||
line_start = Column(Integer, nullable=True)
|
||||
line_end = Column(Integer, nullable=True)
|
||||
description = Column(Text, nullable=True)
|
||||
reference_fix = Column(Text, nullable=True)
|
||||
created_at = Column(DateTime, default=now_utc, nullable=False)
|
||||
|
||||
sample = relationship("Sample", back_populates="defects")
|
||||
|
||||
|
||||
class PromptTemplate(Base):
|
||||
__tablename__ = "prompt_templates"
|
||||
|
||||
id = Column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
|
||||
strategy_id = Column(String(100), nullable=False)
|
||||
level = Column(String(20), nullable=False)
|
||||
created_at = Column(DateTime, default=now_utc, nullable=False)
|
||||
|
||||
versions = relationship(
|
||||
"PromptTemplateVersion",
|
||||
back_populates="template",
|
||||
cascade="all, delete-orphan",
|
||||
order_by="PromptTemplateVersion.version_number",
|
||||
)
|
||||
|
||||
__table_args__ = (UniqueConstraint("strategy_id", "level", name="uix_strategy_level"),)
|
||||
|
||||
|
||||
class PromptTemplateVersion(Base):
|
||||
__tablename__ = "prompt_template_versions"
|
||||
|
||||
id = Column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
|
||||
template_id = Column(UUID(as_uuid=True), ForeignKey("prompt_templates.id"), nullable=False)
|
||||
version_number = Column(Integer, nullable=False)
|
||||
body = Column(Text, nullable=False)
|
||||
variables_schema = Column(JSON, nullable=False, default=dict)
|
||||
created_at = Column(DateTime, default=now_utc, nullable=False)
|
||||
|
||||
template = relationship("PromptTemplate", back_populates="versions")
|
||||
runs = relationship("ExperimentRun", back_populates="template_version")
|
||||
|
||||
|
||||
class Experiment(Base):
|
||||
__tablename__ = "experiments"
|
||||
|
||||
id = Column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
|
||||
name = Column(String(255), nullable=False)
|
||||
models = Column(JSON, nullable=False)
|
||||
levels = Column(JSON, nullable=False)
|
||||
sample_ids = Column(JSON, nullable=False)
|
||||
repeats = Column(Integer, nullable=False, default=3)
|
||||
sampling_params = Column(JSON, nullable=False, default=dict)
|
||||
status = Column(String(20), default="pending", nullable=False)
|
||||
created_at = Column(DateTime, default=now_utc, nullable=False)
|
||||
|
||||
runs = relationship("ExperimentRun", back_populates="experiment", cascade="all, delete-orphan")
|
||||
|
||||
|
||||
class ExperimentRun(Base):
|
||||
__tablename__ = "experiment_runs"
|
||||
|
||||
id = Column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
|
||||
run_id = Column(String(64), unique=True, nullable=False, default=lambda: str(uuid.uuid4()))
|
||||
experiment_id = Column(UUID(as_uuid=True), ForeignKey("experiments.id"), nullable=False)
|
||||
sample_id = Column(UUID(as_uuid=True), ForeignKey("samples.id"), nullable=False)
|
||||
model_id = Column(String(100), nullable=False)
|
||||
template_version_id = Column(UUID(as_uuid=True), ForeignKey("prompt_template_versions.id"), nullable=False)
|
||||
repeat_index = Column(Integer, nullable=False)
|
||||
status = Column(String(20), default="pending", nullable=False)
|
||||
retry_count = Column(Integer, default=0, nullable=False)
|
||||
sampling_params = Column(JSON, nullable=False, default=dict)
|
||||
started_at = Column(DateTime, nullable=True)
|
||||
completed_at = Column(DateTime, nullable=True)
|
||||
created_at = Column(DateTime, default=now_utc, nullable=False)
|
||||
|
||||
experiment = relationship("Experiment", back_populates="runs")
|
||||
sample = relationship("Sample", back_populates="runs")
|
||||
template_version = relationship("PromptTemplateVersion", back_populates="runs")
|
||||
result = relationship("Result", back_populates="run", uselist=False, cascade="all, delete-orphan")
|
||||
|
||||
|
||||
class Result(Base):
|
||||
__tablename__ = "results"
|
||||
|
||||
id = Column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
|
||||
run_id = Column(UUID(as_uuid=True), ForeignKey("experiment_runs.id"), nullable=False, unique=True)
|
||||
raw_output = Column(Text, nullable=True)
|
||||
token_usage = Column(JSON, nullable=True)
|
||||
latency_ms = Column(Float, nullable=True)
|
||||
parsed_findings = Column(JSON, nullable=True)
|
||||
detection_rate = Column(Float, nullable=True)
|
||||
false_positive_rate = Column(Float, nullable=True)
|
||||
coverage_rate = Column(Float, nullable=True)
|
||||
stability_score = Column(Float, nullable=True)
|
||||
likert_score = Column(Integer, nullable=True)
|
||||
created_at = Column(DateTime, default=now_utc, nullable=False)
|
||||
updated_at = Column(DateTime, default=now_utc, onupdate=now_utc, nullable=False)
|
||||
|
||||
run = relationship("ExperimentRun", back_populates="result")
|
||||
@@ -0,0 +1,84 @@
|
||||
"""Default L1/L2/L3 prompt templates."""
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.prompts.service import PromptService
|
||||
|
||||
DEFAULT_TEMPLATES = {
|
||||
("code_review", "L1"): {
|
||||
"body": """You are a code reviewer. Review the following code diff and identify any potential bugs or issues.
|
||||
Only list what is wrong; do not provide locations or fixes.
|
||||
|
||||
Language: {{ language }}
|
||||
Diff:
|
||||
```
|
||||
{{ diff }}
|
||||
```
|
||||
|
||||
Report issues as a plain list.""",
|
||||
"variables_schema": {
|
||||
"language": {"type": "string"},
|
||||
"diff": {"type": "string"},
|
||||
},
|
||||
},
|
||||
("code_review", "L2"): {
|
||||
"body": """You are a code reviewer. Review the following code diff and identify potential bugs or issues.
|
||||
For each issue, provide:
|
||||
1. The defect type (one line)
|
||||
2. The line number range where it occurs
|
||||
3. A brief explanation
|
||||
|
||||
Language: {{ language }}
|
||||
Diff:
|
||||
```
|
||||
{{ diff }}
|
||||
```
|
||||
|
||||
Format each issue as:
|
||||
- Type: <type>
|
||||
Lines: <start>-<end>
|
||||
Explanation: <explanation>""",
|
||||
"variables_schema": {
|
||||
"language": {"type": "string"},
|
||||
"diff": {"type": "string"},
|
||||
},
|
||||
},
|
||||
("code_review", "L3"): {
|
||||
"body": """You are a code reviewer. Review the following code diff and identify potential bugs or issues.
|
||||
For each issue, provide:
|
||||
1. The defect type (one line)
|
||||
2. The line number range where it occurs
|
||||
3. A brief explanation
|
||||
4. A concrete fix suggestion
|
||||
|
||||
Language: {{ language }}
|
||||
Diff:
|
||||
```
|
||||
{{ diff }}
|
||||
```
|
||||
|
||||
Format each issue as:
|
||||
- Type: <type>
|
||||
Lines: <start>-<end>
|
||||
Explanation: <explanation>
|
||||
Fix: <fix>""",
|
||||
"variables_schema": {
|
||||
"language": {"type": "string"},
|
||||
"diff": {"type": "string"},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def seed_default_templates(db: Session) -> None:
|
||||
service = PromptService(db)
|
||||
for (strategy_id, level), data in DEFAULT_TEMPLATES.items():
|
||||
existing = service.list_versions(strategy_id, level)
|
||||
if existing:
|
||||
continue
|
||||
service.create_version(
|
||||
strategy_id=strategy_id,
|
||||
level=level,
|
||||
body=data["body"],
|
||||
variables_schema=data["variables_schema"],
|
||||
)
|
||||
@@ -0,0 +1,113 @@
|
||||
"""Prompt template storage, versioning, and rendering service."""
|
||||
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from jinja2 import BaseLoader, Environment
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.models import PromptTemplate, PromptTemplateVersion
|
||||
|
||||
|
||||
class PromptService:
|
||||
def __init__(self, db: Session):
|
||||
self.db = db
|
||||
self.jinja = Environment(loader=BaseLoader())
|
||||
|
||||
def get_or_create_template(self, strategy_id: str, level: str) -> PromptTemplate:
|
||||
template = (
|
||||
self.db.query(PromptTemplate)
|
||||
.filter_by(strategy_id=strategy_id, level=level)
|
||||
.first()
|
||||
)
|
||||
if not template:
|
||||
template = PromptTemplate(strategy_id=strategy_id, level=level)
|
||||
self.db.add(template)
|
||||
self.db.commit()
|
||||
self.db.refresh(template)
|
||||
return template
|
||||
|
||||
def create_version(
|
||||
self,
|
||||
strategy_id: str,
|
||||
level: str,
|
||||
body: str,
|
||||
variables_schema: Optional[Dict[str, Any]] = None,
|
||||
) -> PromptTemplateVersion:
|
||||
template = self.get_or_create_template(strategy_id, level)
|
||||
next_version = (
|
||||
self.db.query(PromptTemplateVersion)
|
||||
.filter_by(template_id=template.id)
|
||||
.count()
|
||||
+ 1
|
||||
)
|
||||
version = PromptTemplateVersion(
|
||||
template_id=template.id,
|
||||
version_number=next_version,
|
||||
body=body,
|
||||
variables_schema=variables_schema or self._infer_schema(body),
|
||||
)
|
||||
self.db.add(version)
|
||||
self.db.commit()
|
||||
self.db.refresh(version)
|
||||
return version
|
||||
|
||||
def get_version(self, version_id: str) -> Optional[PromptTemplateVersion]:
|
||||
from uuid import UUID
|
||||
|
||||
try:
|
||||
return self.db.query(PromptTemplateVersion).filter_by(id=UUID(version_id)).first()
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
def list_versions(self, strategy_id: str, level: str) -> List[PromptTemplateVersion]:
|
||||
template = (
|
||||
self.db.query(PromptTemplate)
|
||||
.filter_by(strategy_id=strategy_id, level=level)
|
||||
.first()
|
||||
)
|
||||
if not template:
|
||||
return []
|
||||
return (
|
||||
self.db.query(PromptTemplateVersion)
|
||||
.filter_by(template_id=template.id)
|
||||
.order_by(PromptTemplateVersion.version_number)
|
||||
.all()
|
||||
)
|
||||
|
||||
def render(
|
||||
self,
|
||||
strategy_id: str,
|
||||
version: Optional[int],
|
||||
context: Dict[str, Any],
|
||||
) -> str:
|
||||
template = (
|
||||
self.db.query(PromptTemplate)
|
||||
.filter_by(strategy_id=strategy_id)
|
||||
.first()
|
||||
)
|
||||
if not template:
|
||||
raise ValueError(f"Prompt template not found: {strategy_id}")
|
||||
|
||||
query = self.db.query(PromptTemplateVersion).filter_by(template_id=template.id)
|
||||
if version:
|
||||
version_obj = query.filter_by(version_number=version).first()
|
||||
else:
|
||||
version_obj = query.order_by(PromptTemplateVersion.version_number.desc()).first()
|
||||
|
||||
if not version_obj:
|
||||
raise ValueError(f"Prompt version not found: {strategy_id} v{version}")
|
||||
|
||||
jinja_template = self.jinja.from_string(version_obj.body)
|
||||
return jinja_template.render(**context)
|
||||
|
||||
def _infer_schema(self, body: str) -> Dict[str, Any]:
|
||||
"""Infer required variables from Jinja2 template."""
|
||||
from jinja2.meta import find_undeclared_variables
|
||||
|
||||
ast = self.jinja.parse(body)
|
||||
variables = find_undeclared_variables(ast)
|
||||
return {var: {"type": "string"} for var in variables}
|
||||
|
||||
|
||||
def get_prompt_service(db: Session) -> PromptService:
|
||||
return PromptService(db)
|
||||
Reference in New Issue
Block a user