"""API client, credential and least-privilege policy management.""" from __future__ import annotations import base64 import binascii from datetime import datetime, timedelta, timezone import hashlib import hmac import json import os from pathlib import Path import re import secrets import tempfile from typing import Any from urllib.parse import urlsplit import uuid from fastapi import HTTPException from app.config import settings from app.data_platform.mysql_db import ( data_pool_available, get_data_conn, get_provisioner_conn, ) from app.data_platform.mysql_service import ( ensure_platform_registry, list_project_table_entries, ) VALID_ACTIONS = {"metadata", "read", "create", "update", "delete"} _SSH_PUBLIC_KEY_RE = re.compile( r"^(ssh-rsa|ssh-ed25519)\s+([A-Za-z0-9+/]+={0,2})(?:\s+[^\r\n]{1,120})?$" ) _MYSQL_ACCOUNT_RE = re.compile(r"^dba_[a-f0-9]{12}$") def _ssh_field(blob: bytes, offset: int) -> tuple[bytes, int]: if offset + 4 > len(blob): raise ValueError("SSH 公钥内容不完整") length = int.from_bytes(blob[offset : offset + 4], "big") start = offset + 4 end = start + length if length > 4096 or end > len(blob): raise ValueError("SSH 公钥字段长度不正确") return blob[start:end], end def _mysql_account_sql(username: str, *, parameterized: bool = False) -> str: if not _MYSQL_ACCOUNT_RE.fullmatch(username): raise ValueError("DBeaver管理账号格式异常") # PyMySQL/aiomysql use percent-style parameter interpolation. A literal # host wildcard must therefore be doubled only in statements with %s. host = "%%" if parameterized else "%" return f"'{username}'@'{host}'" async def _ready() -> None: if not data_pool_available(): raise HTTPException( 503, "接口中心依赖的数据中心 MySQL 服务未连接", ) await ensure_platform_registry() def _decode_json(value: Any, fallback: Any) -> Any: if value is None: return fallback if isinstance(value, (list, dict)): return value if isinstance(value, (bytes, bytearray)): value = value.decode("utf-8") try: return json.loads(str(value)) except (TypeError, ValueError): return fallback def _safe(value: Any) -> Any: if isinstance(value, datetime): return value.isoformat() if isinstance(value, dict): return {key: _safe(item) for key, item in value.items()} if isinstance(value, list): return [_safe(item) for item in value] return value def _row(row: dict[str, Any] | None) -> dict[str, Any] | None: if row is None: return None result = {key: _safe(value) for key, value in dict(row).items()} for key in ( "actions_json", "readable_fields_json", "writable_fields_json", "row_filter_json", ): if key in result: result[key.removesuffix("_json")] = _decode_json(result.pop(key), [] if key != "row_filter_json" else {}) return result def _api_key_hash(api_key: str) -> str: pepper = settings.interface_api_secret or settings.auth_secret return hmac.new( pepper.encode("utf-8"), api_key.encode("utf-8"), hashlib.sha256, ).hexdigest() def _parse_expiry(value: Any) -> datetime: now = datetime.now(timezone.utc) if value in (None, ""): parsed = now + timedelta(days=settings.interface_api_default_expiry_days) elif isinstance(value, datetime): parsed = value else: try: parsed = datetime.fromisoformat(str(value).replace("Z", "+00:00")) except ValueError as exc: raise ValueError("凭证过期时间格式不正确") from exc if parsed.tzinfo is None: parsed_utc = parsed.replace(tzinfo=timezone.utc) else: parsed_utc = parsed.astimezone(timezone.utc) if parsed_utc <= now: raise ValueError("凭证过期时间必须晚于当前时间") maximum = now + timedelta(days=settings.interface_api_max_expiry_days) if parsed_utc > maximum: raise ValueError( f"接口密钥最长只能签发 {settings.interface_api_max_expiry_days} 天" ) return parsed_utc.replace(tzinfo=None) async def interface_summary() -> dict[str, Any]: await _ready() async with get_data_conn() as conn: async with conn.cursor() as cur: await cur.execute( """ SELECT (SELECT COUNT(*) FROM api_clients WHERE status='active') AS active_clients, (SELECT COUNT(*) FROM api_credentials WHERE revoked_at IS NULL AND (expires_at IS NULL OR expires_at > CURRENT_TIMESTAMP(6))) AS active_credentials, (SELECT COUNT(*) FROM api_policies WHERE status='active') AS active_policies, (SELECT COUNT(*) FROM dbeaver_access_grants WHERE status='active') AS active_dbeaver_grants, (SELECT COUNT(*) FROM api_call_logs WHERE created_at >= CURRENT_TIMESTAMP(6) - INTERVAL 24 HOUR) AS calls_24h, (SELECT COUNT(*) FROM api_call_logs WHERE created_at >= CURRENT_TIMESTAMP(6) - INTERVAL 24 HOUR AND status_code >= 400) AS errors_24h """ ) summary = dict(await cur.fetchone()) mysql_url = urlsplit(settings.data_mysql_url) database_port = ( settings.data_mysql_admin_port or settings.data_mysql_public_port or mysql_url.port or 3306 ) ssh_host = ( settings.data_mysql_ssh_host.strip() or settings.data_mysql_public_host.strip() ) host_bind = settings.mysql_host_bind.strip() or "127.0.0.1" loopback_hosts = {"127.0.0.1", "localhost", "::1"} publicly_bound = host_bind not in loopback_hosts return { **{key: int(value or 0) for key, value in summary.items()}, "data_engine": "MySQL", "public_base_path": "/v1/openapi/data", "direct_database_access": "DBeaver / MySQL account (separate authorization)", "direct_access_enabled": settings.data_mysql_direct_access_enabled, # Legacy aliases retained for existing clients. New clients should # use the explicit SSH and database endpoint fields below. "direct_access_host": ssh_host, "direct_access_port": database_port, "direct_access_transport": settings.data_mysql_direct_transport, "ssh_tunnel_required": settings.data_mysql_ssh_tunnel_required, "ssh_host": ssh_host, "ssh_port": settings.data_mysql_ssh_port, "ssh_username": settings.data_mysql_ssh_username, "ssh_auth_method": settings.data_mysql_ssh_auth_method, "managed_access_enabled": settings.data_mysql_managed_access_enabled, "ssh_key_auto_install": bool( settings.data_mysql_ssh_authorized_keys_file.strip() ), "account_provisioning_ready": bool( settings.data_mysql_provisioner_user.strip() and settings.data_mysql_provisioner_password ), "ssh_host_key_fingerprint": _ssh_host_key_fingerprint(), "database_host": settings.data_mysql_admin_host.strip() or "127.0.0.1", "database_port": database_port, "database_account_policy": settings.data_mysql_admin_account_policy, "mysql_host_bind": host_bind, "mysql_publicly_bound": publicly_bound, "database_audit_enabled": settings.data_mysql_audit_enabled, "application_audit_enabled": True, "sql_console_write_enabled": settings.data_sql_console_write_enabled, "backup_enabled": settings.data_backup_enabled, "backup_encryption_required": settings.data_backup_encryption_required, "backup_retention_days": settings.data_backup_retention_days, } def _normalize_ssh_public_key(value: Any) -> tuple[str, str, str]: candidate = " ".join(str(value or "").strip().split()) match = _SSH_PUBLIC_KEY_RE.fullmatch(candidate) if not match: raise ValueError("SSH 公钥格式不正确,仅支持 RSA 或 ED25519 公钥") key_type, encoded = match.group(1), match.group(2) try: decoded = base64.b64decode(encoded, validate=True) except (binascii.Error, ValueError) as exc: raise ValueError("SSH 公钥内容不是有效的 Base64") from exc if len(decoded) < 48 or len(decoded) > 4096: raise ValueError("SSH 公钥长度不正确") embedded_type, offset = _ssh_field(decoded, 0) if embedded_type.decode("ascii", errors="ignore") != key_type: raise ValueError("SSH 公钥类型与内容不一致") if key_type == "ssh-ed25519": public_bytes, offset = _ssh_field(decoded, offset) if len(public_bytes) != 32: raise ValueError("ED25519 公钥长度不正确") else: exponent, offset = _ssh_field(decoded, offset) modulus, offset = _ssh_field(decoded, offset) exponent_value = int.from_bytes(exponent, "big") modulus_bits = int.from_bytes(modulus, "big").bit_length() if exponent_value < 3 or exponent_value % 2 == 0: raise ValueError("RSA 公钥指数不正确") if not 2048 <= modulus_bits <= 8192: raise ValueError("RSA 公钥必须为 2048–8192 位") if offset != len(decoded): raise ValueError("SSH 公钥包含多余数据") fingerprint = base64.b64encode(hashlib.sha256(decoded).digest()).decode().rstrip("=") return key_type, f"{key_type} {encoded}", f"SHA256:{fingerprint}" def _ssh_host_key_fingerprint() -> str: raw_path = settings.data_mysql_ssh_host_public_key_file.strip() if not raw_path: return "" path = Path(raw_path) try: if not path.is_absolute() or path.is_symlink() or path.stat().st_size > 8192: return "" _key_type, _public_key, fingerprint = _normalize_ssh_public_key( path.read_text(encoding="utf-8") ) return fingerprint except (OSError, UnicodeError, ValueError): return "" def _authorized_keys_path() -> Path: if not settings.data_mysql_managed_access_enabled: raise HTTPException(503, "DBeaver 管理接入尚未启用") raw_path = settings.data_mysql_ssh_authorized_keys_file.strip() if not raw_path: raise HTTPException(503, "服务器尚未配置 DBeaver SSH 公钥文件") path = Path(raw_path) if not path.is_absolute(): raise HTTPException(503, "DBeaver SSH 公钥文件必须使用绝对路径") return path def _write_authorized_keys(rows: list[dict[str, Any]]) -> None: target = _authorized_keys_path() target.parent.mkdir(parents=True, exist_ok=True) if target.is_symlink(): raise RuntimeError("拒绝写入符号链接形式的 SSH 公钥文件") lines = [ "# Managed by the Interface Center. Manual edits will be replaced." ] for item in rows: username = str(item["mysql_username"]) if not _MYSQL_ACCOUNT_RE.fullmatch(username): continue _key_type, public_key, _fingerprint = _normalize_ssh_public_key( item["ssh_public_key"] ) lines.append( 'restrict,port-forwarding,permitopen="127.0.0.1:3307" ' f"{public_key} nianxx:{item['id']}:{username}" ) content = "\n".join(lines) + "\n" file_descriptor, temporary_name = tempfile.mkstemp( prefix=".authorized_keys.", dir=str(target.parent) ) try: with os.fdopen(file_descriptor, "w", encoding="utf-8") as handle: handle.write(content) handle.flush() os.fsync(handle.fileno()) os.chmod(temporary_name, 0o600) os.replace(temporary_name, target) finally: if os.path.exists(temporary_name): os.unlink(temporary_name) async def _refresh_dbeaver_authorized_keys() -> None: async with get_data_conn() as conn: async with conn.cursor() as cur: await cur.execute( """ SELECT id, mysql_username, ssh_public_key FROM dbeaver_access_grants WHERE status='active' ORDER BY created_at """ ) rows = [dict(item) for item in await cur.fetchall()] _write_authorized_keys(rows) async def list_dbeaver_access_grants() -> list[dict[str, Any]]: await _ready() async with get_data_conn() as conn: async with conn.cursor() as cur: await cur.execute( """ SELECT id, display_name, mysql_username, database_id, database_name, permission_level, ssh_key_type, ssh_key_fingerprint, created_by, status, revoked_at, created_at FROM dbeaver_access_grants ORDER BY created_at DESC """ ) return [_row(dict(item)) or {} for item in await cur.fetchall()] async def _write_dbeaver_audit( cur: Any, *, actor: str, action: str, grant_id: str, outcome: str, details: dict[str, Any], ) -> None: await cur.execute( """ INSERT INTO admin_action_logs ( actor, action_name, resource_type, resource_id, outcome, details_json ) VALUES (%s, %s, 'dbeaver_access', %s, %s, %s) """, (actor, action, grant_id, outcome, json.dumps(details, ensure_ascii=False)), ) def _mysql_access_details() -> dict[str, Any]: mysql_url = urlsplit(settings.data_mysql_url) return { "ssh_host": settings.data_mysql_ssh_host.strip(), "ssh_port": settings.data_mysql_ssh_port, "ssh_username": settings.data_mysql_ssh_username, "database_host": settings.data_mysql_admin_host.strip() or "127.0.0.1", "database_port": ( settings.data_mysql_admin_port or settings.data_mysql_public_port or mysql_url.port or 3306 ), } async def issue_dbeaver_access( body: dict[str, Any], actor: str ) -> dict[str, Any]: await _ready() _authorized_keys_path() display_name = str(body.get("display_name") or "").strip() database_id = str(body.get("database_id") or "").strip() permission = str(body.get("permission") or "read").strip().lower() if not display_name: raise ValueError("请输入数据管理员姓名") if len(display_name) > 100: raise ValueError("数据管理员姓名不能超过 100 个字符") if permission not in {"read", "write"}: raise ValueError("DBeaver 权限只能选择只读或读写") key_type, public_key, fingerprint = _normalize_ssh_public_key( body.get("public_key") ) grant_id = str(uuid.uuid4()) mysql_username = f"dba_{uuid.uuid4().hex[:12]}" mysql_password = secrets.token_urlsafe(24) account = _mysql_account_sql(mysql_username) parameterized_account = _mysql_account_sql(mysql_username, parameterized=True) database_name = "" created_account = False try: # Verify the shared key file before creating a database account. This # avoids issuing a credential that cannot be used by the SSH gateway. await _refresh_dbeaver_authorized_keys() 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 AND status='ready' """, (database_id,), ) database = await cur.fetchone() if not database: raise ValueError("目标数据库不存在或尚未就绪") database_name = str(database["database_name"]) if not re.fullmatch(r"[A-Za-z][A-Za-z0-9_]{0,63}", database_name): raise ValueError("目标数据库编码不合法") await cur.execute( """ SELECT 1 FROM dbeaver_access_grants WHERE ssh_key_fingerprint=%s AND status='active' """, (fingerprint,), ) if await cur.fetchone(): raise ValueError("该 SSH 公钥已用于一个有效的管理账号") # CREATE USER and GRANT are intentionally isolated from the normal # Data Center connection. Only this short-lived broker account has # account-provisioning privileges. async with get_provisioner_conn() as provisioner_conn: async with provisioner_conn.cursor() as cur: await cur.execute( f"CREATE USER {parameterized_account} IDENTIFIED BY %s " "PASSWORD EXPIRE INTERVAL 90 DAY " "FAILED_LOGIN_ATTEMPTS 5 PASSWORD_LOCK_TIME 1", (mysql_password,), ) created_account = True privileges = "SELECT, SHOW VIEW" if permission == "write": privileges = "SELECT, INSERT, UPDATE, DELETE, SHOW VIEW" await cur.execute( f"GRANT {privileges} ON `{database_name}`.* TO {account}" ) await provisioner_conn.commit() async with get_data_conn() as conn: async with conn.cursor() as cur: await cur.execute( """ INSERT INTO dbeaver_access_grants ( id, display_name, mysql_username, database_id, database_name, permission_level, ssh_key_type, ssh_public_key, ssh_key_fingerprint, created_by ) VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s) """, ( grant_id, display_name, mysql_username, database_id, database_name, permission, key_type, public_key, fingerprint, actor, ), ) await _write_dbeaver_audit( cur, actor=actor, action="issue_dbeaver_access", grant_id=grant_id, outcome="success", details={ "database_id": database_id, "database_name": database_name, "mysql_username": mysql_username, "permission": permission, "ssh_key_fingerprint": fingerprint, }, ) await conn.commit() await _refresh_dbeaver_authorized_keys() except Exception as exc: if created_account: try: async with get_provisioner_conn() as cleanup_conn: async with cleanup_conn.cursor() as cleanup_cur: await cleanup_cur.execute(f"DROP USER IF EXISTS {account}") await cleanup_conn.commit() async with get_data_conn() as metadata_conn: async with metadata_conn.cursor() as cleanup_cur: await cleanup_cur.execute( "DELETE FROM dbeaver_access_grants WHERE id=%s", (grant_id,), ) await metadata_conn.commit() await _refresh_dbeaver_authorized_keys() except Exception: pass if isinstance(exc, ValueError): raise if isinstance(exc, HTTPException): raise raise HTTPException( 503, "DBeaver账号签发失败,请确认MySQL账号签发权限和SSH网关已完成部署", ) from exc return { "id": grant_id, "display_name": display_name, "mysql_username": mysql_username, "mysql_password": mysql_password, "password_shown_once": True, "database_id": database_id, "database_name": database_name, "permission_level": permission, "ssh_key_type": key_type, "ssh_key_fingerprint": fingerprint, "status": "active", **_mysql_access_details(), } async def revoke_dbeaver_access(grant_id: str, actor: str) -> dict[str, bool]: await _ready() _authorized_keys_path() async with get_data_conn() as conn: async with conn.cursor() as cur: await cur.execute( """ SELECT mysql_username, database_id, database_name, status FROM dbeaver_access_grants WHERE id=%s """, (grant_id,), ) grant = await cur.fetchone() if not grant: raise ValueError("DBeaver管理账号不存在") username = str(grant["mysql_username"]) if not _MYSQL_ACCOUNT_RE.fullmatch(username): raise ValueError("DBeaver管理账号格式异常,已拒绝操作") if str(grant["status"]) == "active": try: async with get_provisioner_conn() as provisioner_conn: async with provisioner_conn.cursor() as cur: await cur.execute( f"DROP USER IF EXISTS {_mysql_account_sql(username)}" ) await provisioner_conn.commit() except Exception as exc: raise HTTPException(503, "MySQL管理账号撤销失败,请检查账号签发服务") from exc async with get_data_conn() as conn: async with conn.cursor() as cur: await cur.execute( """ UPDATE dbeaver_access_grants SET status='revoked', revoked_at=CURRENT_TIMESTAMP(6) WHERE id=%s """, (grant_id,), ) await _write_dbeaver_audit( cur, actor=actor, action="revoke_dbeaver_access", grant_id=grant_id, outcome="success", details={ "database_id": grant["database_id"], "database_name": grant["database_name"], "mysql_username": username, }, ) await conn.commit() try: await _refresh_dbeaver_authorized_keys() except Exception as exc: raise HTTPException(503, "账号已撤销,但SSH公钥同步失败,请重试撤销操作") from exc return {"ok": True} async def interface_catalog() -> list[dict[str, Any]]: await _ready() async with get_data_conn() as conn: async with conn.cursor() as cur: await cur.execute( """ SELECT project_id, display_name, database_name FROM project_databases WHERE status='ready' ORDER BY display_name, project_id """ ) databases = [dict(item) for item in await cur.fetchall()] result = [] for database in databases: entries = await list_project_table_entries(str(database["project_id"])) result.append( { **_safe(database), "tables": [ { "code": definition.code, "label": definition.label, "group": definition.group, "fields": [field.as_dict() for field in definition.fields], } for definition, _origin, _source in entries ], } ) return result async def list_api_clients() -> list[dict[str, Any]]: await _ready() async with get_data_conn() as conn: async with conn.cursor() as cur: await cur.execute( """ SELECT c.*, COUNT(DISTINCT CASE WHEN k.revoked_at IS NULL THEN k.id END) AS credential_count, COUNT(DISTINCT CASE WHEN p.status='active' THEN p.id END) AS policy_count FROM api_clients c LEFT JOIN api_credentials k ON k.client_id=c.id LEFT JOIN api_policies p ON p.client_id=c.id GROUP BY c.id ORDER BY c.created_at DESC """ ) return [_row(dict(item)) or {} for item in await cur.fetchall()] async def create_api_client(body: dict[str, Any], actor: str) -> dict[str, Any]: await _ready() name = str(body.get("name") or "").strip() if not name: raise ValueError("请输入应用或设备名称") if len(name) > 100: raise ValueError("应用或设备名称不能超过 100 个字符") client_id = str(uuid.uuid4()) async with get_data_conn() as conn: async with conn.cursor() as cur: await cur.execute( """ INSERT INTO api_clients (id, name, description, owner, status) VALUES (%s, %s, %s, %s, 'active') """, (client_id, name, str(body.get("description") or "")[:500], actor), ) await cur.execute("SELECT * FROM api_clients WHERE id=%s", (client_id,)) result = _row(await cur.fetchone()) or {} await conn.commit() return result async def update_api_client(client_id: str, body: dict[str, Any]) -> dict[str, Any]: await _ready() updates: list[str] = [] params: list[Any] = [] if "name" in body: name = str(body.get("name") or "").strip() if not name: raise ValueError("请输入应用或设备名称") updates.append("name=%s") params.append(name[:100]) if "description" in body: updates.append("description=%s") params.append(str(body.get("description") or "")[:500]) if "status" in body: status = str(body.get("status") or "") if status not in {"active", "disabled"}: raise ValueError("客户端状态不正确") updates.append("status=%s") params.append(status) if not updates: raise ValueError("没有可修改内容") async with get_data_conn() as conn: async with conn.cursor() as cur: await cur.execute( f"UPDATE api_clients SET {', '.join(updates)} WHERE id=%s", [*params, client_id], ) if cur.rowcount != 1: raise ValueError("接口客户端不存在") if body.get("status") == "disabled": await cur.execute( "UPDATE api_credentials SET revoked_at=CURRENT_TIMESTAMP(6) WHERE client_id=%s AND revoked_at IS NULL", (client_id,), ) await cur.execute( "UPDATE api_policies SET status='disabled' WHERE client_id=%s", (client_id,), ) await cur.execute("SELECT * FROM api_clients WHERE id=%s", (client_id,)) result = _row(await cur.fetchone()) or {} await conn.commit() return result async def delete_api_client(client_id: str) -> dict[str, bool]: await update_api_client(client_id, {"status": "disabled"}) async with get_data_conn() as conn: async with conn.cursor() as cur: await cur.execute( "UPDATE api_policies SET status='disabled' WHERE client_id=%s", (client_id,), ) await conn.commit() return {"ok": True} async def list_api_credentials(client_id: str | None = None) -> list[dict[str, Any]]: await _ready() where = "WHERE k.client_id=%s" if client_id else "" params = (client_id,) if client_id else () async with get_data_conn() as conn: async with conn.cursor() as cur: await cur.execute( f""" SELECT k.*, c.name AS client_name FROM api_credentials k JOIN api_clients c ON c.id=k.client_id {where} ORDER BY k.created_at DESC """, params, ) result = [_row(dict(item)) or {} for item in await cur.fetchall()] for item in result: item.pop("key_hash", None) return result async def issue_api_credential(client_id: str, body: dict[str, Any]) -> dict[str, Any]: await _ready() name = str(body.get("name") or "默认密钥").strip() or "默认密钥" expires_at = _parse_expiry(body.get("expires_at")) credential_id = str(uuid.uuid4()) prefix = f"dc_live_{secrets.token_hex(4)}" api_key = f"{prefix}.{secrets.token_urlsafe(32)}" async with get_data_conn() as conn: async with conn.cursor() as cur: await cur.execute( "SELECT 1 FROM api_clients WHERE id=%s AND status='active'", (client_id,), ) if not await cur.fetchone(): raise ValueError("接口客户端不存在或未启用") await cur.execute( """ INSERT INTO api_credentials ( id, client_id, name, key_prefix, key_hash, expires_at ) VALUES (%s, %s, %s, %s, %s, %s) """, (credential_id, client_id, name[:100], prefix, _api_key_hash(api_key), expires_at), ) await cur.execute("SELECT * FROM api_credentials WHERE id=%s", (credential_id,)) result = _row(await cur.fetchone()) or {} await conn.commit() result.pop("key_hash", None) result["api_key"] = api_key result["shown_once"] = True return result async def revoke_api_credential(credential_id: str) -> dict[str, bool]: await _ready() async with get_data_conn() as conn: async with conn.cursor() as cur: await cur.execute( """ UPDATE api_credentials SET revoked_at=CURRENT_TIMESTAMP(6) WHERE id=%s AND revoked_at IS NULL """, (credential_id,), ) if cur.rowcount != 1: raise ValueError("接口密钥不存在或已撤销") await conn.commit() return {"ok": True} async def list_api_policies(client_id: str | None = None) -> list[dict[str, Any]]: await _ready() where = "WHERE p.client_id=%s" if client_id else "" params = (client_id,) if client_id else () async with get_data_conn() as conn: async with conn.cursor() as cur: await cur.execute( f""" SELECT p.*, c.name AS client_name, d.display_name AS database_name FROM api_policies p JOIN api_clients c ON c.id=p.client_id LEFT JOIN project_databases d ON d.project_id=p.database_id {where} ORDER BY p.created_at DESC """, params, ) return [_row(dict(item)) or {} for item in await cur.fetchall()] async def _validate_policy(body: dict[str, Any]) -> dict[str, Any]: client_id = str(body.get("client_id") or "").strip() database_id = str(body.get("database_id") or "").strip() table_code = str(body.get("table_code") or "*").strip() or "*" actions = sorted({str(value) for value in body.get("actions") or []}) invalid = sorted(set(actions) - VALID_ACTIONS) if not client_id or not database_id: raise ValueError("请选择接口客户端和数据库") if not actions: raise ValueError("请至少授予一个接口动作") if invalid: raise ValueError(f"不支持的接口动作:{', '.join(invalid)}") readable = [str(value) for value in body.get("readable_fields") or ["*"]] writable = [str(value) for value in body.get("writable_fields") or ["*"]] row_filter = body.get("row_filter") or {} if not isinstance(row_filter, dict): raise ValueError("行级数据范围必须是 JSON 对象") status = str(body.get("status") or "active") if status not in {"active", "disabled"}: raise ValueError("权限策略状态不正确") for field_code, expected in row_filter.items(): values = expected if isinstance(expected, list) else [expected] if any(isinstance(value, (dict, list)) for value in values): raise ValueError(f"行级范围 {field_code} 仅支持标量或标量数组") async with get_data_conn() as conn: async with conn.cursor() as cur: await cur.execute("SELECT 1 FROM api_clients WHERE id=%s", (client_id,)) if not await cur.fetchone(): raise ValueError("接口客户端不存在") await cur.execute( "SELECT 1 FROM project_databases WHERE project_id=%s AND status='ready'", (database_id,), ) if not await cur.fetchone(): raise ValueError("数据中心数据库不存在") if table_code == "*": if readable != ["*"] or writable != ["*"] or row_filter: raise ValueError("授权全部表时,字段范围必须为 * 且不能设置跨表行过滤") else: entry = next( ( item for item in await list_project_table_entries(database_id) if item[0].code == table_code ), None, ) if not entry: raise ValueError("授权的数据表不存在") field_codes = {field.code for field in entry[0].fields} for label, fields in (("可读字段", readable), ("可写字段", writable)): unknown = sorted(set(fields) - field_codes - {"*"}) if unknown: raise ValueError(f"{label}不存在:{', '.join(unknown)}") unknown_scope = sorted(set(row_filter) - field_codes) if unknown_scope: raise ValueError(f"行级范围字段不存在:{', '.join(unknown_scope)}") return { "client_id": client_id, "database_id": database_id, "table_code": table_code, "actions": actions, "readable_fields": readable, "writable_fields": writable, "row_filter": row_filter, "status": status, } async def create_api_policy(body: dict[str, Any]) -> dict[str, Any]: await _ready() policy = await _validate_policy(body) policy_id = str(uuid.uuid4()) async with get_data_conn() as conn: async with conn.cursor() as cur: try: await cur.execute( """ INSERT INTO api_policies ( id, client_id, database_id, table_code, actions_json, readable_fields_json, writable_fields_json, row_filter_json, status ) VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s) """, ( policy_id, policy["client_id"], policy["database_id"], policy["table_code"], json.dumps(policy["actions"], ensure_ascii=False), json.dumps(policy["readable_fields"], ensure_ascii=False), json.dumps(policy["writable_fields"], ensure_ascii=False), json.dumps(policy["row_filter"], ensure_ascii=False), policy["status"], ), ) except Exception as exc: raise ValueError("同一客户端、数据库和表只能创建一条权限策略") from exc await cur.execute("SELECT * FROM api_policies WHERE id=%s", (policy_id,)) result = _row(await cur.fetchone()) or {} await conn.commit() return result async def update_api_policy(policy_id: str, body: dict[str, Any]) -> dict[str, Any]: await _ready() async with get_data_conn() as conn: async with conn.cursor() as cur: await cur.execute("SELECT * FROM api_policies WHERE id=%s", (policy_id,)) current = _row(await cur.fetchone()) if not current: raise ValueError("权限策略不存在") merged = { "client_id": body.get("client_id", current["client_id"]), "database_id": body.get("database_id", current["database_id"]), "table_code": body.get("table_code", current["table_code"]), "actions": body.get("actions", current["actions"]), "readable_fields": body.get("readable_fields", current["readable_fields"]), "writable_fields": body.get("writable_fields", current["writable_fields"]), "row_filter": body.get("row_filter", current["row_filter"]), "status": body.get("status", current["status"]), } policy = await _validate_policy(merged) async with get_data_conn() as conn: async with conn.cursor() as cur: await cur.execute( """ UPDATE api_policies SET client_id=%s, database_id=%s, table_code=%s, actions_json=%s, readable_fields_json=%s, writable_fields_json=%s, row_filter_json=%s, status=%s WHERE id=%s """, ( policy["client_id"], policy["database_id"], policy["table_code"], json.dumps(policy["actions"], ensure_ascii=False), json.dumps(policy["readable_fields"], ensure_ascii=False), json.dumps(policy["writable_fields"], ensure_ascii=False), json.dumps(policy["row_filter"], ensure_ascii=False), policy["status"], policy_id, ), ) await cur.execute("SELECT * FROM api_policies WHERE id=%s", (policy_id,)) result = _row(await cur.fetchone()) or {} await conn.commit() return result async def delete_api_policy(policy_id: str) -> dict[str, bool]: await _ready() async with get_data_conn() as conn: async with conn.cursor() as cur: await cur.execute("DELETE FROM api_policies WHERE id=%s", (policy_id,)) if cur.rowcount != 1: raise ValueError("权限策略不存在") await conn.commit() return {"ok": True} async def list_api_call_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 l.*, c.name AS client_name FROM api_call_logs l LEFT JOIN api_clients c ON c.id=l.client_id ORDER BY l.created_at DESC LIMIT %s """, (max(1, min(1000, limit)),), ) return [_row(dict(item)) or {} for item in await cur.fetchall()] async def authenticate_api_key(api_key: str) -> dict[str, Any]: await _ready() if not api_key or "." not in api_key: raise HTTPException(401, "缺少或无效的接口密钥") key_hash = _api_key_hash(api_key) async with get_data_conn() as conn: async with conn.cursor() as cur: await cur.execute( """ SELECT k.id AS credential_id, k.client_id, k.key_prefix, k.expires_at, c.name AS client_name FROM api_credentials k JOIN api_clients c ON c.id=k.client_id WHERE k.key_hash=%s AND k.revoked_at IS NULL AND (k.expires_at IS NULL OR k.expires_at > CURRENT_TIMESTAMP(6)) AND c.status='active' """, (key_hash,), ) identity = _row(await cur.fetchone()) if not identity: raise HTTPException(401, "接口密钥无效、已撤销或已过期") await cur.execute( "UPDATE api_credentials SET last_used_at=CURRENT_TIMESTAMP(6) WHERE id=%s", (identity["credential_id"],), ) await conn.commit() return identity async def resolve_api_policy( client_id: str, database_id: str, table_code: str, action: str, ) -> dict[str, Any]: await _ready() async with get_data_conn() as conn: async with conn.cursor() as cur: await cur.execute( """ SELECT * FROM api_policies WHERE client_id=%s AND database_id=%s AND table_code IN (%s, '*') AND status='active' ORDER BY (table_code=%s) DESC, updated_at DESC """, (client_id, database_id, table_code, table_code), ) policies = [_row(dict(item)) or {} for item in await cur.fetchall()] policy = next((item for item in policies if action in item["actions"]), None) if not policy: raise HTTPException(403, f"当前接口客户端没有 {action} 权限") return policy async def client_catalog(client_id: str) -> list[dict[str, Any]]: policies = [ item for item in await list_api_policies(client_id) if item["status"] == "active" ] catalog = await interface_catalog() result: list[dict[str, Any]] = [] for database in catalog: db_policies = [ item for item in policies if item["database_id"] == database["project_id"] and ({"metadata", "read"} & set(item["actions"])) ] if not db_policies: continue tables = [] for table in database["tables"]: policy = next( (item for item in db_policies if item["table_code"] == table["code"]), next((item for item in db_policies if item["table_code"] == "*"), None), ) if not policy: continue readable = set(policy["readable_fields"]) fields = table["fields"] if "*" in readable else [ field for field in table["fields"] if field["code"] in readable ] tables.append({ **table, "fields": fields, "actions": policy["actions"], }) if tables: result.append({ "project_id": database["project_id"], "display_name": database["display_name"], "tables": tables, }) return result async def write_api_call_log( *, request_id: str, identity: dict[str, Any] | None, method: str, path: str, database_id: str | None, table_code: str | None, action: str | None, status_code: int, duration_ms: float, source_ip: str | None, error_message: str | None = None, ) -> None: if not data_pool_available(): return await ensure_platform_registry() async with get_data_conn() as conn: async with conn.cursor() as cur: await cur.execute( """ INSERT INTO api_call_logs ( request_id, client_id, credential_id, method, path, database_id, table_code, action_name, status_code, duration_ms, source_ip, error_message ) VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s) """, ( request_id, identity.get("client_id") if identity else None, identity.get("credential_id") if identity else None, method, path[:500], database_id, table_code, action, status_code, duration_ms, source_ip, (error_message or "")[:500] or None, ), ) await conn.commit()