feat:新增投顾agent和nl2sqlagent
This commit is contained in:
@@ -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))
|
||||
Reference in New Issue
Block a user