151 lines
5.5 KiB
Python
151 lines
5.5 KiB
Python
"""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"(?<![a-z0-9_]){re.escape(protected)}\s*\.", lowered):
|
||
raise ValueError("控制台禁止访问当前项目之外的数据库空间")
|
||
|
||
for matched_schema in re.findall(r"\b(biz_[a-z0-9_]+)\s*\.", lowered):
|
||
if matched_schema != schema_name.lower():
|
||
raise ValueError("控制台禁止访问其他项目的数据库空间")
|
||
|
||
return statement, keyword
|
||
|
||
|
||
async def execute_project_sql(project_id: str, raw_sql: str) -> 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,
|
||
}
|