feat:新增投顾agent和nl2sqlagent

This commit is contained in:
2026-09-13 16:19:24 +08:00
parent c80c6acac0
commit 163192bf55
122 changed files with 7488 additions and 362 deletions
+407
View File
@@ -0,0 +1,407 @@
"""投顾 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))