first commit

This commit is contained in:
eeymoo
2026-09-19 12:54:45 +08:00
commit 6fc5b64077
126 changed files with 8601 additions and 0 deletions
View File
View File
+80
View File
@@ -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
+111
View File
@@ -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)
+54
View File
@@ -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
+129
View File
@@ -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),
}
+185
View File
@@ -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, []))
+58
View File
@@ -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
+36
View File
@@ -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)}
View File
+24
View File
@@ -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"}
+280
View File
@@ -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()}
View File
+135
View File
@@ -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()
+41
View File
@@ -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()
View File
+128
View File
@@ -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=""))
+94
View File
@@ -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,
}
+121
View File
@@ -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]
+30
View File
@@ -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
+49
View File
@@ -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
+24
View File
@@ -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)
+4
View File
@@ -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"]
+81
View File
@@ -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()
)
+112
View File
@@ -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),
}
+93
View File
@@ -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}"
)
+48
View File
@@ -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())
+57
View File
@@ -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"]
+21
View File
@@ -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",
]
+146
View File
@@ -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")
View File
+84
View File
@@ -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"],
)
+113
View File
@@ -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)