|
# 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},
|
|
)
|