Spaces:
Sleeping
Sleeping
| """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() | |