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
+21
View File
@@ -0,0 +1,21 @@
FROM python:3.11-slim
WORKDIR /app
ENV PYTHONDONTWRITEBYTECODE=1 \
PYTHONUNBUFFERED=1 \
PYTHONPATH=/app
RUN apt-get update && apt-get install -y --no-install-recommends \
git \
libpq-dev \
gcc \
&& rm -rf /var/lib/apt/lists/*
COPY requirements.txt /app/requirements.txt
RUN pip install --no-cache-dir -r requirements.txt
COPY backend /app/backend
COPY backend/alembic.ini /app/alembic.ini
CMD ["sh", "-c", "alembic upgrade head && uvicorn app.api.main:app --host 0.0.0.0 --port 8000"]
+40
View File
@@ -0,0 +1,40 @@
[alembic]
script_location = alembic
prepend_sys_path = .
version_path_separator = os
[post_write_hooks]
[loggers]
keys = root,sqlalchemy,alembic
[handlers]
keys = console
[formatters]
keys = generic
[logger_root]
level = WARN
handlers = console
qualname =
[logger_sqlalchemy]
level = WARN
handlers =
qualname = sqlalchemy.engine
[logger_alembic]
level = INFO
handlers =
qualname = alembic
[handler_console]
class = StreamHandler
args = (sys.stderr,)
level = NOTSET
formatter = generic
[formatter_generic]
format = %(levelname)-5.5s [%(name)s] %(message)s
datefmt = %H:%M:%S
+1
View File
@@ -0,0 +1 @@
Generic single-database configuration with an async dbapi.
+61
View File
@@ -0,0 +1,61 @@
from logging.config import fileConfig
from sqlalchemy import engine_from_config, pool
from alembic import context
import os
import sys
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from app.config import get_settings
from app.db import Base
from app.models import database # noqa: F401
settings = get_settings()
config = context.config
if config.config_file_name is not None:
fileConfig(config.config_file_name)
target_metadata = Base.metadata
def get_url():
return settings.database_url
def run_migrations_offline() -> None:
url = get_url()
context.configure(
url=url,
target_metadata=target_metadata,
literal_binds=True,
dialect_opts={"paramstyle": "named"},
)
with context.begin_transaction():
context.run_migrations()
def run_migrations_online() -> None:
configuration = config.get_section(config.config_ini_section, {})
configuration["sqlalchemy.url"] = get_url()
connectable = engine_from_config(
configuration,
prefix="sqlalchemy.",
poolclass=pool.NullPool,
)
with connectable.connect() as connection:
context.configure(connection=connection, target_metadata=target_metadata)
with context.begin_transaction():
context.run_migrations()
if context.is_offline_mode():
run_migrations_offline()
else:
run_migrations_online()
+26
View File
@@ -0,0 +1,26 @@
"""${message}
Revision ID: ${up_revision}
Revises: ${down_revision | comma,n}
Create Date: ${create_date}
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
${imports if imports else ""}
# revision identifiers, used by Alembic.
revision: str = ${repr(up_revision)}
down_revision: Union[str, None] = ${repr(down_revision)}
branch_labels: Union[str, Sequence[str], None] = ${repr(branch_labels)}
depends_on: Union[str, Sequence[str], None] = ${repr(depends_on)}
def upgrade() -> None:
${upgrades if upgrades else "pass"}
def downgrade() -> None:
${downgrades if downgrades else "pass"}
+121
View File
@@ -0,0 +1,121 @@
"""initial
Revision ID: 0001
Revises:
Create Date: 2026-08-02 00:00:00.000000
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
from sqlalchemy.dialects import postgresql
# revision identifiers, used by Alembic.
revision: str = "0001"
down_revision: Union[str, None] = None
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
op.create_table(
"samples",
sa.Column("id", postgresql.UUID(as_uuid=True), primary_key=True),
sa.Column("repo", sa.String(255), nullable=False),
sa.Column("commit_sha", sa.String(40), nullable=False),
sa.Column("language", sa.String(50), nullable=False),
sa.Column("diff", sa.Text, nullable=False),
sa.Column("before_context", sa.Text, nullable=True),
sa.Column("after_context", sa.Text, nullable=True),
sa.Column("created_at", sa.DateTime, nullable=False),
)
op.create_table(
"defects",
sa.Column("id", postgresql.UUID(as_uuid=True), primary_key=True),
sa.Column("sample_id", postgresql.UUID(as_uuid=True), sa.ForeignKey("samples.id"), nullable=False),
sa.Column("defect_type", sa.String(100), nullable=False),
sa.Column("language", sa.String(50), nullable=False),
sa.Column("line_start", sa.Integer, nullable=True),
sa.Column("line_end", sa.Integer, nullable=True),
sa.Column("description", sa.Text, nullable=True),
sa.Column("reference_fix", sa.Text, nullable=True),
sa.Column("created_at", sa.DateTime, nullable=False),
)
op.create_table(
"prompt_templates",
sa.Column("id", postgresql.UUID(as_uuid=True), primary_key=True),
sa.Column("strategy_id", sa.String(100), nullable=False),
sa.Column("level", sa.String(20), nullable=False),
sa.Column("created_at", sa.DateTime, nullable=False),
sa.UniqueConstraint("strategy_id", "level", name="uix_strategy_level"),
)
op.create_table(
"prompt_template_versions",
sa.Column("id", postgresql.UUID(as_uuid=True), primary_key=True),
sa.Column("template_id", postgresql.UUID(as_uuid=True), sa.ForeignKey("prompt_templates.id"), nullable=False),
sa.Column("version_number", sa.Integer, nullable=False),
sa.Column("body", sa.Text, nullable=False),
sa.Column("variables_schema", sa.JSON, nullable=False),
sa.Column("created_at", sa.DateTime, nullable=False),
)
op.create_table(
"experiments",
sa.Column("id", postgresql.UUID(as_uuid=True), primary_key=True),
sa.Column("name", sa.String(255), nullable=False),
sa.Column("models", sa.JSON, nullable=False),
sa.Column("levels", sa.JSON, nullable=False),
sa.Column("sample_ids", sa.JSON, nullable=False),
sa.Column("repeats", sa.Integer, nullable=False),
sa.Column("sampling_params", sa.JSON, nullable=False),
sa.Column("status", sa.String(20), nullable=False),
sa.Column("created_at", sa.DateTime, nullable=False),
)
op.create_table(
"experiment_runs",
sa.Column("id", postgresql.UUID(as_uuid=True), primary_key=True),
sa.Column("run_id", sa.String(64), unique=True, nullable=False),
sa.Column("experiment_id", postgresql.UUID(as_uuid=True), sa.ForeignKey("experiments.id"), nullable=False),
sa.Column("sample_id", postgresql.UUID(as_uuid=True), sa.ForeignKey("samples.id"), nullable=False),
sa.Column("model_id", sa.String(100), nullable=False),
sa.Column("template_version_id", postgresql.UUID(as_uuid=True), sa.ForeignKey("prompt_template_versions.id"), nullable=False),
sa.Column("repeat_index", sa.Integer, nullable=False),
sa.Column("status", sa.String(20), nullable=False),
sa.Column("retry_count", sa.Integer, nullable=False),
sa.Column("sampling_params", sa.JSON, nullable=False),
sa.Column("started_at", sa.DateTime, nullable=True),
sa.Column("completed_at", sa.DateTime, nullable=True),
sa.Column("created_at", sa.DateTime, nullable=False),
)
op.create_table(
"results",
sa.Column("id", postgresql.UUID(as_uuid=True), primary_key=True),
sa.Column("run_id", postgresql.UUID(as_uuid=True), sa.ForeignKey("experiment_runs.id"), nullable=False, unique=True),
sa.Column("raw_output", sa.Text, nullable=True),
sa.Column("token_usage", sa.JSON, nullable=True),
sa.Column("latency_ms", sa.Float, nullable=True),
sa.Column("parsed_findings", sa.JSON, nullable=True),
sa.Column("detection_rate", sa.Float, nullable=True),
sa.Column("false_positive_rate", sa.Float, nullable=True),
sa.Column("coverage_rate", sa.Float, nullable=True),
sa.Column("stability_score", sa.Float, nullable=True),
sa.Column("likert_score", sa.Integer, nullable=True),
sa.Column("created_at", sa.DateTime, nullable=False),
sa.Column("updated_at", sa.DateTime, nullable=False),
)
def downgrade() -> None:
op.drop_table("results")
op.drop_table("experiment_runs")
op.drop_table("experiments")
op.drop_table("prompt_template_versions")
op.drop_table("prompt_templates")
op.drop_table("defects")
op.drop_table("samples")
+171
View File
@@ -0,0 +1,171 @@
"""Post-experiment analysis: stability, compliance, ANOVA, charts, export.
Usage: python analysis_export.py <experiment_id> <output_dir>
- stability_score: mean pairwise token-Jaccard of the 3 repeat outputs per
(model, level, sample) group, written back to each result row.
- line-number compliance: share of runs whose review reported line numbers.
- ANOVA + paired t-tests on detection_rate across model-level groups.
- Charts (heatmap / boxplot / grouped bar) saved as PNG.
- group_summary.csv + raw run-level export for thesis chapter 5.
"""
import base64
import csv
import json
import re
import sys
import uuid
from collections import defaultdict
from pathlib import Path
from statistics import mean
from app.analysis.service import AnalysisService
from app.db import SessionLocal, init_db
from app.models import ExperimentRun
def token_set(text: str) -> set:
return set(re.findall(r"[a-zA-Z_][a-zA-Z_0-9]*", (text or "").lower()))
def jaccard(a: set, b: set) -> float:
union = a | b
return len(a & b) / len(union) if union else 1.0
def main() -> None:
experiment_id = sys.argv[1]
out_dir = Path(sys.argv[2])
out_dir.mkdir(parents=True, exist_ok=True)
init_db()
db = SessionLocal()
svc = AnalysisService(db)
runs = (
db.query(ExperimentRun)
.filter(ExperimentRun.experiment_id == uuid.UUID(experiment_id), ExperimentRun.status == "done")
.all()
)
# --- stability: token-Jaccard across repeats ---
groups = defaultdict(list)
for run in runs:
level = run.template_version.template.level
groups[(run.model_id, level, str(run.sample_id))].append(run)
for key, group_runs in groups.items():
texts = [r.result.raw_output for r in group_runs if r.result and r.result.raw_output]
if len(texts) < 2:
continue
sets = [token_set(t) for t in texts]
pairs = [(i, j) for i in range(len(sets)) for j in range(i + 1, len(sets))]
score = mean([jaccard(sets[i], sets[j]) for i, j in pairs])
for r in group_runs:
if r.result:
r.result.stability_score = round(score, 4)
db.commit()
# --- per model-level summary ---
summary = defaultdict(lambda: defaultdict(list))
compliance = defaultdict(lambda: [0, 0])
for run in runs:
level = run.template_version.template.level
key = (run.model_id, level)
res = run.result
if not res:
continue
if res.detection_rate is not None:
summary[key]["detection_rate"].append(res.detection_rate)
if res.false_positive_rate is not None:
summary[key]["false_positive_rate"].append(res.false_positive_rate)
if res.coverage_rate is not None:
summary[key]["coverage_rate"].append(res.coverage_rate)
if res.stability_score is not None:
summary[key]["stability"].append(res.stability_score)
pf = res.parsed_findings or {}
verdict = pf.get("judge_verdict") if isinstance(pf, dict) else None
if verdict:
compliance[key][1] += 1
if verdict.get("lines_reported"):
compliance[key][0] += 1
rows = []
for (model, level) in sorted(summary):
g = summary[(model, level)]
c = compliance[(model, level)]
row = {
"model": model,
"level": level,
"n": len(g["detection_rate"]),
"detection_rate": round(mean(g["detection_rate"]), 4) if g["detection_rate"] else None,
"false_positive_rate": round(mean(g["false_positive_rate"]), 4) if g["false_positive_rate"] else None,
"coverage_rate": round(mean(g["coverage_rate"]), 4) if g["coverage_rate"] else None,
"coverage_n": len(g["coverage_rate"]),
"stability": round(mean(g["stability"]), 4) if g["stability"] else None,
"line_compliance": round(c[0] / c[1], 4) if c[1] else None,
}
rows.append(row)
print(row)
with open(out_dir / "group_summary.csv", "w", newline="", encoding="utf-8-sig") as f:
writer = csv.DictWriter(f, fieldnames=list(rows[0].keys()))
writer.writeheader()
writer.writerows(rows)
# --- per-language breakdown ---
lang_summary = defaultdict(list)
for run in runs:
if run.result and run.result.detection_rate is not None:
lang_summary[(run.model_id, run.sample.language)].append(run.result.detection_rate)
lang_rows = [
{"model": m, "language": lang, "n": len(v), "detection_rate": round(mean(v), 4)}
for (m, lang), v in sorted(lang_summary.items())
]
with open(out_dir / "language_summary.csv", "w", newline="", encoding="utf-8-sig") as f:
writer = csv.DictWriter(f, fieldnames=["model", "language", "n", "detection_rate"])
writer.writeheader()
writer.writerows(lang_rows)
# --- ANOVA + paired t-tests ---
stats = {"anova": svc.run_anova(experiment_id), "paired_t_tests": {}}
models = sorted({r.model_id for r in runs})
levels = sorted({r.template_version.template.level for r in runs})
for m in models:
for a, b in [("L1", "L2"), ("L2", "L3"), ("L1", "L3")]:
if a in levels and b in levels:
key = f"{m}:{a} vs {b}"
stats["paired_t_tests"][key] = svc.run_paired_t_test(f"{m}-{a}", f"{m}-{b}", experiment_id)
with open(out_dir / "statistics.json", "w", encoding="utf-8") as f:
json.dump(stats, f, ensure_ascii=False, indent=2)
print("ANOVA:", stats["anova"])
# --- charts ---
charts = svc.generate_charts(experiment_id)
for name, b64 in charts.items():
try:
(out_dir / f"{name}.png").write_bytes(base64.b64decode(b64))
print("chart saved:", name)
except Exception as e: # noqa: BLE001
print(f"chart {name} failed: {e}")
# --- raw run-level export ---
with open(out_dir / "runs.csv", "w", newline="", encoding="utf-8-sig") as f:
writer = csv.writer(f)
writer.writerow(["run_id", "model", "level", "language", "sample_id", "repeat",
"detection_rate", "false_positive_rate", "coverage_rate",
"stability_score", "latency_ms", "total_tokens"])
for run in runs:
res = run.result
writer.writerow([
str(run.id), run.model_id, run.template_version.template.level,
run.sample.language, str(run.sample_id), run.repeat_index,
res.detection_rate, res.false_positive_rate, res.coverage_rate,
res.stability_score, round(res.latency_ms or 0), (res.token_usage or {}).get("total_tokens"),
])
print("exports written to", out_dir)
if __name__ == "__main__":
main()
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)
+23
View File
@@ -0,0 +1,23 @@
# 将 l3_scores.json 的 AI 预评分写回李克特评分表 Excel
import json
from openpyxl import load_workbook
XLSX = r"C:\Users\eeymoo\Documents\毕业论文\实验结果\L3修复建议_李克特评分表.xlsx"
scores = json.load(open("l3_scores.json", encoding="utf-8"))
wb = load_workbook(XLSX)
print("sheets:", wb.sheetnames)
ws = wb["打分表"] if "打分表" in wb.sheetnames else wb.active
# 表头在第 2 行;C 列 run_id、L 列评分、M 列备注
filled = 0
for row in ws.iter_rows(min_row=3):
run_id = row[2].value # C 列 run_id
if run_id in scores:
score, reason = scores[run_id]
row[11].value = score # L 列 评分(1-5)
row[12].value = reason # M 列 备注
filled += 1
wb.save(XLSX)
print(f"filled {filled} rows")
+112
View File
@@ -0,0 +1,112 @@
"""Batch-judge all completed runs of an experiment with the LLM judge.
Usage: python judge_all.py <experiment_id> [judge_model] [concurrency]
For every done run: the judge decides detection + false alarms (semantic,
level-agnostic), while line coverage is computed rule-based from parsed
findings. Idempotent: runs whose result already carries a judge verdict
(judge_verdict key in parsed_findings) are skipped.
"""
import asyncio
import sys
import uuid
from app.analysis.judge import (
build_judge_prompt,
coverage_from_verdict,
parse_verdict,
verdict_to_metrics,
)
from app.analysis.parser import parse_output
from app.db import SessionLocal, init_db
from app.model_adapters.factory import create_adapter
from app.models import ExperimentRun
async def main() -> None:
experiment_id = uuid.UUID(sys.argv[1])
judge_model = sys.argv[2] if len(sys.argv) > 2 else "deepseek"
concurrency = int(sys.argv[3]) if len(sys.argv) > 3 else 5
init_db()
db = SessionLocal()
adapter = create_adapter(judge_model)
runs = (
db.query(ExperimentRun)
.filter(ExperimentRun.experiment_id == experiment_id, ExperimentRun.status == "done")
.all()
)
todo = []
for run in runs:
if not run.result or not run.result.raw_output:
continue
pf = run.result.parsed_findings
if isinstance(pf, dict) and pf.get("judge_verdict"):
continue
todo.append(run)
print(f"{len(todo)} runs to judge (judge={judge_model}, concurrency={concurrency})")
semaphore = asyncio.Semaphore(concurrency)
async def judge_one(run):
async with semaphore:
gt_list = [
{
"defect_type": d.defect_type,
"line_start": d.line_start,
"line_end": d.line_end,
"description": d.description,
"reference_fix": d.reference_fix,
}
for d in run.sample.defects
]
gt = gt_list[0] if gt_list else {}
prompt = build_judge_prompt(run.result.raw_output, gt)
last_err = None
for _ in range(3):
try:
resp = await adapter.chat(prompt, {"temperature": 0.0, "max_tokens": 2048})
verdict = parse_verdict(resp.text)
if verdict:
break
except Exception as e: # noqa: BLE001
last_err = e
await asyncio.sleep(2)
else:
print(f"WARN judge failed for run {run.id}: {last_err or 'unparseable verdict'}")
return None
metrics = verdict_to_metrics(verdict, n_ground_truth=len(gt_list) or 1)
coverage = coverage_from_verdict(verdict, gt_list)
level = run.template_version.template.level
findings = parse_output(run.result.raw_output, level)
run.result.detection_rate = metrics["detection_rate"]
run.result.false_positive_rate = metrics["false_positive_rate"]
run.result.coverage_rate = coverage
run.result.parsed_findings = {
"judge_model": judge_model,
"judge_verdict": verdict,
"findings": [
{
"defect_type": f.defect_type,
"line_start": f.line_start,
"line_end": f.line_end,
"description": f.description,
}
for f in findings
],
}
db.commit()
return verdict
results = await asyncio.gather(*(judge_one(run) for run in todo))
done = sum(1 for r in results if r)
detected = sum(1 for r in results if r and r["detected"])
print(f"judged {done}/{len(todo)}; detected {detected}")
if __name__ == "__main__":
asyncio.run(main())
+1298
View File
File diff suppressed because it is too large Load Diff
+434
View File
@@ -0,0 +1,434 @@
{
"60a6e958-99e5-4d4f-ad82-88548fe4b039": [
3,
"准确指出空指针隐患,修复方向隐含(补空检查),未明示"
],
"8c00b35b-395e-4e42-8118-4589711f6b27": [
3,
"指出 null 直接传入风险,修复隐含,表述略笼统"
],
"25b63c53-9d23-41eb-bef1-82f9dbe2e39d": [
3,
"指出 NPE 风险及与原回退逻辑的差异,修复隐含"
],
"7dd615b3-0f3e-4491-a972-a82ecf74b37c": [
4,
"明确指出原始 < 写法才是正确防护,修复方向明确"
],
"79b3908a-9eb8-48e7-8c42-7784d4cd97d1": [
3,
"指出 <= 变更不当、隐含恢复 <,但越界断言不准确"
],
"241623fc-69fa-4730-bc7b-fdfe75c5c4a6": [
3,
"指向明确、隐含恢复 <,但越界理由不准确"
],
"610e852b-8a21-42eb-b708-7c71a78528a0": [
3,
"精确描述 equals 语义被破坏的后果,未点明运算符变更"
],
"3773ef9e-26df-47f1-ab54-c9e1b6f16e2d": [
4,
"明确点出 || 相对 && 的优先级问题及后果,修复方向明确"
],
"bc64e4de-fe6b-4863-bbc5-c330629e9c39": [
3,
"描述 equals 合约被破坏,未点明 &&→|| 变更"
],
"3a303e33-2510-4316-98f7-2bd318dca4eb": [
4,
"明确指出 || 替换 && 及两个具体后果,修复方向明确"
],
"04468fe8-047e-4cad-aad5-ebcca6b055f9": [
3,
"描述 contains 行为异常,未点明运算符变更"
],
"b9099ae9-230a-4e29-9909-937f482e320a": [
4,
"明确点出 || 短路导致误判,修复方向明确"
],
"fb2a55ae-ef4e-41be-b8e0-7ea3032679da": [
4,
"明确指出 &&→|| 变更及后果,修复方向明确"
],
"56196232-3e5b-4f38-8e74-e60ca6b33e8a": [
4,
"明确指出运算符变更及误用场景,修复方向明确"
],
"f101adbf-3e6a-49db-8bf6-e72f6cbf9899": [
4,
"明确指出 &&→|| 及条件被弱化的后果,修复方向明确"
],
"3bd471c9-466e-4692-a2ef-b9be2e3ead56": [
4,
"精确引用前后代码、指出逻辑反转与 TypeError,修复方向明确"
],
"9bc56114-66e9-46b3-ac1c-37e7f6848295": [
4,
"精确引用 &&→|| 变更及空值风险,修复方向明确"
],
"3666560f-b79f-45cd-a2e1-99f226d32b2d": [
4,
"精确引用变更、指出逻辑反转与异常后果,修复方向明确"
],
"13a2d2d7-9483-4381-a7fa-7e2851a7af6d": [
4,
"明确指出 &&→|| 及误合并后果,修复方向明确"
],
"b80f2457-51ef-45fe-aa3f-c7d1985e0468": [
4,
"明确指出变更及两个具体后果,修复方向明确"
],
"d1838879-aae6-46e1-a483-56b2f268f796": [
4,
"明确指出变更及非对象误合并后果,修复方向明确"
],
"27ebde97-9a09-4811-a8d1-b104329c5949": [
4,
"明确指出变更并分析分支走向,修复方向明确"
],
"efe6536c-f53c-48e9-bbbe-00f056135368": [
4,
"明确指出变更及目标被污染风险,修复方向明确"
],
"e66821dd-175b-4f58-bef4-f2a903679094": [
4,
"明确指出变更及分支逻辑变化,修复方向明确"
],
"f63dee02-7b3f-44e1-9850-b58bd847fe11": [
4,
"明确指出 and→or 及三类输入的具体后果,修复方向明确"
],
"9592df3d-b43d-4b4a-a435-d2807243c7fa": [
4,
"明确指出 chardet 空值守卫被移除及后果,修复方向明确"
],
"7991529a-53e2-4a93-9278-46399e9ad33e": [
2,
"仅提及 chardet 为 None 时 NameError 的表象,未识别守卫移除,修复不明确"
],
"6cf5d834-90f2-4cfa-914f-1e64c3523ef9": [
4,
"明确指出守卫被移除及 target 定义位置变化,修复方向明确"
],
"d9af45af-d57c-4867-8153-9dcc94583b90": [
4,
"明确指出 and→or 及 SSL 上下文被覆盖后果,修复方向明确"
],
"778b3ea1-26c0-482d-bc1b-1e116f1147b2": [
4,
"明确指出变更并详析三个分支影响,修复方向明确"
],
"5781e3d2-b79f-4ef8-b291-ff69d554c06f": [
4,
"明确指出变更及 verify=False 时的矛盾行为,修复方向明确"
],
"33552154-205c-4545-922c-dacc12f4e91c": [
4,
"明确指出空值检查被移除、else 不可达及缩进问题,修复方向明确"
],
"0e8a0198-721f-4e4f-8a63-0bfa9762bfb3": [
4,
"明确指出检查被移除及索引异常后果,修复方向明确"
],
"94b18699-5c10-4f3f-9b5c-33d70b4eb35f": [
4,
"明确指出守卫移除及三项具体后果,修复方向明确"
],
"8b0d0aaa-081f-4589-8c5b-cc5e297bd839": [
3,
"准确指出 null 传入风险,未点明守卫被移除,修复隐含"
],
"6aa57faa-2326-49a5-82fd-6a111468c328": [
4,
"明确提及原有 null 检查及其规避作用,修复方向明确"
],
"099b5e77-ae22-42ea-8546-505ec5c8e0aa": [
3,
"指出 null 未经检查直接传入,修复隐含"
],
"dd40cf99-2c0a-4275-8c90-a9dfe9b731c0": [
3,
"指出 <= 变更及零长写入问题,但定性为冗余偏弱,修复方向隐含"
],
"0b9735ad-ca6b-4601-a1a6-f2d32b64ddd3": [
3,
"指出零长写入副作用,未明确恢复 < 的方向"
],
"1e923a86-07b1-4531-9ceb-d05509079edf": [
3,
"指出 <= 导致的零长写入风险,修复方向隐含"
],
"e05c2805-828d-4e39-8ee1-01b32f154570": [
3,
"准确描述 equals 被破坏及 ClassCast 风险,未点明运算符变更"
],
"d36b6d9d-e9d6-47fe-bae9-f4ff80e76c8b": [
3,
"描述 equals 行为异常与强转风险,未点明 &&→||"
],
"fb551ed2-5bdb-4708-bdce-bf6a64eeaa36": [
4,
"明确指出 &&→|| 变更及两个后果,修复方向明确"
],
"4050c09d-b27e-4ffa-a938-f49057627bd6": [
3,
"详述 contains 行为异常与合约违反,未点明运算符"
],
"8760a46d-1bba-4a76-b8d2-44338b90bc46": [
3,
"指出条件不再要求类型与查找同时成立,语义准确但未点明 ||"
],
"e5a35dbf-b76d-4402-aab8-c7ca17cfe6f5": [
4,
"明确点出 || 短路失效导致强转异常,修复方向明确"
],
"f63431e4-6f3d-4c5a-9aed-c8fca65264d0": [
3,
"准确描述「任一满足即真」的语义变化,未点明运算符"
],
"b38dc15b-94f0-4d5a-b984-839cf485659c": [
3,
"指出不再要求两者同时支持,未点明 &&→||"
],
"b9f4982d-f3bb-4882-b903-b92b241da947": [
4,
"明确指出 &&→|| 变更及运行环境风险,修复方向明确"
],
"e5bc12ac-ac39-43fe-a1af-0d9e12c3ca2e": [
4,
"明确点出 || 短路及头合并被绕过,修复方向明确"
],
"e69d92c0-bdcf-4f0b-9659-38fd7117cd85": [
3,
"详述三种行为后果(含 TypeError 与引用别名),未点明运算符变更"
],
"441de8cd-1697-4155-90fc-1c92059980f1": [
4,
"引用 || 表达式并指出合并缺失与异常,修复方向明确"
],
"b9ce7649-abdd-403e-977b-308470bf1d16": [
3,
"指出仅单方满足即触发合并,未点明 &&→||"
],
"e587a125-3029-4e71-be4f-9c450b2d9300": [
3,
"详析分支不可达与类型错配,未点明运算符"
],
"abb8b86a-bc7f-44cc-b7ae-4748736479e7": [
3,
"指出单方满足即触发及分支失效,未点明运算符"
],
"0824d180-dba6-46e5-9965-5ad0afc5fcc2": [
4,
"明确指出 &&→|| 及三种后果,修复方向明确"
],
"6a9057fb-8f63-4b50-9734-abbd45c6d559": [
4,
"明确点出新 || 条件及语义破坏,修复方向明确"
],
"440f5d47-9cb8-4045-a598-dff736a8ddb7": [
4,
"明确指出 &&→|| 及分支不可达,修复方向明确"
],
"58075bc7-dffa-4179-a710-f9a464a9477a": [
4,
"明确指出 or 替换 and 及解包风险,修复方向明确"
],
"4d1cdddc-6d50-4cef-9070-f6e190ac8f95": [
4,
"引用新 or 条件并指出元组长度校验丢失,修复方向明确"
],
"c623d3cc-77fc-478c-9825-40e47f350111": [
3,
"描述 length-2 字符串误配等后果,未点明 and→or"
],
"b019bc4e-2588-4126-8287-05fc9a8cccc0": [
4,
"点出 or 导致条件放宽及 IndexError 机制,修复方向隐含明确"
],
"61113778-bd49-4f4d-8539-3947b6bb9b08": [
4,
"明确写出 and 改为 or 及三类错误输入后果"
],
"5a764dd0-66ad-4b2c-9d44-2640d94959ca": [
4,
"明确指出 if chardet is not None 守卫被移除及 NameError 后果"
],
"9246543f-13d8-44de-a056-4f96d4cb5dcf": [
4,
"指出守卫移除导致循环无条件执行,并说明映射逻辑丢失"
],
"8435fe89-1c9d-4752-a387-a1cb7c57da15": [
4,
"明确指出 chardet is not None 守卫被移除及失败场景"
],
"7586c40e-cb83-4a86-983f-f59c5e6668c6": [
4,
"引用新 or 条件并精确说明两类误入分支的后果"
],
"aeba23c7-8be9-45cf-a578-00fd40bc6da7": [
3,
"准确描述条件触发范围扩大的后果,未点明 and/or 运算符"
],
"d61481bb-e7cf-40fc-a5bb-778e528d1e8a": [
4,
"明确写出 and 改 or 破坏守卫及正确语义"
],
"44bc02fc-14cf-4557-999e-304a911d8478": [
4,
"明确指出 None 守卫被移除及无条件索引后果"
],
"35357d01-5e38-4edd-aef2-d77bd60569c2": [
4,
"指出 client_cert 不再检查 None 及字符串误索引问题"
],
"28d66204-a406-4739-8622-b5ffdae05c7b": [
4,
"明确指出 None 守卫移除及元组校验缺失后果"
],
"5c699718-681e-4eda-b504-8fa6b6e94530": [
2,
"仅一句笼统的 NPE 提示,未说明成因也无修复指向"
],
"77c54e11-d2a7-4dbf-b5f2-164f249b3550": [
4,
"明确指出守卫移除及 fallback 失效机制"
],
"3ac7d755-afd7-4aa7-a705-6f6cdbb4ad60": [
3,
"指出 NPE 与不可达代码后果,但未点明守卫移除机制"
],
"7ee913bd-7d41-46ee-a154-68d5395052cc": [
3,
"指出边界条件越界(off-by-one)后果,修复仅隐含"
],
"a4d4a3a6-0f15-4f6f-8ef3-0691cf787974": [
3,
"指出 <= 越界后果,修复方向隐含"
],
"5b33523e-6441-4ce0-b206-5916650b1206": [
3,
"描述数组越界后果准确,未给出明确修复"
],
"e439022d-791a-4079-8657-5f7268f05ec9": [
4,
"明确指出 || 替换 && 的运算符变化"
],
"53bd06a1-8d20-4222-a89e-12e4e0ceb8aa": [
4,
"明确指出 && 变 || 及短路语义破坏"
],
"e528d730-288e-434b-bd44-49801536c816": [
4,
"明确点出逻辑运算符替换及后果"
],
"c8cb246f-a361-455d-b9e6-b491ba729bb3": [
4,
"明确指出 && 改 || 及错误语义"
],
"8e5bde72-88c5-4380-b69b-f337858ef808": [
4,
"明确指出 &&→|| 变化及潜在 ClassCastException"
],
"00d028e4-8adb-4ca5-a8d7-a4ede36967d6": [
4,
"明确指出 && 改 || 及类型转换异常后果"
],
"1d33128b-1898-4c8c-af61-2b62de022659": [
4,
"明确指出 && 改 || 导致特性检测错误"
],
"b284228f-9ede-433a-a9db-df43627a2328": [
4,
"明确说明 && 变 || 及正确依赖关系"
],
"478aabc9-8811-4374-ad7d-22082a763a1a": [
4,
"明确指出 || 误用代替 &&"
],
"3314cf51-90c3-4ed0-93e0-b0ae95ef5286": [
3,
"准确描述短路返回值变化及运行时错误后果,未点明运算符"
],
"d55bd2c3-0394-494c-b4fe-18a2191384ca": [
3,
"准确描述 headers 为假时的错误访问后果(裁判未判出但内容命中)"
],
"71cb0124-fc11-4b9e-b99b-b632b3bdfc79": [
3,
"准确描述绕过 merge 与 undefined 访问后果,未点明运算符"
],
"904cd26b-f68f-4792-95be-75928e5df864": [
4,
"明确指出 && 改 || 及单方满足即合并的机制"
],
"5d03a902-90df-4bc4-942d-4b3e7daa07af": [
4,
"明确指出 &&→|| 及错误合并行为"
],
"920a0a2e-adbf-4fec-b9a1-843eb3c88448": [
4,
"明确指出运算符变化并分析分支可达性影响"
],
"b6ed21e1-581e-474d-9c84-13b3eea57c95": [
4,
"明确指出 && 改 || 及冗余分支后果"
],
"0635f620-0e27-4803-8411-580610ebdc10": [
4,
"明确指出 OR 误用代替 AND"
],
"bc8324f6-65ba-41e2-a2e9-85f22fd35b31": [
4,
"明确指出 &&→|| 并列出冗余/不可达代码等四层后果"
],
"ea252238-71e2-41ae-896f-02df20cc20d5": [
4,
"引用含 or 的错误条件表达式并说明 TypeError 机制"
],
"35f3b71f-2151-4dd5-a0c8-0529f1ebe0ea": [
4,
"明确指出逻辑 OR 代替 AND 及错误索引后果"
],
"d1b85c7a-8d30-4140-be90-9c658bcba907": [
4,
"引用含 or 的条件并说明 TypeError 触发路径"
],
"aae79429-3ce1-4256-b407-aae16cbbe011": [
3,
"描述条件块移除后 target 未定义及 chardet 为 None 的后果,未点明守卫机制"
],
"b56f28e4-aeda-4ce1-b675-d59fe2ea442d": [
4,
"明确指出缺少 null check 及 UnboundLocalError 成因"
],
"04e141ba-d32f-4021-ad22-32d44d368207": [
3,
"指出变量未定义导致 NameError 后果,较简略未点明守卫"
],
"a01d49c7-a15f-492a-af64-218316df945c": [
4,
"明确写出 and 改 or 及 ssl_context 误设置机制"
],
"df1d7bc4-2d92-46c7-ab2b-4a59789d828e": [
4,
"引用含 or 的新条件并说明误触发场景"
],
"5573b2cb-1ee2-40ff-a2da-579c1bdb01a5": [
4,
"明确指出 and 改 or 扩大条件适用范围"
],
"510180b2-fbed-4fcf-a49b-773948fef10b": [
4,
"明确指出移除了 client_cert is not None 检查"
],
"4c71b118-0259-4c97-9727-2773d1df106c": [
3,
"准确描述无条件索引与 else 赋 None 后果,未点明守卫移除"
],
"8c550cdc-e3bb-447a-9f6d-ec581cee798b": [
4,
"明确指出不再检查 client_cert 是否为 None 及校验缺失"
]
}
+261
View File
@@ -0,0 +1,261 @@
"""Generate the Likert-5 scoring sheet for L3 fix-suggestion actionability.
Reads the 108 L3 runs (3 models x 12 samples x 3 repeats) from the
experiment database and produces a workbook with:
Sheet 1 封面: title, key facts, sheet index, notes
Sheet 2 评分说明: the 1-5 rubric and scoring guidelines
Sheet 3 打分表: one row per L3 output, dropdown-validated score column
"""
import uuid
from pathlib import Path
from openpyxl import Workbook
from openpyxl.styles import Alignment, Border, Font, PatternFill, Side
from openpyxl.utils import get_column_letter
from openpyxl.worksheet.datavalidation import DataValidation
from app.db import SessionLocal, init_db
from app.models import ExperimentRun
EXPERIMENT_ID = "42a452df-6091-47c3-8632-a9ec29815aaa"
OUTPUT = Path(r"C:\Users\eeymoo\Documents\毕业论文\实验结果\L3修复建议_李克特评分表.xlsx")
GREY_FILL = PatternFill(start_color="F5F5F5", end_color="F5F5F5", fill_type="solid")
HEADER_FILL = PatternFill(start_color="333333", end_color="333333", fill_type="solid")
BLUE_FILL = PatternFill(start_color="E6F0FA", end_color="E6F0FA", fill_type="solid")
THIN = Side(style="thin", color="D0D0D0")
BORDER = Border(left=THIN, right=THIN, top=THIN, bottom=THIN)
def collect_l3_runs():
init_db()
db = SessionLocal()
runs = (
db.query(ExperimentRun)
.filter(ExperimentRun.experiment_id == uuid.UUID(EXPERIMENT_ID), ExperimentRun.status == "done")
.all()
)
l3 = [r for r in runs if r.template_version.template.level == "L3"]
l3.sort(key=lambda r: (r.model_id, r.sample.language, str(r.sample_id), r.repeat_index))
return l3
def style_header(ws, row, cols):
for col in cols:
c = ws.cell(row=row, column=col)
c.font = Font(bold=True, color="FFFFFF", size=11)
c.fill = HEADER_FILL
c.alignment = Alignment(horizontal="center", vertical="center")
def build_cover(wb, n_rows):
ws = wb.active
ws.title = "封面"
ws.sheet_view.showGridLines = False
ws.column_dimensions["A"].width = 3
for col, w in zip("BCDEF", (22, 26, 26, 26, 26)):
ws.column_dimensions[col].width = w
ws.merge_cells("B2:F2")
ws["B2"] = "L3 修复建议可操作性评分表(李克特五级量表)"
ws["B2"].font = Font(size=18, bold=True)
ws["B2"].alignment = Alignment(horizontal="center", vertical="center")
ws.row_dimensions[2].height = 38
ws.merge_cells("B3:F3")
ws["B3"] = "基于大模型的智能代码审查提示词优化研究 · 实验 full-factorial-v1"
ws["B3"].font = Font(size=11, color="666666")
ws["B3"].alignment = Alignment(horizontal="center")
ws["B5"] = "关键信息"
ws["B5"].font = Font(bold=True, size=12)
facts = [
("待评分输出", f"{n_rows} 条(3 模型 × 12 样本 × 3 次重复的 L3 级输出)"),
("评分维度", "建议可操作性(李克特 1–5 级)"),
("评分对象", "模型给出的修复建议,对照 Ground Truth 与参考修复"),
("数据来源", "实验数据库 promptcr.db(experiment: full-factorial-v1)"),
]
r = 6
for k, v in facts:
ws.cell(row=r, column=2, value=k).font = Font(bold=True)
ws.cell(row=r, column=2).fill = BLUE_FILL
ws.cell(row=r, column=2).border = BORDER
ws.merge_cells(start_row=r, start_column=3, end_row=r, end_column=6)
ws.cell(row=r, column=3, value=v).border = BORDER
for col in range(3, 7):
ws.cell(row=r, column=col).border = BORDER
r += 1
ws["B11"] = "工作表索引"
ws["B11"].font = Font(bold=True, size=12)
index = [
("评分说明", "李克特 1–5 级判据与评分注意事项(打钩前必读)"),
("打分表", f"{n_rows} 条 L3 输出逐条打分,评分列为下拉选择 1–5"),
]
r = 12
for name, desc in index:
ws.cell(row=r, column=2, value=name).font = Font(bold=True, color="0066CC")
ws.cell(row=r, column=2).border = BORDER
ws.merge_cells(start_row=r, start_column=3, end_row=r, end_column=6)
ws.cell(row=r, column=3, value=desc)
for col in range(3, 7):
ws.cell(row=r, column=col).border = BORDER
r += 1
ws["B15"] = "备注"
ws["B15"].font = Font(bold=True, size=12)
notes = [
"1. 评分前请先阅读「评分说明」工作表,保持前后判据一致。",
"2. run_id 与实验数据库中的实验单元一一对应,便于回溯原始输出与指标。",
"3. 评分完成后,将按模型 × 级别做描述统计与模糊综合评判(参考亓莱滨方法)。",
"4. 如中途打断,可分多次填写,评分列留空视为未评。",
]
r = 16
for note in notes:
ws.merge_cells(start_row=r, start_column=2, end_row=r, end_column=6)
ws.cell(row=r, column=2, value=note).font = Font(size=10, color="666666")
r += 1
def build_rubric(wb):
ws = wb.create_sheet("评分说明")
ws.sheet_view.showGridLines = False
ws.column_dimensions["A"].width = 3
ws.column_dimensions["B"].width = 8
ws.column_dimensions["C"].width = 18
ws.column_dimensions["D"].width = 80
ws.merge_cells("B2:D2")
ws["B2"] = "建议可操作性 · 李克特五级量表判据"
ws["B2"].font = Font(size=15, bold=True)
ws["B2"].alignment = Alignment(horizontal="center", vertical="center")
ws.row_dimensions[2].height = 30
header_row = 4
ws.cell(row=header_row, column=2, value="分值")
ws.cell(row=header_row, column=3, value="等级")
ws.cell(row=header_row, column=4, value="判据")
style_header(ws, header_row, (2, 3, 4))
rubric = [
(5, "完全可操作", "建议可直接应用:能针对预埋缺陷给出正确的修复方式与具体代码/步骤,无需修改即可采纳。"),
(4, "基本可操作", "建议方向正确、内容具体,仅需少量人工调整(如修正细节、补全边界)即可应用。"),
(3, "部分可操作", "建议方向基本正确,但存在明显缺漏或不够具体,需较多人工补充、验证后才能实施。"),
(2, "可操作性差", "建议笼统模糊(如仅说「建议检查逻辑」「注意空指针」),或与预埋缺陷仅部分相关,难以直接指导修复。"),
(1, "不可操作", "建议与预埋缺陷无关、明显错误、自相矛盾,或未给出任何修复建议。"),
]
r = header_row + 1
for score, label, desc in rubric:
ws.cell(row=r, column=2, value=score).alignment = Alignment(horizontal="center", vertical="center")
ws.cell(row=r, column=2).font = Font(bold=True, size=12, color="0066CC")
ws.cell(row=r, column=3, value=label).font = Font(bold=True)
ws.cell(row=r, column=3).alignment = Alignment(vertical="center")
cell = ws.cell(row=r, column=4, value=desc)
cell.alignment = Alignment(wrap_text=True, vertical="center")
for col in (2, 3, 4):
ws.cell(row=r, column=col).border = BORDER
if score % 2 == 1:
if ws.cell(row=r, column=col).fill.start_color.rgb in (None, "00000000"):
ws.cell(row=r, column=col).fill = GREY_FILL
ws.row_dimensions[r].height = 42
r += 1
r += 1
ws.cell(row=r, column=2, value="评分注意事项").font = Font(bold=True, size=12)
notes = [
"1. 仅评价修复建议的可操作性,不评价缺陷检出是否正确(检出正确性由检出率指标衡量)。",
"2. 若一条输出包含多个问题的建议,以针对预埋缺陷(Ground Truth 列所示)的那条建议为评分对象。",
"3. 对照「参考修复」判断建议的正确性,但措辞不要求一致,语义等价即可。",
"4. 若模型未检出预埋缺陷(建议全部指向其他问题),评 1 分。",
"5. 评分时保持标准前后一致;建议先抽 5 条试评,校准后再正式评分。",
]
r += 1
for note in notes:
ws.merge_cells(start_row=r, start_column=2, end_row=r, end_column=4)
cell = ws.cell(row=r, column=2, value=note)
cell.font = Font(size=10, color="666666")
cell.alignment = Alignment(wrap_text=True, vertical="center")
ws.row_dimensions[r].height = 28
r += 1
def build_scoring(wb, runs):
ws = wb.create_sheet("打分表")
ws.sheet_view.showGridLines = False
headers = [
("序号", 6), ("run_id", 34), ("模型", 10), ("语言", 12), ("仓库", 10),
("commit", 10), ("重复", 6), ("Ground Truth 缺陷", 36), ("参考修复", 36),
("模型 L3 输出(含修复建议)", 70), ("评分(1-5)", 10), ("备注", 14),
]
header_row = 2
for i, (title, width) in enumerate(headers, start=2):
col = get_column_letter(i)
ws.column_dimensions[col].width = width
ws.cell(row=header_row, column=i, value=title)
ws.column_dimensions["A"].width = 2
style_header(ws, header_row, range(2, 2 + len(headers)))
ws.row_dimensions[header_row].height = 24
ws.freeze_panes = "B3"
dv = DataValidation(type="list", formula1='"1,2,3,4,5"', allow_blank=True, showDropDown=False)
dv.error = "请选择 1-5 的整数分值"
dv.errorTitle = "无效评分"
ws.add_data_validation(dv)
r = header_row + 1
for idx, run in enumerate(runs, start=1):
gt = run.sample.defects[0] if run.sample.defects else None
values = [
idx,
str(run.id),
run.model_id,
run.sample.language,
run.sample.repo,
run.sample.commit_sha[:7],
run.repeat_index,
(gt.description or "") if gt else "",
(gt.reference_fix or "") if gt else "",
run.result.raw_output or "",
None,
None,
]
for i, v in enumerate(values, start=2):
cell = ws.cell(row=r, column=i, value=v)
cell.border = BORDER
if i in (8, 9, 10, 11): # GT / 参考修复 / 模型输出 → wrap
cell.alignment = Alignment(wrap_text=True, vertical="top")
else:
cell.alignment = Alignment(horizontal="center", vertical="top")
ws.cell(row=r, column=12).fill = BLUE_FILL # 评分列高亮
dv.add(ws.cell(row=r, column=12))
ws.row_dimensions[r].height = 110
r += 1
ws.auto_filter.ref = f"B{header_row}:M{r - 1}"
def main():
runs = collect_l3_runs()
print(f"L3 runs: {len(runs)}")
wb = Workbook()
build_cover(wb, len(runs))
build_rubric(wb)
build_scoring(wb, runs)
OUTPUT.parent.mkdir(parents=True, exist_ok=True)
wb.save(OUTPUT)
print("saved:", OUTPUT)
# 校验:重开文件确认结构
from openpyxl import load_workbook
wb2 = load_workbook(OUTPUT)
assert wb2.sheetnames == ["封面", "评分说明", "打分表"], wb2.sheetnames
ws = wb2["打分表"]
assert ws.max_row == len(runs) + 2, (ws.max_row, len(runs))
assert ws["L2"].value == "评分(1-5)"
print("verify ok: sheets =", wb2.sheetnames, "| data rows =", ws.max_row - 2)
if __name__ == "__main__":
main()
+5
View File
@@ -0,0 +1,5 @@
[pytest]
pythonpath = .
testpaths = tests
addopts = --cov=app --cov-report=term-missing --cov-report=html --cov-fail-under=50
asyncio_mode = auto
+68
View File
@@ -0,0 +1,68 @@
"""Resume experiment runs filtered by model with custom concurrency.
Usage: python run_by_model.py <experiment_id> <model_id> <concurrency>
Reuses ExperimentRunner internals; safe to re-run (stale running runs are
reset to pending by the caller beforehand via runner.run_experiment normally —
here we handle pending/failed only, plus stale running without result).
"""
import asyncio
import sys
import uuid
from app.db import SessionLocal, init_db
from app.experiments.runner import ExperimentRunner
from app.models import ExperimentRun
async def main() -> None:
experiment_id = uuid.UUID(sys.argv[1])
model_id = sys.argv[2]
concurrency = int(sys.argv[3]) if len(sys.argv) > 3 else 4
init_db()
db = SessionLocal()
runner = ExperimentRunner(db)
# Reclaim runs orphaned in "running" by previously killed processes.
stale = (
db.query(ExperimentRun)
.filter_by(status="running", model_id=model_id)
.filter(~ExperimentRun.result.has())
.all()
)
for run in stale:
run.status = "pending"
db.commit()
runs = (
db.query(ExperimentRun)
.filter(
ExperimentRun.experiment_id == experiment_id,
ExperimentRun.model_id == model_id,
ExperimentRun.status.in_(["pending", "failed"]),
)
.all()
)
print(f"{model_id}: {len(runs)} runs to execute, concurrency={concurrency}")
semaphore = asyncio.Semaphore(concurrency)
async def execute(run):
async with semaphore:
return await runner._execute_run(run)
tasks = [asyncio.create_task(execute(run)) for run in runs]
done = 0
for coro in asyncio.as_completed(tasks):
result = await coro
done += 1
if result["status"] != "done":
print("WARN:", result)
if done % 10 == 0:
print(f"progress: {done}/{len(tasks)}")
print(f"{model_id}: completed {done}/{len(tasks)}")
if __name__ == "__main__":
asyncio.run(main())
+19
View File
@@ -0,0 +1,19 @@
"""Shared pytest fixtures."""
import pytest
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from app.db import Base
@pytest.fixture
def db():
engine = create_engine("sqlite:///:memory:")
Base.metadata.create_all(bind=engine)
Session = sessionmaker(bind=engine)
session = Session()
try:
yield session
finally:
session.close()
+28
View File
@@ -0,0 +1,28 @@
"""Tests for analysis parser and metrics."""
from app.analysis.parser import compare_findings, parse_output
def test_parse_l1_output():
output = "- null pointer\n- off by one"
findings = parse_output(output, "L1")
assert len(findings) == 2
assert findings[0].defect_type == "null pointer"
def test_parse_l2_output():
output = """- Type: null_pointer
Lines: 5-6
Explanation: missing guard"""
findings = parse_output(output, "L2")
assert len(findings) == 1
assert findings[0].line_start == 5
assert findings[0].line_end == 6
def test_compare_findings():
findings = parse_output("- null_pointer", "L1")
gt = [{"id": "1", "defect_type": "null_pointer", "line_start": 5, "line_end": 5}]
metrics = compare_findings(findings, gt)
assert metrics["detection_rate"] == 1.0
assert metrics["coverage_rate"] == 1.0
+61
View File
@@ -0,0 +1,61 @@
"""Tests for analysis service."""
import pytest
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from app.analysis.service import AnalysisService
from app.db import Base
from app.models import Defect, ExperimentRun, PromptTemplate, PromptTemplateVersion, Result, Sample
def get_test_db():
engine = create_engine("sqlite:///:memory:")
Base.metadata.create_all(bind=engine)
Session = sessionmaker(bind=engine)
return Session()
def test_compute_metrics_for_run():
db = get_test_db()
sample = Sample(repo="r", commit_sha="abc", language="python", diff="diff")
defect = Defect(
sample=sample,
defect_type="null_pointer",
language="python",
line_start=5,
line_end=5,
)
db.add(sample)
db.commit()
from uuid import uuid4
template = PromptTemplate(strategy_id="code_review", level="L1")
version = PromptTemplateVersion(
template=template, version_number=1, body="", variables_schema={}
)
db.add(template)
db.commit()
run = ExperimentRun(
id=uuid4(),
experiment_id=uuid4(),
sample_id=sample.id,
model_id="deepseek",
template_version_id=version.id,
repeat_index=1,
status="done",
)
db.add(run)
db.flush()
result = Result(
run_id=run.id,
raw_output="- null_pointer",
)
db.add(result)
db.commit()
service = AnalysisService(db)
metrics = service.compute_metrics_for_run(str(run.id))
assert metrics["detection_rate"] == 1.0
+42
View File
@@ -0,0 +1,42 @@
"""Tests for FastAPI endpoints."""
import pytest
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from app.api.main import app
from app.db import Base, get_db
def get_test_db():
engine = create_engine("sqlite:///:memory:")
Base.metadata.create_all(bind=engine)
Session = sessionmaker(bind=engine)
db = Session()
try:
yield db
finally:
db.close()
app.dependency_overrides[get_db] = get_test_db
client = TestClient(app)
def test_health():
response = client.get("/health")
assert response.status_code == 200
assert response.json()["status"] == "ok"
def test_list_samples_empty():
response = client.get("/api/datasets/samples")
assert response.status_code == 200
assert response.json() == []
def test_list_models():
response = client.get("/api/models")
assert response.status_code == 200
assert "deepseek" in response.json()["models"]
+21
View File
@@ -0,0 +1,21 @@
"""Tests for chart generation."""
from app.analysis.charts import boxplot, grouped_bar, heatmap
def test_heatmap_returns_base64():
data = {"m1": {"L1": 0.5, "L2": 0.8}}
b64 = heatmap(data, title="Test")
assert b64.startswith("iVBOR") or b64.startswith("/9j")
def test_boxplot_returns_base64():
data = {"m1-L1": [0.1, 0.2, 0.3]}
b64 = boxplot(data, title="Test")
assert isinstance(b64, str) and len(b64) > 0
def test_grouped_bar_returns_base64():
data = {"m1": {"L1": 0.5}}
b64 = grouped_bar(data, title="Test")
assert isinstance(b64, str) and len(b64) > 0
+46
View File
@@ -0,0 +1,46 @@
"""Tests for experiment matrix and runner."""
import pytest
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from app.db import Base
from app.experiments.matrix import generate_full_factorial_matrix
from app.model_adapters.factory import create_adapter, list_models
from app.prompts.defaults import seed_default_templates
def get_test_db():
engine = create_engine("sqlite:///:memory:")
Base.metadata.create_all(bind=engine)
Session = sessionmaker(bind=engine)
return Session()
def test_list_models():
assert "deepseek" in list_models()
def test_create_adapter_requires_key():
with pytest.raises(ValueError):
create_adapter("unknown")
def test_generate_matrix():
db = get_test_db()
seed_default_templates(db)
from app.models import Sample
sample = Sample(repo="test", commit_sha="abc", language="python", diff="diff")
db.add(sample)
db.commit()
experiment = generate_full_factorial_matrix(
db,
name="test-exp",
models=["deepseek"],
levels=["L1"],
sample_ids=[str(sample.id)],
repeats=2,
)
assert len(experiment.runs) == 2
+40
View File
@@ -0,0 +1,40 @@
"""Tests for Git parser and candidate selection."""
import os
from pathlib import Path
import pytest
from app.dataset.git_parser import score_commit, select_candidates
def _make_repo(tmp_path: Path) -> str:
repo = tmp_path / "repo"
repo.mkdir()
os.system(
f'cd {repo} && git init -q && git config user.email "test@example.com" && git config user.name "Test"'
)
(repo / "main.py").write_text("print('hello')\n")
os.system(f'cd {repo} && git add . && git commit -q -m "init"')
(repo / "main.py").write_text("print('world')\n")
os.system(f'cd {repo} && git add . && git commit -q -m "update main"')
(repo / "app.java").write_text("class App {}\n")
os.system(f'cd {repo} && git add . && git commit -q -m "add java app"')
return str(repo)
def test_score_commit_prefers_message_and_size(tmp_path):
repo = _make_repo(tmp_path)
from app.dataset.git_parser import list_commits
commits = list_commits(repo)
assert len(commits) >= 2
scores = [score_commit(c) for c in commits]
assert all(0 <= s <= 1 for s in scores)
def test_select_candidates(tmp_path):
repo = _make_repo(tmp_path)
candidates = select_candidates(repo, count=2, languages=["python", "java"])
assert len(candidates) <= 2
assert all(c.repo == "repo" for c in candidates)
+26
View File
@@ -0,0 +1,26 @@
"""Tests for model adapters."""
import pytest
from app.model_adapters.base import ModelAdapter, ChatResponse
from app.model_adapters.providers import DeepSeekAdapter
def test_adapter_payload_and_extraction():
adapter = DeepSeekAdapter(api_key="test", model="deepseek-chat", base_url="https://x")
payload = adapter._build_payload("hello", {"temperature": 0.5, "max_tokens": 100})
assert payload["model"] == "deepseek-chat"
assert payload["messages"][0]["content"] == "hello"
response_text = adapter._extract_text(
{"choices": [{"message": {"content": "hi"}}]}
)
assert response_text == "hi"
def test_adapter_token_usage():
adapter = DeepSeekAdapter(api_key="test", model="m", base_url="https://x")
usage = adapter._extract_token_usage(
{"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}}
)
assert usage["total_tokens"] == 15
+36
View File
@@ -0,0 +1,36 @@
"""Tests for prompt template service."""
import pytest
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from app.db import Base, init_db
from app.prompts.defaults import seed_default_templates
from app.prompts.service import PromptService
def get_test_db():
engine = create_engine("sqlite:///:memory:")
Base.metadata.create_all(bind=engine)
Session = sessionmaker(bind=engine)
return Session()
def test_seed_and_render():
db = get_test_db()
seed_default_templates(db)
service = PromptService(db)
rendered = service.render("code_review", None, {"language": "python", "diff": "+x"})
assert "python" in rendered
assert "+x" in rendered
def test_create_version_increments():
db = get_test_db()
service = PromptService(db)
v1 = service.create_version("test", "L1", "hello {{ name }}")
v2 = service.create_version("test", "L1", "hello {{ name }} v2")
assert v1.version_number == 1
assert v2.version_number == 2
versions = service.list_versions("test", "L1")
assert len(versions) == 2
+44
View File
@@ -0,0 +1,44 @@
"""Tests for mutation rules."""
from app.dataset.rules.python.null_pointer import NoneReferenceRule
from app.dataset.rules.python.boundary_error import BoundaryErrorRule
from app.dataset.rules.python.logic_operator import LogicOperatorRule
from app.dataset.rules.registry import get_registry
def test_python_none_reference():
src = """def process(data):
if data is not None:
x = data.upper()
return x
"""
m = NoneReferenceRule().detect_and_mutate(src)
assert m is not None
assert m.defect_type == "null_pointer"
assert "if data is not None" not in m.mutated_source
def test_python_boundary_error():
src = """def get(items, i):
if i < len(items):
return items[i]
"""
m = BoundaryErrorRule().detect_and_mutate(src)
assert m is not None
assert "<=" in m.mutated_source
def test_python_logic_operator():
src = """def ok(a, b):
return a and b
"""
m = LogicOperatorRule().detect_and_mutate(src)
assert m is not None
assert " or " in m.mutated_source
def test_registry_discovers_rules():
registry = get_registry()
assert "python" in registry.all_rules()
assert "java" in registry.all_rules()
assert "javascript" in registry.all_rules()
+62
View File
@@ -0,0 +1,62 @@
"""Tests for experiment runner."""
from unittest import mock
from uuid import uuid4
import pytest
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from app.db import Base
from app.experiments.runner import ExperimentRunner
from app.model_adapters.base import ChatResponse
from app.models import ExperimentRun, PromptTemplate, PromptTemplateVersion, Result, Sample
def get_test_db():
engine = create_engine("sqlite:///:memory:")
Base.metadata.create_all(bind=engine)
Session = sessionmaker(bind=engine)
return Session()
@pytest.mark.asyncio
async def test_runner_executes_pending_run():
db = get_test_db()
template = PromptTemplate(strategy_id="code_review", level="L1")
version = PromptTemplateVersion(
template=template, version_number=1, body="", variables_schema={}
)
db.add(template)
db.commit()
sample = Sample(repo="r", commit_sha="abc", language="python", diff="d")
db.add(sample)
db.commit()
run = ExperimentRun(
id=uuid4(),
experiment_id=uuid4(),
sample_id=sample.id,
model_id="deepseek",
template_version_id=version.id,
repeat_index=1,
status="pending",
sampling_params={},
)
db.add(run)
db.commit()
runner = ExperimentRunner(db)
with mock.patch("app.experiments.runner.create_adapter") as mock_factory:
mock_adapter = mock.AsyncMock()
mock_adapter.chat.return_value = ChatResponse(
text="- null_pointer", token_usage={}, latency_ms=100.0
)
mock_factory.return_value = mock_adapter
summary = await runner.run_experiment()
assert summary["total"] == 1
db.refresh(run)
assert run.status == "done"
assert run.result is not None
+51
View File
@@ -0,0 +1,51 @@
"""Tests for stability metric."""
from uuid import uuid4
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from app.analysis.stability import compute_stability_score
from app.db import Base
from app.models import ExperimentRun, PromptTemplate, PromptTemplateVersion, Result, Sample
def get_test_db():
engine = create_engine("sqlite:///:memory:")
Base.metadata.create_all(bind=engine)
Session = sessionmaker(bind=engine)
return Session()
def test_stability_perfect():
db = get_test_db()
template = PromptTemplate(strategy_id="code_review", level="L1")
version = PromptTemplateVersion(
template=template, version_number=1, body="", variables_schema={}
)
db.add(template)
db.commit()
sample = Sample(repo="r", commit_sha="abc", language="python", diff="d")
db.add(sample)
db.commit()
experiment_id = uuid4()
for i in range(3):
run = ExperimentRun(
id=uuid4(),
experiment_id=experiment_id,
sample_id=sample.id,
model_id="deepseek",
template_version_id=version.id,
repeat_index=i + 1,
status="done",
)
db.add(run)
db.flush()
result = Result(run_id=run.id, raw_output="- null_pointer")
db.add(result)
db.commit()
score = compute_stability_score(db, str(experiment_id), "deepseek", "L1")
assert score == 1.0