1132 lines
44 KiB
Python
1132 lines
44 KiB
Python
"""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()
|