"""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): """为已脱敏结果补充摘要和安全的基础图表配置。""" return replace( result, summary=await summarize_result( question, result.columns, result.rows, llm_client=llm_client, ), 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: result = await execute_readonly_sql( db, validated_sql, query_id=query_id, trace_id=trace_id, user_id=user.id, masks=permission.get("masks"), ) else: result = await execute_readonly_sql( db, validated_sql, query_id=query_id, trace_id=trace_id, user_id=user.id, masks=permission.get("masks"), ) 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})