forked from retoor/devplacepy
feat: add admin/internal database API with CRUD, read-only query, and natural-language SQL endpoints
Add a new `/dbapi` router package providing a generic database API over `dataset`, restricted to admin sessions, admin API keys, and the internal gateway key. Includes: - `tables.py`: list all tables and inspect table schemas - `crud.py`: full CRUD operations (GET, POST, PATCH, DELETE) with soft-delete awareness, born-live inserts, `?include_deleted`, `.../restore`, and `?hard=true` purge - `query.py`: validated read-only SELECT execution via sqlglot parsing, classification, and EXPLAIN dry-run; async query jobs with WebSocket streaming via `DbApiJobService` - `nl.py`: natural-language-to-SQL conversion using the platform AI gateway with re-prompting until validation passes Also register `DbApiJobService` and `PubSubService` in the service manager, add `DBAPI_DIR` to config data paths, and force cleartext `http://` connections to HTTP/1.1 in `curl_transport` to fix large request failures against uvicorn's HTTP/1.1-only internal gateway.
This commit is contained in:
@@ -29,6 +29,8 @@ CATEGORY_BY_PREFIX: dict[str, str] = {
|
||||
"seo": "tools",
|
||||
"deepsearch": "tools",
|
||||
"ai": "ai",
|
||||
"database": "database",
|
||||
"pubsub": "pubsub",
|
||||
"devii": "devii",
|
||||
"cli": "cli",
|
||||
"reward": "reward",
|
||||
|
||||
@@ -2537,13 +2537,14 @@ async def handle_mentions(dp: DevPlace, answered: set[str]) -> None:
|
||||
reply = await _agent_answer_for_devplace(
|
||||
f"You were @-mentioned in a DevPlace post. The message to you is:\n\n{message}\n\n"
|
||||
f"Do what it asks, then answer it directly. Your reply is posted verbatim as a "
|
||||
f"comment, so write it as botje talking to the user — not as a report about what you did.",
|
||||
f"comment, so write it as botje talking to the user — not as a report about what you did. "
|
||||
f"Keep the entire reply within 1000 characters (the comment length limit).",
|
||||
context=f"Post URL: {DEVPLACE_URL}/posts/{slug}",
|
||||
)
|
||||
|
||||
dp.call(
|
||||
"comments.create",
|
||||
content=reply[:5000],
|
||||
content=reply.strip()[:1000],
|
||||
target_uid=post_uid,
|
||||
target_type="post",
|
||||
)
|
||||
@@ -2597,10 +2598,13 @@ async def handle_dms(dp: DevPlace, answered: set[str]) -> None:
|
||||
try:
|
||||
reply = await _agent_answer_for_devplace(
|
||||
content,
|
||||
context=f"The user @{other_name} (uid {other_uid}) sent you a direct message.",
|
||||
context=(
|
||||
f"The user @{other_name} (uid {other_uid}) sent you a direct message. "
|
||||
f"Keep the entire reply within 2000 characters (the message length limit)."
|
||||
),
|
||||
)
|
||||
|
||||
dp.call("messages.send", content=reply[:5000], receiver_uid=other_uid)
|
||||
dp.call("messages.send", content=reply.strip()[:2000], receiver_uid=other_uid)
|
||||
logger.info("Replied to DM from @%s", other_name)
|
||||
|
||||
except (xmlrpc.client.Fault, Exception) as e:
|
||||
|
||||
@@ -170,8 +170,11 @@ class ContainerService(BaseService):
|
||||
elif ps.state == "created":
|
||||
await backend.start(ps.container_id)
|
||||
_set_status(inst, {"status": store.ST_RUNNING}, reason="start")
|
||||
else:
|
||||
elif status == store.ST_RUNNING:
|
||||
await self._handle_exit(backend, inst, ps)
|
||||
else:
|
||||
await backend.rm(ps.container_id, force=True)
|
||||
await self._launch(backend, inst)
|
||||
|
||||
elif desired == store.DESIRED_STOPPED:
|
||||
if ps is not None and ps.state in ("running", "restarting", "paused"):
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
@@ -0,0 +1,176 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from devplacepy.database import (
|
||||
get_table,
|
||||
purge,
|
||||
restore,
|
||||
soft_delete,
|
||||
text_search_clause,
|
||||
)
|
||||
from devplacepy.utils import generate_uid
|
||||
|
||||
from .policy import soft_delete_aware
|
||||
|
||||
SEARCH_FIELDS = (
|
||||
"title",
|
||||
"description",
|
||||
"content",
|
||||
"name",
|
||||
"username",
|
||||
"email",
|
||||
"body",
|
||||
"summary",
|
||||
"slug",
|
||||
)
|
||||
MAX_SYNC_LIMIT = 500
|
||||
|
||||
|
||||
class DbApiError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
def _now() -> str:
|
||||
return datetime.now(timezone.utc).isoformat()
|
||||
|
||||
|
||||
def column_names(table_name: str) -> list[str]:
|
||||
return list(get_table(table_name).columns)
|
||||
|
||||
|
||||
def schema(table_name: str) -> dict:
|
||||
table = get_table(table_name)
|
||||
columns = []
|
||||
for column in table.table.columns:
|
||||
columns.append({"name": column.name, "type": str(column.type)})
|
||||
return {
|
||||
"table": table_name,
|
||||
"columns": columns,
|
||||
"soft_delete": soft_delete_aware(table_name),
|
||||
"row_count": table.count(),
|
||||
}
|
||||
|
||||
|
||||
def _cursor_field(table) -> str:
|
||||
return "created_at" if table.has_column("created_at") else "id"
|
||||
|
||||
|
||||
def list_rows(
|
||||
table_name: str,
|
||||
*,
|
||||
filters: dict | None = None,
|
||||
comparisons: dict | None = None,
|
||||
search: str = "",
|
||||
before=None,
|
||||
limit: int = 25,
|
||||
include_deleted: bool = False,
|
||||
) -> tuple[list[dict], object]:
|
||||
table = get_table(table_name)
|
||||
cols = set(column_names(table_name))
|
||||
limit = max(1, min(int(limit or 25), MAX_SYNC_LIMIT))
|
||||
cursor_field = _cursor_field(table)
|
||||
clauses = []
|
||||
column_objs = table.table.columns
|
||||
if soft_delete_aware(table_name) and "deleted_at" in cols and not include_deleted:
|
||||
clauses.append(column_objs.deleted_at.is_(None))
|
||||
fields = tuple(field for field in SEARCH_FIELDS if field in cols)
|
||||
search_clause = text_search_clause(table, search, fields=fields) if fields else None
|
||||
if search_clause is not None:
|
||||
clauses.append(search_clause)
|
||||
if before is not None:
|
||||
clauses.append(column_objs[cursor_field] < before)
|
||||
kwargs = {}
|
||||
for key, value in (filters or {}).items():
|
||||
if key in cols:
|
||||
kwargs[key] = value
|
||||
for key, expression in (comparisons or {}).items():
|
||||
if key in cols:
|
||||
kwargs[key] = expression
|
||||
rows = list(
|
||||
table.find(*clauses, order_by=["-" + cursor_field], _limit=limit + 1, **kwargs)
|
||||
)
|
||||
has_more = len(rows) > limit
|
||||
rows = rows[:limit]
|
||||
next_cursor = rows[-1][cursor_field] if has_more and rows else None
|
||||
return [dict(row) for row in rows], next_cursor
|
||||
|
||||
|
||||
def count_rows(table_name: str, filters: dict | None = None) -> int:
|
||||
table = get_table(table_name)
|
||||
cols = set(column_names(table_name))
|
||||
kwargs = {key: value for key, value in (filters or {}).items() if key in cols}
|
||||
return table.count(**kwargs)
|
||||
|
||||
|
||||
def get_row(table_name: str, key: str, value) -> dict | None:
|
||||
_assert_key(table_name, key)
|
||||
row = get_table(table_name).find_one(**{key: value})
|
||||
return dict(row) if row else None
|
||||
|
||||
|
||||
def insert_row(table_name: str, data: dict, actor: str) -> dict:
|
||||
table = get_table(table_name)
|
||||
cols = set(column_names(table_name))
|
||||
row = dict(data or {})
|
||||
unknown = [key for key in row if key not in cols]
|
||||
if unknown:
|
||||
raise DbApiError(f"Unknown columns for {table_name}: {', '.join(sorted(unknown))}")
|
||||
if "uid" in cols and not row.get("uid"):
|
||||
row["uid"] = generate_uid()
|
||||
if "created_at" in cols and not row.get("created_at"):
|
||||
row["created_at"] = _now()
|
||||
if soft_delete_aware(table_name):
|
||||
row.setdefault("deleted_at", None)
|
||||
row.setdefault("deleted_by", None)
|
||||
table.insert(row)
|
||||
if row.get("uid"):
|
||||
return get_row(table_name, "uid", row["uid"]) or row
|
||||
return row
|
||||
|
||||
|
||||
def update_row(table_name: str, key: str, value, data: dict, actor: str) -> dict | None:
|
||||
_assert_key(table_name, key)
|
||||
table = get_table(table_name)
|
||||
cols = set(column_names(table_name))
|
||||
if not table.find_one(**{key: value}):
|
||||
return None
|
||||
payload = {k: v for k, v in (data or {}).items() if k not in ("id", "uid", key)}
|
||||
unknown = [k for k in payload if k not in cols]
|
||||
if unknown:
|
||||
raise DbApiError(f"Unknown columns for {table_name}: {', '.join(sorted(unknown))}")
|
||||
if "updated_at" in cols:
|
||||
payload["updated_at"] = _now()
|
||||
payload[key] = value
|
||||
table.update(payload, [key])
|
||||
return get_row(table_name, key, value)
|
||||
|
||||
|
||||
def delete_row(
|
||||
table_name: str, key: str, value, actor: str, *, hard: bool = False
|
||||
) -> dict | None:
|
||||
_assert_key(table_name, key)
|
||||
table = get_table(table_name)
|
||||
row = table.find_one(**{key: value})
|
||||
if not row:
|
||||
return None
|
||||
if soft_delete_aware(table_name) and not hard:
|
||||
soft_delete(table_name, actor, **{key: value})
|
||||
return {"mode": "soft", "row": dict(row)}
|
||||
purge(table_name, **{key: value})
|
||||
return {"mode": "hard", "row": dict(row)}
|
||||
|
||||
|
||||
def restore_row(table_name: str, key: str, value, actor: str) -> dict | None:
|
||||
_assert_key(table_name, key)
|
||||
if not soft_delete_aware(table_name):
|
||||
raise DbApiError(f"{table_name} does not support restore (no soft delete).")
|
||||
restore(table_name, **{key: value})
|
||||
return get_row(table_name, key, value)
|
||||
|
||||
|
||||
def _assert_key(table_name: str, key: str) -> None:
|
||||
if key not in set(column_names(table_name)):
|
||||
raise DbApiError(f"Unknown key column {key!r} for {table_name}.")
|
||||
@@ -0,0 +1,148 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from devplacepy import stealth
|
||||
from devplacepy.config import INTERNAL_GATEWAY_URL, INTERNAL_MODEL
|
||||
from devplacepy.database import get_setting, internal_gateway_key
|
||||
|
||||
from . import crud
|
||||
from .policy import soft_delete_aware
|
||||
from .validate import validate_select
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
GATEWAY_TIMEOUT_SECONDS = 60.0
|
||||
MAX_TOKENS = 600
|
||||
FENCE = re.compile(r"```(?:sql)?\s*(.+?)```", re.IGNORECASE | re.DOTALL)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Design:
|
||||
question: str
|
||||
table: str
|
||||
sql: str = ""
|
||||
valid: bool = False
|
||||
attempts: int = 0
|
||||
dialect: str = "sqlite"
|
||||
applied_soft_delete: bool = False
|
||||
suspicious: list[str] = field(default_factory=list)
|
||||
tables: list[str] = field(default_factory=list)
|
||||
error: str | None = None
|
||||
|
||||
def as_dict(self) -> dict:
|
||||
return {
|
||||
"question": self.question,
|
||||
"table": self.table,
|
||||
"sql": self.sql,
|
||||
"valid": self.valid,
|
||||
"attempts": self.attempts,
|
||||
"dialect": self.dialect,
|
||||
"applied_soft_delete": self.applied_soft_delete,
|
||||
"suspicious": self.suspicious,
|
||||
"tables": self.tables,
|
||||
"error": self.error,
|
||||
}
|
||||
|
||||
|
||||
def _extract_sql(text: str) -> str:
|
||||
match = FENCE.search(text or "")
|
||||
sql = match.group(1) if match else (text or "")
|
||||
return sql.strip().rstrip(";").strip()
|
||||
|
||||
|
||||
def _system_prompt(table: str, apply_soft_delete: bool, dialect: str) -> str:
|
||||
schema = crud.schema(table)
|
||||
columns = ", ".join(f"{c['name']} ({c['type']})" for c in schema["columns"])
|
||||
examples, _ = crud.list_rows(table, limit=5)
|
||||
sample = json.dumps(examples, ensure_ascii=False, default=str)[:3000]
|
||||
preamble = get_setting("dbapi_nl_system_preamble", "").strip()
|
||||
soft_rule = ""
|
||||
if apply_soft_delete and soft_delete_aware(table):
|
||||
soft_rule = (
|
||||
f"The {table} table uses soft deletes. ALWAYS add `deleted_at IS NULL` to the "
|
||||
"WHERE clause so deleted rows are excluded, UNLESS the user explicitly asks for "
|
||||
"deleted or all rows.\n"
|
||||
)
|
||||
lines = [
|
||||
preamble,
|
||||
"You translate a natural-language request into ONE read-only SQL SELECT statement "
|
||||
f"for a SQLite database (dialect: {dialect}).",
|
||||
f"Target table: {table}",
|
||||
f"Columns: {columns}",
|
||||
f"Example rows (JSON): {sample}",
|
||||
soft_rule,
|
||||
"Rules: return ONLY the SQL, no prose, no explanation, no markdown. Exactly one "
|
||||
"SELECT statement. Never write INSERT, UPDATE, DELETE, DROP, PRAGMA, or ATTACH. "
|
||||
"Prefer an explicit LIMIT when the user does not ask for everything. Dates are ISO "
|
||||
"8601 strings; compare them lexicographically.",
|
||||
]
|
||||
return "\n".join(line for line in lines if line)
|
||||
|
||||
|
||||
async def _complete(messages: list[dict], api_key: str, model: str) -> str:
|
||||
payload = {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
"max_tokens": MAX_TOKENS,
|
||||
"temperature": 0.0,
|
||||
}
|
||||
headers = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"}
|
||||
async with stealth.stealth_async_client(timeout=GATEWAY_TIMEOUT_SECONDS) as client:
|
||||
response = await client.post(INTERNAL_GATEWAY_URL, json=payload, headers=headers)
|
||||
if response.status_code >= 400:
|
||||
raise RuntimeError(f"AI gateway returned {response.status_code}")
|
||||
data = response.json()
|
||||
return (data.get("choices") or [{}])[0].get("message", {}).get("content") or ""
|
||||
|
||||
|
||||
async def design_query(
|
||||
question: str,
|
||||
table: str,
|
||||
*,
|
||||
apply_soft_delete: bool = True,
|
||||
dialect: str = "sqlite",
|
||||
max_attempts: int = 3,
|
||||
api_key: str = "",
|
||||
) -> Design:
|
||||
design = Design(question=question, table=table, dialect=dialect)
|
||||
key = api_key or internal_gateway_key()
|
||||
model = get_setting("dbapi_nl_model", "") or INTERNAL_MODEL
|
||||
messages = [
|
||||
{"role": "system", "content": _system_prompt(table, apply_soft_delete, dialect)},
|
||||
{"role": "user", "content": question},
|
||||
]
|
||||
for attempt in range(1, max(1, max_attempts) + 1):
|
||||
design.attempts = attempt
|
||||
try:
|
||||
raw = await _complete(messages, key, model)
|
||||
except Exception as exc:
|
||||
design.error = f"AI gateway error: {exc}"
|
||||
return design
|
||||
sql = _extract_sql(raw)
|
||||
verdict = validate_select(sql, dialect=dialect)
|
||||
design.sql = verdict.sql
|
||||
design.suspicious = verdict.suspicious
|
||||
design.tables = verdict.tables
|
||||
design.applied_soft_delete = "deleted_at" in verdict.sql.lower()
|
||||
if verdict.valid:
|
||||
design.valid = True
|
||||
design.error = None
|
||||
return design
|
||||
design.error = verdict.error
|
||||
messages.append({"role": "assistant", "content": sql})
|
||||
messages.append(
|
||||
{
|
||||
"role": "user",
|
||||
"content": (
|
||||
f"That query is invalid: {verdict.error}. Return a corrected single "
|
||||
"SELECT statement only, no prose."
|
||||
),
|
||||
}
|
||||
)
|
||||
return design
|
||||
@@ -0,0 +1,103 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
|
||||
from devplacepy.database import (
|
||||
SOFT_DELETE_TABLES,
|
||||
db,
|
||||
get_setting,
|
||||
internal_gateway_key,
|
||||
)
|
||||
from devplacepy.utils import _user_from_api_key, get_current_user, is_admin
|
||||
|
||||
TABLE_NAME = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
|
||||
DEFAULT_DENY = {"sessions", "password_resets", "cache_state"}
|
||||
|
||||
|
||||
class DbApiDenied(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class DbApiBadTable(Exception):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Caller:
|
||||
kind: str
|
||||
uid: str
|
||||
username: str
|
||||
api_key: str
|
||||
|
||||
@property
|
||||
def is_internal(self) -> bool:
|
||||
return self.kind == "internal"
|
||||
|
||||
|
||||
def deny_tables() -> set[str]:
|
||||
raw = get_setting("dbapi_deny_tables", "")
|
||||
extra = {name.strip() for name in raw.split(",") if name.strip()}
|
||||
return DEFAULT_DENY | extra
|
||||
|
||||
|
||||
def _bearer_key(request) -> str:
|
||||
key = request.headers.get("x-api-key", "").strip()
|
||||
if key:
|
||||
return key
|
||||
scheme, _, credentials = request.headers.get("authorization", "").partition(" ")
|
||||
if scheme.lower() == "bearer":
|
||||
return credentials.strip()
|
||||
return ""
|
||||
|
||||
|
||||
def caller_for(request) -> Caller | None:
|
||||
user = get_current_user(request)
|
||||
if user and is_admin(user):
|
||||
return Caller(
|
||||
kind="admin",
|
||||
uid=user["uid"],
|
||||
username=user.get("username", ""),
|
||||
api_key=user.get("api_key", "") or "",
|
||||
)
|
||||
key = _bearer_key(request)
|
||||
if key and key == internal_gateway_key():
|
||||
return Caller(kind="internal", uid="internal", username="internal", api_key=key)
|
||||
if key:
|
||||
keyed = _user_from_api_key(key)
|
||||
if keyed and is_admin(keyed):
|
||||
return Caller(
|
||||
kind="admin",
|
||||
uid=keyed["uid"],
|
||||
username=keyed.get("username", ""),
|
||||
api_key=keyed.get("api_key", "") or "",
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
def require_caller(request) -> Caller:
|
||||
caller = caller_for(request)
|
||||
if caller is None:
|
||||
raise DbApiDenied("Administrator or internal-key access required.")
|
||||
return caller
|
||||
|
||||
|
||||
def soft_delete_aware(table: str) -> bool:
|
||||
return table in SOFT_DELETE_TABLES
|
||||
|
||||
|
||||
def assert_table(name: str) -> str:
|
||||
if not name or not TABLE_NAME.match(name):
|
||||
raise DbApiBadTable(f"Invalid table name: {name!r}")
|
||||
if name in deny_tables():
|
||||
raise DbApiBadTable(f"Table {name!r} is not accessible through the database API.")
|
||||
if name not in db.tables:
|
||||
raise DbApiBadTable(f"Unknown table: {name!r}")
|
||||
return name
|
||||
|
||||
|
||||
def allowed_tables() -> list[str]:
|
||||
blocked = deny_tables()
|
||||
return sorted(name for name in db.tables if name not in blocked)
|
||||
@@ -0,0 +1,43 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
|
||||
MAX_BUFFER = 2000
|
||||
|
||||
|
||||
class ProgressHub:
|
||||
def __init__(self) -> None:
|
||||
self._subscribers: dict[str, set[asyncio.Queue]] = {}
|
||||
self._buffers: dict[str, list[dict]] = {}
|
||||
|
||||
def subscribe(self, uid: str) -> asyncio.Queue:
|
||||
queue: asyncio.Queue = asyncio.Queue()
|
||||
self._subscribers.setdefault(uid, set()).add(queue)
|
||||
return queue
|
||||
|
||||
def unsubscribe(self, uid: str, queue: asyncio.Queue) -> None:
|
||||
listeners = self._subscribers.get(uid)
|
||||
if not listeners:
|
||||
return
|
||||
listeners.discard(queue)
|
||||
if not listeners:
|
||||
self._subscribers.pop(uid, None)
|
||||
|
||||
def publish(self, uid: str, frame: dict) -> None:
|
||||
buffer = self._buffers.setdefault(uid, [])
|
||||
buffer.append(frame)
|
||||
if len(buffer) > MAX_BUFFER:
|
||||
del buffer[: len(buffer) - MAX_BUFFER]
|
||||
for queue in self._subscribers.get(uid, set()):
|
||||
queue.put_nowait(frame)
|
||||
|
||||
def snapshot(self, uid: str) -> list[dict]:
|
||||
return list(self._buffers.get(uid, []))
|
||||
|
||||
def clear(self, uid: str) -> None:
|
||||
self._buffers.pop(uid, None)
|
||||
|
||||
|
||||
hub = ProgressHub()
|
||||
@@ -0,0 +1,153 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import shutil
|
||||
import sqlite3
|
||||
from pathlib import Path
|
||||
|
||||
from devplacepy.config import DBAPI_DIR
|
||||
from devplacepy.database import get_int_setting
|
||||
from devplacepy.services.base import ConfigField
|
||||
from devplacepy.services.jobs.base import JobService
|
||||
|
||||
from .progress import hub
|
||||
from .validate import _database_path, validate_select
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
BATCH_SIZE = 200
|
||||
DEFAULT_MAX_ROWS = 5000
|
||||
|
||||
|
||||
class DbApiJobService(JobService):
|
||||
kind = "dbquery"
|
||||
title = "Database API"
|
||||
description = (
|
||||
"Runs validated read-only SQL SELECT queries off the request path and streams the "
|
||||
"result rows over a websocket. Backs the asynchronous /dbapi/query/async endpoint."
|
||||
)
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(name="dbquery", interval_seconds=2)
|
||||
self.config_fields.extend(
|
||||
[
|
||||
ConfigField(
|
||||
"dbapi_max_rows",
|
||||
"Max result rows",
|
||||
type="int",
|
||||
default=DEFAULT_MAX_ROWS,
|
||||
minimum=1,
|
||||
help="Hard cap on rows returned by a query (sync and async).",
|
||||
group="Database API",
|
||||
),
|
||||
ConfigField(
|
||||
"dbapi_nl_model",
|
||||
"NL-to-SQL model",
|
||||
type="str",
|
||||
default="",
|
||||
help="Model used to design SQL from natural language. Blank uses molodetz.",
|
||||
group="Database API",
|
||||
),
|
||||
ConfigField(
|
||||
"dbapi_nl_system_preamble",
|
||||
"NL-to-SQL preamble",
|
||||
type="text",
|
||||
default="",
|
||||
help="Optional operator text prepended to the NL-to-SQL system prompt.",
|
||||
group="Database API",
|
||||
),
|
||||
ConfigField(
|
||||
"dbapi_deny_tables",
|
||||
"Denied tables",
|
||||
type="str",
|
||||
default="",
|
||||
help="Comma separated extra tables to hide from the database API.",
|
||||
group="Database API",
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
def result_dir(self, uid: str) -> Path:
|
||||
return DBAPI_DIR / uid
|
||||
|
||||
async def process(self, job: dict) -> dict:
|
||||
uid = job["uid"]
|
||||
payload = job.get("payload", {})
|
||||
sql = payload.get("sql", "")
|
||||
verdict = validate_select(sql)
|
||||
if not verdict.valid:
|
||||
hub.publish(uid, {"type": "failed", "message": verdict.error or "invalid query"})
|
||||
hub.clear(uid)
|
||||
raise RuntimeError(verdict.error or "invalid query")
|
||||
|
||||
max_rows = max(1, get_int_setting("dbapi_max_rows", DEFAULT_MAX_ROWS))
|
||||
rows, truncated = await self._stream(uid, verdict.sql, max_rows)
|
||||
|
||||
output_dir = self.result_dir(uid)
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
result = {
|
||||
"sql": verdict.sql,
|
||||
"row_count": len(rows),
|
||||
"truncated": truncated,
|
||||
"suspicious": verdict.suspicious,
|
||||
"rows": rows,
|
||||
}
|
||||
(output_dir / "result.json").write_text(
|
||||
json.dumps(result, default=str), encoding="utf-8"
|
||||
)
|
||||
hub.publish(
|
||||
uid,
|
||||
{
|
||||
"type": "done",
|
||||
"row_count": len(rows),
|
||||
"truncated": truncated,
|
||||
"result_url": f"/dbapi/query/{uid}/result",
|
||||
},
|
||||
)
|
||||
hub.clear(uid)
|
||||
return {
|
||||
"sql": verdict.sql,
|
||||
"row_count": len(rows),
|
||||
"truncated": truncated,
|
||||
"suspicious": verdict.suspicious,
|
||||
"result_url": f"/dbapi/query/{uid}/result",
|
||||
"bytes_out": len(json.dumps(result, default=str)),
|
||||
"item_count": len(rows),
|
||||
}
|
||||
|
||||
async def _stream(self, uid: str, sql: str, max_rows: int):
|
||||
connection = sqlite3.connect(
|
||||
f"file:{_database_path()}?mode=ro", uri=True, timeout=30
|
||||
)
|
||||
rows: list[dict] = []
|
||||
truncated = False
|
||||
try:
|
||||
connection.row_factory = sqlite3.Row
|
||||
connection.execute("PRAGMA query_only=ON")
|
||||
cursor = connection.execute(sql)
|
||||
while True:
|
||||
batch = cursor.fetchmany(BATCH_SIZE)
|
||||
if not batch:
|
||||
break
|
||||
for raw in batch:
|
||||
if len(rows) >= max_rows:
|
||||
truncated = True
|
||||
break
|
||||
rows.append(dict(raw))
|
||||
hub.publish(
|
||||
uid, {"type": "progress", "row_count": len(rows)}
|
||||
)
|
||||
if truncated:
|
||||
break
|
||||
await asyncio.sleep(0)
|
||||
finally:
|
||||
connection.close()
|
||||
return rows, truncated
|
||||
|
||||
def cleanup(self, job: dict) -> None:
|
||||
hub.clear(job["uid"])
|
||||
shutil.rmtree(self.result_dir(job["uid"]), ignore_errors=True)
|
||||
@@ -0,0 +1,156 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
import sqlite3
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import sqlglot
|
||||
from sqlglot import exp
|
||||
|
||||
from devplacepy.config import DATABASE_URL
|
||||
|
||||
UNSAFE = re.compile(r"\b(attach|detach|pragma|vacuum|reindex)\b", re.IGNORECASE)
|
||||
|
||||
_STATEMENT_TYPES = {
|
||||
exp.Select: "select",
|
||||
exp.Union: "select",
|
||||
exp.Insert: "insert",
|
||||
exp.Update: "update",
|
||||
exp.Delete: "delete",
|
||||
exp.Create: "ddl",
|
||||
exp.Drop: "ddl",
|
||||
exp.Alter: "ddl",
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class Verdict:
|
||||
sql: str
|
||||
statement_type: str = "other"
|
||||
is_select: bool = False
|
||||
valid: bool = False
|
||||
tables: list[str] = field(default_factory=list)
|
||||
has_where: bool = False
|
||||
has_join: bool = False
|
||||
has_limit: bool = False
|
||||
suspicious: list[str] = field(default_factory=list)
|
||||
error: str | None = None
|
||||
|
||||
def as_dict(self) -> dict:
|
||||
return {
|
||||
"sql": self.sql,
|
||||
"statement_type": self.statement_type,
|
||||
"is_select": self.is_select,
|
||||
"valid": self.valid,
|
||||
"tables": self.tables,
|
||||
"has_where": self.has_where,
|
||||
"has_join": self.has_join,
|
||||
"has_limit": self.has_limit,
|
||||
"suspicious": self.suspicious,
|
||||
"error": self.error,
|
||||
}
|
||||
|
||||
|
||||
def _database_path() -> str:
|
||||
return DATABASE_URL.split("sqlite:///", 1)[-1]
|
||||
|
||||
|
||||
def _statement_type(node: exp.Expression) -> str:
|
||||
for kind, label in _STATEMENT_TYPES.items():
|
||||
if isinstance(node, kind):
|
||||
return label
|
||||
return "other"
|
||||
|
||||
|
||||
def classify(sql: str, dialect: str = "sqlite") -> Verdict:
|
||||
text = (sql or "").strip().rstrip(";").strip()
|
||||
verdict = Verdict(sql=text)
|
||||
if not text:
|
||||
verdict.error = "Empty query."
|
||||
return verdict
|
||||
try:
|
||||
statements = [node for node in sqlglot.parse(text, read=dialect) if node]
|
||||
except Exception as exc:
|
||||
verdict.error = f"Could not parse SQL: {str(exc).splitlines()[0][:200]}"
|
||||
return verdict
|
||||
if not statements:
|
||||
verdict.error = "No statement found."
|
||||
return verdict
|
||||
node = statements[0]
|
||||
verdict.statement_type = _statement_type(node)
|
||||
verdict.is_select = verdict.statement_type == "select"
|
||||
verdict.tables = sorted({t.name for t in node.find_all(exp.Table) if t.name})
|
||||
verdict.has_where = node.find(exp.Where) is not None
|
||||
verdict.has_join = node.find(exp.Join) is not None
|
||||
verdict.has_limit = node.find(exp.Limit) is not None
|
||||
if len(statements) > 1:
|
||||
verdict.suspicious.append("Multiple statements in one query are not allowed.")
|
||||
if UNSAFE.search(text):
|
||||
verdict.suspicious.append("ATTACH/PRAGMA/VACUUM style statements are not allowed.")
|
||||
verdict.is_select = False
|
||||
if verdict.is_select and not (
|
||||
verdict.has_where or verdict.has_join or verdict.has_limit
|
||||
):
|
||||
verdict.suspicious.append(
|
||||
"SELECT has no WHERE, JOIN, or LIMIT and may return an entire table."
|
||||
)
|
||||
return verdict
|
||||
|
||||
|
||||
def dry_run(sql: str) -> tuple[bool, str | None]:
|
||||
text = (sql or "").strip().rstrip(";").strip()
|
||||
connection = sqlite3.connect(f"file:{_database_path()}?mode=ro", uri=True, timeout=5)
|
||||
try:
|
||||
connection.execute("PRAGMA query_only=ON")
|
||||
connection.execute(f"EXPLAIN {text}")
|
||||
return True, None
|
||||
except sqlite3.Error as exc:
|
||||
return False, str(exc)[:300]
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
|
||||
def validate_select(sql: str, dialect: str = "sqlite") -> Verdict:
|
||||
verdict = classify(sql, dialect=dialect)
|
||||
if verdict.error:
|
||||
return verdict
|
||||
if len(verdict.suspicious) and any(
|
||||
"Multiple statements" in note or "not allowed" in note
|
||||
for note in verdict.suspicious
|
||||
):
|
||||
verdict.error = verdict.suspicious[0]
|
||||
return verdict
|
||||
if not verdict.is_select:
|
||||
verdict.error = (
|
||||
f"Only SELECT queries run through query(); this is a {verdict.statement_type} "
|
||||
"statement. Use the structured /dbapi/{table} CRUD routes to change data."
|
||||
)
|
||||
return verdict
|
||||
ok, error = dry_run(verdict.sql)
|
||||
verdict.valid = ok
|
||||
if not ok:
|
||||
verdict.error = error
|
||||
return verdict
|
||||
|
||||
|
||||
def run_select(sql: str, limit: int) -> tuple[list[dict], bool]:
|
||||
text = sql.strip().rstrip(";").strip()
|
||||
connection = sqlite3.connect(f"file:{_database_path()}?mode=ro", uri=True, timeout=30)
|
||||
try:
|
||||
connection.row_factory = sqlite3.Row
|
||||
connection.execute("PRAGMA query_only=ON")
|
||||
cursor = connection.execute(text)
|
||||
rows = [dict(row) for row in cursor.fetchmany(limit + 1)]
|
||||
truncated = len(rows) > limit
|
||||
return rows[:limit], truncated
|
||||
finally:
|
||||
connection.close()
|
||||
|
||||
|
||||
def transpile(sql: str, read: str = "sqlite", write: str = "sqlite") -> str:
|
||||
try:
|
||||
return sqlglot.transpile(sql, read=read, write=write)[0]
|
||||
except Exception:
|
||||
return sql
|
||||
@@ -1400,6 +1400,175 @@ ACTIONS: tuple[Action, ...] = (
|
||||
params=(path("uid", "Attachment uid to restore."),),
|
||||
requires_admin=True,
|
||||
),
|
||||
Action(
|
||||
name="db_list_tables",
|
||||
method="GET",
|
||||
path="/dbapi/tables",
|
||||
summary="List database tables exposed by the database API (admin only)",
|
||||
description=(
|
||||
"Returns every table reachable through the database API with its row count and "
|
||||
"whether it uses soft deletes. Use this to discover what data exists before "
|
||||
"querying or designing a query."
|
||||
),
|
||||
params=(),
|
||||
requires_admin=True,
|
||||
),
|
||||
Action(
|
||||
name="db_table_schema",
|
||||
method="GET",
|
||||
path="/dbapi/{table}/schema",
|
||||
summary="Show a table's columns and types (admin only)",
|
||||
description="Returns the column names, types, row count, and soft-delete flag for one table.",
|
||||
params=(path("table", "Table name (from db_list_tables)."),),
|
||||
requires_admin=True,
|
||||
),
|
||||
Action(
|
||||
name="db_list_rows",
|
||||
method="GET",
|
||||
path="/dbapi/{table}",
|
||||
summary="List rows of a table with keyset pagination (admin only)",
|
||||
description=(
|
||||
"Browses rows newest-first. Soft-deleted rows are excluded unless include_deleted is "
|
||||
"true. For filtered or joined questions prefer db_query or db_design_query."
|
||||
),
|
||||
params=(
|
||||
path("table", "Table name."),
|
||||
query("limit", "Maximum rows (1-500, default 25)."),
|
||||
query("search", "Free-text search over common text columns."),
|
||||
query("before", "Keyset cursor: return rows older than this created_at/id value."),
|
||||
Param(
|
||||
name="include_deleted",
|
||||
location="query",
|
||||
description="Include soft-deleted rows.",
|
||||
required=False,
|
||||
type="boolean",
|
||||
),
|
||||
),
|
||||
requires_admin=True,
|
||||
),
|
||||
Action(
|
||||
name="db_get_row",
|
||||
method="GET",
|
||||
path="/dbapi/{table}/{key}/{value}",
|
||||
summary="Fetch one row by a key column (admin only)",
|
||||
description="Returns a single row where key column equals value (key is usually 'uid').",
|
||||
params=(
|
||||
path("table", "Table name."),
|
||||
path("key", "Key column to match (usually 'uid')."),
|
||||
path("value", "Value of the key column."),
|
||||
),
|
||||
requires_admin=True,
|
||||
),
|
||||
Action(
|
||||
name="db_query",
|
||||
method="POST",
|
||||
path="/dbapi/query",
|
||||
summary="Run a read-only SQL SELECT and return rows (admin only)",
|
||||
description=(
|
||||
"Executes a SINGLE validated SELECT statement read-only and returns the rows. Only "
|
||||
"SELECT is allowed; INSERT/UPDATE/DELETE/DDL are rejected (use db_insert_row, "
|
||||
"db_update_row, db_delete_row for changes). The response may include a 'suspicious' "
|
||||
"list (e.g. a SELECT with no WHERE/JOIN/LIMIT that scans a whole table); when present, "
|
||||
"surface that warning to the user before trusting the results."
|
||||
),
|
||||
params=(
|
||||
body("sql", "A single SELECT statement.", required=True),
|
||||
body("dialect", "Optional source SQL dialect (default sqlite)."),
|
||||
),
|
||||
requires_admin=True,
|
||||
read_only=True,
|
||||
),
|
||||
Action(
|
||||
name="db_design_query",
|
||||
method="POST",
|
||||
path="/dbapi/nl",
|
||||
summary="Design a SQL SELECT from a natural-language question (admin only)",
|
||||
description=(
|
||||
"Turns a plain-language question about one table into a validated read-only SELECT. "
|
||||
"It auto-adds 'deleted_at IS NULL' for soft-delete tables unless apply_soft_delete is "
|
||||
"false. By default it only returns the SQL; pass execute=true to also run it read-only "
|
||||
"and return rows. Show the user the SQL and any 'suspicious' notes."
|
||||
),
|
||||
params=(
|
||||
body("question", "The natural-language request.", required=True),
|
||||
body("table", "Target table the question is about.", required=True),
|
||||
Param(
|
||||
name="apply_soft_delete",
|
||||
location="body",
|
||||
description="Add deleted_at IS NULL for soft-delete tables (default true).",
|
||||
required=False,
|
||||
type="boolean",
|
||||
),
|
||||
body("dialect", "Optional target SQL dialect (default sqlite)."),
|
||||
Param(
|
||||
name="execute",
|
||||
location="body",
|
||||
description="Also run the validated query read-only and return rows.",
|
||||
required=False,
|
||||
type="boolean",
|
||||
),
|
||||
),
|
||||
requires_admin=True,
|
||||
read_only=True,
|
||||
),
|
||||
Action(
|
||||
name="db_insert_row",
|
||||
method="POST",
|
||||
path="/dbapi/{table}",
|
||||
summary="Insert a row into a table (admin only, confirmation required)",
|
||||
description=(
|
||||
"Inserts a new row. Pass the column values as a JSON object string in values_json. "
|
||||
"Soft-delete columns and uid/created_at are filled automatically. Requires confirmation."
|
||||
),
|
||||
params=(
|
||||
path("table", "Table name."),
|
||||
body("values_json", "JSON object of column:value pairs for the new row.", required=True),
|
||||
confirm(),
|
||||
),
|
||||
requires_admin=True,
|
||||
),
|
||||
Action(
|
||||
name="db_update_row",
|
||||
method="PATCH",
|
||||
path="/dbapi/{table}/{key}/{value}",
|
||||
summary="Update a row in a table (admin only, confirmation required)",
|
||||
description=(
|
||||
"Updates the row where key equals value. Pass the changed columns as a JSON object "
|
||||
"string in values_json. uid and id cannot be changed. Requires confirmation."
|
||||
),
|
||||
params=(
|
||||
path("table", "Table name."),
|
||||
path("key", "Key column to match (usually 'uid')."),
|
||||
path("value", "Value of the key column."),
|
||||
body("values_json", "JSON object of column:value pairs to change.", required=True),
|
||||
confirm(),
|
||||
),
|
||||
requires_admin=True,
|
||||
),
|
||||
Action(
|
||||
name="db_delete_row",
|
||||
method="DELETE",
|
||||
path="/dbapi/{table}/{key}/{value}",
|
||||
summary="Delete a row from a table (admin only, confirmation required)",
|
||||
description=(
|
||||
"Soft-deletes the row where key equals value (restorable). Pass hard=true to "
|
||||
"PERMANENTLY purge it (or for tables without soft delete). Requires confirmation."
|
||||
),
|
||||
params=(
|
||||
path("table", "Table name."),
|
||||
path("key", "Key column to match (usually 'uid')."),
|
||||
path("value", "Value of the key column."),
|
||||
Param(
|
||||
name="hard",
|
||||
location="query",
|
||||
description="Permanently purge instead of soft delete.",
|
||||
required=False,
|
||||
type="boolean",
|
||||
),
|
||||
confirm(),
|
||||
),
|
||||
requires_admin=True,
|
||||
),
|
||||
)
|
||||
|
||||
PLATFORM_CATALOG = Catalog(actions=ACTIONS)
|
||||
|
||||
@@ -51,6 +51,9 @@ CONFIRM_REQUIRED = {
|
||||
"admin_reset_guest_ai_quota",
|
||||
"admin_reset_user_ai_quota",
|
||||
"notification_reset",
|
||||
"db_insert_row",
|
||||
"db_update_row",
|
||||
"db_delete_row",
|
||||
}
|
||||
|
||||
CONDITIONAL_CONFIRM = {
|
||||
@@ -212,6 +215,26 @@ def confirmation_error(name: str, arguments: dict[str, Any]) -> ToolInputError |
|
||||
f"such as rm, dd, truncate, or drop): {command!r}. Show the user the exact command, get "
|
||||
"explicit confirmation, then call again with confirm=true."
|
||||
)
|
||||
if name == "db_insert_row":
|
||||
table = str(arguments.get("table", "")).strip() or "(unspecified)"
|
||||
return ToolInputError(
|
||||
f"This writes a new row directly into the '{table}' table. Show the user the exact "
|
||||
"table and values, get explicit confirmation, then call again with confirm=true."
|
||||
)
|
||||
if name == "db_update_row":
|
||||
table = str(arguments.get("table", "")).strip() or "(unspecified)"
|
||||
return ToolInputError(
|
||||
f"This updates an existing row in the '{table}' table directly. Show the user the "
|
||||
"exact row and new values, get explicit confirmation, then call again with confirm=true."
|
||||
)
|
||||
if name == "db_delete_row":
|
||||
table = str(arguments.get("table", "")).strip() or "(unspecified)"
|
||||
hard = str(arguments.get("hard", "")).strip().lower() in ("true", "1", "yes", "on")
|
||||
kind = "PERMANENTLY purges" if hard else "soft-deletes"
|
||||
return ToolInputError(
|
||||
f"This {kind} a row in the '{table}' table. Show the user the exact row, get explicit "
|
||||
"confirmation, then call again with confirm=true."
|
||||
)
|
||||
if name in CONFIRM_REQUIRED:
|
||||
return ToolInputError(
|
||||
"This removes the item as a soft delete: it disappears from every surface and is only "
|
||||
@@ -547,7 +570,7 @@ class Dispatcher:
|
||||
key = self._file_key(arguments)
|
||||
if key is not None:
|
||||
self._read_files.add(key)
|
||||
if action.method in MUTATING_METHODS:
|
||||
if action.method in MUTATING_METHODS and not action.is_read_only:
|
||||
record_mutation(action.name)
|
||||
store = get_store()
|
||||
if store is not None:
|
||||
|
||||
@@ -31,7 +31,10 @@ PLAN_VIOLATION = (
|
||||
VERIFICATION_GATE = (
|
||||
"[verification-gate] You produced a final answer after performing changes without calling "
|
||||
"verify(). Confirm the change took effect and call verify() with a summary. If verification "
|
||||
"truly does not apply, reply explicitly starting with: 'No verification applicable: <reason>'."
|
||||
"truly does not apply (for example the action only read data and changed nothing), do NOT "
|
||||
"reply with a bare disclaimer: give the user the full answer they asked for - including any "
|
||||
"rows or data you retrieved - and you may note 'No verification applicable: <reason>' at the "
|
||||
"end. Never drop the requested data."
|
||||
)
|
||||
ITERATION_LIMIT_MESSAGE = "[stopped] Maximum iterations reached without a final answer."
|
||||
|
||||
|
||||
@@ -0,0 +1,12 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
from .hub import pubsub
|
||||
from .service import PubSubService
|
||||
|
||||
__all__ = ["pubsub", "PubSubService", "publish"]
|
||||
|
||||
|
||||
async def publish(topic: str, data) -> int:
|
||||
return await pubsub.publish(
|
||||
topic, {"type": "message", "topic": topic, "data": data}
|
||||
)
|
||||
@@ -0,0 +1,63 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def topic_matches(pattern: str, topic: str) -> bool:
|
||||
if pattern == topic or pattern == "*":
|
||||
return True
|
||||
if pattern.endswith(".*"):
|
||||
prefix = pattern[:-1]
|
||||
return topic == pattern[:-2] or topic.startswith(prefix)
|
||||
return False
|
||||
|
||||
|
||||
class PubSubHub:
|
||||
def __init__(self) -> None:
|
||||
self._subscriptions: dict[str, set] = {}
|
||||
|
||||
def subscribe(self, pattern: str, socket) -> None:
|
||||
self._subscriptions.setdefault(pattern, set()).add(socket)
|
||||
|
||||
def unsubscribe(self, pattern: str, socket) -> None:
|
||||
sockets = self._subscriptions.get(pattern)
|
||||
if not sockets:
|
||||
return
|
||||
sockets.discard(socket)
|
||||
if not sockets:
|
||||
self._subscriptions.pop(pattern, None)
|
||||
|
||||
def drop_socket(self, socket) -> None:
|
||||
for pattern in list(self._subscriptions):
|
||||
self.unsubscribe(pattern, socket)
|
||||
|
||||
def _targets(self, topic: str) -> set:
|
||||
targets: set = set()
|
||||
for pattern, sockets in self._subscriptions.items():
|
||||
if topic_matches(pattern, topic):
|
||||
targets.update(sockets)
|
||||
return targets
|
||||
|
||||
async def publish(self, topic: str, frame: dict) -> int:
|
||||
targets = self._targets(topic)
|
||||
delivered = 0
|
||||
for socket in targets:
|
||||
try:
|
||||
await socket.send_json(frame)
|
||||
delivered += 1
|
||||
except Exception:
|
||||
logger.debug("pubsub dropping dead socket for %s", topic)
|
||||
return delivered
|
||||
|
||||
def topics(self) -> list[dict]:
|
||||
return [
|
||||
{"topic": pattern, "subscribers": len(sockets)}
|
||||
for pattern, sockets in sorted(self._subscriptions.items())
|
||||
]
|
||||
|
||||
|
||||
pubsub = PubSubHub()
|
||||
@@ -0,0 +1,76 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
|
||||
from devplacepy.database import get_setting
|
||||
from devplacepy.services.dbapi.policy import caller_for as _dbapi_caller
|
||||
from devplacepy.utils import get_current_user, is_admin
|
||||
|
||||
TOPIC_NAME = re.compile(r"^[A-Za-z0-9_.*-]{1,128}$")
|
||||
MAX_PAYLOAD_BYTES = 64 * 1024
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Actor:
|
||||
kind: str
|
||||
uid: str
|
||||
username: str
|
||||
|
||||
@property
|
||||
def privileged(self) -> bool:
|
||||
return self.kind in ("admin", "internal")
|
||||
|
||||
|
||||
def guests_enabled() -> bool:
|
||||
return get_setting("pubsub_allow_guests", "0") == "1"
|
||||
|
||||
|
||||
def resolve_actor(scope) -> Actor:
|
||||
caller = _dbapi_caller(scope)
|
||||
if caller is not None:
|
||||
return Actor(kind=caller.kind, uid=caller.uid, username=caller.username)
|
||||
user = get_current_user(scope)
|
||||
if user:
|
||||
kind = "admin" if is_admin(user) else "user"
|
||||
return Actor(kind=kind, uid=user["uid"], username=user.get("username", ""))
|
||||
return Actor(kind="guest", uid="", username="guest")
|
||||
|
||||
|
||||
def valid_topic(topic: str) -> bool:
|
||||
return bool(topic and TOPIC_NAME.match(topic))
|
||||
|
||||
|
||||
def _own_namespace(actor: Actor, topic: str) -> bool:
|
||||
if not actor.uid:
|
||||
return False
|
||||
base = f"user.{actor.uid}"
|
||||
return topic == base or topic.startswith(base + ".")
|
||||
|
||||
|
||||
def _is_public(topic: str) -> bool:
|
||||
return topic == "public" or topic.startswith("public.") or topic.startswith("public")
|
||||
|
||||
|
||||
def can_subscribe(actor: Actor, topic: str) -> bool:
|
||||
if actor.privileged:
|
||||
return True
|
||||
if topic == "*":
|
||||
return False
|
||||
if topic == "public" or topic.startswith("public."):
|
||||
return actor.kind != "guest" or guests_enabled()
|
||||
if actor.kind == "user" and _own_namespace(actor, topic):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def can_publish(actor: Actor, topic: str) -> bool:
|
||||
if actor.privileged:
|
||||
return True
|
||||
if actor.kind == "guest":
|
||||
return False
|
||||
if actor.kind == "user" and _own_namespace(actor, topic):
|
||||
return True
|
||||
return False
|
||||
@@ -0,0 +1,36 @@
|
||||
# retoor <retoor@molodetz.nl>
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from devplacepy.services.base import BaseService, ConfigField
|
||||
|
||||
from .hub import pubsub
|
||||
|
||||
|
||||
class PubSubService(BaseService):
|
||||
title = "Pub/Sub"
|
||||
description = (
|
||||
"In-process publish/subscribe bus over websockets. Lets the frontend, services and "
|
||||
"Devii broadcast and receive messages on named topics. Database-free and ephemeral; "
|
||||
"served on the service lock owner so every subscriber converges on one worker."
|
||||
)
|
||||
default_enabled = True
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(name="pubsub", interval_seconds=3600)
|
||||
self.config_fields = [
|
||||
ConfigField(
|
||||
"pubsub_allow_guests",
|
||||
"Allow guests",
|
||||
type="bool",
|
||||
default=False,
|
||||
help="Allow unauthenticated guests to subscribe to public.* topics.",
|
||||
group="Pub/Sub",
|
||||
),
|
||||
]
|
||||
|
||||
async def run_once(self) -> None:
|
||||
return None
|
||||
|
||||
def collect_metrics(self) -> dict:
|
||||
return {"topics": len(pubsub.topics())}
|
||||
Reference in New Issue
Block a user