first commit
This commit is contained in:
@@ -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()}
|
||||
Reference in New Issue
Block a user