Files

2338 lines
92 KiB
Python
Raw Permalink 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.
"""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),
}