first commit

This commit is contained in:
eeymoo
2026-09-19 12:54:45 +08:00
commit 6fc5b64077
126 changed files with 8601 additions and 0 deletions
View File
+24
View File
@@ -0,0 +1,24 @@
"""FastAPI application."""
from contextlib import asynccontextmanager
from fastapi import Depends, FastAPI, HTTPException
from sqlalchemy.orm import Session
from app.api import router
from app.db import get_db, init_db
@asynccontextmanager
async def lifespan(app: FastAPI):
init_db()
yield
app = FastAPI(title="PromptCR-Lab API", lifespan=lifespan)
app.include_router(router.api_router, prefix="/api")
@app.get("/health")
def health_check():
return {"status": "ok"}
+280
View File
@@ -0,0 +1,280 @@
"""FastAPI RESTful routers."""
from typing import Any, Dict, List, Optional
from uuid import UUID
from fastapi import APIRouter, Depends, HTTPException
from pydantic import BaseModel, Field
from sqlalchemy.orm import Session
from app.analysis.service import AnalysisService
from app.db import get_db
from app.dataset.builder import build_dataset_from_repo
from app.dataset.git_parser import select_candidates
from app.experiments.matrix import generate_full_factorial_matrix
from app.experiments.runner import ExperimentRunner
from app.model_adapters.factory import create_adapter, list_models
from app.models import Experiment, ExperimentRun, PromptTemplate, PromptTemplateVersion, Sample
from app.prompts.defaults import seed_default_templates
from app.prompts.service import PromptService
api_router = APIRouter()
# ---- Dataset ----
class DatasetBuildRequest(BaseModel):
repo_path: str
count: int = 12
languages: List[str] = None
@api_router.post("/datasets/build")
def build_dataset(req: DatasetBuildRequest, db: Session = Depends(get_db)):
languages = req.languages or ["python", "java", "javascript"]
samples = build_dataset_from_repo(db, req.repo_path, req.count, languages)
return {
"count": len(samples),
"samples": [
{"id": str(s.id), "repo": s.repo, "commit_sha": s.commit_sha, "language": s.language}
for s in samples
],
}
@api_router.get("/datasets/samples")
def list_samples(db: Session = Depends(get_db)):
samples = db.query(Sample).all()
return [
{
"id": str(s.id),
"repo": s.repo,
"commit_sha": s.commit_sha,
"language": s.language,
"defect_count": len(s.defects),
}
for s in samples
]
@api_router.get("/datasets/samples/{sample_id}")
def get_sample(sample_id: str, db: Session = Depends(get_db)):
try:
sample = db.query(Sample).filter_by(id=UUID(sample_id)).first()
except ValueError:
raise HTTPException(status_code=400, detail="Invalid UUID")
if not sample:
raise HTTPException(status_code=404, detail="Sample not found")
return {
"id": str(sample.id),
"repo": sample.repo,
"commit_sha": sample.commit_sha,
"language": sample.language,
"diff": sample.diff,
"defects": [
{
"id": str(d.id),
"defect_type": d.defect_type,
"line_start": d.line_start,
"line_end": d.line_end,
"reference_fix": d.reference_fix,
}
for d in sample.defects
],
}
# ---- Prompts ----
class PromptVersionCreate(BaseModel):
strategy_id: str
level: str
body: str
variables_schema: Optional[Dict[str, Any]] = None
@api_router.get("/prompts")
def list_prompts(db: Session = Depends(get_db)):
templates = db.query(PromptTemplate).all()
return [
{
"id": str(t.id),
"strategy_id": t.strategy_id,
"level": t.level,
"version_count": len(t.versions),
}
for t in templates
]
@api_router.get("/prompts/{strategy_id}/{level}/versions")
def list_prompt_versions(strategy_id: str, level: str, db: Session = Depends(get_db)):
service = PromptService(db)
return [
{
"id": str(v.id),
"version_number": v.version_number,
"body": v.body,
"variables_schema": v.variables_schema,
"created_at": v.created_at.isoformat() if v.created_at else None,
}
for v in service.list_versions(strategy_id, level)
]
@api_router.post("/prompts/versions")
def create_prompt_version(req: PromptVersionCreate, db: Session = Depends(get_db)):
service = PromptService(db)
version = service.create_version(
req.strategy_id, req.level, req.body, req.variables_schema
)
return {
"id": str(version.id),
"version_number": version.version_number,
"template_id": str(version.template_id),
}
# ---- Experiments ----
class ExperimentCreate(BaseModel):
name: str
models: List[str]
levels: List[str]
sample_ids: List[str]
repeats: int = 3
sampling_params: Optional[Dict[str, Any]] = None
@api_router.post("/experiments")
def create_experiment(req: ExperimentCreate, db: Session = Depends(get_db)):
seed_default_templates(db)
experiment = generate_full_factorial_matrix(
db,
name=req.name,
models=req.models,
levels=req.levels,
sample_ids=req.sample_ids,
repeats=req.repeats,
sampling_params=req.sampling_params,
)
return {
"id": str(experiment.id),
"name": experiment.name,
"status": experiment.status,
"run_count": len(experiment.runs),
}
@api_router.get("/experiments")
def list_experiments(db: Session = Depends(get_db)):
experiments = db.query(Experiment).all()
return [
{
"id": str(e.id),
"name": e.name,
"status": e.status,
"run_count": len(e.runs),
}
for e in experiments
]
@api_router.post("/experiments/{experiment_id}/run")
async def run_experiment(experiment_id: str, db: Session = Depends(get_db)):
runner = ExperimentRunner(db)
summary = await runner.run_experiment(experiment_id=experiment_id)
return summary
@api_router.get("/experiments/{experiment_id}/runs")
def get_experiment_runs(experiment_id: str, db: Session = Depends(get_db)):
try:
runs = db.query(ExperimentRun).filter_by(experiment_id=UUID(experiment_id)).all()
except ValueError:
raise HTTPException(status_code=400, detail="Invalid UUID")
return [
{
"id": str(r.id),
"run_id": r.run_id,
"model_id": r.model_id,
"level": r.template_version.template.level,
"sample_id": str(r.sample_id),
"repeat_index": r.repeat_index,
"status": r.status,
"retry_count": r.retry_count,
}
for r in runs
]
@api_router.get("/experiments/{experiment_id}/runs/{run_id}")
def get_run(run_id: str, db: Session = Depends(get_db)):
try:
run = db.query(ExperimentRun).filter_by(id=UUID(run_id)).first()
except ValueError:
raise HTTPException(status_code=400, detail="Invalid UUID")
if not run:
raise HTTPException(status_code=404, detail="Run not found")
return {
"id": str(run.id),
"run_id": run.run_id,
"model_id": run.model_id,
"level": run.template_version.template.level,
"sample_id": str(run.sample_id),
"status": run.status,
"raw_output": run.result.raw_output if run.result else None,
"latency_ms": run.result.latency_ms if run.result else None,
}
# ---- Analysis ----
@api_router.post("/analysis/{run_id}/metrics")
def compute_run_metrics(run_id: str, db: Session = Depends(get_db)):
service = AnalysisService(db)
return service.compute_metrics_for_run(run_id)
@api_router.get("/analysis/{experiment_id}/aggregate")
def aggregate_metrics(experiment_id: str, db: Session = Depends(get_db)):
service = AnalysisService(db)
return service.aggregate_metrics(experiment_id)
@api_router.get("/analysis/{experiment_id}/charts")
def get_charts(experiment_id: str, db: Session = Depends(get_db)):
service = AnalysisService(db)
return service.generate_charts(experiment_id)
@api_router.get("/analysis/likert")
def likert_aggregation(db: Session = Depends(get_db)):
service = AnalysisService(db)
return service.likert_aggregation()
@api_router.post("/analysis/likert/{run_id}")
def save_likert(run_id: str, score: int, db: Session = Depends(get_db)):
from app.analysis.likert import save_likert_score
result = save_likert_score(db, run_id, score)
return {"run_id": run_id, "likert_score": result.likert_score}
@api_router.get("/analysis/{experiment_id}/anova")
def experiment_anova(experiment_id: str, db: Session = Depends(get_db)):
service = AnalysisService(db)
return service.run_anova(experiment_id)
# ---- Models ----
@api_router.get("/models")
def list_available_models():
return {"models": list_models()}