281 lines
8.1 KiB
Python
281 lines
8.1 KiB
Python
"""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()}
|