2338 lines
92 KiB
Python
2338 lines
92 KiB
Python
"""MySQL implementation of the Data Center service contract.
|
||
|
||
The React Data Center keeps using the existing ``/data-platform`` endpoints.
|
||
Only their persistence implementation changes: every logical database now maps
|
||
to a real MySQL database, while catalog, API-permission and audit metadata live
|
||
in the protected MySQL control database.
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import csv
|
||
from datetime import date, datetime, timezone
|
||
from decimal import Decimal
|
||
import hashlib
|
||
import io
|
||
import json
|
||
import re
|
||
import time
|
||
from typing import Any, Iterable
|
||
import uuid
|
||
import warnings
|
||
|
||
from fastapi import HTTPException
|
||
from pymysql import MySQLError
|
||
|
||
from app.config import settings
|
||
from app.data_platform.mysql_db import (
|
||
control_database_name,
|
||
data_pool_available,
|
||
data_pool_generation,
|
||
get_data_conn,
|
||
)
|
||
from app.data_platform.registry import TABLE_REGISTRY, FieldDefinition, TableDefinition
|
||
from app.data_platform.schema import (
|
||
IDENTIFIER_PATTERN,
|
||
SYSTEM_COLUMN_CODES,
|
||
_custom_definition_from_payload,
|
||
_field_row,
|
||
_structure_field_from_payload,
|
||
custom_definition_from_create_sql,
|
||
)
|
||
|
||
|
||
MAX_CSV_BYTES = 50 * 1024 * 1024
|
||
MAX_IMPORT_ROWS = 50_000
|
||
MAX_EXPORT_ROWS = 200_000
|
||
PREVIEW_ROWS = 5
|
||
MAX_REPORTED_ERRORS = 100
|
||
MAX_SQL_LENGTH = 50_000
|
||
MAX_RESULT_ROWS = 500
|
||
STATEMENT_TIMEOUT_MS = 15_000
|
||
|
||
_registry_generation = -1
|
||
_MYSQL_IDENTIFIER = re.compile(r"^[A-Za-z][A-Za-z0-9_]{0,63}$")
|
||
_MYSQL_SYSTEM_DATABASES = {"information_schema", "mysql", "performance_schema", "sys"}
|
||
_ALLOWED_SQL = {"select", "with", "insert", "update", "delete", "explain", "show", "describe", "desc"}
|
||
_FORBIDDEN_SQL = re.compile(
|
||
r"\b(?:call|grant|revoke|prepare|execute|deallocate|load|outfile|dumpfile|"
|
||
r"create|alter|drop|truncate|rename|lock|unlock|handler|install|uninstall|use)\b",
|
||
re.IGNORECASE,
|
||
)
|
||
|
||
|
||
def _require_data_center() -> None:
|
||
if not data_pool_available():
|
||
raise HTTPException(
|
||
503,
|
||
"数据中心 MySQL 服务未连接,请检查 DATA_MYSQL_URL 与 MySQL 服务状态",
|
||
)
|
||
|
||
|
||
def _q(identifier: str) -> str:
|
||
if not _MYSQL_IDENTIFIER.fullmatch(identifier):
|
||
raise ValueError(f"不合法的 MySQL 标识符:{identifier}")
|
||
return f"`{identifier}`"
|
||
|
||
|
||
def _decode_json(value: Any, fallback: Any) -> Any:
|
||
if value is None:
|
||
return fallback
|
||
if isinstance(value, (dict, list)):
|
||
return value
|
||
if isinstance(value, (bytes, bytearray)):
|
||
value = value.decode("utf-8")
|
||
try:
|
||
return json.loads(str(value))
|
||
except (TypeError, ValueError, json.JSONDecodeError):
|
||
return fallback
|
||
|
||
|
||
def _json_safe(value: Any) -> Any:
|
||
if isinstance(value, Decimal):
|
||
return float(value)
|
||
if isinstance(value, uuid.UUID):
|
||
return str(value)
|
||
if isinstance(value, (date, datetime)):
|
||
return value.isoformat()
|
||
if isinstance(value, (bytes, bytearray, memoryview)):
|
||
return bytes(value).hex()
|
||
if isinstance(value, dict):
|
||
return {key: _json_safe(item) for key, item in value.items()}
|
||
if isinstance(value, (list, tuple)):
|
||
return [_json_safe(item) for item in value]
|
||
return value
|
||
|
||
|
||
def _row_json(row: dict[str, Any] | None) -> dict[str, Any] | None:
|
||
if row is None:
|
||
return None
|
||
return {key: _json_safe(value) for key, value in dict(row).items()}
|
||
|
||
|
||
def project_database_name(project_id: str) -> str:
|
||
"""Use the user-entered database code as the physical MySQL name."""
|
||
|
||
return project_id
|
||
|
||
|
||
def _managed_index_name(prefix: str, column_code: str) -> str:
|
||
candidate = f"{prefix}_{column_code}"
|
||
if len(candidate) <= 64:
|
||
return candidate
|
||
digest = hashlib.sha1(candidate.encode("utf-8")).hexdigest()[:8]
|
||
return f"{candidate[:55]}_{digest}"
|
||
|
||
|
||
def _mysql_type(sql_type: str, *, unique: bool = False) -> str:
|
||
value = re.sub(r"\s+", " ", sql_type.strip().upper())
|
||
if value == "UUID":
|
||
return "CHAR(36)"
|
||
if value in {"JSON", "JSONB"}:
|
||
return "JSON"
|
||
if value.startswith("TIMESTAMPTZ") or "WITH TIME ZONE" in value:
|
||
precision = re.search(r"\((\d)\)", value)
|
||
return f"DATETIME({precision.group(1)})" if precision else "DATETIME(6)"
|
||
if value.startswith("TIMESTAMP"):
|
||
precision = re.search(r"\((\d)\)", value)
|
||
return f"DATETIME({precision.group(1)})" if precision else "DATETIME(6)"
|
||
if value in {"DOUBLE PRECISION", "FLOAT8"}:
|
||
return "DOUBLE"
|
||
if value in {"REAL", "FLOAT4"}:
|
||
return "FLOAT"
|
||
if value in {"BOOLEAN", "BOOL"}:
|
||
return "BOOLEAN"
|
||
if value in {"INTEGER", "INT", "INT4", "SERIAL"}:
|
||
return "INT"
|
||
if value in {"SMALLINT", "INT2", "SMALLSERIAL"}:
|
||
return "SMALLINT"
|
||
if value in {"BIGINT", "INT8", "BIGSERIAL"}:
|
||
return "BIGINT"
|
||
numeric = re.fullmatch(r"(?:NUMERIC|DECIMAL)(\s*\([^)]*\))?", value)
|
||
if numeric:
|
||
return f"DECIMAL{numeric.group(1) or '(18,4)'}".replace(" ", "")
|
||
varchar = re.fullmatch(r"(?:CHARACTER VARYING|VARCHAR)\s*(\(\d+\))?", value)
|
||
if varchar:
|
||
length = int((varchar.group(1) or "(255)")[1:-1])
|
||
if not 1 <= length <= 16_383:
|
||
raise ValueError("VARCHAR 长度必须在 1–16383 之间")
|
||
if unique and length > 512:
|
||
raise ValueError("唯一文本字段长度不能超过 512,以兼容 utf8mb4 索引")
|
||
return f"VARCHAR({length})"
|
||
char = re.fullmatch(r"(?:CHARACTER|CHAR)\s*(\(\d+\))?", value)
|
||
if char:
|
||
length = int((char.group(1) or "(1)")[1:-1])
|
||
if not 1 <= length <= 255:
|
||
raise ValueError("CHAR 长度必须在 1–255 之间")
|
||
return f"CHAR({length})"
|
||
if value == "TEXT":
|
||
# MySQL cannot create an unrestricted UNIQUE index on TEXT.
|
||
return "VARCHAR(512)" if unique else "TEXT"
|
||
if value == "DATE":
|
||
return "DATE"
|
||
raise ValueError(f"暂不支持的 MySQL 字段类型:{sql_type}")
|
||
|
||
|
||
def _mysql_default(default_sql: str | None, mysql_type: str) -> str | None:
|
||
if not default_sql:
|
||
return None
|
||
value = re.sub(
|
||
r"\s*::\s*(?:text|varchar|json|jsonb|uuid|date|timestamptz)\s*$",
|
||
"",
|
||
default_sql.strip(),
|
||
flags=re.IGNORECASE,
|
||
)
|
||
lowered = value.lower()
|
||
if lowered in {"now()", "current_timestamp"}:
|
||
return "CURRENT_TIMESTAMP(6)"
|
||
if lowered == "gen_random_uuid()":
|
||
return "(UUID())"
|
||
if lowered == "true":
|
||
return "1"
|
||
if lowered == "false":
|
||
return "0"
|
||
if mysql_type == "JSON":
|
||
if value == "'{}'":
|
||
return "(JSON_OBJECT())"
|
||
if value == "'[]'":
|
||
return "(JSON_ARRAY())"
|
||
return f"({value})"
|
||
if mysql_type in {"TEXT", "BLOB"}:
|
||
return f"({value})"
|
||
return value
|
||
|
||
|
||
def _column_clause(field: FieldDefinition, *, include_name: bool = True) -> str:
|
||
mysql_type = _mysql_type(field.sql_type, unique=field.unique)
|
||
parts = [_q(field.code), mysql_type] if include_name else [mysql_type]
|
||
parts.append("NOT NULL" if field.required else "NULL")
|
||
default = _mysql_default(field.default_sql, mysql_type)
|
||
if default is not None:
|
||
parts.extend(("DEFAULT", default))
|
||
return " ".join(parts)
|
||
|
||
|
||
async def ensure_platform_registry() -> bool:
|
||
"""Create the protected MySQL catalog used by both new centers."""
|
||
|
||
global _registry_generation
|
||
generation = data_pool_generation()
|
||
if _registry_generation == generation:
|
||
return True
|
||
if not data_pool_available():
|
||
return False
|
||
statements = (
|
||
"""
|
||
CREATE TABLE IF NOT EXISTS project_databases (
|
||
project_id VARCHAR(63) PRIMARY KEY,
|
||
tenant_id VARCHAR(100) NOT NULL,
|
||
display_name VARCHAR(100) NOT NULL,
|
||
database_name VARCHAR(64) NOT NULL UNIQUE,
|
||
schema_name VARCHAR(64) NOT NULL UNIQUE,
|
||
engine VARCHAR(16) NOT NULL DEFAULT 'mysql',
|
||
status VARCHAR(24) NOT NULL DEFAULT 'ready',
|
||
created_at DATETIME(6) NOT NULL DEFAULT CURRENT_TIMESTAMP(6),
|
||
updated_at DATETIME(6) NOT NULL DEFAULT CURRENT_TIMESTAMP(6)
|
||
ON UPDATE CURRENT_TIMESTAMP(6)
|
||
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_0900_ai_ci
|
||
""",
|
||
"""
|
||
CREATE TABLE IF NOT EXISTS project_table_definitions (
|
||
project_id VARCHAR(63) NOT NULL,
|
||
source_code VARCHAR(63) NOT NULL,
|
||
table_code VARCHAR(63) NOT NULL,
|
||
label VARCHAR(100) NOT NULL,
|
||
group_name VARCHAR(100) NOT NULL DEFAULT '自定义',
|
||
description TEXT NOT NULL,
|
||
fields_json JSON NOT NULL,
|
||
origin VARCHAR(16) NOT NULL DEFAULT 'custom',
|
||
display_order INT NOT NULL DEFAULT 1000,
|
||
allow_create BOOLEAN NOT NULL DEFAULT TRUE,
|
||
allow_update BOOLEAN NOT NULL DEFAULT TRUE,
|
||
allow_delete BOOLEAN NOT NULL DEFAULT TRUE,
|
||
status VARCHAR(16) NOT NULL DEFAULT 'active',
|
||
created_by VARCHAR(191) NOT NULL,
|
||
created_at DATETIME(6) NOT NULL DEFAULT CURRENT_TIMESTAMP(6),
|
||
updated_at DATETIME(6) NOT NULL DEFAULT CURRENT_TIMESTAMP(6)
|
||
ON UPDATE CURRENT_TIMESTAMP(6),
|
||
PRIMARY KEY (project_id, source_code),
|
||
UNIQUE KEY uq_project_table_code (project_id, table_code),
|
||
KEY idx_project_table_status (project_id, status, display_order)
|
||
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_0900_ai_ci
|
||
""",
|
||
"""
|
||
CREATE TABLE IF NOT EXISTS data_change_logs (
|
||
id BIGINT UNSIGNED NOT NULL AUTO_INCREMENT PRIMARY KEY,
|
||
tenant_id VARCHAR(100) NOT NULL,
|
||
project_id VARCHAR(63) NOT NULL,
|
||
table_code VARCHAR(63) NOT NULL,
|
||
record_id CHAR(36) NOT NULL,
|
||
operation VARCHAR(24) NOT NULL,
|
||
before_data JSON NULL,
|
||
after_data JSON NULL,
|
||
actor VARCHAR(191) NOT NULL,
|
||
created_at DATETIME(6) NOT NULL DEFAULT CURRENT_TIMESTAMP(6),
|
||
KEY idx_data_change_record (project_id, table_code, record_id),
|
||
KEY idx_data_change_created (created_at)
|
||
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_0900_ai_ci
|
||
""",
|
||
"""
|
||
CREATE TABLE IF NOT EXISTS dbeaver_access_grants (
|
||
id CHAR(36) PRIMARY KEY,
|
||
display_name VARCHAR(100) NOT NULL,
|
||
mysql_username VARCHAR(32) NOT NULL UNIQUE,
|
||
database_id VARCHAR(63) NOT NULL,
|
||
database_name VARCHAR(64) NOT NULL,
|
||
permission_level VARCHAR(16) NOT NULL,
|
||
ssh_key_type VARCHAR(32) NOT NULL,
|
||
ssh_public_key TEXT NOT NULL,
|
||
ssh_key_fingerprint VARCHAR(100) NOT NULL,
|
||
created_by VARCHAR(191) NOT NULL,
|
||
status VARCHAR(16) NOT NULL DEFAULT 'active',
|
||
revoked_at DATETIME(6) NULL,
|
||
created_at DATETIME(6) NOT NULL DEFAULT CURRENT_TIMESTAMP(6),
|
||
KEY idx_dbeaver_access_status (status, created_at),
|
||
KEY idx_dbeaver_access_database (database_id, status),
|
||
KEY idx_dbeaver_access_fingerprint (ssh_key_fingerprint, status)
|
||
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_0900_ai_ci
|
||
""",
|
||
"""
|
||
CREATE TABLE IF NOT EXISTS api_clients (
|
||
id CHAR(36) PRIMARY KEY,
|
||
name VARCHAR(100) NOT NULL,
|
||
description VARCHAR(500) NOT NULL DEFAULT '',
|
||
owner VARCHAR(191) NOT NULL,
|
||
status VARCHAR(16) NOT NULL DEFAULT 'active',
|
||
created_at DATETIME(6) NOT NULL DEFAULT CURRENT_TIMESTAMP(6),
|
||
updated_at DATETIME(6) NOT NULL DEFAULT CURRENT_TIMESTAMP(6)
|
||
ON UPDATE CURRENT_TIMESTAMP(6),
|
||
KEY idx_api_clients_status (status)
|
||
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_0900_ai_ci
|
||
""",
|
||
"""
|
||
CREATE TABLE IF NOT EXISTS api_credentials (
|
||
id CHAR(36) PRIMARY KEY,
|
||
client_id CHAR(36) NOT NULL,
|
||
name VARCHAR(100) NOT NULL,
|
||
key_prefix VARCHAR(24) NOT NULL UNIQUE,
|
||
key_hash CHAR(64) NOT NULL UNIQUE,
|
||
expires_at DATETIME(6) NULL,
|
||
last_used_at DATETIME(6) NULL,
|
||
revoked_at DATETIME(6) NULL,
|
||
created_at DATETIME(6) NOT NULL DEFAULT CURRENT_TIMESTAMP(6),
|
||
KEY idx_api_credentials_client (client_id, revoked_at)
|
||
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_0900_ai_ci
|
||
""",
|
||
"""
|
||
CREATE TABLE IF NOT EXISTS api_policies (
|
||
id CHAR(36) PRIMARY KEY,
|
||
client_id CHAR(36) NOT NULL,
|
||
database_id VARCHAR(63) NOT NULL,
|
||
table_code VARCHAR(63) NOT NULL DEFAULT '*',
|
||
actions_json JSON NOT NULL,
|
||
readable_fields_json JSON NOT NULL,
|
||
writable_fields_json JSON NOT NULL,
|
||
row_filter_json JSON NOT NULL,
|
||
status VARCHAR(16) NOT NULL DEFAULT 'active',
|
||
created_at DATETIME(6) NOT NULL DEFAULT CURRENT_TIMESTAMP(6),
|
||
updated_at DATETIME(6) NOT NULL DEFAULT CURRENT_TIMESTAMP(6)
|
||
ON UPDATE CURRENT_TIMESTAMP(6),
|
||
UNIQUE KEY uq_api_policy_scope (client_id, database_id, table_code),
|
||
KEY idx_api_policy_client (client_id, status)
|
||
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_0900_ai_ci
|
||
""",
|
||
"""
|
||
CREATE TABLE IF NOT EXISTS api_call_logs (
|
||
id BIGINT UNSIGNED NOT NULL AUTO_INCREMENT PRIMARY KEY,
|
||
request_id CHAR(36) NOT NULL,
|
||
client_id CHAR(36) NULL,
|
||
credential_id CHAR(36) NULL,
|
||
method VARCHAR(12) NOT NULL,
|
||
path VARCHAR(500) NOT NULL,
|
||
database_id VARCHAR(63) NULL,
|
||
table_code VARCHAR(63) NULL,
|
||
action_name VARCHAR(24) NULL,
|
||
status_code INT NOT NULL,
|
||
duration_ms DECIMAL(12,2) NOT NULL DEFAULT 0,
|
||
source_ip VARCHAR(64) NULL,
|
||
error_message VARCHAR(500) NULL,
|
||
created_at DATETIME(6) NOT NULL DEFAULT CURRENT_TIMESTAMP(6),
|
||
KEY idx_api_call_created (created_at),
|
||
KEY idx_api_call_client (client_id, created_at)
|
||
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_0900_ai_ci
|
||
""",
|
||
"""
|
||
CREATE TABLE IF NOT EXISTS admin_action_logs (
|
||
id BIGINT UNSIGNED NOT NULL AUTO_INCREMENT PRIMARY KEY,
|
||
actor VARCHAR(191) NOT NULL,
|
||
action_name VARCHAR(64) NOT NULL,
|
||
resource_type VARCHAR(64) NOT NULL,
|
||
resource_id VARCHAR(191) NOT NULL,
|
||
outcome VARCHAR(16) NOT NULL,
|
||
statement_hash CHAR(64) NULL,
|
||
affected_rows BIGINT NOT NULL DEFAULT 0,
|
||
source_ip VARCHAR(64) NULL,
|
||
details_json JSON NOT NULL,
|
||
created_at DATETIME(6) NOT NULL DEFAULT CURRENT_TIMESTAMP(6),
|
||
KEY idx_admin_action_created (created_at),
|
||
KEY idx_admin_action_actor (actor, created_at),
|
||
KEY idx_admin_action_resource (resource_type, resource_id, created_at)
|
||
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_0900_ai_ci
|
||
""",
|
||
)
|
||
async with get_data_conn() as conn:
|
||
async with conn.cursor() as cur:
|
||
for statement in statements:
|
||
with warnings.catch_warnings():
|
||
warnings.filterwarnings(
|
||
"ignore",
|
||
message=r"Table '.*' already exists",
|
||
)
|
||
await cur.execute(statement)
|
||
await conn.commit()
|
||
_registry_generation = generation
|
||
return True
|
||
|
||
|
||
async def _ready() -> None:
|
||
_require_data_center()
|
||
await ensure_platform_registry()
|
||
|
||
|
||
def _table_from_row(row: dict[str, Any]) -> TableDefinition:
|
||
fields = tuple(
|
||
FieldDefinition(
|
||
code=str(item["code"]),
|
||
label=str(item["label"]),
|
||
sql_type=str(item["sql_type"]),
|
||
data_type=str(item.get("data_type") or "text"),
|
||
required=bool(item.get("required")),
|
||
searchable=bool(item.get("searchable")),
|
||
sortable=bool(item.get("sortable", True)),
|
||
editable=bool(item.get("editable", True)),
|
||
visible_in_list=bool(item.get("visible_in_list", True)),
|
||
default_sql=str(item["default_sql"]) if item.get("default_sql") else None,
|
||
unique=bool(item.get("unique")),
|
||
options=tuple(str(value) for value in item.get("options") or ()),
|
||
)
|
||
for item in _decode_json(row.get("fields_json"), [])
|
||
)
|
||
return TableDefinition(
|
||
code=str(row["table_code"]),
|
||
label=str(row["label"]),
|
||
group=str(row.get("group_name") or "自定义"),
|
||
description=str(row.get("description") or ""),
|
||
fields=fields,
|
||
allow_create=bool(row.get("allow_create", True)),
|
||
allow_update=bool(row.get("allow_update", True)),
|
||
allow_delete=bool(row.get("allow_delete", True)),
|
||
)
|
||
|
||
|
||
async def get_project_database(project_id: str) -> dict[str, Any] | None:
|
||
await _ready()
|
||
async with get_data_conn() as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
"SELECT * FROM project_databases WHERE project_id=%s AND status='ready'",
|
||
(project_id,),
|
||
)
|
||
return _row_json(await cur.fetchone())
|
||
|
||
|
||
async def _database_or_404(project_id: str) -> dict[str, Any]:
|
||
database = await get_project_database(project_id)
|
||
if database:
|
||
return database
|
||
raise HTTPException(404, "关系数据库不存在或未就绪")
|
||
|
||
|
||
async def list_project_table_entries(
|
||
project_id: str,
|
||
) -> tuple[tuple[TableDefinition, str, str], ...]:
|
||
await _ready()
|
||
async with get_data_conn() as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
"""
|
||
SELECT * FROM project_table_definitions
|
||
WHERE project_id=%s AND status='active'
|
||
ORDER BY display_order, created_at, table_code
|
||
""",
|
||
(project_id,),
|
||
)
|
||
rows = await cur.fetchall()
|
||
return tuple(
|
||
(_table_from_row(dict(row)), str(row["origin"]), str(row["source_code"]))
|
||
for row in rows
|
||
)
|
||
|
||
|
||
async def list_project_table_definitions(project_id: str) -> tuple[TableDefinition, ...]:
|
||
return tuple(item[0] for item in await list_project_table_entries(project_id))
|
||
|
||
|
||
async def resolve_project_table_entry(
|
||
project_id: str,
|
||
table_code: str,
|
||
) -> tuple[TableDefinition, str, str] | None:
|
||
return next(
|
||
(entry for entry in await list_project_table_entries(project_id) if entry[0].code == table_code),
|
||
None,
|
||
)
|
||
|
||
|
||
async def _table_or_404(project_id: str, table_code: str) -> TableDefinition:
|
||
entry = await resolve_project_table_entry(project_id, table_code)
|
||
if entry:
|
||
return entry[0]
|
||
raise HTTPException(404, "数据表未注册")
|
||
|
||
|
||
async def _create_physical_table(
|
||
database_name: str,
|
||
definition: TableDefinition,
|
||
) -> None:
|
||
business_columns = [_column_clause(field) for field in definition.fields]
|
||
columns = [
|
||
"`id` CHAR(36) NOT NULL PRIMARY KEY",
|
||
"`tenant_id` VARCHAR(100) NOT NULL",
|
||
"`project_id` VARCHAR(63) NOT NULL",
|
||
*business_columns,
|
||
"`created_at` DATETIME(6) NOT NULL DEFAULT CURRENT_TIMESTAMP(6)",
|
||
"`updated_at` DATETIME(6) NOT NULL DEFAULT CURRENT_TIMESTAMP(6) ON UPDATE CURRENT_TIMESTAMP(6)",
|
||
"`deleted_at` DATETIME(6) NULL",
|
||
"`deleted_by` VARCHAR(191) NULL",
|
||
"KEY `idx_updated_at` (`updated_at`)",
|
||
"KEY `idx_project_active` (`project_id`, `deleted_at`)",
|
||
]
|
||
for field in definition.fields:
|
||
if field.unique:
|
||
columns.append(
|
||
f"UNIQUE KEY {_q(_managed_index_name('uq', field.code))} ({_q(field.code)})"
|
||
)
|
||
async with get_data_conn(database_name) as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
f"CREATE TABLE IF NOT EXISTS {_q(definition.code)} "
|
||
f"({', '.join(columns)}) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 "
|
||
"COLLATE=utf8mb4_0900_ai_ci"
|
||
)
|
||
await conn.commit()
|
||
|
||
|
||
async def ensure_project_database(
|
||
project_id: str,
|
||
tenant_id: str,
|
||
display_name: str,
|
||
) -> dict[str, Any]:
|
||
await _ready()
|
||
if not IDENTIFIER_PATTERN.fullmatch(project_id):
|
||
raise ValueError("数据库编码格式不正确")
|
||
next_name = display_name.strip()
|
||
if not next_name:
|
||
raise ValueError("请输入数据库名称")
|
||
if len(next_name) > 100:
|
||
raise ValueError("数据库名称不能超过 100 个字符")
|
||
database_name = project_database_name(project_id)
|
||
if database_name.lower() in {
|
||
*_MYSQL_SYSTEM_DATABASES,
|
||
control_database_name().lower(),
|
||
}:
|
||
raise ValueError("数据库编码不能与 MySQL 系统数据库或数据中心控制库重名")
|
||
|
||
async with get_data_conn() as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
"SELECT database_name FROM project_databases WHERE project_id=%s",
|
||
(project_id,),
|
||
)
|
||
existing = await cur.fetchone()
|
||
if existing:
|
||
# Preserve the catalogued physical name during idempotent
|
||
# recovery. Legacy prefixed names are migrated separately.
|
||
database_name = str(existing["database_name"])
|
||
await cur.execute(
|
||
f"CREATE DATABASE IF NOT EXISTS {_q(database_name)} "
|
||
"CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci"
|
||
)
|
||
await cur.execute(
|
||
"""
|
||
INSERT INTO project_databases (
|
||
project_id, tenant_id, display_name, database_name,
|
||
schema_name, engine, status
|
||
) VALUES (%s, %s, %s, %s, %s, 'mysql', 'provisioning')
|
||
ON DUPLICATE KEY UPDATE
|
||
tenant_id=%s, display_name=%s,
|
||
status='provisioning', updated_at=CURRENT_TIMESTAMP(6)
|
||
""",
|
||
(
|
||
project_id,
|
||
tenant_id,
|
||
next_name,
|
||
database_name,
|
||
database_name,
|
||
tenant_id,
|
||
next_name,
|
||
),
|
||
)
|
||
await conn.commit()
|
||
|
||
try:
|
||
# A newly created Data Center database is intentionally empty. Existing
|
||
# registered definitions are provisioned only for idempotent recovery
|
||
# of an already catalogued database; legacy templates are never injected
|
||
# into a new database here.
|
||
for definition in await list_project_table_definitions(project_id):
|
||
await _create_physical_table(database_name, definition)
|
||
except Exception:
|
||
async with get_data_conn() as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
"UPDATE project_databases SET status='provision_failed' WHERE project_id=%s",
|
||
(project_id,),
|
||
)
|
||
await conn.commit()
|
||
raise
|
||
|
||
async with get_data_conn() as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
"UPDATE project_databases SET status='ready', updated_at=CURRENT_TIMESTAMP(6) WHERE project_id=%s",
|
||
(project_id,),
|
||
)
|
||
await conn.commit()
|
||
database = await get_project_database(project_id)
|
||
if not database:
|
||
raise ValueError("MySQL 数据库创建失败")
|
||
return database
|
||
|
||
|
||
async def rename_project_database(project_id: str, display_name: str) -> dict[str, Any]:
|
||
await _ready()
|
||
next_name = display_name.strip()
|
||
if not next_name:
|
||
raise ValueError("请输入数据库名称")
|
||
if len(next_name) > 100:
|
||
raise ValueError("数据库名称不能超过 100 个字符")
|
||
async with get_data_conn() as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
"""
|
||
UPDATE project_databases
|
||
SET display_name=%s, updated_at=CURRENT_TIMESTAMP(6)
|
||
WHERE project_id=%s AND status='ready'
|
||
""",
|
||
(next_name, project_id),
|
||
)
|
||
if cur.rowcount != 1:
|
||
raise ValueError("关系数据库不存在或未就绪")
|
||
await conn.commit()
|
||
database = await get_project_database(project_id)
|
||
return database or {}
|
||
|
||
|
||
async def delete_project_database(project_id: str, confirm_name: str) -> dict[str, str]:
|
||
await _ready()
|
||
database = await get_project_database(project_id)
|
||
if not database:
|
||
raise ValueError("关系数据库不存在或未就绪")
|
||
if confirm_name.strip() not in {
|
||
str(database["display_name"]),
|
||
str(database["project_id"]),
|
||
str(database["database_name"]),
|
||
}:
|
||
raise ValueError("确认名称不正确,未执行删除")
|
||
|
||
async with get_data_conn() as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(f"DROP DATABASE {_q(str(database['database_name']))}")
|
||
await cur.execute(
|
||
"DELETE FROM project_table_definitions WHERE project_id=%s",
|
||
(project_id,),
|
||
)
|
||
await cur.execute(
|
||
"UPDATE api_policies SET status='disabled' WHERE database_id=%s",
|
||
(project_id,),
|
||
)
|
||
await cur.execute("DELETE FROM project_databases WHERE project_id=%s", (project_id,))
|
||
await conn.commit()
|
||
return {"project_id": project_id, "status": "deleted"}
|
||
|
||
|
||
async def _create_custom_table_definition(
|
||
project_id: str,
|
||
definition: TableDefinition,
|
||
field_rows: list[dict[str, Any]],
|
||
actor: str,
|
||
) -> TableDefinition:
|
||
database = await _database_or_404(project_id)
|
||
if any(item.code == definition.code for item in await list_project_table_definitions(project_id)):
|
||
raise ValueError("当前数据库中已存在同名数据表")
|
||
await _create_physical_table(str(database["database_name"]), definition)
|
||
try:
|
||
async with get_data_conn() as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
"""
|
||
INSERT INTO project_table_definitions (
|
||
project_id, source_code, table_code, label, group_name,
|
||
description, fields_json, origin, display_order,
|
||
allow_create, allow_update, allow_delete, created_by
|
||
) VALUES (%s, %s, %s, %s, %s, %s, %s, 'custom', 1000, %s, %s, %s, %s)
|
||
""",
|
||
(
|
||
project_id,
|
||
definition.code,
|
||
definition.code,
|
||
definition.label,
|
||
definition.group,
|
||
definition.description,
|
||
json.dumps(field_rows, ensure_ascii=False),
|
||
int(definition.allow_create),
|
||
int(definition.allow_update),
|
||
int(definition.allow_delete),
|
||
actor,
|
||
),
|
||
)
|
||
await conn.commit()
|
||
except Exception:
|
||
async with get_data_conn(str(database["database_name"])) as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(f"DROP TABLE IF EXISTS {_q(definition.code)}")
|
||
await conn.commit()
|
||
raise
|
||
return definition
|
||
|
||
|
||
async def create_custom_table(
|
||
project_id: str,
|
||
body: dict[str, Any],
|
||
actor: str,
|
||
) -> TableDefinition:
|
||
definition, rows = _custom_definition_from_payload(body)
|
||
return await _create_custom_table_definition(project_id, definition, rows, actor)
|
||
|
||
|
||
async def create_custom_table_from_sql(
|
||
project_id: str,
|
||
body: dict[str, Any],
|
||
actor: str,
|
||
) -> TableDefinition:
|
||
definition, rows = custom_definition_from_create_sql(body)
|
||
return await _create_custom_table_definition(project_id, definition, rows, actor)
|
||
|
||
|
||
async def _persist_fields(
|
||
project_id: str,
|
||
source_code: str,
|
||
fields: Iterable[FieldDefinition],
|
||
) -> None:
|
||
async with get_data_conn() as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
"""
|
||
UPDATE project_table_definitions
|
||
SET fields_json=%s, updated_at=CURRENT_TIMESTAMP(6)
|
||
WHERE project_id=%s AND source_code=%s AND status='active'
|
||
""",
|
||
(
|
||
json.dumps([_field_row(field) for field in fields], ensure_ascii=False),
|
||
project_id,
|
||
source_code,
|
||
),
|
||
)
|
||
if cur.rowcount != 1:
|
||
raise ValueError("数据表字段元数据不存在")
|
||
await conn.commit()
|
||
|
||
|
||
async def rename_custom_table(
|
||
project_id: str,
|
||
table_code: str,
|
||
new_code: str,
|
||
new_label: str,
|
||
) -> TableDefinition:
|
||
database = await _database_or_404(project_id)
|
||
next_code = new_code.strip()
|
||
next_label = new_label.strip()
|
||
if not IDENTIFIER_PATTERN.fullmatch(next_code):
|
||
raise ValueError("表名必须以小写字母开头,只能包含小写字母、数字和下划线,长度为 2–63 位")
|
||
if next_code in {"data_change_logs", "graph_sync_queue"}:
|
||
raise ValueError("该表名已被系统占用")
|
||
if not next_label:
|
||
raise ValueError("请输入数据表名称")
|
||
entry = await resolve_project_table_entry(project_id, table_code)
|
||
if not entry:
|
||
raise ValueError("数据表不存在")
|
||
current, _origin, source_code = entry
|
||
if next_code in TABLE_REGISTRY and next_code != source_code:
|
||
raise ValueError("该表名已被系统占用")
|
||
if next_code != table_code and any(
|
||
item.code == next_code for item in await list_project_table_definitions(project_id)
|
||
):
|
||
raise ValueError("当前数据库中已存在同名数据表")
|
||
|
||
if next_code != table_code:
|
||
async with get_data_conn(str(database["database_name"])) as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
f"RENAME TABLE {_q(table_code)} TO {_q(next_code)}"
|
||
)
|
||
await conn.commit()
|
||
async with get_data_conn() as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
"""
|
||
UPDATE project_table_definitions
|
||
SET table_code=%s, label=%s, updated_at=CURRENT_TIMESTAMP(6)
|
||
WHERE project_id=%s AND source_code=%s AND status='active'
|
||
""",
|
||
(next_code, next_label, project_id, source_code),
|
||
)
|
||
await cur.execute(
|
||
"UPDATE data_change_logs SET table_code=%s WHERE project_id=%s AND table_code=%s",
|
||
(next_code, project_id, table_code),
|
||
)
|
||
await conn.commit()
|
||
return TableDefinition(
|
||
code=next_code,
|
||
label=next_label,
|
||
group=current.group,
|
||
description=current.description,
|
||
fields=current.fields,
|
||
allow_create=current.allow_create,
|
||
allow_update=current.allow_update,
|
||
allow_delete=current.allow_delete,
|
||
)
|
||
|
||
|
||
async def delete_custom_table(
|
||
project_id: str,
|
||
table_code: str,
|
||
confirm_name: str,
|
||
) -> dict[str, str]:
|
||
database = await _database_or_404(project_id)
|
||
entry = await resolve_project_table_entry(project_id, table_code)
|
||
if not entry:
|
||
raise ValueError("数据表不存在")
|
||
current, _origin, source_code = entry
|
||
if confirm_name.strip() not in {table_code, current.label}:
|
||
raise ValueError("确认名称不正确,未执行删除")
|
||
async with get_data_conn(str(database["database_name"])) as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(f"DROP TABLE {_q(table_code)}")
|
||
await conn.commit()
|
||
async with get_data_conn() as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
"""
|
||
UPDATE project_table_definitions
|
||
SET status='deleted', updated_at=CURRENT_TIMESTAMP(6)
|
||
WHERE project_id=%s AND source_code=%s
|
||
""",
|
||
(project_id, source_code),
|
||
)
|
||
await cur.execute(
|
||
"UPDATE api_policies SET status='disabled' WHERE database_id=%s AND table_code=%s",
|
||
(project_id, table_code),
|
||
)
|
||
await conn.commit()
|
||
return {"project_id": project_id, "table_code": table_code, "status": "deleted"}
|
||
|
||
|
||
async def create_table_column(
|
||
project_id: str,
|
||
table_code: str,
|
||
body: dict[str, Any],
|
||
) -> TableDefinition:
|
||
database = await _database_or_404(project_id)
|
||
entry = await resolve_project_table_entry(project_id, table_code)
|
||
if not entry:
|
||
raise ValueError("数据表不存在")
|
||
current, _origin, source_code = entry
|
||
definition = _structure_field_from_payload(body)
|
||
if any(field.code == definition.code for field in current.fields):
|
||
raise ValueError("当前数据表已存在同名字段")
|
||
if definition.required and not definition.default_sql:
|
||
raise ValueError("新增非空字段必须设置默认值,避免现有记录无法补值")
|
||
try:
|
||
async with get_data_conn(str(database["database_name"])) as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
f"ALTER TABLE {_q(table_code)} ADD COLUMN {_column_clause(definition)}"
|
||
)
|
||
if definition.unique:
|
||
await cur.execute(
|
||
f"ALTER TABLE {_q(table_code)} ADD UNIQUE KEY "
|
||
f"{_q(_managed_index_name('uq', definition.code))} ({_q(definition.code)})"
|
||
)
|
||
await conn.commit()
|
||
except MySQLError as exc:
|
||
raise ValueError(f"MySQL 拒绝结构修改:{exc}") from exc
|
||
next_fields = (*current.fields, definition)
|
||
await _persist_fields(project_id, source_code, next_fields)
|
||
return TableDefinition(
|
||
current.code, current.label, current.group, current.description, next_fields,
|
||
current.allow_create, current.allow_update, current.allow_delete,
|
||
)
|
||
|
||
|
||
async def _single_column_unique_indexes(
|
||
database_name: str,
|
||
table_code: str,
|
||
column_code: str,
|
||
) -> list[str]:
|
||
async with get_data_conn(database_name) as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
"""
|
||
SELECT index_name AS index_name
|
||
FROM information_schema.statistics
|
||
WHERE table_schema=%s AND table_name=%s AND column_name=%s
|
||
AND non_unique=0 AND index_name <> 'PRIMARY'
|
||
GROUP BY index_name
|
||
HAVING COUNT(*)=1
|
||
""",
|
||
(database_name, table_code, column_code),
|
||
)
|
||
return [str(row["index_name"]) for row in await cur.fetchall()]
|
||
|
||
|
||
async def update_table_column(
|
||
project_id: str,
|
||
table_code: str,
|
||
column_code: str,
|
||
body: dict[str, Any],
|
||
) -> TableDefinition:
|
||
if column_code in SYSTEM_COLUMN_CODES:
|
||
raise ValueError("系统字段不可修改")
|
||
database = await _database_or_404(project_id)
|
||
entry = await resolve_project_table_entry(project_id, table_code)
|
||
if not entry:
|
||
raise ValueError("数据表不存在")
|
||
current, _origin, source_code = entry
|
||
current_field = next((field for field in current.fields if field.code == column_code), None)
|
||
if not current_field:
|
||
raise ValueError("业务字段不存在或不可修改")
|
||
definition = _structure_field_from_payload(body, current=current_field)
|
||
if definition.code != column_code and any(
|
||
field.code == definition.code for field in current.fields
|
||
):
|
||
raise ValueError("当前数据表已存在同名字段")
|
||
database_name = str(database["database_name"])
|
||
unique_indexes = await _single_column_unique_indexes(database_name, table_code, column_code)
|
||
try:
|
||
async with get_data_conn(database_name) as conn:
|
||
async with conn.cursor() as cur:
|
||
for index_name in unique_indexes:
|
||
await cur.execute(
|
||
f"ALTER TABLE {_q(table_code)} DROP INDEX {_q(index_name)}"
|
||
)
|
||
await cur.execute(
|
||
f"ALTER TABLE {_q(table_code)} CHANGE COLUMN {_q(column_code)} {_column_clause(definition)}"
|
||
)
|
||
if definition.unique:
|
||
await cur.execute(
|
||
f"ALTER TABLE {_q(table_code)} ADD UNIQUE KEY "
|
||
f"{_q(_managed_index_name('uq', definition.code))} ({_q(definition.code)})"
|
||
)
|
||
await conn.commit()
|
||
except MySQLError as exc:
|
||
raise ValueError(f"MySQL 拒绝结构修改:{exc}") from exc
|
||
next_fields = tuple(
|
||
definition if field.code == column_code else field for field in current.fields
|
||
)
|
||
await _persist_fields(project_id, source_code, next_fields)
|
||
return TableDefinition(
|
||
current.code, current.label, current.group, current.description, next_fields,
|
||
current.allow_create, current.allow_update, current.allow_delete,
|
||
)
|
||
|
||
|
||
async def delete_table_column(
|
||
project_id: str,
|
||
table_code: str,
|
||
column_code: str,
|
||
confirm_name: str,
|
||
) -> TableDefinition:
|
||
if column_code in SYSTEM_COLUMN_CODES:
|
||
raise ValueError("系统字段不可删除")
|
||
if confirm_name.strip() != column_code:
|
||
raise ValueError("确认字段名不正确,未执行删除")
|
||
database = await _database_or_404(project_id)
|
||
entry = await resolve_project_table_entry(project_id, table_code)
|
||
if not entry:
|
||
raise ValueError("数据表不存在")
|
||
current, _origin, source_code = entry
|
||
if not any(field.code == column_code for field in current.fields):
|
||
raise ValueError("业务字段不存在或不可删除")
|
||
next_fields = tuple(field for field in current.fields if field.code != column_code)
|
||
if not next_fields:
|
||
raise ValueError("数据表至少需要保留一个业务字段")
|
||
try:
|
||
async with get_data_conn(str(database["database_name"])) as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
f"ALTER TABLE {_q(table_code)} DROP COLUMN {_q(column_code)}"
|
||
)
|
||
await conn.commit()
|
||
except MySQLError as exc:
|
||
raise ValueError(f"MySQL 拒绝结构修改:{exc}") from exc
|
||
await _persist_fields(project_id, source_code, next_fields)
|
||
return TableDefinition(
|
||
current.code, current.label, current.group, current.description, next_fields,
|
||
current.allow_create, current.allow_update, current.allow_delete,
|
||
)
|
||
|
||
|
||
async def list_databases() -> list[dict[str, Any]]:
|
||
await _ready()
|
||
async with get_data_conn() as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
"""
|
||
SELECT * FROM project_databases
|
||
WHERE status='ready'
|
||
ORDER BY CASE WHEN project_id='yunyou_libo' THEN 0 ELSE 1 END,
|
||
created_at DESC
|
||
"""
|
||
)
|
||
databases = [dict(row) for row in await cur.fetchall()]
|
||
|
||
result: list[dict[str, Any]] = []
|
||
for database in databases:
|
||
project_id = str(database["project_id"])
|
||
database_name = str(database["database_name"])
|
||
definitions = await list_project_table_definitions(project_id)
|
||
record_count = 0
|
||
latest_updates: list[datetime] = []
|
||
async with get_data_conn(database_name) as conn:
|
||
async with conn.cursor() as cur:
|
||
for definition in definitions:
|
||
await cur.execute(
|
||
f"SELECT COUNT(*) AS count, MAX(`updated_at`) AS updated_at "
|
||
f"FROM {_q(definition.code)} WHERE `deleted_at` IS NULL"
|
||
)
|
||
summary = await cur.fetchone()
|
||
record_count += int(summary["count"] or 0)
|
||
if summary.get("updated_at"):
|
||
latest_updates.append(summary["updated_at"])
|
||
row = _row_json(database) or {}
|
||
row.update(
|
||
{
|
||
"table_count": len(definitions),
|
||
"record_count": record_count,
|
||
"updated_at": _json_safe(max(latest_updates))
|
||
if latest_updates
|
||
else _json_safe(database.get("updated_at")),
|
||
}
|
||
)
|
||
result.append(row)
|
||
return result
|
||
|
||
|
||
async def list_tables(project_id: str) -> dict[str, Any]:
|
||
database = await _database_or_404(project_id)
|
||
entries = await list_project_table_entries(project_id)
|
||
tables: list[dict[str, Any]] = []
|
||
async with get_data_conn(str(database["database_name"])) as conn:
|
||
async with conn.cursor() as cur:
|
||
for definition, origin, _source_code in entries:
|
||
await cur.execute(
|
||
f"SELECT COUNT(*) AS count, MAX(`updated_at`) AS updated_at "
|
||
f"FROM {_q(definition.code)} WHERE `deleted_at` IS NULL"
|
||
)
|
||
summary = await cur.fetchone()
|
||
item = definition.as_dict()
|
||
item.update(
|
||
{
|
||
"record_count": int(summary["count"] or 0),
|
||
"updated_at": _json_safe(summary.get("updated_at")),
|
||
"is_custom": origin == "custom",
|
||
"origin": origin,
|
||
}
|
||
)
|
||
tables.append(item)
|
||
return {"database": _row_json(database), "tables": tables}
|
||
|
||
|
||
async def inspect_table(project_id: str, table_code: str) -> dict[str, Any]:
|
||
"""Return the physical MySQL structure and foreign-key relationships."""
|
||
|
||
table = await _table_or_404(project_id, table_code)
|
||
database = await _database_or_404(project_id)
|
||
database_name = str(database["database_name"])
|
||
definitions = await list_project_table_definitions(project_id)
|
||
table_labels = {definition.code: definition.label for definition in definitions}
|
||
field_definitions = {field.code: field for field in table.fields}
|
||
system_labels = {
|
||
"id": "记录 ID",
|
||
"tenant_id": "租户 ID",
|
||
"project_id": "数据库编码",
|
||
"created_at": "创建时间",
|
||
"updated_at": "更新时间",
|
||
"deleted_at": "删除时间",
|
||
"deleted_by": "删除人",
|
||
}
|
||
|
||
async with get_data_conn(database_name) as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
"""
|
||
SELECT ordinal_position AS ordinal_position,
|
||
column_name AS name, column_type AS sql_type,
|
||
is_nullable AS is_nullable, column_default AS default_value,
|
||
column_comment AS comment, column_key, extra
|
||
FROM information_schema.columns
|
||
WHERE table_schema=%s AND table_name=%s
|
||
ORDER BY ordinal_position
|
||
""",
|
||
(database_name, table_code),
|
||
)
|
||
physical_columns = [dict(row) for row in await cur.fetchall()]
|
||
|
||
await cur.execute(
|
||
"""
|
||
SELECT tc.constraint_name AS constraint_name,
|
||
tc.constraint_type AS constraint_type,
|
||
GROUP_CONCAT(kcu.column_name ORDER BY kcu.ordinal_position) AS columns_csv
|
||
FROM information_schema.table_constraints tc
|
||
LEFT JOIN information_schema.key_column_usage kcu
|
||
ON kcu.constraint_schema=tc.constraint_schema
|
||
AND kcu.table_name=tc.table_name
|
||
AND kcu.constraint_name=tc.constraint_name
|
||
WHERE tc.table_schema=%s AND tc.table_name=%s
|
||
GROUP BY tc.constraint_name, tc.constraint_type
|
||
ORDER BY FIELD(tc.constraint_type, 'PRIMARY KEY', 'FOREIGN KEY', 'UNIQUE'),
|
||
tc.constraint_name
|
||
""",
|
||
(database_name, table_code),
|
||
)
|
||
constraint_rows = [dict(row) for row in await cur.fetchall()]
|
||
|
||
await cur.execute(
|
||
"""
|
||
SELECT index_name AS index_name, non_unique AS non_unique,
|
||
GROUP_CONCAT(column_name ORDER BY seq_in_index) AS columns_csv,
|
||
index_type AS index_type
|
||
FROM information_schema.statistics
|
||
WHERE table_schema=%s AND table_name=%s
|
||
GROUP BY index_name, non_unique, index_type
|
||
ORDER BY (index_name='PRIMARY') DESC, non_unique, index_name
|
||
""",
|
||
(database_name, table_code),
|
||
)
|
||
index_rows = [dict(row) for row in await cur.fetchall()]
|
||
|
||
await cur.execute(
|
||
"""
|
||
SELECT kcu.constraint_name AS name,
|
||
kcu.table_schema AS source_schema,
|
||
kcu.table_name AS source_table,
|
||
kcu.column_name AS source_column,
|
||
kcu.referenced_table_schema AS target_schema,
|
||
kcu.referenced_table_name AS target_table,
|
||
kcu.referenced_column_name AS target_column,
|
||
rc.update_rule AS on_update,
|
||
rc.delete_rule AS on_delete
|
||
FROM information_schema.key_column_usage kcu
|
||
JOIN information_schema.referential_constraints rc
|
||
ON rc.constraint_schema=kcu.constraint_schema
|
||
AND rc.constraint_name=kcu.constraint_name
|
||
WHERE kcu.referenced_table_name IS NOT NULL
|
||
AND (
|
||
(kcu.table_schema=%s AND kcu.table_name=%s)
|
||
OR (kcu.referenced_table_schema=%s AND kcu.referenced_table_name=%s)
|
||
)
|
||
ORDER BY kcu.constraint_name, kcu.ordinal_position
|
||
""",
|
||
(database_name, table_code, database_name, table_code),
|
||
)
|
||
relationship_rows = [dict(row) for row in await cur.fetchall()]
|
||
|
||
type_names = {
|
||
"PRIMARY KEY": "primary_key",
|
||
"UNIQUE": "unique",
|
||
"FOREIGN KEY": "foreign_key",
|
||
"CHECK": "check",
|
||
}
|
||
constraints: list[dict[str, Any]] = []
|
||
for row in constraint_rows:
|
||
columns = str(row.get("columns_csv") or "").split(",")
|
||
constraint_type = type_names.get(
|
||
str(row["constraint_type"]).upper(),
|
||
str(row["constraint_type"]).lower().replace(" ", "_"),
|
||
)
|
||
constraints.append(
|
||
{
|
||
"name": str(row["constraint_name"]),
|
||
"type": constraint_type,
|
||
"columns": [value for value in columns if value],
|
||
"definition": f"{row['constraint_type']} ({', '.join(columns)})",
|
||
}
|
||
)
|
||
primary_columns = {
|
||
column
|
||
for constraint in constraints
|
||
if constraint["type"] == "primary_key"
|
||
for column in constraint["columns"]
|
||
}
|
||
unique_columns = {
|
||
column
|
||
for constraint in constraints
|
||
if constraint["type"] in {"primary_key", "unique"}
|
||
and len(constraint["columns"]) == 1
|
||
for column in constraint["columns"]
|
||
}
|
||
foreign_columns = {
|
||
column
|
||
for constraint in constraints
|
||
if constraint["type"] == "foreign_key"
|
||
for column in constraint["columns"]
|
||
}
|
||
|
||
columns: list[dict[str, Any]] = []
|
||
for row in physical_columns:
|
||
code = str(row["name"])
|
||
definition = field_definitions.get(code)
|
||
columns.append(
|
||
{
|
||
"ordinal_position": int(row["ordinal_position"]),
|
||
"name": code,
|
||
"label": definition.label if definition else system_labels.get(code, code),
|
||
"data_type": definition.data_type if definition else "system",
|
||
"sql_type": str(row["sql_type"]),
|
||
"nullable": str(row["is_nullable"]).upper() == "YES",
|
||
"default": _json_safe(row.get("default_value")),
|
||
"primary_key": code in primary_columns,
|
||
"unique": code in unique_columns,
|
||
"foreign_key": code in foreign_columns,
|
||
"editable": bool(definition.editable) if definition else False,
|
||
"structure_editable": bool(definition) and code not in system_labels,
|
||
"searchable": bool(definition.searchable) if definition else False,
|
||
"sortable": bool(definition.sortable) if definition else False,
|
||
"visible_in_list": bool(definition.visible_in_list) if definition else False,
|
||
"comment": row.get("comment") or None,
|
||
}
|
||
)
|
||
|
||
indexes = []
|
||
for row in index_rows:
|
||
index_columns = [value for value in str(row.get("columns_csv") or "").split(",") if value]
|
||
name = str(row["index_name"])
|
||
indexes.append(
|
||
{
|
||
"name": name,
|
||
"unique": not bool(row["non_unique"]),
|
||
"primary": name == "PRIMARY",
|
||
"definition": (
|
||
f"{'UNIQUE ' if not row['non_unique'] else ''}INDEX {name} "
|
||
f"({', '.join(index_columns)}) USING {row['index_type']}"
|
||
),
|
||
}
|
||
)
|
||
|
||
grouped_relationships: dict[tuple[str, str, str], dict[str, Any]] = {}
|
||
for row in relationship_rows:
|
||
key = (str(row["name"]), str(row["source_table"]), str(row["target_table"]))
|
||
relationship = grouped_relationships.setdefault(
|
||
key,
|
||
{
|
||
"name": str(row["name"]),
|
||
"direction": "outgoing" if str(row["source_table"]) == table_code else "incoming",
|
||
"source_schema": str(row["source_schema"]),
|
||
"source_table": str(row["source_table"]),
|
||
"source_table_label": table_labels.get(str(row["source_table"]), str(row["source_table"])),
|
||
"source_columns": [],
|
||
"target_schema": str(row["target_schema"]),
|
||
"target_table": str(row["target_table"]),
|
||
"target_table_label": table_labels.get(str(row["target_table"]), str(row["target_table"])),
|
||
"target_columns": [],
|
||
"on_update": str(row["on_update"]),
|
||
"on_delete": str(row["on_delete"]),
|
||
},
|
||
)
|
||
relationship["source_columns"].append(str(row["source_column"]))
|
||
relationship["target_columns"].append(str(row["target_column"]))
|
||
|
||
return {
|
||
"table": {
|
||
"schema_name": database_name,
|
||
"table_code": table.code,
|
||
"label": table.label,
|
||
"description": table.description,
|
||
"column_count": len(columns),
|
||
"index_count": len(indexes),
|
||
"constraint_count": len(constraints),
|
||
"primary_key": list(primary_columns),
|
||
},
|
||
"columns": columns,
|
||
"constraints": constraints,
|
||
"indexes": indexes,
|
||
"relationships": list(grouped_relationships.values()),
|
||
}
|
||
|
||
|
||
def _coerce_value(field: FieldDefinition, value: Any) -> Any:
|
||
if value in ("", None):
|
||
return None
|
||
if field.data_type == "json":
|
||
if isinstance(value, str):
|
||
try:
|
||
value = json.loads(value)
|
||
except json.JSONDecodeError as exc:
|
||
raise HTTPException(400, f"{field.label}不是合法 JSON") from exc
|
||
return json.dumps(value, ensure_ascii=False, separators=(",", ":"))
|
||
if field.data_type == "boolean":
|
||
if isinstance(value, bool):
|
||
return int(value)
|
||
lowered = str(value).lower()
|
||
if lowered in {"true", "1", "yes", "是", "启用"}:
|
||
return 1
|
||
if lowered in {"false", "0", "no", "否", "禁用"}:
|
||
return 0
|
||
raise HTTPException(400, f"{field.label}必须是布尔值")
|
||
if field.data_type == "number":
|
||
if any(token in field.sql_type.upper() for token in ("INT", "SERIAL")):
|
||
try:
|
||
return int(value)
|
||
except (TypeError, ValueError) as exc:
|
||
raise HTTPException(400, f"{field.label}必须是整数") from exc
|
||
try:
|
||
return Decimal(str(value))
|
||
except Exception as exc:
|
||
raise HTTPException(400, f"{field.label}必须是数字") from exc
|
||
if field.data_type == "date":
|
||
try:
|
||
return date.fromisoformat(str(value)).isoformat()
|
||
except ValueError as exc:
|
||
raise HTTPException(400, f"{field.label}必须使用 YYYY-MM-DD 格式") from exc
|
||
if field.data_type == "datetime":
|
||
parsed = str(value).replace("Z", "+00:00")
|
||
try:
|
||
return datetime.fromisoformat(parsed).replace(tzinfo=None)
|
||
except ValueError as exc:
|
||
raise HTTPException(400, f"{field.label}不是合法日期时间") from exc
|
||
if field.options and str(value) not in field.options:
|
||
raise HTTPException(400, f"{field.label}可选值为:{', '.join(field.options)}")
|
||
return value
|
||
|
||
|
||
def _validated_payload(
|
||
table: TableDefinition,
|
||
body: dict[str, Any],
|
||
*,
|
||
create: bool,
|
||
) -> dict[str, Any]:
|
||
definitions = {field.code: field for field in table.fields}
|
||
unknown = sorted(set(body) - set(definitions))
|
||
if unknown:
|
||
raise HTTPException(400, f"不允许的字段:{', '.join(unknown)}")
|
||
payload: dict[str, Any] = {}
|
||
for code, value in body.items():
|
||
field = definitions[code]
|
||
if field.editable:
|
||
payload[code] = _coerce_value(field, value)
|
||
if create:
|
||
missing = [
|
||
field.label
|
||
for field in table.fields
|
||
if field.required and payload.get(field.code) in (None, "")
|
||
]
|
||
if missing:
|
||
raise HTTPException(400, f"缺少必填字段:{', '.join(missing)}")
|
||
return payload
|
||
|
||
|
||
def _where_scope(
|
||
table: TableDefinition,
|
||
row_filter: dict[str, Any] | None,
|
||
) -> tuple[list[str], list[Any]]:
|
||
if not row_filter:
|
||
return [], []
|
||
definitions = {field.code: field for field in table.fields}
|
||
clauses: list[str] = []
|
||
params: list[Any] = []
|
||
for code, expected in row_filter.items():
|
||
field = definitions.get(code)
|
||
if not field:
|
||
raise HTTPException(403, f"接口行级范围字段不存在:{code}")
|
||
if isinstance(expected, list):
|
||
if not expected:
|
||
clauses.append("1=0")
|
||
continue
|
||
values = [_coerce_value(field, value) for value in expected]
|
||
clauses.append(f"{_q(code)} IN ({', '.join(['%s'] * len(values))})")
|
||
params.extend(values)
|
||
elif expected is None:
|
||
clauses.append(f"{_q(code)} IS NULL")
|
||
else:
|
||
clauses.append(f"{_q(code)}=%s")
|
||
params.append(_coerce_value(field, expected))
|
||
return clauses, params
|
||
|
||
|
||
def _payload_in_scope(
|
||
table: TableDefinition,
|
||
payload: dict[str, Any],
|
||
row_filter: dict[str, Any] | None,
|
||
*,
|
||
inject_missing: bool,
|
||
) -> dict[str, Any]:
|
||
if not row_filter:
|
||
return payload
|
||
definitions = {field.code: field for field in table.fields}
|
||
scoped = dict(payload)
|
||
for code, expected in row_filter.items():
|
||
field = definitions.get(code)
|
||
if not field:
|
||
raise HTTPException(403, f"接口行级范围字段不存在:{code}")
|
||
if isinstance(expected, list):
|
||
allowed = [_coerce_value(field, value) for value in expected]
|
||
if code not in scoped or scoped[code] not in allowed:
|
||
raise HTTPException(403, f"字段 {code} 超出接口授权的数据范围")
|
||
continue
|
||
normalized = _coerce_value(field, expected) if expected is not None else None
|
||
if code not in scoped and inject_missing:
|
||
scoped[code] = normalized
|
||
elif code in scoped and scoped[code] != normalized:
|
||
raise HTTPException(403, f"字段 {code} 超出接口授权的数据范围")
|
||
return scoped
|
||
|
||
|
||
def _selected_columns(
|
||
table: TableDefinition,
|
||
allowed_fields: set[str] | None = None,
|
||
) -> list[str]:
|
||
business = [field.code for field in table.fields]
|
||
if allowed_fields is not None and "*" not in allowed_fields:
|
||
business = [code for code in business if code in allowed_fields]
|
||
return ["id", *business, "created_at", "updated_at"]
|
||
|
||
|
||
def _decode_record(table: TableDefinition, row: dict[str, Any] | None) -> dict[str, Any] | None:
|
||
if row is None:
|
||
return None
|
||
result = _row_json(row) or {}
|
||
for field in table.fields:
|
||
if field.data_type == "json" and field.code in result and result[field.code] is not None:
|
||
result[field.code] = _decode_json(result[field.code], result[field.code])
|
||
if field.data_type == "boolean" and field.code in result and result[field.code] is not None:
|
||
result[field.code] = bool(result[field.code])
|
||
return result
|
||
|
||
|
||
async def list_records(
|
||
project_id: str,
|
||
table_code: str,
|
||
*,
|
||
page: int,
|
||
page_size: int,
|
||
search: str | None,
|
||
sort_field: str | None,
|
||
sort_order: str,
|
||
allowed_fields: set[str] | None = None,
|
||
row_filter: dict[str, Any] | None = None,
|
||
) -> dict[str, Any]:
|
||
table = await _table_or_404(project_id, table_code)
|
||
database = await _database_or_404(project_id)
|
||
page = max(1, page)
|
||
page_size = max(1, min(5000, page_size))
|
||
offset = (page - 1) * page_size
|
||
readable = {field.code for field in table.fields}
|
||
if allowed_fields is not None and "*" not in allowed_fields:
|
||
readable &= allowed_fields
|
||
allowed_sort = {
|
||
"id",
|
||
"created_at",
|
||
"updated_at",
|
||
*(field.code for field in table.fields if field.sortable and field.code in readable),
|
||
}
|
||
order_field = sort_field if sort_field in allowed_sort else "updated_at"
|
||
order_keyword = "ASC" if sort_order.lower() == "asc" else "DESC"
|
||
where_parts = ["`deleted_at` IS NULL"]
|
||
params: list[Any] = []
|
||
scope_parts, scope_params = _where_scope(table, row_filter)
|
||
where_parts.extend(scope_parts)
|
||
params.extend(scope_params)
|
||
searchable = [
|
||
field for field in table.fields if field.searchable and field.code in readable
|
||
]
|
||
if search and searchable:
|
||
pattern = f"%{search.strip()}%"
|
||
where_parts.append(
|
||
"(" + " OR ".join(
|
||
f"CAST({_q(field.code)} AS CHAR) LIKE %s" for field in searchable
|
||
) + ")"
|
||
)
|
||
params.extend(pattern for _ in searchable)
|
||
where_clause = " AND ".join(where_parts)
|
||
columns = _selected_columns(table, allowed_fields)
|
||
|
||
async with get_data_conn(str(database["database_name"])) as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
f"SELECT COUNT(*) AS count FROM {_q(table.code)} WHERE {where_clause}",
|
||
params,
|
||
)
|
||
total = int((await cur.fetchone())["count"])
|
||
await cur.execute(
|
||
f"SELECT {', '.join(_q(column) for column in columns)} "
|
||
f"FROM {_q(table.code)} WHERE {where_clause} "
|
||
f"ORDER BY {_q(order_field)} {order_keyword} LIMIT %s OFFSET %s",
|
||
[*params, page_size, offset],
|
||
)
|
||
rows = [_decode_record(table, dict(row)) for row in await cur.fetchall()]
|
||
table_payload = table.as_dict()
|
||
if allowed_fields is not None and "*" not in allowed_fields:
|
||
table_payload["fields"] = [
|
||
field for field in table_payload["fields"]
|
||
if field["code"] in {"id", "created_at", "updated_at", *allowed_fields}
|
||
]
|
||
return {
|
||
"table": table_payload,
|
||
"items": rows,
|
||
"page": page,
|
||
"page_size": page_size,
|
||
"total": total,
|
||
}
|
||
|
||
|
||
async def _write_audit(
|
||
cur: Any,
|
||
*,
|
||
tenant_id: str,
|
||
project_id: str,
|
||
table_code: str,
|
||
record_id: str,
|
||
operation: str,
|
||
before_data: dict[str, Any] | None,
|
||
after_data: dict[str, Any] | None,
|
||
actor: str,
|
||
) -> None:
|
||
await cur.execute(
|
||
f"""
|
||
INSERT INTO {_q(control_database_name())}.data_change_logs (
|
||
tenant_id, project_id, table_code, record_id, operation,
|
||
before_data, after_data, actor
|
||
) VALUES (%s, %s, %s, %s, %s, %s, %s, %s)
|
||
""",
|
||
(
|
||
tenant_id,
|
||
project_id,
|
||
table_code,
|
||
record_id,
|
||
operation,
|
||
json.dumps(before_data, ensure_ascii=False) if before_data is not None else None,
|
||
json.dumps(after_data, ensure_ascii=False) if after_data is not None else None,
|
||
actor,
|
||
),
|
||
)
|
||
|
||
|
||
async def _write_admin_action_audit(
|
||
cur: Any,
|
||
*,
|
||
actor: str,
|
||
action_name: str,
|
||
resource_type: str,
|
||
resource_id: str,
|
||
outcome: str,
|
||
statement_hash: str | None = None,
|
||
affected_rows: int = 0,
|
||
source_ip: str | None = None,
|
||
details: dict[str, Any] | None = None,
|
||
) -> None:
|
||
await cur.execute(
|
||
f"""
|
||
INSERT INTO {_q(control_database_name())}.admin_action_logs (
|
||
actor, action_name, resource_type, resource_id, outcome,
|
||
statement_hash, affected_rows, source_ip, details_json
|
||
) VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s)
|
||
""",
|
||
(
|
||
actor[:191],
|
||
action_name[:64],
|
||
resource_type[:64],
|
||
resource_id[:191],
|
||
outcome[:16],
|
||
statement_hash,
|
||
max(0, int(affected_rows)),
|
||
(source_ip or "")[:64] or None,
|
||
json.dumps(details or {}, ensure_ascii=False),
|
||
),
|
||
)
|
||
|
||
|
||
async def _record_admin_action_safely(**payload: Any) -> None:
|
||
try:
|
||
async with get_data_conn() as conn:
|
||
async with conn.cursor() as cur:
|
||
await _write_admin_action_audit(cur, **payload)
|
||
await conn.commit()
|
||
except Exception:
|
||
# The original security decision or SQL error must remain visible even
|
||
# if the audit store is temporarily unavailable.
|
||
return
|
||
|
||
|
||
async def list_admin_action_logs(limit: int = 200) -> list[dict[str, Any]]:
|
||
await _ready()
|
||
async with get_data_conn() as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
"""
|
||
SELECT id, actor, action_name, resource_type, resource_id,
|
||
outcome, statement_hash, affected_rows, source_ip,
|
||
details_json, created_at
|
||
FROM admin_action_logs
|
||
ORDER BY created_at DESC
|
||
LIMIT %s
|
||
""",
|
||
(max(1, min(1000, limit)),),
|
||
)
|
||
rows = []
|
||
for item in await cur.fetchall():
|
||
row = _row_json(dict(item)) or {}
|
||
row["details"] = _decode_json(row.pop("details_json", None), {})
|
||
rows.append(row)
|
||
return rows
|
||
|
||
|
||
async def _select_record(
|
||
cur: Any,
|
||
table: TableDefinition,
|
||
record_id: str,
|
||
*,
|
||
row_filter: dict[str, Any] | None = None,
|
||
) -> dict[str, Any] | None:
|
||
where = ["`id`=%s", "`deleted_at` IS NULL"]
|
||
params: list[Any] = [record_id]
|
||
scope, scope_params = _where_scope(table, row_filter)
|
||
where.extend(scope)
|
||
params.extend(scope_params)
|
||
columns = _selected_columns(table)
|
||
await cur.execute(
|
||
f"SELECT {', '.join(_q(column) for column in columns)} "
|
||
f"FROM {_q(table.code)} WHERE {' AND '.join(where)}",
|
||
params,
|
||
)
|
||
return _decode_record(table, await cur.fetchone())
|
||
|
||
|
||
async def create_record(
|
||
project_id: str,
|
||
table_code: str,
|
||
body: dict[str, Any],
|
||
actor: str,
|
||
*,
|
||
row_filter: dict[str, Any] | None = None,
|
||
) -> dict[str, Any]:
|
||
table = await _table_or_404(project_id, table_code)
|
||
if not table.allow_create:
|
||
raise HTTPException(403, "该表不允许新增")
|
||
database = await _database_or_404(project_id)
|
||
payload = _payload_in_scope(
|
||
table,
|
||
_validated_payload(table, body, create=True),
|
||
row_filter,
|
||
inject_missing=True,
|
||
)
|
||
record_id = str(uuid.uuid4())
|
||
columns = ["id", "tenant_id", "project_id", *payload.keys()]
|
||
values = [record_id, database["tenant_id"], project_id, *payload.values()]
|
||
async with get_data_conn(str(database["database_name"])) as conn:
|
||
async with conn.cursor() as cur:
|
||
try:
|
||
await cur.execute(
|
||
f"INSERT INTO {_q(table.code)} "
|
||
f"({', '.join(_q(column) for column in columns)}) "
|
||
f"VALUES ({', '.join(['%s'] * len(values))})",
|
||
values,
|
||
)
|
||
row = await _select_record(cur, table, record_id)
|
||
await _write_audit(
|
||
cur,
|
||
tenant_id=str(database["tenant_id"]),
|
||
project_id=project_id,
|
||
table_code=table_code,
|
||
record_id=record_id,
|
||
operation="create",
|
||
before_data=None,
|
||
after_data=row,
|
||
actor=actor,
|
||
)
|
||
await conn.commit()
|
||
except MySQLError as exc:
|
||
raise HTTPException(400, f"MySQL 数据约束校验失败:{exc}") from exc
|
||
return row or {}
|
||
|
||
|
||
async def update_record(
|
||
project_id: str,
|
||
table_code: str,
|
||
record_id: str,
|
||
body: dict[str, Any],
|
||
actor: str,
|
||
*,
|
||
row_filter: dict[str, Any] | None = None,
|
||
) -> dict[str, Any]:
|
||
table = await _table_or_404(project_id, table_code)
|
||
if not table.allow_update:
|
||
raise HTTPException(403, "该表不允许修改")
|
||
try:
|
||
str(uuid.UUID(record_id))
|
||
except ValueError as exc:
|
||
raise HTTPException(400, "记录 ID 格式错误") from exc
|
||
database = await _database_or_404(project_id)
|
||
payload = _validated_payload(table, body, create=False)
|
||
if not payload:
|
||
raise HTTPException(400, "没有可修改字段")
|
||
async with get_data_conn(str(database["database_name"])) as conn:
|
||
async with conn.cursor() as cur:
|
||
before = await _select_record(cur, table, record_id, row_filter=row_filter)
|
||
if not before:
|
||
raise HTTPException(404, "记录不存在或超出接口授权范围")
|
||
merged = {**before, **payload}
|
||
_payload_in_scope(table, merged, row_filter, inject_missing=False)
|
||
assignments = [f"{_q(column)}=%s" for column in payload]
|
||
if any(field.code == "version" for field in table.fields):
|
||
assignments.append("`version`=`version`+1")
|
||
assignments.append("`updated_at`=CURRENT_TIMESTAMP(6)")
|
||
try:
|
||
await cur.execute(
|
||
f"UPDATE {_q(table.code)} SET {', '.join(assignments)} WHERE `id`=%s",
|
||
[*payload.values(), record_id],
|
||
)
|
||
after = await _select_record(cur, table, record_id)
|
||
await _write_audit(
|
||
cur,
|
||
tenant_id=str(database["tenant_id"]),
|
||
project_id=project_id,
|
||
table_code=table_code,
|
||
record_id=record_id,
|
||
operation="update",
|
||
before_data=before,
|
||
after_data=after,
|
||
actor=actor,
|
||
)
|
||
await conn.commit()
|
||
except MySQLError as exc:
|
||
raise HTTPException(400, f"MySQL 数据约束校验失败:{exc}") from exc
|
||
return after or {}
|
||
|
||
|
||
async def delete_record(
|
||
project_id: str,
|
||
table_code: str,
|
||
record_id: str,
|
||
actor: str,
|
||
*,
|
||
row_filter: dict[str, Any] | None = None,
|
||
) -> dict[str, bool]:
|
||
table = await _table_or_404(project_id, table_code)
|
||
if not table.allow_delete:
|
||
raise HTTPException(403, "该表不允许删除")
|
||
try:
|
||
str(uuid.UUID(record_id))
|
||
except ValueError as exc:
|
||
raise HTTPException(400, "记录 ID 格式错误") from exc
|
||
database = await _database_or_404(project_id)
|
||
async with get_data_conn(str(database["database_name"])) as conn:
|
||
async with conn.cursor() as cur:
|
||
before = await _select_record(cur, table, record_id, row_filter=row_filter)
|
||
if not before:
|
||
raise HTTPException(404, "记录不存在或超出接口授权范围")
|
||
await cur.execute(
|
||
f"UPDATE {_q(table.code)} SET `deleted_at`=CURRENT_TIMESTAMP(6), "
|
||
"`deleted_by`=%s, `updated_at`=CURRENT_TIMESTAMP(6) WHERE `id`=%s",
|
||
(actor, record_id),
|
||
)
|
||
await _write_audit(
|
||
cur,
|
||
tenant_id=str(database["tenant_id"]),
|
||
project_id=project_id,
|
||
table_code=table_code,
|
||
record_id=record_id,
|
||
operation="delete",
|
||
before_data=before,
|
||
after_data=None,
|
||
actor=actor,
|
||
)
|
||
await conn.commit()
|
||
return {"ok": True}
|
||
|
||
|
||
def _decode_csv(content: bytes) -> tuple[str, str]:
|
||
if not content:
|
||
raise HTTPException(400, "CSV 文件为空")
|
||
if len(content) > MAX_CSV_BYTES:
|
||
raise HTTPException(413, "CSV 文件不能超过 50 MB")
|
||
for encoding in ("utf-8-sig", "utf-8", "gb18030"):
|
||
try:
|
||
return content.decode(encoding), encoding
|
||
except UnicodeDecodeError:
|
||
continue
|
||
raise HTTPException(400, "CSV 编码无法识别,请使用 UTF-8 或 GB18030")
|
||
|
||
|
||
def _csv_reader(text: str) -> csv.DictReader:
|
||
try:
|
||
dialect = csv.Sniffer().sniff(text[:8192], delimiters=",;\t|")
|
||
except csv.Error:
|
||
dialect = csv.excel
|
||
return csv.DictReader(io.StringIO(text, newline=""), dialect=dialect)
|
||
|
||
|
||
def _normalize_header(value: str | None) -> str:
|
||
return str(value or "").replace("\ufeff", "").strip()
|
||
|
||
|
||
_SYSTEM_CSV_HEADERS = {
|
||
"id", "记录 id", "记录id", "tenant_id", "租户 id", "租户id",
|
||
"project_id", "数据库编码", "created_at", "创建时间", "updated_at", "更新时间",
|
||
"deleted_at", "deleted_by",
|
||
}
|
||
|
||
|
||
def _parse_csv(
|
||
table: TableDefinition,
|
||
content: bytes,
|
||
file_name: str,
|
||
) -> tuple[dict[str, Any], list[dict[str, Any]], list[str]]:
|
||
text, encoding = _decode_csv(content)
|
||
reader = _csv_reader(text)
|
||
headers = [_normalize_header(header) for header in (reader.fieldnames or [])]
|
||
by_code = {field.code.lower(): field for field in table.fields}
|
||
by_label = {field.label: field for field in table.fields}
|
||
mappings: list[dict[str, Any]] = []
|
||
mapped: dict[str, FieldDefinition] = {}
|
||
mapped_codes: set[str] = set()
|
||
schema_errors: list[str] = []
|
||
if not headers or not any(headers):
|
||
schema_errors.append("CSV 缺少表头")
|
||
for header in headers:
|
||
field = by_code.get(header.lower()) or by_label.get(header)
|
||
if field and field.code not in mapped_codes:
|
||
mapped[header] = field
|
||
mapped_codes.add(field.code)
|
||
mappings.append({
|
||
"source_header": header,
|
||
"field_code": field.code,
|
||
"field_label": field.label,
|
||
"status": "mapped",
|
||
})
|
||
elif field:
|
||
mappings.append({
|
||
"source_header": header,
|
||
"field_code": field.code,
|
||
"field_label": field.label,
|
||
"status": "duplicate",
|
||
})
|
||
schema_errors.append(f"字段“{field.label}”在 CSV 中重复映射")
|
||
elif header.lower() in _SYSTEM_CSV_HEADERS:
|
||
mappings.append({
|
||
"source_header": header,
|
||
"field_code": None,
|
||
"field_label": "平台自动生成",
|
||
"status": "ignored",
|
||
})
|
||
else:
|
||
mappings.append({
|
||
"source_header": header,
|
||
"field_code": None,
|
||
"field_label": None,
|
||
"status": "unknown",
|
||
})
|
||
schema_errors.append(f"无法识别 CSV 列“{header}”")
|
||
missing = [
|
||
field.label for field in table.fields if field.required and field.code not in mapped_codes
|
||
]
|
||
if missing:
|
||
schema_errors.append(f"缺少必填列:{', '.join(missing)}")
|
||
if not mapped:
|
||
schema_errors.append("CSV 没有可导入的业务字段")
|
||
|
||
rows: list[dict[str, Any]] = []
|
||
errors: list[dict[str, Any]] = []
|
||
invalid_count = 0
|
||
total_rows = 0
|
||
for physical_row, source in enumerate(reader, start=2):
|
||
normalized = {
|
||
_normalize_header(key): value for key, value in source.items() if key is not None
|
||
}
|
||
if None in source and source[None]:
|
||
invalid_count += 1
|
||
if len(errors) < MAX_REPORTED_ERRORS:
|
||
errors.append({"row": physical_row, "field": "", "kind": "invalid", "message": "该行列数多于表头"})
|
||
continue
|
||
if not any(str(value or "").strip() for value in normalized.values()):
|
||
continue
|
||
total_rows += 1
|
||
if total_rows > MAX_IMPORT_ROWS:
|
||
raise HTTPException(400, f"单次最多导入 {MAX_IMPORT_ROWS:,} 条记录")
|
||
raw: dict[str, Any] = {}
|
||
row_error: str | None = None
|
||
row_field = ""
|
||
for header, field in mapped.items():
|
||
value = str(normalized.get(header) or "").strip()
|
||
if not value:
|
||
continue
|
||
try:
|
||
raw[field.code] = _coerce_value(field, value)
|
||
except HTTPException as exc:
|
||
row_error = str(exc.detail)
|
||
row_field = field.label
|
||
break
|
||
if row_error is None:
|
||
try:
|
||
payload = _validated_payload(table, raw, create=True)
|
||
except HTTPException as exc:
|
||
row_error = str(exc.detail)
|
||
payload = {}
|
||
else:
|
||
payload = {}
|
||
if row_error or not payload:
|
||
invalid_count += 1
|
||
if len(errors) < MAX_REPORTED_ERRORS:
|
||
errors.append({
|
||
"row": physical_row,
|
||
"field": row_field,
|
||
"kind": "invalid",
|
||
"message": row_error or "该行没有可导入的业务字段",
|
||
})
|
||
continue
|
||
rows.append({"row": physical_row, "payload": payload})
|
||
if total_rows == 0:
|
||
schema_errors.append("CSV 没有可导入的数据行")
|
||
preview = {
|
||
"file_name": file_name,
|
||
"encoding": encoding,
|
||
"total_rows": total_rows,
|
||
"valid_rows": len(rows),
|
||
"invalid_rows": invalid_count,
|
||
"duplicate_rows": 0,
|
||
"file_duplicate_rows": 0,
|
||
"existing_duplicate_rows": 0,
|
||
"skipped_rows": invalid_count,
|
||
"columns": mappings,
|
||
"schema_errors": schema_errors,
|
||
"errors": errors,
|
||
"errors_truncated": invalid_count > len(errors),
|
||
"preview_rows": [],
|
||
"can_import": not schema_errors and bool(rows),
|
||
}
|
||
return preview, rows, [field.code for field in mapped.values()]
|
||
|
||
|
||
def _signature_value(value: Any) -> tuple[str, Any]:
|
||
if value is None:
|
||
return ("null", None)
|
||
if isinstance(value, bool):
|
||
return ("boolean", value)
|
||
if isinstance(value, (int, float, Decimal)):
|
||
return ("number", format(Decimal(str(value)).normalize(), "f"))
|
||
if isinstance(value, str):
|
||
stripped = value.strip()
|
||
if stripped.startswith(("{", "[")):
|
||
try:
|
||
return ("json", json.dumps(json.loads(stripped), ensure_ascii=False, sort_keys=True, separators=(",", ":")))
|
||
except json.JSONDecodeError:
|
||
pass
|
||
return ("text", stripped.casefold())
|
||
return ("json", json.dumps(_json_safe(value), ensure_ascii=False, sort_keys=True, separators=(",", ":")))
|
||
|
||
|
||
def _payload_signature(payload: dict[str, Any], fields: list[str]) -> tuple[Any, ...]:
|
||
return tuple((field, _signature_value(payload.get(field))) for field in fields)
|
||
|
||
|
||
async def _prepare_csv_import(
|
||
project_id: str,
|
||
table: TableDefinition,
|
||
database: dict[str, Any],
|
||
content: bytes,
|
||
file_name: str,
|
||
) -> tuple[dict[str, Any], list[dict[str, Any]]]:
|
||
preview, candidates, dedupe_fields = _parse_csv(table, content, file_name)
|
||
if preview["schema_errors"] or not candidates or not dedupe_fields:
|
||
preview["can_import"] = False
|
||
return preview, []
|
||
existing: set[tuple[Any, ...]] = set()
|
||
async with get_data_conn(str(database["database_name"])) as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
f"SELECT {', '.join(_q(code) for code in dedupe_fields)} "
|
||
f"FROM {_q(table.code)} WHERE `deleted_at` IS NULL"
|
||
)
|
||
for row in await cur.fetchall():
|
||
existing.add(_payload_signature(dict(row), dedupe_fields))
|
||
valid: list[dict[str, Any]] = []
|
||
seen: dict[tuple[Any, ...], int] = {}
|
||
duplicate_errors: list[dict[str, Any]] = []
|
||
file_duplicates = 0
|
||
existing_duplicates = 0
|
||
for candidate in candidates:
|
||
signature = _payload_signature(candidate["payload"], dedupe_fields)
|
||
row_number = int(candidate["row"])
|
||
if signature in existing:
|
||
existing_duplicates += 1
|
||
duplicate_errors.append({"row": row_number, "field": "", "kind": "duplicate", "message": "与当前数据表中的数据重复,已跳过"})
|
||
elif signature in seen:
|
||
file_duplicates += 1
|
||
duplicate_errors.append({"row": row_number, "field": "", "kind": "duplicate", "message": f"与 CSV 第 {seen[signature]} 行重复,已跳过"})
|
||
else:
|
||
seen[signature] = row_number
|
||
valid.append(candidate)
|
||
duplicate_count = file_duplicates + existing_duplicates
|
||
slots = max(0, MAX_REPORTED_ERRORS - len(preview["errors"]))
|
||
preview["errors"].extend(duplicate_errors[:slots])
|
||
preview["errors"].sort(key=lambda item: int(item["row"]))
|
||
preview.update({
|
||
"valid_rows": len(valid),
|
||
"duplicate_rows": duplicate_count,
|
||
"file_duplicate_rows": file_duplicates,
|
||
"existing_duplicate_rows": existing_duplicates,
|
||
"skipped_rows": int(preview["invalid_rows"]) + duplicate_count,
|
||
"errors_truncated": int(preview["invalid_rows"]) + duplicate_count > len(preview["errors"]),
|
||
"preview_rows": [_json_safe(item["payload"]) for item in valid[:PREVIEW_ROWS]],
|
||
"can_import": bool(valid),
|
||
})
|
||
return preview, valid
|
||
|
||
|
||
async def preview_csv_import(
|
||
project_id: str,
|
||
table_code: str,
|
||
content: bytes,
|
||
file_name: str,
|
||
) -> dict[str, Any]:
|
||
table = await _table_or_404(project_id, table_code)
|
||
if not table.allow_create:
|
||
raise HTTPException(403, "该表不允许导入新增记录")
|
||
database = await _database_or_404(project_id)
|
||
preview, _ = await _prepare_csv_import(project_id, table, database, content, file_name)
|
||
return preview
|
||
|
||
|
||
async def import_csv_records(
|
||
project_id: str,
|
||
table_code: str,
|
||
content: bytes,
|
||
file_name: str,
|
||
actor: str,
|
||
) -> dict[str, Any]:
|
||
table = await _table_or_404(project_id, table_code)
|
||
if not table.allow_create:
|
||
raise HTTPException(403, "该表不允许导入新增记录")
|
||
database = await _database_or_404(project_id)
|
||
preview, rows = await _prepare_csv_import(project_id, table, database, content, file_name)
|
||
if not preview["can_import"]:
|
||
detail = (preview["schema_errors"] or ["CSV 没有可导入的新数据"])[0]
|
||
raise HTTPException(400, f"CSV 校验失败:{detail}")
|
||
imported = 0
|
||
failed = 0
|
||
failed_errors: list[dict[str, Any]] = []
|
||
async with get_data_conn(str(database["database_name"])) as conn:
|
||
async with conn.cursor() as cur:
|
||
for candidate in rows:
|
||
payload = candidate["payload"]
|
||
record_id = str(uuid.uuid4())
|
||
columns = ["id", "tenant_id", "project_id", *payload.keys()]
|
||
values = [record_id, database["tenant_id"], project_id, *payload.values()]
|
||
try:
|
||
await cur.execute(
|
||
f"INSERT INTO {_q(table.code)} ({', '.join(_q(column) for column in columns)}) "
|
||
f"VALUES ({', '.join(['%s'] * len(values))})",
|
||
values,
|
||
)
|
||
imported += 1
|
||
await _write_audit(
|
||
cur,
|
||
tenant_id=str(database["tenant_id"]),
|
||
project_id=project_id,
|
||
table_code=table_code,
|
||
record_id=record_id,
|
||
operation="import",
|
||
before_data=None,
|
||
after_data={"id": record_id, **_json_safe(payload)},
|
||
actor=actor,
|
||
)
|
||
except MySQLError:
|
||
failed += 1
|
||
if len(failed_errors) < MAX_REPORTED_ERRORS:
|
||
failed_errors.append({"row": int(candidate["row"]), "field": "", "kind": "failed", "message": "数据库约束校验失败,已跳过"})
|
||
await conn.commit()
|
||
return {
|
||
"ok": True,
|
||
"file_name": file_name,
|
||
"imported_count": imported,
|
||
"invalid_count": int(preview["invalid_rows"]),
|
||
"duplicate_count": int(preview["duplicate_rows"]),
|
||
"failed_count": failed,
|
||
"skipped_count": int(preview["skipped_rows"]) + failed,
|
||
"errors": failed_errors,
|
||
"table_code": table.code,
|
||
}
|
||
|
||
|
||
def _csv_value(value: Any) -> str:
|
||
value = _json_safe(value)
|
||
if value is None:
|
||
return ""
|
||
if isinstance(value, bool):
|
||
return "true" if value else "false"
|
||
if isinstance(value, (dict, list)):
|
||
return json.dumps(value, ensure_ascii=False, separators=(",", ":"))
|
||
return str(value)
|
||
|
||
|
||
async def export_csv_records(
|
||
project_id: str,
|
||
table_code: str,
|
||
search: str | None = None,
|
||
) -> tuple[bytes, str, int]:
|
||
table = await _table_or_404(project_id, table_code)
|
||
database = await _database_or_404(project_id)
|
||
columns = _selected_columns(table)
|
||
headers = ["记录 ID", *(field.label for field in table.fields), "创建时间", "更新时间"]
|
||
where = ["`deleted_at` IS NULL"]
|
||
params: list[Any] = []
|
||
searchable = [field for field in table.fields if field.searchable]
|
||
if search and searchable:
|
||
where.append("(" + " OR ".join(f"CAST({_q(field.code)} AS CHAR) LIKE %s" for field in searchable) + ")")
|
||
params.extend(f"%{search.strip()}%" for _ in searchable)
|
||
where_clause = " AND ".join(where)
|
||
async with get_data_conn(str(database["database_name"])) as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
f"SELECT COUNT(*) AS count FROM {_q(table.code)} WHERE {where_clause}",
|
||
params,
|
||
)
|
||
total = int((await cur.fetchone())["count"])
|
||
if total > MAX_EXPORT_ROWS:
|
||
raise HTTPException(400, f"当前数据量为 {total:,} 条,单次最多导出 {MAX_EXPORT_ROWS:,} 条,请先搜索筛选")
|
||
await cur.execute(
|
||
f"SELECT {', '.join(_q(column) for column in columns)} "
|
||
f"FROM {_q(table.code)} WHERE {where_clause} ORDER BY `created_at` ASC",
|
||
params,
|
||
)
|
||
rows = [dict(row) for row in await cur.fetchall()]
|
||
output = io.StringIO(newline="")
|
||
writer = csv.writer(output, lineterminator="\n")
|
||
writer.writerow(headers)
|
||
for row in rows:
|
||
writer.writerow([_csv_value(row.get(column)) for column in columns])
|
||
content = ("\ufeff" + output.getvalue()).encode("utf-8")
|
||
timestamp = datetime.now(timezone.utc).strftime("%Y%m%d_%H%M%S")
|
||
return content, f"{table.code}_{timestamp}.csv", total
|
||
|
||
|
||
def validate_console_sql(raw_sql: str, database_name: str) -> tuple[str, str]:
|
||
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_SQL:
|
||
raise ValueError("控制台仅支持受控查询;建表或修改结构请使用页面中的表结构功能")
|
||
if _FORBIDDEN_SQL.search(statement):
|
||
raise ValueError("SQL 包含控制台不允许执行的结构、文件或权限操作")
|
||
if keyword in {"insert", "update", "delete"} and not settings.data_sql_console_write_enabled:
|
||
raise ValueError("SQL 控制台当前为只读模式;数据修改请使用数据表页面,以保留完整变更记录")
|
||
if keyword == "delete":
|
||
raise ValueError("SQL 控制台禁止物理删除;请使用数据表页面执行可审计的软删除")
|
||
if keyword == "update":
|
||
if not re.search(r"\bwhere\b", statement, re.IGNORECASE):
|
||
raise ValueError("UPDATE 必须包含明确的 WHERE 条件")
|
||
if re.search(r"\bwhere\s+(?:1\s*=\s*1|true)\b", statement, re.IGNORECASE):
|
||
raise ValueError("UPDATE 禁止使用恒真 WHERE 条件")
|
||
if keyword == "with" and re.search(
|
||
r"\b(?:insert|update|delete)\b",
|
||
statement,
|
||
re.IGNORECASE,
|
||
):
|
||
raise ValueError("WITH 在 SQL 控制台中仅允许只读查询")
|
||
if re.search(
|
||
r"\b(?:sleep|benchmark|get_lock|release_lock|load_file)\s*\(",
|
||
statement,
|
||
re.IGNORECASE,
|
||
):
|
||
raise ValueError("SQL 包含控制台禁止调用的高风险函数")
|
||
if re.search(r"\bfor\s+update\b|\block\s+in\s+share\s+mode\b", statement, re.IGNORECASE):
|
||
raise ValueError("SQL 控制台禁止执行可能长期持有数据锁的查询")
|
||
if keyword == "select" and re.search(r"\binto\b", statement, re.IGNORECASE):
|
||
raise ValueError("SQL 控制台禁止 SELECT INTO")
|
||
lowered = statement.lower()
|
||
if keyword == "show" and not re.match(
|
||
r"^\s*show\s+(?:full\s+)?(?:tables|columns|fields|index|indexes|keys|create\s+table)\b",
|
||
statement,
|
||
re.IGNORECASE,
|
||
):
|
||
raise ValueError("SHOW 仅允许查看当前数据库的表、字段、索引或建表语句")
|
||
qualified_objects = re.findall(
|
||
r"\b(?:from|join|update|into)\s+`?([a-z][a-z0-9_]*)`?\s*\.",
|
||
lowered,
|
||
re.IGNORECASE,
|
||
)
|
||
if keyword in {"describe", "desc"}:
|
||
described = re.match(
|
||
r"^\s*(?:describe|desc)\s+`?([a-z][a-z0-9_]*)`?\s*\.",
|
||
lowered,
|
||
re.IGNORECASE,
|
||
)
|
||
if described:
|
||
qualified_objects.append(described.group(1))
|
||
for qualifier in qualified_objects:
|
||
if qualifier.lower() != database_name.lower():
|
||
raise ValueError("控制台禁止访问其他业务数据库")
|
||
if keyword == "show":
|
||
show_tables = re.match(r"^\s*show\s+(?:full\s+)?tables\b", lowered)
|
||
show_scopes = re.findall(r"\b(?:from|in)\s+`?([a-z][a-z0-9_]*)`?", lowered)
|
||
explicit_scope = (
|
||
show_scopes[-1]
|
||
if show_tables and show_scopes
|
||
else show_scopes[-1] if len(show_scopes) >= 2 else None
|
||
)
|
||
if explicit_scope and explicit_scope.lower() != database_name.lower():
|
||
raise ValueError("控制台禁止查看其他业务数据库")
|
||
protected = {
|
||
"information_schema", "mysql", "performance_schema", "sys",
|
||
control_database_name().lower(),
|
||
}
|
||
for name in protected:
|
||
if re.search(rf"(?<![a-z0-9_])`?{re.escape(name)}`?\s*\.", lowered):
|
||
raise ValueError("控制台禁止访问数据中心控制库或系统数据库")
|
||
for other in re.findall(r"\b(dc_[a-z0-9_]+)\s*\.", lowered):
|
||
if other != database_name.lower():
|
||
raise ValueError("控制台禁止访问其他业务数据库")
|
||
return statement, keyword
|
||
|
||
|
||
async def execute_project_sql(
|
||
project_id: str,
|
||
raw_sql: str,
|
||
actor: str = "unknown",
|
||
source_ip: str | None = None,
|
||
) -> dict[str, Any]:
|
||
database = await _database_or_404(project_id)
|
||
database_name = str(database["database_name"])
|
||
raw_hash = hashlib.sha256(raw_sql.strip().encode("utf-8")).hexdigest()
|
||
try:
|
||
statement, keyword = validate_console_sql(raw_sql, database_name)
|
||
except ValueError as exc:
|
||
await _record_admin_action_safely(
|
||
actor=actor,
|
||
action_name="sql_console",
|
||
resource_type="database",
|
||
resource_id=project_id,
|
||
outcome="rejected",
|
||
statement_hash=raw_hash,
|
||
source_ip=source_ip,
|
||
details={"reason": str(exc)[:300]},
|
||
)
|
||
raise
|
||
started_at = time.perf_counter()
|
||
async with get_data_conn(database_name) as conn:
|
||
try:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
"SELECT @@SESSION.sql_mode AS sql_mode, "
|
||
"@@SESSION.MAX_EXECUTION_TIME AS max_execution_time"
|
||
)
|
||
session_state = await cur.fetchone()
|
||
try:
|
||
# Existing frontend-generated SQL quotes identifiers with
|
||
# ". ANSI_QUOTES preserves that API contract on MySQL.
|
||
await cur.execute(
|
||
"SET SESSION sql_mode=CONCAT_WS(',', @@sql_mode, 'ANSI_QUOTES')"
|
||
)
|
||
await cur.execute(
|
||
"SET SESSION MAX_EXECUTION_TIME=%s",
|
||
(STATEMENT_TIMEOUT_MS,),
|
||
)
|
||
await cur.execute(statement)
|
||
columns: list[str] = []
|
||
rows: list[dict[str, Any]] = []
|
||
truncated = False
|
||
if cur.description:
|
||
columns = [str(item[0]) for item in cur.description]
|
||
fetched = await cur.fetchmany(MAX_RESULT_ROWS + 1)
|
||
truncated = len(fetched) > MAX_RESULT_ROWS
|
||
rows = [
|
||
{key: _json_safe(value) for key, value in dict(row).items()}
|
||
for row in fetched[:MAX_RESULT_ROWS]
|
||
]
|
||
affected_rows = max(int(cur.rowcount or 0), 0)
|
||
command = keyword.upper()
|
||
finally:
|
||
await cur.execute(
|
||
"SET SESSION sql_mode=%s",
|
||
(session_state["sql_mode"],),
|
||
)
|
||
await cur.execute(
|
||
"SET SESSION MAX_EXECUTION_TIME=%s",
|
||
(session_state["max_execution_time"],),
|
||
)
|
||
await _write_admin_action_audit(
|
||
cur,
|
||
actor=actor,
|
||
action_name="sql_console",
|
||
resource_type="database",
|
||
resource_id=project_id,
|
||
outcome="success",
|
||
statement_hash=raw_hash,
|
||
affected_rows=affected_rows,
|
||
source_ip=source_ip,
|
||
details={
|
||
"statement_type": keyword,
|
||
"returned_rows": len(rows),
|
||
"truncated": truncated,
|
||
},
|
||
)
|
||
await conn.commit()
|
||
except MySQLError as exc:
|
||
await conn.rollback()
|
||
try:
|
||
async with conn.cursor() as cur:
|
||
await _write_admin_action_audit(
|
||
cur,
|
||
actor=actor,
|
||
action_name="sql_console",
|
||
resource_type="database",
|
||
resource_id=project_id,
|
||
outcome="failed",
|
||
statement_hash=raw_hash,
|
||
source_ip=source_ip,
|
||
details={"error": str(exc)[:300]},
|
||
)
|
||
await conn.commit()
|
||
except Exception:
|
||
await conn.rollback()
|
||
raise ValueError(f"MySQL 执行失败:{exc}") from exc
|
||
return {
|
||
"database": {
|
||
"project_id": project_id,
|
||
"display_name": database["display_name"],
|
||
"database_name": database_name,
|
||
"schema_name": database_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": round((time.perf_counter() - started_at) * 1000, 2),
|
||
}
|