feat:新增投顾agent和nl2sqlagent

This commit is contained in:
2026-09-13 16:19:24 +08:00
parent c80c6acac0
commit 163192bf55
122 changed files with 7488 additions and 362 deletions
+407
View File
@@ -0,0 +1,407 @@
"""投顾 Agent HTTP 契约入口与本地业务编排。"""
from __future__ import annotations
import json
from typing import Literal
from fastapi import APIRouter, BackgroundTasks, Depends, Query, Request
from fastapi.responses import StreamingResponse
from pydantic import ValidationError
from sqlalchemy.ext.asyncio import AsyncSession
from agent.advisor_agent.auth import ensure_customer_access
from agent.advisor_agent.intent.fund_analysis import build_fund_analysis
from agent.advisor_agent.intent.talk_script import build_talk_script
from agent.advisor_agent.llm import generate_text
from agent.advisor_agent.intent.generation_flow import (
generate_rebalance_draft,
generate_recommendation_draft,
)
from agent.advisor_agent.protocol import agent_failure, agent_success
from common.common_const import (
AGENT_INTENT_RECOMMEND,
CUSTOMER_REL_STATUS_SIGNED,
ERR_CODE_DRAFT_NOT_FOUND,
ERR_CODE_FORBIDDEN_CUSTOMER,
ERR_CODE_LLM_ERROR,
ERR_CODE_NOT_SIGNED_REBALANCE,
DRAFT_STATUS_DISCARDED,
DRAFT_STATUS_DRAFT,
SSE_EVENT_TYPE_DONE,
SSE_EVENT_TYPE_ERROR,
SSE_EVENT_TYPE_META,
SSE_EVENT_TYPE_TEXT,
)
from api.deps import audited_advisor
from config.deps import get_db
from config.database import mysql, redis as redis_db
from model.sys_user import SysUser
from repositories.advisor_draft import AdvisorDraftRepo
from service.advisor_agent.context import (
load_fund_analysis_context,
load_customer_risk,
load_rebalance_context,
load_recommendation_context,
)
from service.event_publisher import publish_event
from schemas.advisor_agent import (
AdvisorDraftOperateReq,
AdvisorDraftSaveReq,
AdvisorChatReq,
AdvisorFundAnalysisReq,
AdvisorRebalanceRunReq,
AdvisorTalkScriptReq,
)
from service.advisor_agent.draft import (
detail_draft,
discard_draft,
ensure_draft_owner,
get_draft as get_draft_service,
list_drafts as list_drafts_service,
save_draft as save_draft_service,
)
from utils.exceptions import ApiError
from utils.request_id import get_request_id, new_request_id
from utils.logger import get_logger
router = APIRouter(prefix="/advisor-agent")
logger = get_logger("advisor_agent.router")
_NOT_READY_CODE = ERR_CODE_LLM_ERROR
_NOT_READY_MESSAGE = "投顾 Agent 核心能力尚未初始化"
_DRAFT_NOT_FOUND_CODE = ERR_CODE_DRAFT_NOT_FOUND
_DRAFT_NOT_FOUND_MESSAGE = "草稿不存在或者已废弃"
def _trace_id(request: Request) -> str:
return request.headers.get("X-Trace-Id") or get_request_id() or new_request_id()
def _not_ready(request: Request):
return agent_failure(_NOT_READY_CODE, _NOT_READY_MESSAGE, trace_id=_trace_id(request))
def _advisor_runtime(request: Request):
app = request.scope.get("app")
return getattr(getattr(app, "state", None), "advisor_agent_runtime", None)
async def _recall_advisor_memories(
request: Request, *, customer_id: int, query: str
) -> list[dict]:
"""从应用共享 runtime 读取记忆;未装配或故障时安全降级为空。"""
app = request.scope.get("app")
state = getattr(app, "state", None)
runtime = getattr(state, "advisor_agent_runtime", None)
provider = getattr(runtime, "memory_provider", None)
if provider is None:
return []
try:
return await provider.recall(customer_id=customer_id, query=query)
except Exception:
logger.warning(
"advisor memory recall failed: customer_id=%s", customer_id, exc_info=True
)
return []
async def _ensure_draft_access(db, draft, advisor_id: int) -> None:
ensure_draft_owner(draft, advisor_id=advisor_id)
await ensure_customer_access(
db,
advisor_id=advisor_id,
customer_id=draft.customer_id,
)
async def _run_rebalance_background(
*, advisor_id: int, customer_id: int, trace_id: str
) -> None:
"""后台任务使用独立会话,避免请求返回后复用已关闭的请求会话。"""
try:
async with mysql.get_session_factory()() as task_db:
relation = await ensure_customer_access(
task_db, advisor_id=advisor_id, customer_id=customer_id
)
if relation.status != CUSTOMER_REL_STATUS_SIGNED:
return
context = await load_rebalance_context(task_db, customer_id=customer_id)
if context is None:
return
await generate_rebalance_draft(
draft_repo=AdvisorDraftRepo(task_db),
publish=lambda **kwargs: publish_event(redis_db.client(), task_db, **kwargs),
advisor_id=advisor_id,
trace_id=trace_id,
**context,
)
except Exception:
logger.warning("advisor rebalance background task failed", exc_info=True)
@router.post("/chat/stream")
async def chat_stream(
request: Request,
body: dict,
user: SysUser = Depends(audited_advisor),
db: AsyncSession = Depends(get_db),
):
trace_id = _trace_id(request)
try:
chat_request = AdvisorChatReq.model_validate(body)
except ValidationError:
payload = agent_failure(
ERR_CODE_FORBIDDEN_CUSTOMER,
"对话请求缺少有效客户范围或参数",
trace_id=trace_id,
)
else:
customer_id = chat_request.customer_id
relation = await ensure_customer_access(
db, advisor_id=user.id, customer_id=int(customer_id)
)
if chat_request.intent == AGENT_INTENT_RECOMMEND:
memories = await _recall_advisor_memories(
request,
customer_id=int(customer_id),
query=chat_request.query or "",
)
runtime = _advisor_runtime(request)
context = await load_recommendation_context(
db, customer_id=int(customer_id)
)
if context is not None:
draft = await generate_recommendation_draft(
draft_repo=AdvisorDraftRepo(db),
advisor_id=user.id,
relation_status=relation.status,
trace_id=trace_id,
memories=memories,
llm_client=getattr(runtime, "llm_client", None),
**context,
)
async def events():
for event in (
{
"type": SSE_EVENT_TYPE_META,
"draft_id": draft.draft_id,
"intent": draft.intent,
"status": draft.status,
},
{"type": SSE_EVENT_TYPE_TEXT, "content": draft.content},
{"type": SSE_EVENT_TYPE_DONE, "draft_id": draft.draft_id},
):
yield f"data: {json.dumps(event, ensure_ascii=False)}\n\n"
return StreamingResponse(
events(),
media_type="text/event-stream",
headers={"X-Trace-Id": trace_id},
)
payload = agent_failure(_NOT_READY_CODE, _NOT_READY_MESSAGE, trace_id=trace_id)
async def events():
yield f"data: {json.dumps({'type': SSE_EVENT_TYPE_ERROR, **payload}, ensure_ascii=False)}\n\n"
return StreamingResponse(
events(),
media_type="text/event-stream",
headers={"X-Trace-Id": trace_id},
)
@router.get("/draft/list")
async def list_drafts(
request: Request,
advisor_id: int | None = Query(default=None),
customer_id: int | None = Query(default=None),
status: Literal[DRAFT_STATUS_DRAFT, DRAFT_STATUS_DISCARDED] | None = Query(default=None),
page: int = Query(default=1, ge=1),
page_size: int = Query(default=20, ge=1, le=100),
user: SysUser = Depends(audited_advisor),
db: AsyncSession = Depends(get_db),
):
if advisor_id is not None and advisor_id != user.id:
raise ApiError(ERR_CODE_FORBIDDEN_CUSTOMER, "无权操作该客户数据")
if customer_id is not None:
await ensure_customer_access(db, advisor_id=user.id, customer_id=customer_id)
result = await list_drafts_service(
AdvisorDraftRepo(db),
advisor_id=user.id,
customer_id=customer_id,
status=status,
page=page,
page_size=page_size,
)
return agent_success(result, trace_id=_trace_id(request))
@router.get("/draft/{draft_id}")
async def get_draft(
draft_id: str,
request: Request,
user: SysUser = Depends(audited_advisor),
db: AsyncSession = Depends(get_db),
):
draft = await get_draft_service(AdvisorDraftRepo(db), draft_id)
await _ensure_draft_access(db, draft, user.id)
return agent_success(detail_draft(draft), trace_id=_trace_id(request))
@router.put("/draft/{draft_id}/save")
async def save_draft(
draft_id: str,
request: Request,
body: AdvisorDraftSaveReq,
user: SysUser = Depends(audited_advisor),
db: AsyncSession = Depends(get_db),
):
repo = AdvisorDraftRepo(db)
draft = await get_draft_service(repo, draft_id)
await _ensure_draft_access(db, draft, user.id)
customer_risk = await load_customer_risk(db, customer_id=draft.customer_id)
result = await save_draft_service(
repo,
draft_id,
title=body.title,
content=body.content,
structured_data=body.structured_data,
customer_risk=customer_risk,
)
return agent_success(result, trace_id=_trace_id(request))
@router.post("/draft/{draft_id}/operate")
async def operate_draft(
draft_id: str,
request: Request,
body: AdvisorDraftOperateReq,
user: SysUser = Depends(audited_advisor),
db: AsyncSession = Depends(get_db),
):
if body.operation != "discard":
raise ApiError(_DRAFT_NOT_FOUND_CODE, "不支持的草稿操作")
repo = AdvisorDraftRepo(db)
draft = await get_draft_service(repo, draft_id)
await _ensure_draft_access(db, draft, user.id)
result = await discard_draft(repo, draft_id)
return agent_success(detail_draft(result), trace_id=_trace_id(request))
@router.post("/rebalance/run")
async def run_rebalance(
request: Request,
body: AdvisorRebalanceRunReq,
user: SysUser = Depends(audited_advisor),
db: AsyncSession = Depends(get_db),
background_tasks: BackgroundTasks = None,
):
customer_id = body.customer_id
relation = await ensure_customer_access(
db, advisor_id=user.id, customer_id=int(customer_id)
)
if relation.status != CUSTOMER_REL_STATUS_SIGNED:
return agent_failure(
ERR_CODE_NOT_SIGNED_REBALANCE,
"客户尚未签约,禁止生成调仓草稿",
trace_id=_trace_id(request),
)
if background_tasks is None:
background_tasks = BackgroundTasks()
background_tasks.add_task(
_run_rebalance_background,
advisor_id=user.id,
customer_id=customer_id,
trace_id=_trace_id(request),
)
return agent_success(
{"accepted": True, "status": "queued"},
trace_id=_trace_id(request),
)
@router.post("/fund-analysis")
async def fund_analysis(
request: Request,
body: AdvisorFundAnalysisReq,
user: SysUser = Depends(audited_advisor),
db: AsyncSession = Depends(get_db),
):
memories: list[dict] = []
if body.customer_id is not None:
await ensure_customer_access(
db, advisor_id=user.id, customer_id=body.customer_id
)
memories = await _recall_advisor_memories(
request,
customer_id=body.customer_id,
query=f"基金分析 {','.join(body.fund_codes)}",
)
fund = body.fund
performance = body.performance
if fund is None or performance is None:
contexts = await load_fund_analysis_context(
db,
fund_codes=[str(code) for code in body.fund_codes],
)
if not contexts:
return _not_ready(request)
if len(contexts) == 1:
fund = contexts[0]["fund"]
performance = contexts[0]["performance"]
else:
return agent_success(
{"items": [build_fund_analysis(item["fund"], item["performance"]) for item in contexts]},
trace_id=_trace_id(request),
)
result = build_fund_analysis(fund, performance)
runtime = _advisor_runtime(request)
if getattr(runtime, "llm_client", None) is not None:
result["analysis_text"] = await generate_text(
runtime.llm_client,
system_prompt="你是基金投顾助手,只基于给定历史数据生成谨慎的内部分析,不承诺收益。",
user_prompt=json.dumps(
{"analysis": result, "customer_memories": memories},
ensure_ascii=False,
),
fallback=lambda: result["analysis_text"],
timeout=5.0,
)
return agent_success(result, trace_id=_trace_id(request))
@router.post("/generate-talk-script")
async def generate_talk_script(
request: Request,
body: AdvisorTalkScriptReq,
user: SysUser = Depends(audited_advisor),
db: AsyncSession = Depends(get_db),
):
await ensure_customer_access(db, advisor_id=user.id, customer_id=body.customer_id)
memories = await _recall_advisor_memories(
request,
customer_id=body.customer_id,
query=f"沟通话术 {body.scene_type}",
)
try:
result = build_talk_script(
body.scene_type,
customer_name=body.customer_name,
)
except ValueError as exc:
raise ApiError(ERR_CODE_LLM_ERROR, str(exc)) from exc
runtime = _advisor_runtime(request)
if getattr(runtime, "llm_client", None) is not None:
result["content"] = await generate_text(
runtime.llm_client,
system_prompt="你是合规的基金投顾助手,只生成谨慎沟通话术,不承诺收益、不代客交易。",
user_prompt=json.dumps(
{"script": result, "customer_memories": memories},
ensure_ascii=False,
),
fallback=lambda: result["content"],
timeout=5.0,
)
return agent_success(result, trace_id=_trace_id(request))
+25
View File
@@ -0,0 +1,25 @@
"""基础设施健康检查路由。"""
from __future__ import annotations
from fastapi import APIRouter
from fastapi.responses import JSONResponse
from config import database
from utils.response import fail, success
router = APIRouter()
@router.get("/health/ready")
async def ready():
"""按需探测四库,并返回适合负载均衡器使用的 HTTP 就绪状态。"""
databases = await database.check_ready_detail()
is_ready = all(item.get("status") == "ok" for item in databases.values())
payload = {"ready": is_ready, "databases": databases}
if is_ready:
return JSONResponse(status_code=200, content=success(payload).model_dump())
return JSONResponse(
status_code=503,
content=fail(503, "基础设施尚未就绪", payload).model_dump(),
)
+614
View File
@@ -0,0 +1,614 @@
"""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
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)
rows = await Nl2SqlPermissionRepo(db).list_query_history(
user.id,
limit=page_size,
offset=(page - 1) * page_size,
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})
+323
View File
@@ -0,0 +1,323 @@
"""NL2SQL 管理接口。"""
from __future__ import annotations
import json
import time
from datetime import datetime, timedelta, timezone
from fastapi import APIRouter, Depends
from fastapi.responses import PlainTextResponse
from sqlalchemy.ext.asyncio import AsyncSession
from api.deps import get_current_user
from config.deps import get_db, get_redis
from model.sys_user import SysUser
from nl2sql.audit import write_nl2sql_audit_safely
from nl2sql.metrics import load_history_metrics_safely, query_metrics, render_prometheus
from nl2sql.runtime_config import runtime_config
from nl2sql.job_history import list_job_history, record_job_history_safely
from nl2sql.jobs import (
run_consistency_check,
run_history_cleanup,
run_metadata_sync,
run_vector_cleanup,
)
from nl2sql.semantics import get_semantic_catalog_info, refresh_semantic_catalog
from repositories.nl2sql_permission import Nl2SqlPermissionRepo
from schemas.nl2sql_admin import (
ColumnPermissionCreateReq,
ColumnPermissionUpdateReq,
RoleCreateReq,
RoleUpdateReq,
SensitiveFieldCreateReq,
SensitiveFieldUpdateReq,
TablePermissionCreateReq,
TablePermissionUpdateReq,
MaintenanceJobReq,
RuntimeConfigUpdateReq,
)
from service.nl2sql.admin_service import (
column_permission_payload,
create_role,
role_payload,
sensitive_field_payload,
table_permission_payload,
validate_mask_type,
validate_table_permission,
)
from utils.exceptions import NotFoundError, ParamError
from utils.request_id import get_request_id, new_request_id
from utils.response import success
from api.routers.nl2sql import ensure_query_admin
router = APIRouter()
async def _audit(
db,
user: SysUser,
action: str,
target: str,
detail: dict | None = None,
status: str = "成功",
):
await write_nl2sql_audit_safely(
db,
user_id=user.id,
username=user.username,
action=action,
target=target,
trace_id=get_request_id() or new_request_id(),
detail=detail,
status=status,
)
@router.get("/nl2sql/admin/roles")
async def list_roles(user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)):
ensure_query_admin(user)
rows = await Nl2SqlPermissionRepo(db).list_roles(include_inactive=True)
return success([role_payload(row) for row in rows])
@router.post("/nl2sql/admin/roles")
async def add_role(body: RoleCreateReq, user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)):
ensure_query_admin(user)
try:
role = await create_role(db, body.model_dump())
except ValueError as exc:
raise ParamError(str(exc)) from exc
await _audit(db, user, "permission_role_create", str(role.id), {"role_code": role.role_code})
return success(role_payload(role))
@router.patch("/nl2sql/admin/roles/{role_id}")
async def edit_role(role_id: int, body: RoleUpdateReq, user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)):
ensure_query_admin(user)
role = await Nl2SqlPermissionRepo(db).update_role(role_id, **body.model_dump(exclude_none=True))
if role is None:
raise NotFoundError("NL2SQL 角色不存在")
await _audit(db, user, "permission_role_update", str(role_id))
return success(role_payload(role))
@router.get("/nl2sql/admin/roles/{role_id}/tables")
async def list_tables(role_id: int, user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)):
ensure_query_admin(user)
rows = await Nl2SqlPermissionRepo(db).list_role_table_permissions(role_id, include_inactive=True)
return success([table_permission_payload(row) for row in rows])
@router.post("/nl2sql/admin/roles/{role_id}/tables")
async def add_table(role_id: int, body: TablePermissionCreateReq, user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)):
ensure_query_admin(user)
if await Nl2SqlPermissionRepo(db).get_role(role_id) is None:
raise NotFoundError("NL2SQL 角色不存在")
try:
values = validate_table_permission({**body.model_dump(), "role_id": role_id})
item = await Nl2SqlPermissionRepo(db).add_table_permission(role_id=role_id, **values)
except ValueError as exc:
raise ParamError(str(exc)) from exc
await _audit(db, user, "permission_table_create", str(item.id), {"table_name": item.table_name})
return success(table_permission_payload(item))
@router.patch("/nl2sql/admin/table-permissions/{permission_id}")
async def edit_table(permission_id: int, body: TablePermissionUpdateReq, user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)):
ensure_query_admin(user)
item = await Nl2SqlPermissionRepo(db).update_table_permission(permission_id, **body.model_dump(exclude_none=True))
if item is None:
raise NotFoundError("NL2SQL 表权限不存在")
await _audit(db, user, "permission_table_update", str(permission_id))
return success(table_permission_payload(item))
@router.get("/nl2sql/admin/roles/{role_id}/columns")
async def list_columns(role_id: int, user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)):
ensure_query_admin(user)
rows = await Nl2SqlPermissionRepo(db).list_role_column_permissions(role_id, include_inactive=True)
return success([column_permission_payload(row) for row in rows])
@router.post("/nl2sql/admin/roles/{role_id}/columns")
async def add_column(role_id: int, body: ColumnPermissionCreateReq, user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)):
ensure_query_admin(user)
if await Nl2SqlPermissionRepo(db).get_role(role_id) is None:
raise NotFoundError("NL2SQL 角色不存在")
if body.access_mode == "mask":
try:
validate_mask_type(body.mask_type)
except ValueError as exc:
raise ParamError(str(exc)) from exc
item = await Nl2SqlPermissionRepo(db).add_column_permission(role_id=role_id, **body.model_dump())
await _audit(db, user, "permission_column_create", str(item.id), {"table_name": item.table_name})
return success(column_permission_payload(item))
@router.patch("/nl2sql/admin/column-permissions/{permission_id}")
async def edit_column(permission_id: int, body: ColumnPermissionUpdateReq, user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)):
ensure_query_admin(user)
values = body.model_dump(exclude_none=True)
if "mask_type" in values:
try:
validate_mask_type(values["mask_type"])
except ValueError as exc:
raise ParamError(str(exc)) from exc
item = await Nl2SqlPermissionRepo(db).update_column_permission(permission_id, **values)
if item is None:
raise NotFoundError("NL2SQL 字段权限不存在")
await _audit(db, user, "permission_column_update", str(permission_id))
return success(column_permission_payload(item))
@router.get("/nl2sql/admin/sensitive-fields")
async def list_sensitive(user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)):
ensure_query_admin(user)
rows = await Nl2SqlPermissionRepo(db).list_sensitive_fields_admin(include_inactive=True)
return success([sensitive_field_payload(row) for row in rows])
@router.post("/nl2sql/admin/sensitive-fields")
async def add_sensitive(body: SensitiveFieldCreateReq, user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)):
ensure_query_admin(user)
item = await Nl2SqlPermissionRepo(db).add_sensitive_field(**body.model_dump())
await _audit(db, user, "permission_sensitive_create", str(item.id), {"table_name": item.table_name})
return success(sensitive_field_payload(item))
@router.patch("/nl2sql/admin/sensitive-fields/{field_id}")
async def edit_sensitive(field_id: int, body: SensitiveFieldUpdateReq, user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)):
ensure_query_admin(user)
item = await Nl2SqlPermissionRepo(db).update_sensitive_field(field_id, **body.model_dump(exclude_none=True))
if item is None:
raise NotFoundError("NL2SQL 敏感字段不存在")
await _audit(db, user, "permission_sensitive_update", str(field_id))
return success(sensitive_field_payload(item))
@router.get("/nl2sql/admin/metrics")
async def admin_metrics(user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)):
ensure_query_admin(user)
return success({
"runtime": query_metrics.snapshot(),
"history": await load_history_metrics_safely(db),
})
@router.get("/nl2sql/admin/runtime-config")
async def get_runtime_config(user: SysUser = Depends(get_current_user)):
"""管理员查看当前进程内 NL2SQL 运行参数。"""
ensure_query_admin(user)
return success(runtime_config.model_dump())
@router.patch("/nl2sql/admin/runtime-config")
async def update_runtime_config(
body: RuntimeConfigUpdateReq,
user: SysUser = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""管理员更新当前进程内 NL2SQL 非敏感运行参数。"""
ensure_query_admin(user)
try:
values = runtime_config.update(**body.model_dump(exclude_none=True))
except ValueError as exc:
raise ParamError("运行参数不合法") from exc
await _audit(db, user, "runtime_config_update", "nl2sql", {"fields": sorted(body.model_dump(exclude_none=True))})
return success(values)
@router.get("/nl2sql/admin/metrics/prometheus", response_class=PlainTextResponse)
async def admin_metrics_prometheus(user: SysUser = Depends(get_current_user)):
"""管理员读取聚合 Prometheus 指标,不返回查询明细。"""
ensure_query_admin(user)
return render_prometheus()
@router.post("/nl2sql/admin/jobs")
async def run_admin_job(
body: MaintenanceJobReq,
user: SysUser = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
redis=Depends(get_redis),
):
"""管理员手动执行一个幂等运维任务并保存执行历史。"""
ensure_query_admin(user)
started = time.perf_counter()
if body.task == "metadata_sync":
result = await run_metadata_sync(redis=redis)
elif body.task == "vector_cleanup":
result = await run_vector_cleanup(redis=redis)
elif body.task == "consistency_check":
from scripts.check_nl2sql_consistency import collect_consistency
result = await run_consistency_check(redis=redis, worker=collect_consistency)
else:
before = datetime.now(timezone.utc).replace(tzinfo=None) - timedelta(days=body.before_days)
result = await run_history_cleanup(db, before, redis=redis)
await record_job_history_safely(
db,
result,
elapsed_ms=(time.perf_counter() - started) * 1000,
parameter_summary={"before_days": body.before_days} if body.task == "history_cleanup" else {},
)
await _audit(db, user, "maintenance_job_run", result.name, {"status": result.status})
return success({
"name": result.name,
"status": result.status,
"attempts": result.attempts,
"detail": result.detail,
"error_type": result.error_type,
})
@router.get("/nl2sql/admin/jobs/history")
async def admin_job_history(
page: int = 1,
page_size: int = 20,
status: str | None = None,
user: SysUser = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
"""管理员查询运维任务执行历史摘要。"""
ensure_query_admin(user)
return success(await list_job_history(db, page=page, page_size=page_size, status=status))
@router.get("/nl2sql/admin/semantics")
async def admin_semantics(user: SysUser = Depends(get_current_user)):
"""管理员查看当前语义目录版本和规模摘要。"""
ensure_query_admin(user)
return success(get_semantic_catalog_info())
@router.post("/nl2sql/admin/semantics/refresh")
async def refresh_semantics(user: SysUser = Depends(get_current_user), db: AsyncSession = Depends(get_db)):
"""管理员刷新语义目录缓存,目录无效时继续使用旧缓存。"""
ensure_query_admin(user)
try:
result = refresh_semantic_catalog()
except (OSError, ValueError, json.JSONDecodeError) as exc:
await _audit(
db,
user,
"semantic_catalog_refresh",
"default",
{"error_type": type(exc).__name__},
status="失败",
)
raise ParamError("语义目录刷新失败") from exc
await _audit(
db,
user,
"semantic_catalog_refresh",
"default",
{
"version": result["version"],
"previous_version": result["previous_version"],
"changed": result["changed"],
"digest": result["digest"],
},
)
return success(result)