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

1132 lines
44 KiB
Python
Raw Blame History

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