Files
Mutual_Fund/api/routers/nl2sql.py
T
2026-09-14 10:57:48 +08:00

631 lines
22 KiB
Python

"""NL2SQL 查询 HTTP 接口。"""
from __future__ import annotations
import csv
import json
from dataclasses import asdict, replace
from io import StringIO
from uuid import uuid4
from datetime import datetime
from fastapi import APIRouter, Depends, Query, Request
from fastapi.responses import Response, StreamingResponse
from sqlalchemy import text
from sqlalchemy.ext.asyncio import AsyncSession
from api.deps import get_current_user
from config import database
from config.deps import get_db, get_milvus, get_redis
from config.settings import settings
from model.sys_user import SysUser
from nl2sql.contracts import DataQueryRequest
from nl2sql.cache import (
DEFAULT_CACHE_TTL,
build_cache_key,
cache_get,
cache_set,
invalidate_tables,
result_from_cache,
result_to_cache,
)
from nl2sql.audit import write_nl2sql_audit_safely
from nl2sql.diagnostics import diagnostic_registry
from nl2sql.executor import execute_readonly_sql
from nl2sql.explain import build_explain_sql
from nl2sql.history import archive_query_safely
from nl2sql.health import check_nl2sql_health
from nl2sql.query_experience import (
build_clarification,
render_csv,
)
from nl2sql.retrieval import retrieve_metadata
from nl2sql.result import build_chart_config, summarize_result
from nl2sql.runtime import kill_mysql_query, query_runtime_registry
from nl2sql.runtime_config import runtime_config
from nl2sql.session_context import SessionContextStore, build_conversation_context
from nl2sql.limits import QueryLimiter
from nl2sql.metrics import query_metrics
from nl2sql.schema import load_authoritative_schema
from nl2sql.supervisor import SessionBusyError, SessionLock
from nl2sql.sql_security import SqlSecurityError, validate_select_sql
from repositories.nl2sql_permission import Nl2SqlPermissionRepo
from schemas.nl2sql import DataCacheInvalidateReq, DataExplainReq, DataKillReq, DataQueryReq
from service.nl2sql.query_service import QueryServiceError, query as build_query
from service.nl2sql.permission_service import load_query_permission
from tool.llm import llm
from utils.exceptions import ForbiddenError, NotFoundError, ParamError
from utils.request_id import get_request_id, new_request_id
from utils.response import success
from utils.pagination import normalize_pagination
router = APIRouter()
def ensure_query_employee(user: SysUser) -> None:
"""NL2SQL 仅允许已登录员工账号使用。"""
if user.user_type != "EMPLOYEE":
raise ForbiddenError("仅登录员工可以使用数据查询")
def ensure_query_admin(user: SysUser) -> None:
"""只允许系统管理员管理运行中查询。"""
if user.user_type == "ADMIN":
return
if user.user_type == "EMPLOYEE" and user.employee_role in {"ADMIN", "系统管理员"}:
return
raise ForbiddenError("仅管理员可以管理运行中查询")
def build_history_payload(*, query_id, request_data, user, trace_id, status, result, query_result=None):
"""构造查询历史字段,明确排除查询结果行。"""
return {
"query_id": query_id,
"user_id": user.id,
"session_id": request_data.session_id,
"caller_agent": request_data.caller_agent,
"question": request_data.question,
"generated_sql": getattr(result, "sql", None),
"access_tables": getattr(result, "access_tables", set()),
"status": status,
"row_count": getattr(query_result or result, "row_count", 0),
"truncated": getattr(query_result or result, "truncated", False),
"elapsed_ms": getattr(query_result or result, "elapsed_ms", None),
"trace_id": trace_id,
}
def history_payload(history) -> dict:
"""将查询历史模型转换为不含结果行和连接信息的响应。"""
return {
"query_id": history.query_id,
"user_id": history.user_id,
"session_id": history.session_id,
"caller_agent": history.caller_agent,
"question": history.question,
"generated_sql": history.generated_sql,
"access_tables": history.access_tables or [],
"status": history.status,
"error_code": history.error_code,
"error_message": history.error_message,
"row_count": history.row_count,
"truncated": history.truncated,
"elapsed_ms": history.elapsed_ms,
"trace_id": history.trace_id,
"create_time": history.create_time,
}
async def enrich_query_result(question: str, result, *, llm_client):
"""为已脱敏结果补充摘要和安全的基础图表配置。"""
summary = await summarize_result(
question,
result.columns,
result.rows,
llm_client=llm_client,
)
return replace(
result,
summary=summary,
answer=summary,
chart=build_chart_config(result.columns, result.rows),
)
async def _load_permission(db: AsyncSession, user: SysUser) -> dict:
return await load_query_permission(db, user.id)
def format_sse_event(event: str, payload: dict) -> str:
"""将结构化事件编码为标准 SSE 文本。"""
return f"event: {event}\ndata: {json.dumps(payload, ensure_ascii=False)}\n\n"
@router.post("/nl2sql/query/stream")
async def stream_query_data(
body: DataQueryReq,
request: Request,
user: SysUser = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
milvus=Depends(get_milvus),
redis=Depends(get_redis),
):
"""以 SSE 返回查询进度和脱敏统计,不改变原查询链路。"""
ensure_query_employee(user)
if body.output_format != "json":
raise ParamError("流式查询只支持 JSON 输出")
stream_id = uuid4().hex
trace_id = request.headers.get("X-Trace-Id") or get_request_id() or new_request_id()
async def events():
yield format_sse_event(
"started",
{"query_id": stream_id, "trace_id": trace_id},
)
try:
response = await query_data(body, request, user, db, milvus, redis)
payload = response.model_dump() if hasattr(response, "model_dump") else {}
data = payload.get("data") or {}
yield format_sse_event(
"completed",
{
"query_id": data.get("query_id", stream_id),
"trace_id": data.get("trace_id", trace_id),
"row_count": data.get("row_count", 0),
"truncated": bool(data.get("truncated", False)),
},
)
except Exception as exc: # noqa: BLE001 流式错误只返回异常类型
yield format_sse_event(
"failed",
{
"query_id": stream_id,
"trace_id": trace_id,
"error_type": type(exc).__name__,
},
)
return StreamingResponse(
events(),
media_type="text/event-stream",
headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"},
)
@router.post("/nl2sql/query")
async def query_data(
body: DataQueryReq,
request: Request,
user: SysUser = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
milvus=Depends(get_milvus),
redis=Depends(get_redis),
):
"""执行员工自然语言查询并返回数据结果。"""
ensure_query_employee(user)
trace_id = request.headers.get("X-Trace-Id") or get_request_id() or new_request_id()
query_id = uuid4().hex
clarification = build_clarification(body.question)
if clarification is not None:
return success({"status": "clarification_required", "clarification": clarification})
permission = await _load_permission(db, user)
limiter = QueryLimiter(redis)
quota_acquired = await limiter.acquire(
user.id,
daily_quota=permission.get("daily_quota", 0),
max_concurrent=1,
rate_limit=30,
)
if not quota_acquired:
query_metrics.record_rate_limited()
await write_nl2sql_audit_safely(
db,
user_id=user.id,
username=user.username,
action="limit_rejected",
target=query_id,
trace_id=trace_id,
detail={"question_length": len(body.question)},
status="失败",
)
raise ParamError("查询配额或并发限制已达上限")
session_lock = SessionLock(redis, body.session_id) if body.session_id else None
if session_lock is not None:
try:
await session_lock.acquire()
except SessionBusyError as exc:
raise ParamError("同一会话已有查询运行") from exc
async def permission_loader(_user_id: int):
return permission
async def metadata_retriever(question: str):
return await retrieve_metadata(question, milvus, top_k=runtime_config.retrieval_top_k)
async def schema_loader(table_names: set[str], _permission: dict):
return await load_authoritative_schema(
db,
database=settings.mysql.database,
candidate_tables=table_names,
)
request_contract = DataQueryRequest(
question=body.question,
user_id=user.id,
trace_id=trace_id,
session_id=body.session_id,
caller_agent=body.caller_agent,
data_scope=body.data_scope,
max_rows=min(body.max_rows or runtime_config.max_rows, runtime_config.max_rows),
include_sql=body.include_sql,
page=body.page,
page_size=body.page_size,
sort_by=body.sort_by,
sort_order=body.sort_order,
)
context_store = SessionContextStore(redis)
conversation_context = build_conversation_context(
await context_store.load(user.id, body.session_id)
)
validated_sql = None
cache_hit = False
try:
validated_sql = await build_query(
request_contract,
permission_loader=permission_loader,
metadata_retriever=metadata_retriever,
schema_loader=schema_loader,
conversation_context=conversation_context,
)
cache_key = build_cache_key(validated_sql.sql, permission=permission)
cached_payload = await cache_get(redis, cache_key)
if cached_payload is not None:
try:
cached_result = result_from_cache(cached_payload)
except Exception: # noqa: BLE001 缓存内容损坏时回退数据库
cached_result = None
if cached_result is not None:
cache_hit = True
await write_nl2sql_audit_safely(
db,
user_id=user.id,
username=user.username,
action="cache_hit",
target=query_id,
trace_id=trace_id,
detail={"tables": sorted(validated_sql.access_tables)},
)
result = replace(
cached_result,
query_id=query_id,
trace_id=trace_id,
warnings=[*cached_result.warnings, "cache_hit"],
)
else:
parameters = (
{"advisor_id": user.id}
if ":advisor_id" in validated_sql.sql
else None
)
result = await execute_readonly_sql(
db,
validated_sql,
query_id=query_id,
trace_id=trace_id,
user_id=user.id,
masks=permission.get("masks"),
parameters=parameters,
)
else:
parameters = (
{"advisor_id": user.id}
if ":advisor_id" in validated_sql.sql
else None
)
result = await execute_readonly_sql(
db,
validated_sql,
query_id=query_id,
trace_id=trace_id,
user_id=user.id,
masks=permission.get("masks"),
parameters=parameters,
)
if result.summary is None:
result = await enrich_query_result(body.question, result, llm_client=llm)
await cache_set(
redis,
cache_key,
result_to_cache(result),
ttl=runtime_config.cache_ttl,
access_tables=validated_sql.access_tables,
)
await write_nl2sql_audit_safely(
db,
user_id=user.id,
username=user.username,
action="query",
target=query_id,
trace_id=trace_id,
detail={"tables": sorted(validated_sql.access_tables), "row_count": result.row_count},
)
await context_store.append(user.id, body.session_id, body.question, "success")
diagnostic_registry.record(
query_id=query_id,
status="success",
access_tables=validated_sql.access_tables,
sql=validated_sql.sql,
model_elapsed_ms=result.elapsed_ms,
)
query_metrics.record(
status="success",
elapsed_ms=result.elapsed_ms,
cache_hit=cache_hit,
)
except Exception as exc: # noqa: BLE001 查询失败统一归档后再转业务异常
await archive_query_safely(
db,
query_id=query_id,
user_id=user.id,
question=body.question,
generated_sql=getattr(validated_sql, "sql", None),
access_tables=getattr(validated_sql, "access_tables", set()),
status="failed",
error_message="查询处理失败",
trace_id=trace_id,
session_id=body.session_id,
caller_agent=body.caller_agent,
)
await write_nl2sql_audit_safely(
db,
user_id=user.id,
username=user.username,
action="query_failed",
target=query_id,
trace_id=trace_id,
detail={"error_type": type(exc).__name__},
status="失败",
)
diagnostic_registry.record(
query_id=query_id,
status="failed",
access_tables=getattr(validated_sql, "access_tables", set()),
sql=getattr(validated_sql, "sql", None),
security_rule=type(exc).__name__,
)
await limiter.release(user.id)
query_metrics.record(
status="timeout" if type(exc).__name__ == "QueryExecutionError" and "超时" in str(exc) else "failed",
failure_reason=type(exc).__name__,
)
if session_lock is not None:
await session_lock.release()
if isinstance(exc, QueryServiceError):
raise ParamError(str(exc)) from exc
raise
await archive_query_safely(
db,
**build_history_payload(
query_id=query_id,
request_data=body,
user=user,
trace_id=trace_id,
status="success",
result=validated_sql,
query_result=result,
),
)
if not body.include_sql:
result = replace(result, sql=None)
if session_lock is not None:
await session_lock.release()
await limiter.release(user.id)
if body.output_format == "csv":
await write_nl2sql_audit_safely(
db,
user_id=user.id,
username=user.username,
action="result_export",
target=query_id,
trace_id=trace_id,
detail={"format": "csv", "row_count": result.row_count},
)
return Response(
content="\ufeff" + render_csv(result.columns, result.rows),
media_type="text/csv; charset=utf-8",
headers={"Content-Disposition": f'attachment; filename="nl2sql-{query_id}.csv"'},
)
return success(asdict(result))
@router.get("/nl2sql/health")
async def nl2sql_health(user: SysUser = Depends(get_current_user)):
"""返回 NL2SQL 依赖状态,仅对已登录员工开放。"""
ensure_query_employee(user)
result = await check_nl2sql_health(
{
"mysql": database.mysql.check_health,
"redis": database.redis.check_health,
"milvus": database.milvus.check_health,
"llm": llm.check_health,
}
)
return success(result)
@router.post("/nl2sql/query/explain")
async def explain_query(
body: DataExplainReq,
user: SysUser = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""对当前员工有权限的 SELECT 返回执行计划。"""
ensure_query_employee(user)
permission = await _load_permission(db, user)
max_rows = permission.get("max_rows") or 1000
try:
validated = validate_select_sql(
body.sql,
authorized_tables=permission.get("tables", set()),
authorized_columns=permission.get("columns"),
max_rows=max_rows,
)
result = await db.execute(text(build_explain_sql(validated)))
return success([dict(row) for row in result.mappings().all()])
except SqlSecurityError as exc:
raise ForbiddenError("SQL 未通过安全校验") from exc
@router.get("/nl2sql/query-history")
async def list_query_history(
page: int = Query(1, ge=1),
page_size: int = Query(20, ge=1, le=100),
status: str | None = Query(None, max_length=16),
start_time: datetime | None = Query(None),
end_time: datetime | None = Query(None),
user: SysUser = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""分页读取当前员工自己的查询历史。"""
ensure_query_employee(user)
page, page_size, offset = normalize_pagination(page, page_size)
rows = await Nl2SqlPermissionRepo(db).list_query_history(
user.id,
limit=page_size,
offset=offset,
status=status,
start_time=start_time,
end_time=end_time,
)
return success({"page": page, "page_size": page_size, "items": [history_payload(row) for row in rows]})
@router.get("/nl2sql/query-history/{query_id}")
async def get_query_history(
query_id: str,
user: SysUser = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""读取当前员工自己的查询历史详情。"""
ensure_query_employee(user)
history = await Nl2SqlPermissionRepo(db).get_query_history(user.id, query_id)
if history is None:
raise NotFoundError("查询历史不存在")
return success(history_payload(history))
@router.get("/nl2sql/query/running")
async def list_running_queries(user: SysUser = Depends(get_current_user)):
"""管理员查看当前进程登记的运行中查询。"""
ensure_query_admin(user)
return success([item.to_dict() for item in query_runtime_registry.list_active()])
@router.post("/nl2sql/query/kill")
async def kill_running_query(
body: DataKillReq,
user: SysUser = Depends(get_current_user),
):
"""管理员尝试中止运行中查询,并返回是否已验证。"""
ensure_query_admin(user)
item = query_runtime_registry.get(body.query_id)
if item is None or item.status != "running" or item.connection_id is None:
return success({"query_id": body.query_id, "verified": False})
try:
verified = await kill_mysql_query(item.connection_id)
except Exception: # noqa: BLE001 中止失败返回未验证,不暴露连接异常
verified = False
if verified:
await query_runtime_registry.mark_killed(body.query_id)
return success({"query_id": body.query_id, "verified": verified})
@router.get("/nl2sql/query-history/{query_id}/export")
async def export_query_history(
query_id: str,
user: SysUser = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""导出查询元数据 CSV,不导出结果行、凭据或连接信息。"""
ensure_query_employee(user)
history = await Nl2SqlPermissionRepo(db).get_query_history(user.id, query_id)
if history is None:
raise NotFoundError("查询历史不存在")
await write_nl2sql_audit_safely(
db,
user_id=user.id,
username=user.username,
action="history_export",
target=query_id,
trace_id=get_request_id() or new_request_id(),
detail={"format": "csv"},
)
output = StringIO()
fieldnames = [
"query_id", "question", "generated_sql", "access_tables", "status",
"row_count", "truncated", "elapsed_ms", "trace_id", "create_time",
]
writer = csv.DictWriter(output, fieldnames=fieldnames)
writer.writeheader()
payload = history_payload(history)
writer.writerow({
"query_id": payload["query_id"],
"question": payload["question"],
"generated_sql": payload["generated_sql"] or "",
"access_tables": ",".join(payload["access_tables"]),
"status": payload["status"],
"row_count": payload["row_count"],
"truncated": payload["truncated"],
"elapsed_ms": payload["elapsed_ms"] or "",
"trace_id": payload["trace_id"] or "",
"create_time": payload["create_time"] or "",
})
return Response(
content="\ufeff" + output.getvalue(),
media_type="text/csv; charset=utf-8",
headers={"Content-Disposition": f'attachment; filename="nl2sql-{query_id}.csv"'},
)
@router.get("/nl2sql/query/diagnostics")
async def query_diagnostics(user: SysUser = Depends(get_current_user)):
"""管理员查看脱敏运行诊断信息。"""
ensure_query_admin(user)
health = await check_nl2sql_health(
{
"mysql": database.mysql.check_health,
"redis": database.redis.check_health,
"milvus": database.milvus.check_health,
"llm": llm.check_health,
}
)
return success({
"health": health,
"running": [item.to_dict() for item in query_runtime_registry.list_active()],
"metrics": query_metrics.snapshot(),
"recent": diagnostic_registry.list_recent(),
})
@router.post("/nl2sql/cache/invalidate")
async def invalidate_query_cache(
body: DataCacheInvalidateReq,
user: SysUser = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
redis=Depends(get_redis),
):
"""管理员按表失效查询结果缓存。"""
ensure_query_admin(user)
deleted = await invalidate_tables(redis, set(body.table_names))
await write_nl2sql_audit_safely(
db,
user_id=user.id,
username=user.username,
action="cache_invalidate",
target=",".join(body.table_names),
trace_id=get_request_id() or new_request_id(),
detail={"deleted": deleted},
)
return success({"table_names": sorted(set(body.table_names)), "deleted": deleted})