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