"""投顾 Agent HTTP 契约入口与本地业务编排。""" from __future__ import annotations import json import re 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 from repositories.customer_relation import CustomerRelationRepo 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) 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 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, body: AdvisorChatReq, user: SysUser = Depends(audited_advisor), db: AsyncSession = Depends(get_db), ): trace_id = _trace_id(request) 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, ) 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}, ) else: relation = await ensure_customer_access( db, advisor_id=user.id, customer_id=int(customer_id) ) if inferred_intent == AGENT_INTENT_RECOMMEND: memories = await _recall_advisor_memories( request, customer_id=int(customer_id), query=chat_request.query, ) 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=10, 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))