Files
Cloud-Tour-to-Libo/app/data_platform/interface_service.py
T
xuelong 3dd5731751 feat: streamline platform and secure data access
- move the relational data center to MySQL and a standalone workbench\n- add Interface Center API credentials, policies, logs, and DBeaver SSH guidance\n- harden authentication and deployment while retiring unused management surfaces
2026-08-25 02:06:28 -07:00

705 lines
27 KiB
Python

"""API client, credential and least-privilege policy management."""
from __future__ import annotations
from datetime import datetime, timedelta, timezone
import hashlib
import hmac
import json
import secrets
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
from app.data_platform.mysql_service import (
ensure_platform_registry,
list_project_table_entries,
)
VALID_ACTIONS = {"metadata", "read", "create", "update", "delete"}
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 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_auth_method": settings.data_mysql_ssh_auth_method,
"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,
}
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()