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