"""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()}