diff --git a/.env.example b/.env.example index 363c381..f45eced 100644 --- a/.env.example +++ b/.env.example @@ -68,4 +68,4 @@ LLM_MAX_TOKENS=1024 LLM_TIMEOUT=30 LLM_MAX_RETRIES=3 LLM_RETRY_BACKOFF_SEC=1 -LLM_FALLBACK_CHAT_MODEL= # 备用模型:主模型失败自动切换(留空则不启用) \ No newline at end of file +LLM_FALLBACK_CHAT_MODEL= # 备用模型:主模型失败自动切换(留空则不启用) diff --git a/agent/advisor_agent/data_query.py b/agent/advisor_agent/data_query.py new file mode 100644 index 0000000..b70e5a2 --- /dev/null +++ b/agent/advisor_agent/data_query.py @@ -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 diff --git a/agent/advisor_agent/intent/recognizer.py b/agent/advisor_agent/intent/recognizer.py new file mode 100644 index 0000000..b6d4907 --- /dev/null +++ b/agent/advisor_agent/intent/recognizer.py @@ -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"] diff --git a/api/advisor/__init__.py b/api/advisor/__init__.py index bff65c2..0bb8232 100644 --- a/api/advisor/__init__.py +++ b/api/advisor/__init__.py @@ -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) diff --git a/api/advisor/data_query.py b/api/advisor/data_query.py new file mode 100644 index 0000000..f3850b2 --- /dev/null +++ b/api/advisor/data_query.py @@ -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) diff --git a/api/router.py b/api/router.py index 7d7772e..f99603a 100644 --- a/api/router.py +++ b/api/router.py @@ -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, diff --git a/api/routers/advisor_agent.py b/api/routers/advisor_agent.py index aae2121..0542a78 100644 --- a/api/routers/advisor_agent.py +++ b/api/routers/advisor_agent.py @@ -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, diff --git a/common/common_const.py b/common/common_const.py index 7313563..fb64ed3 100644 --- a/common/common_const.py +++ b/common/common_const.py @@ -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" diff --git a/common_const.py b/common_const.py index ab7d717..48e9977 100644 --- a/common_const.py +++ b/common_const.py @@ -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) diff --git a/schemas/advisor.py b/schemas/advisor.py index 25c8b07..a096d89 100644 --- a/schemas/advisor.py +++ b/schemas/advisor.py @@ -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): """生成沟通话术草稿(同步,超时由工作台降级为「稍后重试」)。""" diff --git a/schemas/advisor_agent.py b/schemas/advisor_agent.py index f48a238..9d106e3 100644 --- a/schemas/advisor_agent.py +++ b/schemas/advisor_agent.py @@ -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) diff --git a/service/advisor/agent_client.py b/service/advisor/agent_client.py index d7e20bd..416f272 100644 --- a/service/advisor/agent_client.py +++ b/service/advisor/agent_client.py @@ -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 diff --git a/service/advisor/data_query.py b/service/advisor/data_query.py new file mode 100644 index 0000000..fca55e4 --- /dev/null +++ b/service/advisor/data_query.py @@ -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