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
+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