feat:投顾agent接入nl2sql能力

This commit is contained in:
2026-09-13 18:24:44 +08:00
parent 43330042b8
commit 3f5fc1b9f0
13 changed files with 339 additions and 4 deletions
+2 -1
View File
@@ -1,10 +1,11 @@
"""投顾工作台路由聚合(前缀 /api/advisor,在 api/router.py 以 /api 挂载)。"""
from fastapi import APIRouter
from api.advisor import audit, customers, dashboard, diagnosis, drafts, report, todos, visits
from api.advisor import audit, customers, dashboard, data_query, diagnosis, drafts, report, todos, visits
router = APIRouter(prefix="/advisor", tags=["投顾工作台"])
router.include_router(dashboard.router)
router.include_router(data_query.router)
router.include_router(customers.router)
router.include_router(diagnosis.router)
router.include_router(drafts.router)
+31
View File
@@ -0,0 +1,31 @@
"""投顾工作台客户数据查询代理路由。"""
from fastapi import APIRouter, Depends, Request
from sqlalchemy.ext.asyncio import AsyncSession
from api.advisor._auth import extract_auth
from api.deps import require_advisor
from config.deps import get_db
from model.sys_user import SysUser
from schemas.advisor import AdvisorDataQueryReq
from service.advisor.data_query import query_customer_data
from utils.response import success
router = APIRouter()
@router.post("/data-query", summary="查询当前客户数据(代理投顾Agent NL2SQL)")
async def data_query(
req: AdvisorDataQueryReq,
request: Request,
user: SysUser = Depends(require_advisor),
db: AsyncSession = Depends(get_db),
):
auth, trace_id = extract_auth(request)
data = await query_customer_data(
db,
user,
auth_header=auth,
trace_id=trace_id,
req=req,
)
return success(data)
+2 -1
View File
@@ -7,7 +7,7 @@ from api.chat import client_agent, customer_agent, knowledge
from api.routers import product, questionnaire
from api.routers import account, auth, holdings, purchase, redeem, risk, work_order
from api.routers import advisor_agent, health, nl2sql, nl2sql_admin
from api.advisor import audit, customers, dashboard, diagnosis, drafts, report, todos, visits
from api.advisor import audit, customers, dashboard, data_query, diagnosis, drafts, report, todos, visits
api_router = APIRouter()
api_router.include_router(auth.router, prefix="/api", tags=["认证"])
@@ -25,6 +25,7 @@ api_router.include_router(questionnaire.router, prefix="/api", tags=["问卷"])
api_router.include_router(advisor_agent.router, prefix="/api", tags=["投顾Agent"])
for workbench_router in (
dashboard.router,
data_query.router,
customers.router,
diagnosis.router,
drafts.router,
+78 -1
View File
@@ -10,7 +10,9 @@ from pydantic import ValidationError
from sqlalchemy.ext.asyncio import AsyncSession
from agent.advisor_agent.auth import ensure_customer_access
from agent.advisor_agent.data_query import execute_advisor_data_query
from agent.advisor_agent.intent.fund_analysis import build_fund_analysis
from agent.advisor_agent.intent.recognizer import recognize_advisor_intent
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 (
@@ -19,6 +21,7 @@ from agent.advisor_agent.intent.generation_flow import (
)
from agent.advisor_agent.protocol import agent_failure, agent_success
from common.common_const import (
AGENT_INTENT_DATA_QUERY,
AGENT_INTENT_RECOMMEND,
CUSTOMER_REL_STATUS_SIGNED,
ERR_CODE_DRAFT_NOT_FOUND,
@@ -51,7 +54,9 @@ from schemas.advisor_agent import (
AdvisorFundAnalysisReq,
AdvisorRebalanceRunReq,
AdvisorTalkScriptReq,
AdvisorDataQueryReq,
)
from service.nl2sql.query_service import QueryServiceError
from service.advisor_agent.draft import (
detail_draft,
discard_draft,
@@ -158,10 +163,14 @@ async def chat_stream(
)
else:
customer_id = chat_request.customer_id
resolved_intent = recognize_advisor_intent(
chat_request.query,
chat_request.intent,
)
relation = await ensure_customer_access(
db, advisor_id=user.id, customer_id=int(customer_id)
)
if chat_request.intent == AGENT_INTENT_RECOMMEND:
if resolved_intent == AGENT_INTENT_RECOMMEND:
memories = await _recall_advisor_memories(
request,
customer_id=int(customer_id),
@@ -200,6 +209,41 @@ async def chat_stream(
media_type="text/event-stream",
headers={"X-Trace-Id": trace_id},
)
if resolved_intent == AGENT_INTENT_DATA_QUERY:
if not chat_request.query or not chat_request.query.strip():
payload = agent_failure(
ERR_CODE_FORBIDDEN_CUSTOMER,
"查询问题不能为空",
trace_id=trace_id,
)
else:
try:
result = await execute_advisor_data_query(
db,
advisor_id=user.id,
customer_id=int(customer_id),
question=chat_request.query,
trace_id=trace_id,
llm_client=getattr(_advisor_runtime(request), "llm_client", None),
)
except QueryServiceError:
payload = agent_failure(
ERR_CODE_LLM_ERROR,
"客户数据查询失败,请稍后重试",
trace_id=trace_id,
)
else:
async def events():
yield f"data: {json.dumps({'type': SSE_EVENT_TYPE_META, 'intent': AGENT_INTENT_DATA_QUERY, 'query_id': result.get('query_id'), 'trace_id': trace_id}, ensure_ascii=False)}\n\n"
if result.get("summary"):
yield f"data: {json.dumps({'type': SSE_EVENT_TYPE_TEXT, 'content': result['summary']}, ensure_ascii=False)}\n\n"
yield f"data: {json.dumps({'type': SSE_EVENT_TYPE_DONE, 'query_id': result.get('query_id')}, 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():
@@ -212,6 +256,39 @@ async def chat_stream(
)
@router.post("/data-query")
async def advisor_data_query(
request: Request,
body: AdvisorDataQueryReq,
user: SysUser = Depends(audited_advisor),
db: AsyncSession = Depends(get_db),
):
"""查询当前投顾选中客户的数据,不返回 SQL,也不生成草稿。"""
trace_id = _trace_id(request)
try:
result = await execute_advisor_data_query(
db,
advisor_id=user.id,
customer_id=body.customer_id,
question=body.question,
trace_id=trace_id,
session_id=body.session_id,
max_rows=body.max_rows,
page=body.page,
page_size=body.page_size,
sort_by=body.sort_by,
sort_order=body.sort_order,
)
except QueryServiceError:
return agent_failure(
ERR_CODE_LLM_ERROR,
"客户数据查询失败,请稍后重试",
trace_id=trace_id,
)
result.pop("sql", None)
return agent_success(result, trace_id=trace_id)
@router.get("/draft/list")
async def list_drafts(
request: Request,