# retoor <retoor@molodetz.nl>
import json
import logging
from fastapi import APIRouter, Request, WebSocket, WebSocketDisconnect
from fastapi.responses import JSONResponse
from devplacepy.config import DBAPI_DIR
from devplacepy.database import get_int_setting
from devplacepy.schemas import DbQueryJobOut, DbQueryOut
from devplacepy.services.dbapi import policy
from devplacepy.services.dbapi.progress import hub
from devplacepy.services.dbapi.validate import run_select, validate_select
from devplacepy.services.jobs import queue
from devplacepy.services.manager import service_manager
from ._shared import error, read_body, require_dbapi_caller
logger = logging.getLogger(__name__)
router = APIRouter()
TERMINAL_STATES = (queue.DONE, queue.FAILED)
DEFAULT_MAX_ROWS = 5000
def _owner(caller) -> tuple[str, str]:
return "user", caller.uid
def _max_rows() -> int:
return max(1, get_int_setting("dbapi_max_rows", DEFAULT_MAX_ROWS))
@router.post("/query")
async def dbapi_query(request: Request):
caller = require_dbapi_caller(request)
body = await read_body(request)
sql = str(body.get("sql", "")).strip()
dialect = str(body.get("dialect", "sqlite")).strip() or "sqlite"
if not sql:
return error(400, "Provide a 'sql' SELECT statement.")
verdict = validate_select(sql, dialect=dialect)
if not verdict.valid:
status = 409 if not verdict.is_select else 400
hint = (
" The database API is read-only; data cannot be changed through it."
if not verdict.is_select
else ""
)
return JSONResponse(
DbQueryOut(
sql=verdict.sql,
valid=False,
statement_type=verdict.statement_type,
suspicious=verdict.suspicious,
error=(verdict.error or "Invalid query.") + hint,
).model_dump(mode="json"),
status_code=status,
)
rows, truncated = run_select(verdict.sql, _max_rows())
_audit_query(request, caller, verdict.sql, len(rows))
return JSONResponse(
DbQueryOut(
sql=verdict.sql,
valid=True,
statement_type="select",
rows=rows,
row_count=len(rows),
truncated=truncated,
suspicious=verdict.suspicious,
).model_dump(mode="json")
)
@router.post("/query/async")
async def dbapi_query_async(request: Request):
caller = require_dbapi_caller(request)
body = await read_body(request)
sql = str(body.get("sql", "")).strip()
if not sql:
return error(400, "Provide a 'sql' SELECT statement.")
verdict = validate_select(sql)
if not verdict.valid:
status = 409 if not verdict.is_select else 400
return JSONResponse(
DbQueryOut(
sql=verdict.sql,
valid=False,
statement_type=verdict.statement_type,
suspicious=verdict.suspicious,
error=verdict.error or "Invalid query.",
).model_dump(mode="json"),
status_code=status,
)
owner_kind, owner_id = _owner(caller)
uid = queue.enqueue(
"dbquery", {"sql": verdict.sql}, owner_kind, owner_id, "dbquery"
)
_audit_query(request, caller, verdict.sql, None, event="database.query.async")
return JSONResponse(
{
"uid": uid,
"status_url": f"/dbapi/query/{uid}",
"ws_url": f"/dbapi/query/{uid}/ws",
}
)
@router.get("/query/{uid}")
async def dbapi_query_status(request: Request, uid: str):
require_dbapi_caller(request)
job = queue.get_job(uid)
if not job or job.get("kind") != "dbquery":
return error(404, "Query job not found")
result = job.get("result", {})
done = job.get("status") == queue.DONE
return JSONResponse(
DbQueryJobOut(
uid=uid,
kind="dbquery",
status=job.get("status", ""),
ws_url=f"/dbapi/query/{uid}/ws",
result_url=f"/dbapi/query/{uid}/result" if done else None,
row_count=result.get("row_count") if done else None,
truncated=result.get("truncated") if done else None,
error=job.get("error") or None,
created_at=job.get("created_at"),
completed_at=job.get("completed_at") or None,
).model_dump(mode="json")
)
@router.get("/query/{uid}/result")
async def dbapi_query_result(request: Request, uid: str):
require_dbapi_caller(request)
job = queue.get_job(uid)
if not job or job.get("kind") != "dbquery":
return error(404, "Query job not found")
path = (DBAPI_DIR / uid / "result.json").resolve()
if not path.is_relative_to(DBAPI_DIR.resolve()) or not path.is_file():
return error(404, "Result not available")
queue.touch_job(uid, get_int_setting("dbquery_retention_seconds", 604800))
return JSONResponse(json.loads(path.read_text(encoding="utf-8")))
@router.websocket("/query/{uid}/ws")
async def dbapi_query_ws(websocket: WebSocket, uid: str):
await websocket.accept()
if not service_manager.owns_lock():
await websocket.close(code=4013)
return
if policy.caller_for(websocket) is None:
await websocket.close(code=1008)
return
job = queue.get_job(uid)
if not job or job.get("kind") != "dbquery":
await websocket.close(code=1008)
return
listener = hub.subscribe(uid)
try:
for frame in hub.snapshot(uid):
await websocket.send_json(frame)
current = queue.get_job(uid)
if current and current.get("status") in TERMINAL_STATES:
done = current.get("status") == queue.DONE
result = current.get("result", {})
await websocket.send_json(
{
"type": "done" if done else "failed",
"status": current.get("status"),
"row_count": result.get("row_count"),
"truncated": result.get("truncated"),
"result_url": f"/dbapi/query/{uid}/result" if done else None,
"error": current.get("error") or None,
}
)
return
while True:
frame = await listener.get()
await websocket.send_json(frame)
if frame.get("type") in ("done", "failed"):
break
except WebSocketDisconnect:
pass
except Exception:
logger.exception("dbquery websocket loop failed for %s", uid)
finally:
hub.unsubscribe(uid, listener)
def _audit_query(request, caller, sql, row_count, event="database.query"):
from devplacepy.services.audit import record as audit
audit.record(
request,
event,
summary=f"{caller.username or caller.kind} ran a database query",
metadata={"sql": sql[:500], "row_count": row_count, "caller": caller.kind},
)