feat:新增投顾agent和nl2sqlagent
This commit is contained in:
@@ -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))
|
||||
@@ -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(),
|
||||
)
|
||||
@@ -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})
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user