408 lines
14 KiB
Python
408 lines
14 KiB
Python
"""投顾 Agent HTTP 契约入口与本地业务编排。"""
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
from typing import Literal
|
|
|
|
from fastapi import APIRouter, BackgroundTasks, Depends, Query, Request
|
|
from fastapi.responses import StreamingResponse
|
|
from pydantic import ValidationError
|
|
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 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)
|
|
|
|
|
|
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: dict,
|
|
user: SysUser = Depends(audited_advisor),
|
|
db: AsyncSession = Depends(get_db),
|
|
):
|
|
trace_id = _trace_id(request)
|
|
try:
|
|
chat_request = AdvisorChatReq.model_validate(body)
|
|
except ValidationError:
|
|
payload = agent_failure(
|
|
ERR_CODE_FORBIDDEN_CUSTOMER,
|
|
"对话请求缺少有效客户范围或参数",
|
|
trace_id=trace_id,
|
|
)
|
|
else:
|
|
customer_id = chat_request.customer_id
|
|
relation = await ensure_customer_access(
|
|
db, advisor_id=user.id, customer_id=int(customer_id)
|
|
)
|
|
if chat_request.intent == AGENT_INTENT_RECOMMEND:
|
|
memories = await _recall_advisor_memories(
|
|
request,
|
|
customer_id=int(customer_id),
|
|
query=chat_request.query or "",
|
|
)
|
|
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))
|