"""Project-scoped SQL console execution for the relational data center.""" from __future__ import annotations import re import time import uuid from decimal import Decimal from typing import Any from fastapi import HTTPException from psycopg import Error, sql from app.config import settings from app.data_platform.schema import get_project_database from app.db import get_conn MAX_SQL_LENGTH = 50_000 MAX_RESULT_ROWS = 500 STATEMENT_TIMEOUT_MS = 15_000 ALLOWED_STATEMENTS = {"select", "with", "insert", "update", "delete", "explain"} FORBIDDEN_SQL_PATTERN = re.compile( r"\b(?:" r"copy|call|do|grant|revoke|prepare|execute|deallocate|listen|unlisten|notify|" r"vacuum|analyze|cluster|reindex|refresh|reset|show|" r"create|alter|drop|truncate|comment|security|lock" r")\b", re.IGNORECASE, ) def _json_safe(value: Any) -> Any: if isinstance(value, Decimal): return float(value) if isinstance(value, uuid.UUID): return str(value) if hasattr(value, "isoformat"): return value.isoformat() if isinstance(value, (bytes, bytearray, memoryview)): return bytes(value).hex() return value def validate_console_sql(raw_sql: str, schema_name: str) -> tuple[str, str]: """Return a safe, single DML/query statement and its leading keyword.""" statement = raw_sql.strip() if not statement: raise ValueError("请输入要执行的 SQL") if len(statement) > MAX_SQL_LENGTH: raise ValueError("单次 SQL 不能超过 50,000 个字符") if "--" in statement or "/*" in statement or "*/" in statement: raise ValueError("控制台暂不支持 SQL 注释,请删除注释后重试") if statement.endswith(";"): statement = statement[:-1].rstrip() if ";" in statement: raise ValueError("一次只能执行一条 SQL 语句") match = re.match(r"^\s*([a-z]+)\b", statement, re.IGNORECASE) keyword = match.group(1).lower() if match else "" if keyword not in ALLOWED_STATEMENTS: raise ValueError( "控制台仅支持 SELECT、WITH、INSERT、UPDATE、DELETE 和 EXPLAIN;" "建表或修改结构请使用“新建数据表 / 表结构”" ) if keyword == "with" and not re.search( r"\b(?:select|insert|update|delete)\b", statement, re.IGNORECASE ): raise ValueError("WITH 语句必须包含 SELECT、INSERT、UPDATE 或 DELETE") if FORBIDDEN_SQL_PATTERN.search(statement): raise ValueError("SQL 包含控制台不允许执行的结构或权限操作") lowered = statement.lower() protected_schemas = { "public", "information_schema", "pg_catalog", settings.db_schema.lower(), } for protected in protected_schemas: if re.search(rf"(? dict[str, Any]: database = await get_project_database(project_id) if not database: raise HTTPException(404, "关系数据库不存在或已收藏") schema_name = str(database["schema_name"]) statement, keyword = validate_console_sql(raw_sql, schema_name) started_at = time.perf_counter() async with get_conn() as conn: try: async with conn.cursor() as cur: await cur.execute( "SELECT set_config('statement_timeout', %s, true)", (str(STATEMENT_TIMEOUT_MS),), ) await cur.execute( sql.SQL("SET LOCAL search_path TO {}, pg_catalog").format( sql.Identifier(schema_name) ) ) await cur.execute(statement) columns: list[str] = [] rows: list[dict[str, Any]] = [] truncated = False if cur.description: columns = [item.name for item in cur.description] fetched = await cur.fetchmany(MAX_RESULT_ROWS + 1) truncated = len(fetched) > MAX_RESULT_ROWS for row in fetched[:MAX_RESULT_ROWS]: rows.append({key: _json_safe(value) for key, value in dict(row).items()}) affected_rows = max(int(cur.rowcount or 0), 0) command = str(cur.statusmessage or keyword.upper()) await conn.commit() except Error as exc: await conn.rollback() detail = getattr(exc, "diag", None) primary = getattr(detail, "message_primary", None) if detail else None raise ValueError(f"PostgreSQL 执行失败:{primary or str(exc)}") from exc duration_ms = round((time.perf_counter() - started_at) * 1000, 2) return { "database": { "project_id": project_id, "display_name": database["display_name"], "database_name": database["database_name"], "schema_name": schema_name, }, "statement_type": keyword, "command": command, "columns": columns, "rows": rows, "returned_rows": len(rows), "affected_rows": affected_rows, "truncated": truncated, "max_result_rows": MAX_RESULT_ROWS, "duration_ms": duration_ms, }