antern-bot / db.py
johnpitteera's picture
Upload folder using huggingface_hub
2dd2de0 verified
Raw
History Blame Contribute Delete
7.74 kB
"""Database access layer.
Two responsibilities:
1. Introspect the schema once at startup (cached) so the bot can be told
what tables/columns/relationships exist.
2. Execute model-generated SQL safely: READ-ONLY, single statement, row-capped.
The read-only guard is the core safety mechanism. We connect with a single
login, so we cannot rely on database-level permissions; instead every query is
validated to be a single SELECT/WITH statement before it ever reaches the server.
"""
from __future__ import annotations
import datetime
import decimal
import re
import struct
import uuid
from collections import OrderedDict
import pyodbc
import config
# SQL Server's datetimeoffset (ODBC type -155) isn't decoded by pyodbc natively;
# without this converter, selecting any CreatedDate/ModifiedDate/DeletedDate
# column raises "ODBC SQL type -155 is not yet supported". Decode the 20-byte
# SQL_SS_TIMESTAMPOFFSET struct into a tz-aware datetime.
SQL_SS_TIMESTAMPOFFSET = -155
def _decode_datetimeoffset(raw: bytes):
try:
y, mo, d, h, mi, s, frac, tzh, tzm = struct.unpack("<6hI2h", raw)
return datetime.datetime(
y, mo, d, h, mi, s, frac // 1000,
datetime.timezone(datetime.timedelta(hours=tzh, minutes=tzm)),
)
except Exception:
return None
class UnsafeQueryError(Exception):
"""Raised when a query fails the read-only safety checks."""
# Whole-word keywords that must never appear in a query. SELECT INTO (which
# creates a table) is covered by the INTO entry.
_FORBIDDEN = [
"INSERT", "UPDATE", "DELETE", "DROP", "ALTER", "CREATE", "TRUNCATE",
"MERGE", "EXEC", "EXECUTE", "GRANT", "REVOKE", "DENY", "BACKUP",
"RESTORE", "INTO", "SHUTDOWN", "RECONFIGURE", "WAITFOR",
]
_FORBIDDEN_RE = re.compile(r"\b(" + "|".join(_FORBIDDEN) + r")\b", re.IGNORECASE)
# Stored-procedure prefixes (sp_, xp_) used for system access.
_PROC_RE = re.compile(r"\b(sp_|xp_)\w+", re.IGNORECASE)
_COMMENT_BLOCK = re.compile(r"/\*.*?\*/", re.DOTALL)
_COMMENT_LINE = re.compile(r"--[^\n]*")
def _connect() -> pyodbc.Connection:
conn = pyodbc.connect(config.connection_string(), timeout=15)
conn.timeout = config.QUERY_TIMEOUT # query (command) timeout in seconds
conn.add_output_converter(SQL_SS_TIMESTAMPOFFSET, _decode_datetimeoffset)
return conn
def _strip_comments(sql: str) -> str:
sql = _COMMENT_BLOCK.sub(" ", sql)
sql = _COMMENT_LINE.sub(" ", sql)
return sql
def validate_readonly(sql: str) -> str:
"""Validate that `sql` is a single read-only statement. Returns the cleaned
SQL (no trailing semicolon) or raises UnsafeQueryError."""
if not sql or not sql.strip():
raise UnsafeQueryError("Empty query.")
cleaned = _strip_comments(sql).strip()
# Disallow stacked statements: a semicolon is only allowed as the very
# last character.
body = cleaned[:-1] if cleaned.endswith(";") else cleaned
if ";" in body:
raise UnsafeQueryError(
"Multiple statements are not allowed. Send a single SELECT query."
)
first = body.lstrip().split(None, 1)[0].upper() if body.strip() else ""
if first not in ("SELECT", "WITH"):
raise UnsafeQueryError(
"Only SELECT (or WITH ... SELECT) queries are allowed."
)
if _FORBIDDEN_RE.search(body):
bad = _FORBIDDEN_RE.search(body).group(1).upper()
raise UnsafeQueryError(f"Disallowed keyword in query: {bad}.")
if _PROC_RE.search(body):
raise UnsafeQueryError("Stored-procedure calls are not allowed.")
return body
def _jsonify(value):
"""Convert SQL Server values into JSON-serialisable Python values."""
if value is None:
return None
if isinstance(value, (datetime.datetime, datetime.date, datetime.time)):
return value.isoformat()
if isinstance(value, decimal.Decimal):
# keep integers as ints, others as float
return int(value) if value == value.to_integral_value() else float(value)
if isinstance(value, uuid.UUID):
return str(value)
if isinstance(value, (bytes, bytearray)):
return value.hex()
return value
def run_query(sql: str) -> dict:
"""Validate and execute a read-only query.
Returns a dict: {columns: [...], rows: [[...]], row_count, truncated}.
"""
body = validate_readonly(sql)
conn = _connect()
try:
cur = conn.cursor()
cur.execute(body)
if cur.description is None:
return {"columns": [], "rows": [], "row_count": 0, "truncated": False}
columns = [d[0] for d in cur.description]
cap = config.MAX_RESULT_ROWS
raw = cur.fetchmany(cap + 1)
truncated = len(raw) > cap
raw = raw[:cap]
rows = [[_jsonify(v) for v in row] for row in raw]
return {
"columns": columns,
"rows": rows,
"row_count": len(rows),
"truncated": truncated,
}
finally:
conn.close()
def introspect_schema() -> str:
"""Build a compact text description of the schema for the system prompt:
every table with its columns, plus foreign-key relationships."""
conn = _connect()
try:
cur = conn.cursor()
cur.execute(
"""
SELECT t.TABLE_NAME, c.COLUMN_NAME, c.DATA_TYPE
FROM INFORMATION_SCHEMA.TABLES t
JOIN INFORMATION_SCHEMA.COLUMNS c
ON t.TABLE_NAME = c.TABLE_NAME AND t.TABLE_SCHEMA = c.TABLE_SCHEMA
WHERE t.TABLE_TYPE = 'BASE TABLE'
ORDER BY t.TABLE_NAME, c.ORDINAL_POSITION
"""
)
# Boilerplate audit columns present on nearly every table — omitted from
# the listing to save prompt tokens (IsActive is kept; it's meaningful).
AUDIT_COLS = {
"CreatedBy", "CreatedDate", "ModifiedBy", "ModifiedDate",
"DeletedBy", "DeletedDate",
}
tables: "OrderedDict[str, list[str]]" = OrderedDict()
for tname, cname, dtype in cur.fetchall():
if cname in AUDIT_COLS:
continue
tables.setdefault(tname, []).append(f"{cname} {dtype}")
# Foreign keys for join hints.
cur.execute(
"""
SELECT
fk_tab.name AS fk_table, fk_col.name AS fk_column,
pk_tab.name AS pk_table, pk_col.name AS pk_column
FROM sys.foreign_key_columns fkc
JOIN sys.tables fk_tab ON fkc.parent_object_id = fk_tab.object_id
JOIN sys.columns fk_col
ON fkc.parent_object_id = fk_col.object_id
AND fkc.parent_column_id = fk_col.column_id
JOIN sys.tables pk_tab ON fkc.referenced_object_id = pk_tab.object_id
JOIN sys.columns pk_col
ON fkc.referenced_object_id = pk_col.object_id
AND fkc.referenced_column_id = pk_col.column_id
ORDER BY fk_tab.name, fk_col.name
"""
)
fks = [
f"{r.fk_table}.{r.fk_column} -> {r.pk_table}.{r.pk_column}"
for r in cur.fetchall()
]
finally:
conn.close()
lines = ["# Tables (table(column type, ...))", ""]
for tname, cols in tables.items():
lines.append(f"{tname}({', '.join(cols)})")
if fks:
lines.append("")
lines.append("# Foreign keys (from -> to)")
lines.append("")
lines.extend(fks)
return "\n".join(lines)
def ping() -> str:
"""Quick connectivity check; returns the server version."""
conn = _connect()
try:
cur = conn.cursor()
cur.execute("SELECT @@VERSION")
return cur.fetchone()[0]
finally:
conn.close()