Files
Cloud-Tour-to-Libo/app/data_platform/sql_console.py
T

151 lines
5.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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,
}