feat:投顾agent接入nl2sql能力
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user