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
+1 -1
View File
@@ -169,7 +169,7 @@ class AdvisorAgentClient:
) -> dict:
return await self._request(
"POST", f"/draft/{draft_id}/operate", auth_header=auth_header,
trace_id=trace_id, json={"action": action},
trace_id=trace_id, json={"operation": action},
)
async def rebalance_run(
+4 -2
View File
@@ -29,6 +29,7 @@ async def list_ledger(
user: SysUser,
*,
action: str | None = None,
customer_id: int | None = None,
keyword: str | None = None,
start: datetime | None = None,
end: datetime | None = None,
@@ -37,7 +38,7 @@ async def list_ledger(
) -> dict:
repo = AuditLogRepo(db)
items = await repo.list_by_advisor(
user_id=user.id, action=action, keyword=keyword, start=start, end=end,
user_id=user.id, action=action, customer_id=customer_id, keyword=keyword, start=start, end=end,
limit=page_size, offset=(page - 1) * page_size,
)
total = await repo.count_by_advisor(
@@ -62,6 +63,7 @@ async def export_ledger(
user: SysUser,
*,
action: str | None = None,
customer_id: int | None = None,
keyword: str | None = None,
start: datetime | None = None,
end: datetime | None = None,
@@ -69,7 +71,7 @@ async def export_ledger(
"""导出本人审计台账为 CSV 文本(全量,不限分页大小,上限 10000 条)。"""
repo = AuditLogRepo(db)
items = await repo.list_by_advisor(
user_id=user.id, action=action, keyword=keyword, start=start, end=end,
user_id=user.id, action=action, customer_id=customer_id, keyword=keyword, start=start, end=end,
limit=10000, offset=0,
)
return _to_csv([_audit_item(a) for a in items])
+4 -1
View File
@@ -138,7 +138,10 @@ async def get_customer_reports(
await ensure_customer_owned(db, user.id, customer_id)
repo = AdvisorReportRepo(db)
items = await repo.list_by_customer(
customer_id, limit=page_size, offset=(page - 1) * page_size
customer_id,
advisor_id=user.id,
limit=page_size,
offset=(page - 1) * page_size,
)
return {
"total": await repo.count_by_advisor(advisor_id=user.id, customer_id=customer_id),
+18 -1
View File
@@ -21,6 +21,17 @@ from repositories.advisor_visit_record import AdvisorVisitRecordRepo
from repositories.customer_relation import CustomerRelationRepo
def _todo_summary(todo) -> dict:
return {
"id": todo.id,
"todo_type": todo.todo_type,
"customer_id": todo.customer_id,
"priority": todo.priority,
"status": todo.status,
"due_at": todo.due_at.isoformat() if todo.due_at else None,
}
async def get_dashboard(db: AsyncSession, user: SysUser) -> dict:
relation_repo = CustomerRelationRepo(db)
todo_repo = AdvisorTodoRepo(db)
@@ -28,6 +39,9 @@ async def get_dashboard(db: AsyncSession, user: SysUser) -> dict:
visit_repo = AdvisorVisitRecordRepo(db)
pending = await todo_repo.count_by_advisor(advisor_id=user.id, status=TODO_STATUS_PENDING)
pending_items = await todo_repo.list_by_advisor(
advisor_id=user.id, status=TODO_STATUS_PENDING, limit=10, offset=0
)
total = await relation_repo.count_customer_rows(advisor_id=user.id)
signed = await relation_repo.count_customer_rows(
advisor_id=user.id, status=CUSTOMER_REL_STATUS_SIGNED
@@ -54,7 +68,10 @@ async def get_dashboard(db: AsyncSession, user: SysUser) -> dict:
unknown += cnt
return {
"todos": {"pending": pending},
"todos": {
"pending": pending,
"items": [_todo_summary(todo) for todo in pending_items],
},
"overview": {
"total_customers": total,
"signed_customers": signed,
+46 -11
View File
@@ -29,6 +29,7 @@ from repositories.advisor_report import AdvisorReportRepo
from repositories.product import ProductRepo
from repositories.risk_assessment import CustomerProfileRepo
from repositories.sensitive_word import SensitiveWordRepo
from repositories.sys_message import SysMessageRepo
from schemas.advisor import DraftSaveReq, RebalanceRunReq, TalkScriptReq
from service.advisor.agent_client import get_agent_client
from service.advisor.permissions import ensure_customer_owned, require_owned_relation
@@ -60,20 +61,30 @@ def _build_report_from_detail(detail: dict, advisor_id: int) -> AdvisorReport:
)
def _extract_buy_codes(suggestions: Any) -> list[str]:
"""从建议清单提取「低配申购」产品代码(发送终审适当性校验对象)。
def _extract_buy_codes(detail: Any) -> list[str]:
"""从 Agent 草稿详情提取发送终审需要校验的产品代码。
假设:结构为 dict,申购侧键见 _BUY_KEYS;项为 {product_code} 或 {code}/{fund_code}。
解析失败返回空列表(视为无结构化产品建议,适当性空过,不误拦)。
兼容旧的顶层 ``suggestions``,以及 Agent 当前的 ``structured_data``:调仓只取
``buy``,推荐取 ``items``。解析失败返回空列表,交由其它终审规则继续处理。
"""
if not isinstance(suggestions, dict):
if not isinstance(detail, dict):
return []
structured_data = detail.get("structured_data")
data = structured_data if isinstance(structured_data, dict) else detail
suggestions = data.get("suggestions")
if isinstance(suggestions, dict):
data = suggestions
items: list | None = None
for key in _BUY_KEYS:
value = suggestions.get(key)
value = data.get(key)
if isinstance(value, list):
items = value
break
if items is None and isinstance(data.get("items"), list):
items = data["items"]
codes: list[str] = []
for it in items or []:
if isinstance(it, dict):
@@ -104,7 +115,7 @@ async def _resolve_product_risks(
result = await get_agent_client().draft_detail(
draft_id, auth_header=auth_header, trace_id=trace_id
)
codes = _extract_buy_codes((result["data"] or {}).get("suggestions"))
codes = _extract_buy_codes(result["data"] or {})
product_repo = ProductRepo(db)
risks: list[str | None] = []
for code in codes:
@@ -214,7 +225,11 @@ async def save_draft(
# 1) 先取草稿确认归属(避免对无权限草稿执行写操作),并拿到 intent/customer_id
detail = await get_draft(db, user, auth_header=auth_header, trace_id=trace_id, draft_id=draft_id)
# 2) 调 Agent 保存(Agent 重新适当性校验,违规 40020;缺免责仅告警不阻断)
payload = {"title": req.title, "content": req.content, "suggestions": req.suggestions}
payload = {
"title": req.title,
"content": req.content,
"structured_data": req.suggestions,
}
result = await get_agent_client().draft_save(
draft_id, payload, auth_header=auth_header, trace_id=trace_id
)
@@ -262,6 +277,9 @@ async def send_draft(
elif report.advisor_id != user.id:
raise ForbiddenError("无权操作该客户数据")
if report.send_status == REPORT_SEND_STATUS_DISCARDED:
raise ParamError("已废弃的报告不可发送")
# 幂等:已发送直接返回(重复点击不重复发站内信)
if report.send_status == REPORT_SEND_STATUS_SENT:
return {**_sent_payload(report), "duplicated": True}
@@ -283,6 +301,19 @@ async def send_draft(
sensitive_words=sensitive_words,
)
# 提交结果不确定后重试时,先复用已经写入的站内信,避免重复触达客户。
existing_message = await SysMessageRepo(db).get_by_biz_id(
report.report_id, user_id=report.customer_id
)
if existing_message is not None:
report.send_status = REPORT_SEND_STATUS_SENT
report.send_time = report.send_time or existing_message.create_time
report.send_by = report.send_by or user.id
report.msg_id = existing_message.id
db.add(report)
await db.commit()
return {**_sent_payload(report), "duplicated": True}
# 写站内信 + 报告置 sent,同事务(失败可重试,避免假送达)
msg_type = INTENT_TO_MSG_TYPE.get(report.intent, MSG_TYPE_RECOMMEND)
message = SysMessage(
@@ -297,9 +328,13 @@ async def send_draft(
report.send_by = user.id
db.add(message)
db.add(report) # 已跟踪对象时无副作用,新对象时入 session
await db.flush() # 生成 message.id,供回填 msg_id
report.msg_id = message.id
await db.commit()
try:
await db.flush() # 生成 message.id,供回填 msg_id
report.msg_id = message.id
await db.commit()
except Exception:
await db.rollback()
raise
return _sent_payload(report)
+75 -21
View File
@@ -18,20 +18,31 @@ from common_const import (
TODO_SOURCE_AGENT_EVENT,
TODO_TYPE_NEW_REBALANCE_DRAFT,
)
from config.database.redis import client as redis_client
from config.settings import settings
from model.advisor_todo import AdvisorTodo
from repositories.advisor_todo import AdvisorTodoRepo
from repositories.event_log import EventLogRepo
from service.memory.profile import CustomerProfileMemory
logger = logging.getLogger("service.advisor.event_consumer")
def _payload_log_summary(payload: dict) -> str:
"""日志只记录字段名,避免事件扩展后误打客户敏感值。"""
return "{" + ", ".join(sorted(str(key) for key in payload)) + "}"
async def handle_rebalance_created(db, payload: dict) -> None:
"""消费 rebalance_draft_created → 生成待办(uk_todo 三元组去重,幂等)。"""
advisor_id = payload.get("advisor_id")
customer_id = payload.get("customer_id")
draft_id = payload.get("draft_id")
if not advisor_id:
logger.warning("rebalance_draft_created 缺少 advisor_id,忽略: %s", payload)
if not advisor_id or customer_id is None or not draft_id:
logger.warning(
"rebalance_draft_created 缺少 advisor_id/customer_id/draft_id,忽略字段: %s",
_payload_log_summary(payload),
)
return
repo = AdvisorTodoRepo(db)
existing = await repo.get_by_unique(TODO_TYPE_NEW_REBALANCE_DRAFT, customer_id, draft_id)
@@ -47,16 +58,23 @@ async def handle_rebalance_created(db, payload: dict) -> None:
await repo.add(todo) # add 内 commit + refresh
async def handle_profile_update(db, payload: dict) -> None:
"""消费 profile_update → 失效该客户 Redis 画像缓存(不落库,画像只读)。
async def handle_profile_update(db, payload: dict, *, redis=None) -> None:
"""消费 profile_update,失效该客户 Redis 画像缓存。"""
customer_id = payload.get("customer_id")
if customer_id is None:
logger.warning("profile_update 缺少 customer_id,忽略字段: %s", _payload_log_summary(payload))
return
注:V1.0 工作台尚未建设 profile:{customer_id} 画像热缓存读取,此处仅占位;后续
360 读画像接入缓存后生效。事件仍需标记已消费。
"""
return
warnings = await CustomerProfileMemory(redis=redis or redis_client()).invalidate(
int(customer_id)
)
for warning in warnings:
logger.warning("profile_update cache invalidation warning: %s", warning)
async def process_event(db, *, event_id: str | None, event_name: str | None, payload: dict) -> None:
async def process_event(
db, *, event_id: str | None, event_name: str | None, payload: dict, redis=None
) -> None:
"""处理单条事件(幂等入口,实时订阅与补拉共用)。"""
# 幂等兜底:event_log 已消费则跳过(实时订阅与补拉并发时防重)
if event_id:
@@ -67,7 +85,7 @@ async def process_event(db, *, event_id: str | None, event_name: str | None, pay
if event_name == EVENT_REBALANCE_DRAFT_CREATED:
await handle_rebalance_created(db, payload)
elif event_name == EVENT_PROFILE_UPDATE:
await handle_profile_update(db, payload)
await handle_profile_update(db, payload, redis=redis)
if event_id:
await EventLogRepo(db).mark_consumed(event_id)
@@ -92,12 +110,16 @@ def parse_message(raw: str) -> tuple[str | None, str | None, dict] | None:
return event_id, event_name, payload
async def pull_pending_events(db) -> int:
async def pull_pending_events(db, *, redis=None) -> int:
"""补拉 event_log 中未消费的投顾域事件(Pub/Sub 丢消息兜底),返回处理条数。"""
events = await EventLogRepo(db).list_pending(ADVISOR_EVENTS, limit=100)
for ev in events:
await process_event(
db, event_id=ev.event_id, event_name=ev.event_name, payload=ev.payload or {}
db,
event_id=ev.event_id,
event_name=ev.event_name,
payload=ev.payload or {},
redis=redis,
)
return len(events)
@@ -105,9 +127,34 @@ async def pull_pending_events(db) -> int:
class EventConsumer:
"""Redis 订阅循环(后台任务,dev 默认关闭,避免 --reload 重复订阅)。"""
def __init__(self, redis):
def __init__(self, redis, *, retry_interval_sec: float | None = None):
self.redis = redis
self._task: asyncio.Task | None = None
self._retry_task: asyncio.Task | None = None
self.retry_interval_sec = (
settings.advisor.event_retry_interval_sec
if retry_interval_sec is None
else retry_interval_sec
)
async def retry_pending_once(self) -> int:
"""补拉并处理一批尚未消费事件;异常保留 pending,供下一轮重试。"""
from config.database.mysql import get_session_factory
async with get_session_factory()() as session:
return await pull_pending_events(session, redis=self.redis)
async def _retry_loop(self) -> None:
while True:
await asyncio.sleep(self.retry_interval_sec)
try:
count = await self.retry_pending_once()
if count:
logger.info("advisor event pending retry processed: %s", count)
except asyncio.CancelledError:
raise
except Exception:
logger.exception("advisor event pending retry failed")
async def _run(self) -> None:
from config.database.mysql import get_session_factory
@@ -128,7 +175,11 @@ class EventConsumer:
async with get_session_factory()() as session:
try:
await process_event(
session, event_id=event_id, event_name=event_name, payload=payload
session,
event_id=event_id,
event_name=event_name,
payload=payload,
redis=self.redis,
)
except Exception:
logger.exception("process event failed: %s %s", event_name, event_id)
@@ -139,12 +190,15 @@ class EventConsumer:
def start(self) -> None:
if self._task is None:
self._task = asyncio.create_task(self._run())
self._retry_task = asyncio.create_task(self._retry_loop())
async def stop(self) -> None:
if self._task is not None:
self._task.cancel()
try:
await self._task
except asyncio.CancelledError:
pass
self._task = None
for task in (self._task, self._retry_task):
if task is not None:
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
self._task = None
self._retry_task = None
+44 -2
View File
@@ -11,9 +11,10 @@ from model.advisor_visit_record import AdvisorVisitRecord
from model.sys_user import SysUser
from repositories.advisor_visit_record import AdvisorVisitRecordRepo
from repositories.sys_message import SysMessageRepo
from schemas.advisor import VisitCreateReq
from schemas.advisor import VisitCreateReq, VisitUpdateReq
from service.advisor.permissions import ensure_customer_owned
from utils.exceptions import ParamError
from service.advisor.audit_writer import write_audit
from utils.exceptions import NotFoundError, ParamError
# 合规话术库(内置,标准化投教/市场解读/调仓沟通)。
# 假设:V1.0 内置常量,后续可迁移 sys_config 运营化;内容不含承诺收益等敏感词。
@@ -63,9 +64,50 @@ async def create_visit(db: AsyncSession, user: SysUser, req: VisitCreateReq) ->
audio_url=req.audio_url,
)
record = await AdvisorVisitRecordRepo(db).add(record)
await write_audit(
db,
user_id=user.id,
username=user.username,
module="advisor",
action="visit_create",
target=str(record.id),
detail={"customer_id": req.customer_id},
)
return {"visit_id": record.id}
async def get_visit(db: AsyncSession, user: SysUser, visit_id: int) -> dict:
record = await AdvisorVisitRecordRepo(db).get_by_advisor(visit_id, user.id)
if record is None:
raise NotFoundError("回访记录不存在")
return _visit_item(record)
async def update_visit(
db: AsyncSession, user: SysUser, visit_id: int, req: VisitUpdateReq
) -> dict:
changes = req.model_dump(exclude_unset=True)
if not changes:
raise ParamError("至少提供一项回访记录修改内容")
record = await AdvisorVisitRecordRepo(db).get_by_advisor(visit_id, user.id)
if record is None:
raise NotFoundError("回访记录不存在")
for field, value in changes.items():
setattr(record, field, value)
await db.commit()
await db.refresh(record)
await write_audit(
db,
user_id=user.id,
username=user.username,
module="advisor",
action="visit_update",
target=str(record.id),
detail={"customer_id": record.customer_id, "fields": list(changes)},
)
return _visit_item(record)
def list_talk_templates() -> list[dict]:
"""合规话术库(内置,投顾参考;不自动发送)。"""
return TALK_TEMPLATES