Files
Mutual_Fund/api/routers/advisor_agent.py
T

511 lines
18 KiB
Python
Raw Normal View History

2026-09-13 16:19:24 +08:00
"""投顾 Agent HTTP 契约入口与本地业务编排。"""
from __future__ import annotations
import json
2026-09-13 21:25:04 +08:00
import re
2026-09-13 16:19:24 +08:00
from typing import Literal
from fastapi import APIRouter, BackgroundTasks, Depends, Query, Request
from fastapi.responses import StreamingResponse
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
2026-09-13 21:25:04 +08:00
from repositories.customer_relation import CustomerRelationRepo
2026-09-13 16:19:24 +08:00
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)
2026-09-13 21:25:04 +08:00
def _infer_chat_intent(query: str) -> str | None:
"""从自然语言问题推断投顾意图;无法确定时保留通用问答。"""
if any(word in query for word in ("调仓", "再平衡", "组合偏离")):
return "rebalance"
if any(word in query for word in ("沟通话术", "怎么和客户说", "解释给客户")):
return "dialogue-script"
if any(word in query for word in ("基金分析", "分析这只基金", "分析产品")):
return "fund_analysis"
if any(word in query for word in ("推荐", "产品建议", "买什么基金", "适合的基金")):
return AGENT_INTENT_RECOMMEND
return None
async def _resolve_customer_from_query(db, *, advisor_id: int, query: str) -> tuple[int | None, str | None]:
"""解析问题中的客户编号或姓名,并限制在当前投顾客户范围内。"""
number_match = re.search(r"(?:客户|用户)\s*[#编号号:]?\s*(\d+)", query)
relation_repo = CustomerRelationRepo(db)
if number_match:
customer_id = int(number_match.group(1))
relation = await relation_repo.get_active_relation(
customer_id=customer_id,
advisor_id=advisor_id,
)
if relation is None:
return None, "问题中的客户不在当前投顾的授权范围内"
return customer_id, None
rows = await relation_repo.list_customer_rows(advisor_id=advisor_id, limit=100)
matched = {
int(account.id)
for _relation, account, _profile in rows
if account.real_name and account.real_name in query
}
if len(matched) == 1:
return next(iter(matched)), None
if len(matched) > 1:
return None, "问题中的客户姓名无法唯一确定,请补充客户编号"
if "客户" in query or "用户" in query:
return None, "请在问题中补充客户编号或客户姓名"
return None, None
2026-09-13 16:19:24 +08:00
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,
2026-09-13 21:25:04 +08:00
body: AdvisorChatReq,
2026-09-13 16:19:24 +08:00
user: SysUser = Depends(audited_advisor),
db: AsyncSession = Depends(get_db),
):
trace_id = _trace_id(request)
2026-09-13 21:25:04 +08:00
chat_request = body
customer_id = chat_request.customer_id
inferred_intent = chat_request.intent or _infer_chat_intent(chat_request.query)
# 请求体只传问题时,从问题中解析客户;解析结果仍必须经过投顾关系授权校验。
if customer_id is None:
customer_id, resolve_error = await _resolve_customer_from_query(
db,
advisor_id=user.id,
query=chat_request.query,
2026-09-13 16:19:24 +08:00
)
2026-09-13 21:25:04 +08:00
if resolve_error:
payload = agent_failure(
ERR_CODE_FORBIDDEN_CUSTOMER,
resolve_error,
trace_id=trace_id,
)
async def resolve_error_events():
yield f"data: {json.dumps({'type': SSE_EVENT_TYPE_ERROR, **payload}, ensure_ascii=False)}\n\n"
return StreamingResponse(
resolve_error_events(),
media_type="text/event-stream",
headers={"X-Trace-Id": trace_id},
)
else:
payload = None
# 不带客户编号时只提供通用基金问答,不读取客户画像,也不生成个性化草稿。
if customer_id is None:
if inferred_intent in {
AGENT_INTENT_RECOMMEND,
"rebalance",
"fund_analysis",
"dialogue-script",
}:
if payload is None:
payload = agent_failure(
ERR_CODE_FORBIDDEN_CUSTOMER,
"个性化投顾分析需要在问题中明确客户编号或姓名",
trace_id=trace_id,
)
else:
runtime = _advisor_runtime(request)
llm_client = getattr(runtime, "llm_client", None)
if llm_client is None:
answer = "已收到问题。当前未配置通用投顾模型,请选择客户后使用个性化分析,或联系管理员配置 Agent 服务。"
else:
answer = await generate_text(
llm_client,
system_prompt="你是基金投顾助手,只回答通用基金知识和产品分析问题,不读取或推断任何客户信息,不承诺收益,不代客交易。",
user_prompt=chat_request.query,
fallback=lambda: "当前模型暂时不可用,请稍后重试。",
timeout=5.0,
)
async def events():
for event in (
{"type": SSE_EVENT_TYPE_META, "intent": "general_question"},
{"type": SSE_EVENT_TYPE_TEXT, "content": answer},
{"type": SSE_EVENT_TYPE_DONE},
):
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},
)
2026-09-13 16:19:24 +08:00
else:
relation = await ensure_customer_access(
db, advisor_id=user.id, customer_id=int(customer_id)
)
2026-09-13 21:25:04 +08:00
if inferred_intent == AGENT_INTENT_RECOMMEND:
2026-09-13 16:19:24 +08:00
memories = await _recall_advisor_memories(
request,
customer_id=int(customer_id),
2026-09-13 21:25:04 +08:00
query=chat_request.query,
2026-09-13 16:19:24 +08:00
)
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))