Files
Mutual_Fund/service/advisor/event_consumer.py

205 lines
7.6 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Redis 事件消费者(投顾域事件:rebalance_draft_created / profile_update)。
双写原则(common_const §9):Agent 双写 event_log + Redis Pub/Sub;Pub/Sub 仅实时通知,
消费以 event_id 幂等,补拉靠 event_log 扫描兜底。process_event 是幂等入口,实时订阅与
补拉两条路径共用。
"""
from __future__ import annotations
import asyncio
import json
import logging
from common_const import (
ADVISOR_EVENTS,
EVENT_PROFILE_UPDATE,
EVENT_REBALANCE_DRAFT_CREATED,
EVENT_STATUS_CONSUMED,
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 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)
if existing is not None:
return # 已存在,幂等跳过
todo = AdvisorTodo(
todo_type=TODO_TYPE_NEW_REBALANCE_DRAFT,
customer_id=customer_id,
advisor_id=int(advisor_id),
source=TODO_SOURCE_AGENT_EVENT,
biz_id=str(draft_id) if draft_id else None,
)
await repo.add(todo) # add 内 commit + refresh
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
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, redis=None
) -> None:
"""处理单条事件(幂等入口,实时订阅与补拉共用)。"""
# 幂等兜底:event_log 已消费则跳过(实时订阅与补拉并发时防重)
if event_id:
existing = await EventLogRepo(db).get_by_event_id(event_id)
if existing is not None and existing.status == EVENT_STATUS_CONSUMED:
return
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, redis=redis)
if event_id:
await EventLogRepo(db).mark_consumed(event_id)
def parse_message(raw: str) -> tuple[str | None, str | None, dict] | None:
"""解析 Redis Pub/Sub 消息为 (event_id, event_name, payload)。
容忍两种结构:带内层 payload 的标准结构;无内层 payload 时把外层整体当 payload。
"""
try:
msg = json.loads(raw)
except (ValueError, TypeError):
return None
if not isinstance(msg, dict):
return None
event_name = msg.get("event_name")
event_id = msg.get("event_id") or msg.get("id")
payload = msg.get("payload")
if not isinstance(payload, dict):
payload = msg
return event_id, event_name, payload
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 {},
redis=redis,
)
return len(events)
class EventConsumer:
"""Redis 订阅循环(后台任务,dev 默认关闭,避免 --reload 重复订阅)。"""
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
pubsub = self.redis.pubsub()
await pubsub.subscribe(*ADVISOR_EVENTS)
try:
async for message in pubsub.listen():
if message.get("type") != "message":
continue
parsed = parse_message(message.get("data"))
if parsed is None:
continue
event_id, event_name, payload = parsed
if event_name not in ADVISOR_EVENTS:
continue
# 每事件开独立 session,避免长事务占用连接
async with get_session_factory()() as session:
try:
await process_event(
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)
finally:
await pubsub.unsubscribe(*ADVISOR_EVENTS)
await pubsub.aclose()
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:
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