sage-api / server /app.py
hiasgnpsadgd's picture
Deploy current MIMIC-only SAGE backend
de0f30b verified
Raw
History Blame Contribute Delete
13.5 kB
"""FastAPI app for the local SAGE WebUI and service API."""
from __future__ import annotations
import asyncio
import json
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
from fastapi import FastAPI, HTTPException
from fastapi.responses import FileResponse, HTMLResponse, StreamingResponse
from fastapi.staticfiles import StaticFiles
from .jobs import JobManager
from .models import RunRequest
from .runtime_config import hosted_mode_enabled
from .store import RunStore
APP_DIR = Path(__file__).resolve().parent
STATIC_DIR = APP_DIR / "static"
HOSTED_MODE = hosted_mode_enabled()
store = RunStore()
jobs = JobManager(store)
app = FastAPI(
title="SAGE Local Service",
version="0.1.0",
docs_url=None if HOSTED_MODE else "/docs",
redoc_url=None if HOSTED_MODE else "/redoc",
openapi_url=None if HOSTED_MODE else "/openapi.json",
)
if not HOSTED_MODE:
app.mount("/static", StaticFiles(directory=STATIC_DIR), name="static")
@app.get("/", include_in_schema=not HOSTED_MODE)
def index() -> Any:
if HOSTED_MODE:
return {"service": "sage", "status": "ok"}
return HTMLResponse((STATIC_DIR / "index.html").read_text(encoding="utf-8"))
@app.get("/api/health")
def health() -> dict[str, str]:
return {"status": "ok"}
@app.post("/api/runs")
def create_run(request: RunRequest) -> dict:
if not request.request_text:
raise HTTPException(status_code=400, detail="The WebUI service accepts request_text for end-to-end runs.")
end_stage = 2 if request.require_human_confirmation else 3
request = request.model_copy(update={"input_path": None, "start_stage": 1, "end_stage": end_stage})
try:
return jobs.submit(request).model_dump(mode="json")
except FileExistsError as exc:
raise HTTPException(status_code=409, detail="run_id already exists") from exc
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
@app.get("/api/runs")
def list_runs() -> list[dict]:
return [run_payload_with_prompt(status.run_id) for status in store.list_runs()]
@app.get("/api/runs/{run_id}")
def get_run(run_id: str) -> dict:
try:
store.read(run_id)
return run_payload_with_prompt(run_id)
except FileNotFoundError as exc:
raise HTTPException(status_code=404, detail="run not found") from exc
@app.post("/api/runs/{run_id}/cancel")
def cancel_run(run_id: str) -> dict:
try:
return jobs.cancel(run_id).model_dump(mode="json")
except FileNotFoundError as exc:
raise HTTPException(status_code=404, detail="run not found") from exc
@app.post("/api/runs/{run_id}/retry")
def retry_run(run_id: str) -> dict:
try:
return jobs.retry(run_id).model_dump(mode="json")
except FileNotFoundError as exc:
raise HTTPException(status_code=404, detail="run not found") from exc
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
@app.post("/api/runs/{run_id}/generate-sql")
def generate_sql(run_id: str) -> dict:
try:
return jobs.continue_to_stage3(run_id).model_dump(mode="json")
except FileNotFoundError as exc:
raise HTTPException(status_code=404, detail="run not found") from exc
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
except RuntimeError as exc:
raise HTTPException(status_code=409, detail=str(exc)) from exc
@app.get("/api/runs/{run_id}/human-confirmation")
def get_human_confirmation(run_id: str) -> dict:
try:
store.read(run_id)
for relative_path in (
"stage_02_clause_retrieval/human_confirmation.review.json",
"stage_02_clause_retrieval/human_confirmation.json",
):
try:
payload = store.read_json(run_id, relative_path)
payload["_artifact_path"] = relative_path
return payload
except FileNotFoundError:
continue
except FileNotFoundError as exc:
raise HTTPException(status_code=404, detail="run not found") from exc
raise HTTPException(status_code=404, detail="human confirmation artifact not found")
@app.put("/api/runs/{run_id}/human-confirmation")
def save_human_confirmation(run_id: str, payload: dict[str, Any]) -> dict:
try:
store.read(run_id)
base = store.read_json(run_id, "stage_02_clause_retrieval/human_confirmation.json")
normalized = normalize_human_confirmation_payload(base=base, payload=payload)
store.write_json(run_id, "stage_02_clause_retrieval/human_confirmation.review.json", normalized)
store.append_event(run_id, "Human confirmation checklist saved.", stage="stage_02")
return normalized
except FileNotFoundError as exc:
raise HTTPException(status_code=404, detail="human confirmation artifact not found") from exc
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
@app.delete("/api/runs/{run_id}", status_code=204)
def delete_run(run_id: str) -> None:
try:
jobs.delete(run_id)
except FileNotFoundError as exc:
raise HTTPException(status_code=404, detail="run not found") from exc
except RuntimeError as exc:
raise HTTPException(status_code=409, detail=str(exc)) from exc
@app.get("/api/runs/{run_id}/artifacts")
def list_artifacts(run_id: str) -> list[dict]:
try:
store.read(run_id)
return [artifact.model_dump(mode="json") for artifact in store.list_artifacts(run_id)]
except FileNotFoundError as exc:
raise HTTPException(status_code=404, detail="run not found") from exc
@app.get("/api/runs/{run_id}/lineage")
def get_lineage(run_id: str) -> dict:
"""Per-clause provenance: request sentence -> Stage 1 intent -> Stage 2 binding ->
Stage 3 SQL, for the clickable audit path + stage-production plot in the UI."""
from .lineage import build_lineage
try:
store.read(run_id)
except FileNotFoundError as exc:
raise HTTPException(status_code=404, detail="run not found") from exc
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
return build_lineage(store, run_id)
@app.get("/api/runs/{run_id}/artifacts/{artifact_path:path}")
def get_artifact(run_id: str, artifact_path: str) -> FileResponse:
try:
path = store.artifact_path(run_id, artifact_path)
except FileNotFoundError as exc:
raise HTTPException(status_code=404, detail="artifact not found") from exc
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
return FileResponse(path)
@app.get("/api/runs/{run_id}/events")
async def run_events(run_id: str) -> StreamingResponse:
async def stream():
emitted = 0
terminal_states = {"succeeded", "failed", "cancelled"}
while True:
try:
status = store.read(run_id)
except FileNotFoundError:
yield sse({"error": "run not found"})
break
events = status.events[emitted:]
for item in events:
yield sse(item)
emitted += len(events)
if status.state in terminal_states:
yield sse({"level": "info", "message": f"Run {status.state}.", "state": status.state})
break
await asyncio.sleep(1.0)
return StreamingResponse(stream(), media_type="text/event-stream")
def sse(payload: dict) -> str:
return f"data: {json.dumps(payload, sort_keys=True)}\n\n"
def run_payload_with_prompt(run_id: str) -> dict:
status = store.read(run_id)
payload = status.model_dump(mode="json")
prompt = read_original_prompt(run_id, payload.get("artifacts") or {})
payload["original_prompt"] = prompt
payload["original_prompt_preview"] = prompt_preview(prompt)
return payload
def read_original_prompt(run_id: str, artifacts: dict) -> str:
input_request = artifacts.get("input_request")
if input_request:
try:
return store.read_text(run_id, input_request)
except FileNotFoundError:
pass
try:
request = store.read_json(run_id, "request.json")
except FileNotFoundError:
return ""
return str(request.get("request_text") or "")
def prompt_preview(prompt: str, limit: int = 220) -> str:
compact = " ".join(prompt.split())
if len(compact) <= limit:
return compact
return compact[: limit - 1].rstrip() + "..."
def normalize_human_confirmation_payload(*, base: dict[str, Any], payload: dict[str, Any]) -> dict[str, Any]:
incoming = {
str(item.get("review_item_id") or ""): item
for item in payload.get("review_items") or []
if isinstance(item, dict)
}
normalized = dict(base)
review_items = []
for item in base.get("review_items") or []:
next_item = dict(item)
item_id = str(item.get("review_item_id") or "")
update = incoming.get(item_id) or {}
valid_row_uids = valid_review_row_uids(item)
review = normalized_review_from_payload(update) if update else normalized_review_from_payload(item)
selected = clean_row_uid_list(review.get("selected_row_uids"), valid_row_uids=valid_row_uids)
repair = review.get("retrieval_repair") if isinstance(review.get("retrieval_repair"), dict) else {}
repair_instructions = str(repair.get("instructions") or review.get("notes") or "")[:2000].strip()
review_status = "verified" if selected and not repair_instructions else "needs_retrieval_repair"
repair_requested = review_status == "needs_retrieval_repair"
next_item.pop("human_ranking", None)
next_item["human_review"] = {
"review_status": review_status,
"selected_row_uids": selected,
"retrieval_selected_row_uids": clean_row_uid_list(
review.get("retrieval_selected_row_uids"),
valid_row_uids=valid_row_uids,
),
"retrieval_repair": {
"requested": repair_requested,
"instructions": repair_instructions,
},
"reviewed_at": utc_now(),
}
review_items.append(next_item)
normalized["review_items"] = review_items
normalized["schema_version"] = "human_confirmation.v2"
statuses = [(item.get("human_review") or {}).get("review_status") for item in review_items]
normalized["status"] = (
"retrieval_repair_requested"
if any(status == "needs_retrieval_repair" for status in statuses)
else "reviewed"
if review_items and all(status == "verified" for status in statuses)
else "pending"
if review_items
else "empty"
)
normalized["updated_at"] = utc_now()
return normalized
def normalized_review_from_payload(item: dict[str, Any]) -> dict[str, Any]:
review = item.get("human_review")
if isinstance(review, dict):
return dict(review)
ranking = item.get("human_ranking")
if isinstance(ranking, dict):
status = canonical_review_status(ranking.get("review_status"))
return {
"review_status": status,
"selected_row_uids": ranking.get("preferred_row_uids") or [],
"retrieval_selected_row_uids": [],
"notes": str(ranking.get("notes") or ""),
"retrieval_repair": {
"requested": status == "needs_retrieval_repair",
"instructions": str(ranking.get("notes") or ""),
},
"reviewed_at": ranking.get("reviewed_at"),
}
return {
"review_status": "pending",
"selected_row_uids": [],
"retrieval_selected_row_uids": [],
"notes": "",
"retrieval_repair": {"requested": False, "instructions": ""},
"reviewed_at": None,
}
def canonical_review_status(value: Any) -> str:
status = str(value or "pending")
if status == "accepted":
return "verified"
if status in {"rejected", "needs_revision"}:
return "needs_retrieval_repair"
return status
def clean_row_uid_list(value: Any, *, valid_row_uids: set[str]) -> list[str]:
if not isinstance(value, list):
return []
result: list[str] = []
seen: set[str] = set()
for item in value:
text = str(item or "").strip()
if text and text in valid_row_uids and text not in seen:
result.append(text)
seen.add(text)
return result
def valid_review_row_uids(item: dict[str, Any]) -> set[str]:
valid: set[str] = set()
for candidate in item.get("candidates") or []:
if not isinstance(candidate, dict):
continue
row_uid = str(candidate.get("row_uid") or "")
if row_uid:
valid.add(row_uid)
for child in candidate.get("child_candidates") or []:
child_uid = str((child or {}).get("row_uid") or "")
if child_uid:
valid.add(child_uid)
for key in ("source_row_uids", "included_source_row_uids", "child_row_uids", "included_child_row_uids"):
for value in candidate.get(key) or []:
text = str(value or "").strip()
if text:
valid.add(text)
return valid
def utc_now() -> str:
return datetime.now(timezone.utc).isoformat()