feat:投顾agent接入nl2sql能力
This commit is contained in:
+1
-1
@@ -68,4 +68,4 @@ LLM_MAX_TOKENS=1024
|
||||
LLM_TIMEOUT=30
|
||||
LLM_MAX_RETRIES=3
|
||||
LLM_RETRY_BACKOFF_SEC=1
|
||||
LLM_FALLBACK_CHAT_MODEL= # 备用模型:主模型失败自动切换(留空则不启用)
|
||||
LLM_FALLBACK_CHAT_MODEL= # 备用模型:主模型失败自动切换(留空则不启用)
|
||||
|
||||
@@ -0,0 +1,98 @@
|
||||
"""投顾 Agent 的 NL2SQL 数据查询适配层。"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import asdict
|
||||
from typing import Any
|
||||
|
||||
from agent.advisor_agent.auth import ensure_customer_access
|
||||
from agent.data_query.agent import DataQueryAgent
|
||||
from config import database
|
||||
from config.settings import settings
|
||||
from nl2sql.contracts import DataQueryRequest, DataQueryResult
|
||||
from nl2sql.retrieval import retrieve_metadata
|
||||
from nl2sql.runtime_config import runtime_config
|
||||
from nl2sql.schema import load_authoritative_schema
|
||||
from service.nl2sql.permission_service import load_query_permission
|
||||
from service.nl2sql.query_service import QueryServiceError
|
||||
from tool.llm import llm as default_llm
|
||||
|
||||
|
||||
async def execute_advisor_data_query(
|
||||
db,
|
||||
*,
|
||||
advisor_id: int,
|
||||
customer_id: int,
|
||||
question: str,
|
||||
trace_id: str,
|
||||
session_id: str | None = None,
|
||||
data_scope: dict[str, Any] | None = None,
|
||||
max_rows: int | None = None,
|
||||
page: int = 1,
|
||||
page_size: int = 100,
|
||||
sort_by: str | None = None,
|
||||
sort_order: str = "asc",
|
||||
milvus=None,
|
||||
redis=None,
|
||||
llm_client=None,
|
||||
query_agent=None,
|
||||
) -> dict[str, Any]:
|
||||
"""在当前投顾和选中客户范围内执行只读自然语言查询。
|
||||
|
||||
``data_scope`` 即使由调用方传入也不会被信任,服务端始终覆盖为当前
|
||||
``customer_id``,避免投顾借助 NL2SQL 查询其他客户数据。
|
||||
"""
|
||||
await ensure_customer_access(db, advisor_id=advisor_id, customer_id=customer_id)
|
||||
permission = await load_query_permission(db, advisor_id)
|
||||
if not permission.get("can_query", False):
|
||||
raise QueryServiceError("当前投顾没有 NL2SQL 查询权限")
|
||||
|
||||
milvus = milvus or database.milvus.client()
|
||||
redis = redis or database.redis.client()
|
||||
llm_client = llm_client or default_llm
|
||||
request = DataQueryRequest(
|
||||
question=question,
|
||||
user_id=advisor_id,
|
||||
trace_id=trace_id,
|
||||
session_id=session_id,
|
||||
caller_agent="advisor_agent",
|
||||
data_scope={"customer_ids": [customer_id]},
|
||||
max_rows=min(max_rows or runtime_config.max_rows, runtime_config.max_rows),
|
||||
include_sql=False,
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
sort_by=sort_by,
|
||||
sort_order=sort_order,
|
||||
)
|
||||
|
||||
async def permission_loader(_user_id: int):
|
||||
return permission
|
||||
|
||||
async def metadata_retriever(query: str):
|
||||
return await retrieve_metadata(
|
||||
query,
|
||||
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,
|
||||
)
|
||||
|
||||
result: DataQueryResult = await (query_agent or DataQueryAgent()).query(
|
||||
request,
|
||||
session=db,
|
||||
permission_loader=permission_loader,
|
||||
metadata_retriever=metadata_retriever,
|
||||
schema_loader=schema_loader,
|
||||
llm_client=llm_client,
|
||||
summary_llm=llm_client,
|
||||
masks=permission.get("masks"),
|
||||
redis=redis,
|
||||
)
|
||||
payload = asdict(result)
|
||||
payload["sql"] = None
|
||||
payload["customer_id"] = customer_id
|
||||
return payload
|
||||
@@ -0,0 +1,54 @@
|
||||
"""投顾聊天入口的轻量意图识别。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
|
||||
from common.common_const import AGENT_INTENT_DATA_QUERY
|
||||
|
||||
|
||||
_QUERY_ACTIONS = (
|
||||
"查询",
|
||||
"查一下",
|
||||
"查看",
|
||||
"统计",
|
||||
"列出",
|
||||
"显示",
|
||||
"多少",
|
||||
"有哪些",
|
||||
"明细",
|
||||
)
|
||||
_DATA_TERMS = (
|
||||
"持仓",
|
||||
"资产",
|
||||
"收益",
|
||||
"交易记录",
|
||||
"申购",
|
||||
"赎回",
|
||||
"余额",
|
||||
"市值",
|
||||
"份额",
|
||||
"客户数据",
|
||||
"账户",
|
||||
)
|
||||
_NON_QUERY_INTENTS = ("推荐", "调仓", "再平衡", "话术", "沟通")
|
||||
|
||||
|
||||
def recognize_advisor_intent(query: str | None, explicit_intent: str | None = None) -> str | None:
|
||||
"""返回当前投顾聊天应使用的意图;无法判断时返回 ``None``。
|
||||
|
||||
数据查询采用保守规则:必须命中查询动作或客户数据表达,且不能明显是
|
||||
推荐、调仓或话术请求,避免把生成类请求送入 NL2SQL。
|
||||
"""
|
||||
if explicit_intent:
|
||||
return explicit_intent
|
||||
text = re.sub(r"\s+", "", query or "")
|
||||
if not text or any(term in text for term in _NON_QUERY_INTENTS):
|
||||
return None
|
||||
has_action = any(term in text for term in _QUERY_ACTIONS)
|
||||
has_data = any(term in text for term in _DATA_TERMS)
|
||||
if has_data and (has_action or "客户" in text or "近一年" in text or "本月" in text):
|
||||
return AGENT_INTENT_DATA_QUERY
|
||||
return None
|
||||
|
||||
|
||||
__all__ = ["recognize_advisor_intent"]
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -17,6 +17,7 @@ AGENT_INTENT_RECOMMEND = "recommend"
|
||||
AGENT_INTENT_REBALANCE = "rebalance"
|
||||
AGENT_INTENT_FUND_ANALYSIS = "fund_analysis"
|
||||
AGENT_INTENT_DIALOGUE_SCRIPT = "dialogue-script"
|
||||
AGENT_INTENT_DATA_QUERY = "data_query"
|
||||
|
||||
TALK_SCENE_RISK_BLOCK_ORDER = "risk_block_order"
|
||||
TALK_SCENE_MARKET_FLUCTUATION = "market_fluctuation"
|
||||
|
||||
@@ -39,6 +39,7 @@ AGENT_INTENT_RECOMMEND = "recommend" # 基金推荐
|
||||
AGENT_INTENT_REBALANCE = "rebalance" # 持仓诊断 & 调仓再平衡
|
||||
AGENT_INTENT_FUND_ANALYSIS = "fund_analysis" # 基金深度分析
|
||||
AGENT_INTENT_DIALOGUE_SCRIPT = "dialogue-script" # 生成沟通话术
|
||||
AGENT_INTENT_DATA_QUERY = "data_query" # 客户数据查询
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# §5 沟通话术场景(generate-talk-script 入参 scene_type)
|
||||
|
||||
@@ -33,6 +33,19 @@ class RebalanceRunReq(BaseModel):
|
||||
customer_id: int = Field(gt=0)
|
||||
|
||||
|
||||
class AdvisorDataQueryReq(BaseModel):
|
||||
"""工作台代理 Agent 的客户数据查询请求。"""
|
||||
|
||||
customer_id: int = Field(gt=0)
|
||||
question: str = Field(min_length=1, max_length=2000)
|
||||
session_id: str | None = Field(default=None, max_length=64)
|
||||
max_rows: int | None = Field(default=None, gt=0, le=10000)
|
||||
page: int = Field(default=1, ge=1, le=100000)
|
||||
page_size: int = Field(default=100, ge=1, le=10000)
|
||||
sort_by: str | None = Field(default=None, max_length=128)
|
||||
sort_order: Literal["asc", "desc"] = "asc"
|
||||
|
||||
|
||||
class TalkScriptReq(BaseModel):
|
||||
"""生成沟通话术草稿(同步,超时由工作台降级为「稍后重试」)。"""
|
||||
|
||||
|
||||
@@ -7,6 +7,7 @@ from pydantic import BaseModel, Field
|
||||
|
||||
from common.common_const import (
|
||||
AGENT_INTENT_DIALOGUE_SCRIPT,
|
||||
AGENT_INTENT_DATA_QUERY,
|
||||
AGENT_INTENT_FUND_ANALYSIS,
|
||||
AGENT_INTENT_REBALANCE,
|
||||
AGENT_INTENT_RECOMMEND,
|
||||
@@ -38,10 +39,22 @@ class AdvisorChatReq(BaseModel):
|
||||
AGENT_INTENT_REBALANCE,
|
||||
AGENT_INTENT_FUND_ANALYSIS,
|
||||
AGENT_INTENT_DIALOGUE_SCRIPT,
|
||||
AGENT_INTENT_DATA_QUERY,
|
||||
] | None = None
|
||||
query: str | None = Field(default=None, max_length=4000)
|
||||
|
||||
|
||||
class AdvisorDataQueryReq(BaseModel):
|
||||
customer_id: int = Field(gt=0)
|
||||
question: str = Field(min_length=1, max_length=2000)
|
||||
session_id: str | None = Field(default=None, max_length=64)
|
||||
max_rows: int | None = Field(default=None, gt=0, le=10000)
|
||||
page: int = Field(default=1, ge=1, le=100000)
|
||||
page_size: int = Field(default=100, ge=1, le=10000)
|
||||
sort_by: str | None = Field(default=None, max_length=128)
|
||||
sort_order: Literal["asc", "desc"] = "asc"
|
||||
|
||||
|
||||
class AdvisorFundAnalysisReq(BaseModel):
|
||||
customer_id: int | None = Field(default=None, gt=0)
|
||||
fund_codes: list[str] = Field(min_length=1)
|
||||
|
||||
@@ -188,6 +188,19 @@ class AdvisorAgentClient:
|
||||
trace_id=trace_id, json={"customer_id": customer_id, "scene_type": scene_type},
|
||||
)
|
||||
|
||||
async def data_query(
|
||||
self,
|
||||
payload: dict,
|
||||
*,
|
||||
auth_header: str,
|
||||
trace_id: str,
|
||||
) -> dict:
|
||||
"""代理当前投顾选中客户的数据查询,不向工作台暴露 SQL。"""
|
||||
return await self._request(
|
||||
"POST", "/data-query", auth_header=auth_header,
|
||||
trace_id=trace_id, json=payload,
|
||||
)
|
||||
|
||||
|
||||
_client: AdvisorAgentClient | None = None
|
||||
|
||||
|
||||
@@ -0,0 +1,32 @@
|
||||
"""工作台到投顾 Agent 的客户数据查询代理。"""
|
||||
from __future__ import annotations
|
||||
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from model.sys_user import SysUser
|
||||
from schemas.advisor import AdvisorDataQueryReq
|
||||
from service.advisor.agent_client import get_agent_client
|
||||
from service.advisor.permissions import ensure_customer_owned
|
||||
|
||||
|
||||
async def query_customer_data(
|
||||
db: AsyncSession,
|
||||
user: SysUser,
|
||||
*,
|
||||
auth_header: str,
|
||||
trace_id: str,
|
||||
req: AdvisorDataQueryReq,
|
||||
) -> dict:
|
||||
"""先做工作台归属校验,再透传到 Agent;customer_id 不由 Agent 客户端覆盖。"""
|
||||
await ensure_customer_owned(db, user.id, req.customer_id)
|
||||
payload = req.model_dump(exclude_none=True)
|
||||
result = await get_agent_client().data_query(
|
||||
payload,
|
||||
auth_header=auth_header,
|
||||
trace_id=trace_id,
|
||||
)
|
||||
data = result["data"] or {}
|
||||
data.pop("sql", None)
|
||||
if result["warning"]:
|
||||
data["warning"] = result["warning"]
|
||||
return data
|
||||
Reference in New Issue
Block a user