WechatOnCloud/bridge/woc_bridge/messaging/batch_worker.py

500 lines
22 KiB
Python
Raw Normal View History

"""群发消息后台协程:串行复用 FlowOrchestrator + SendQueue。
设计原则
- 串行发送不并行避免破坏 UI 操作互斥
- 复用 orchestrator.send_text 的幂等 + 熔断 + 会话缓存
- 风控规避间隔随机化 + 内容个性化 + 单批200 + 日频次3 + 批间隔2 小时
三层幂等
- 批级client_request_idNone 时自动生成 batch_id
- 单条内部按 f"{batch_id}:{to_wxid}" 作为 client_request_id 传给 orchestrator
- orchestrator复用 IdemCacheTTL=300s
"""
from __future__ import annotations
import asyncio
import logging
import random
import time
import uuid
from datetime import datetime, timedelta
from typing import Optional, Any
from woc_bridge.models import BridgeError, BatchSendRequest, BatchSendStatus, BatchItemResult
logger = logging.getLogger("woc-bridge")
# 群发单条间隔抖动base=3 ± 1 → [2, 4]
# orchestrator.send_text 内部已通过 send_queue 串行化并施加 send_delay_ms默认 3000
# 本 worker 在其之上额外施加 2-4s 随机抖动,使得相邻两条群发的总间隔呈现随机性,
# 降低被风控识别为机器行为的概率。当 worker 抖动 ≥ send_queue 延时时,
# send_queue 的延时被吸收worker 主导节奏)。
_BATCH_JITTER_MIN_SEC = 2.0
_BATCH_JITTER_MAX_SEC = 4.0
# 单批发送耗时样本上限(用于均值估算剩余时间)
_RECENT_SEND_SAMPLE_MAX = 10
class BatchSendWorker:
"""群发消息后台协程。"""
def __init__(
self,
orchestrator, # FlowOrchestrator
send_queue, # SendQueue
db_reader, # DbReader用于 nickname 查询)
config, # BridgeConfig含 batch_max_per_batch / batch_daily_limit / batch_min_interval_sec
) -> None:
self._orchestrator = orchestrator
self._send_queue = send_queue
self._db_reader = db_reader
self._config = config
# 任务存储batch_id -> 任务运行态
# 运行态字段req, status, items, task_obj, start_time, end_time, cancel_requested
self._tasks: dict[str, dict] = {}
self._active_tasks: dict[str, asyncio.Task] = {}
# 风控状态
self._daily_count: int = 0
self._daily_reset_time: float = self._next_midnight_timestamp() # 下一次重置时间(次日 0 点)
self._last_batch_end_time: float = 0.0 # 上一批次结束时间monotonic
# 统计:最近 N 条单批发送耗时(秒)
self._recent_send_sec: list[float] = []
# 锁:保护 _tasks / _active_tasks / 风控状态
self._lock = asyncio.Lock()
# ------------------------------------------------------------------
# 公共 API
# ------------------------------------------------------------------
async def submit_batch(self, req: BatchSendRequest) -> str:
"""提交群发任务,返回 batch_id。
Raises:
BridgeError: 超过单批上限 / 日频次 / 间隔限制 / 重复 batch_id
"""
async with self._lock:
# 1. 风控校验(可能抛 RATE_LIMITED
self._reset_daily_count_if_needed()
self._check_risk_control(req)
# 2. 生成 batch_id批级幂等
batch_id = req.client_request_id or f"batch_{uuid.uuid4().hex[:16]}"
if batch_id in self._tasks:
raise BridgeError(
code="INVALID_PARAMS",
message=f"batch_id={batch_id} 已存在(批级幂等命中)",
http_status=409,
)
# 3. targets 去重(保序)
seen = set()
unique_targets: list[str] = []
for wxid in req.targets:
if wxid not in seen:
seen.add(wxid)
unique_targets.append(wxid)
# 4. 初始化任务运行态
now_mono = time.monotonic()
items = [
BatchItemResult(to_wxid=wxid, status="pending")
for wxid in unique_targets
]
self._tasks[batch_id] = {
"req": req,
"status": "pending",
"items": items,
"task_obj": None,
"start_time": now_mono,
"end_time": None,
"cancel_requested": False,
"consecutive_fail": 0,
}
# P0 修复:风控状态在 submit_batch 持锁时立即更新,防止并发提交绕过。
# 原实现在 _run_batch 的 finally 块(批次结束)才更新 _daily_count 和
# _last_batch_end_timesubmit_batch 释放锁后 _run_batch 才运行,
# 用户可在 1 秒内并发提交 3 批 × 200 = 600 条完全绕过风控。
# 现在提交时立即递增 _daily_count 并更新 _last_batch_end_time
# _run_batch 的 finally 块不再重复递增cancelled 也不回退,因配额已消耗)。
self._daily_count += 1
self._last_batch_end_time = now_mono
# 5. 启动后台 _run_batch 协程
task = asyncio.create_task(self._run_batch(batch_id, req, unique_targets))
self._active_tasks[batch_id] = task
self._tasks[batch_id]["task_obj"] = task
logger.info(
"[batch_worker] 提交群发 batch_id=%s targets=%d content_len=%d",
batch_id, len(unique_targets), len(req.content),
)
return batch_id
async def _run_batch(
self,
batch_id: str,
req: BatchSendRequest,
targets: list[str],
) -> None:
"""后台串行发送整批。
P0 修复整个 for 循环包裹在 try/except CancelledError + try/finally
确保 cancel/异常时收尾代码_daily_count++ / _active_tasks 清理 /
_last_batch_end_time 更新 / end_time 设置必定执行避免任务永久卡在 running
cancel 中断时剩余 pending 项标记为 skipped
"""
state = self._tasks[batch_id]
state["status"] = "running"
items: list[BatchItemResult] = state["items"]
total = len(items)
abort_reason: Optional[str] = None
try:
for idx, item in enumerate(items):
# 取消检查
if state["cancel_requested"]:
# 剩余全部标记 skipped
for remain in items[idx:]:
if remain.status == "pending":
remain.status = "skipped"
state["status"] = "cancelled"
logger.info("[batch_worker] batch_id=%s 已取消(已处理 %d/%d", batch_id, idx, total)
break
# P0 修复记录当前正在处理的条目索引cancel 中断时该条目
# 的 send_text 可能已在 send_queue worker 中执行future 取消不
# 中断 worker消息可能已实际发送不应标记 skipped。
state["_current_idx"] = idx
to_wxid = item.to_wxid
t_send_start = time.monotonic()
try:
# a. 查 nickname 替换 {nickname}
content = await self._personalize_content(req.content, to_wxid)
# b. 调 orchestrator.send_text单条幂等 ID = f"{batch_id}:{to_wxid}"
result = await self._orchestrator.send_text(
to_wxid=to_wxid,
content=content,
client_request_id=f"{batch_id}:{to_wxid}",
)
# c. 记录结果
if result.success and result.verified:
item.status = "success"
item.verified = True
state["consecutive_fail"] = 0
elif result.success and not result.verified and result.skipped:
# 幂等命中跳过:算成功(已发送过)
item.status = "skipped"
item.verified = result.verified
state["consecutive_fail"] = 0
elif result.success and not result.verified:
# 发送成功但校验失败:仍计为 success消息已发
item.status = "success"
item.verified = False
state["consecutive_fail"] = 0
else:
item.status = "failed"
item.error_code = result.error_code or "SEND_FAILED"
item.error_message = result.error or "发送失败"
state["consecutive_fail"] += 1
except BridgeError as exc:
item.status = "failed"
item.error_code = exc.code
item.error_message = exc.message
state["consecutive_fail"] += 1
logger.warning(
"[batch_worker] batch_id=%s to=%s BridgeError code=%s",
batch_id, to_wxid, exc.code,
)
except Exception as exc: # 兜底
item.status = "failed"
item.error_code = "SEND_FAILED"
item.error_message = str(exc)
state["consecutive_fail"] += 1
logger.exception(
"[batch_worker] batch_id=%s to=%s 异常",
batch_id, to_wxid,
)
# 记录单条耗时样本(用于剩余时间估算)
send_sec = time.monotonic() - t_send_start
self._recent_send_sec.append(send_sec)
if len(self._recent_send_sec) > _RECENT_SEND_SAMPLE_MAX:
self._recent_send_sec = self._recent_send_sec[-_RECENT_SEND_SAMPLE_MAX:]
# d. abort_on_consecutive_fail 触发:整批 failed剩余 skipped
if state["consecutive_fail"] >= req.abort_on_consecutive_fail:
abort_reason = (
f"连续失败 {state['consecutive_fail']} 次,"
f"达到阈值 {req.abort_on_consecutive_fail},整批中止"
)
logger.warning(
"[batch_worker] batch_id=%s 触发中止:%s", batch_id, abort_reason
)
for remain in items[idx + 1:]:
if remain.status == "pending":
remain.status = "skipped"
state["status"] = "failed"
state["abort_reason"] = abort_reason
break
# e. 间隔抖动(非最后一条)
if idx < total - 1 and not state["cancel_requested"]:
jitter = random.uniform(_BATCH_JITTER_MIN_SEC, _BATCH_JITTER_MAX_SEC)
await asyncio.sleep(jitter)
else:
# for 循环正常结束(未 break / 未 cancel
if state["status"] not in ("cancelled", "failed"):
state["status"] = "completed"
except asyncio.CancelledError:
# P0 修复cancel_batch 调用 task.cancel() 在 await 点抛 CancelledError
# 剩余 pending 项标记 skipped状态置 cancelled。
# 不 re-raise 以确保 finally 收尾代码正常执行且 cancel_batch 的 await task 正常返回。
#
# P0 修复cancel 误标 skipped当前正在发送的条目_current_idx
# 的 send_text 可能已在 send_queue worker 中实际执行future 取消不
# 中断 worker消息可能已发出。将该条目标为 "unknown" 而非 "skipped"
# 避免数据不一致(用户看到 skipped 但接收方收到了消息)。
current_idx = state.pop("_current_idx", None)
for i, item in enumerate(items):
if item.status == "pending":
if current_idx is not None and i == current_idx:
# 正在发送中被中断,状态不确定
item.status = "unknown"
item.error_message = "发送被取消,实际发送状态未知"
else:
item.status = "skipped"
if state["status"] not in ("completed", "failed"):
state["status"] = "cancelled"
state["abort_reason"] = "用户取消"
logger.info(
"[batch_worker] batch_id=%s 被 cancel 中断(剩余项标 skipped当前项标 unknown",
batch_id,
)
finally:
# P0 修复:收尾代码移入 finally确保 cancel/异常/正常结束都执行
# 注意_daily_count 和 _last_batch_end_time 已在 submit_batch 中更新,
# 此处不再重复递增(避免双重计数)。
state["end_time"] = time.monotonic()
async with self._lock:
# _last_batch_end_time 更新为实际结束时间(更精确)
self._last_batch_end_time = state["end_time"]
self._active_tasks.pop(batch_id, None)
logger.info(
"[batch_worker] batch_id=%s 结束 status=%s success=%d failed=%d skipped=%d 耗时=%.1fs",
batch_id, state["status"],
sum(1 for it in items if it.status == "success"),
sum(1 for it in items if it.status == "failed"),
sum(1 for it in items if it.status == "skipped"),
state["end_time"] - state["start_time"],
)
async def get_status(self, batch_id: str) -> Optional[BatchSendStatus]:
"""查询任务状态。"""
async with self._lock:
state = self._tasks.get(batch_id)
if state is None:
return None
return self._build_status(batch_id, state)
async def cancel_batch(self, batch_id: str) -> bool:
"""取消群发任务。
Returns:
True=已标记取消False=任务不存在或已结束
"""
async with self._lock:
state = self._tasks.get(batch_id)
if state is None:
return False
if state["status"] in ("completed", "failed", "cancelled"):
return False
state["cancel_requested"] = True
task = self._active_tasks.get(batch_id)
# 在锁外 await避免阻塞其他查询
if task is not None and not task.done():
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
except Exception:
pass
return True
async def list_batches(self, limit: int = 20) -> list[BatchSendStatus]:
"""列出最近 N 条群发任务(按提交时间倒序)。"""
async with self._lock:
# 按 start_time 倒序
sorted_ids = sorted(
self._tasks.keys(),
key=lambda bid: self._tasks[bid]["start_time"],
reverse=True,
)[:limit]
return [
self._build_status(bid, self._tasks[bid])
for bid in sorted_ids
]
async def stop(self) -> None:
"""停止所有运行中任务lifespan finally 调用)。
取消所有 _active_tasks 中的 asyncio.Task等待最多 5s
"""
async with self._lock:
running = list(self._active_tasks.items())
for batch_id, task in running:
if not task.done():
task.cancel()
if not running:
return
# 等待所有任务落地(带总超时 5s
try:
await asyncio.wait_for(
asyncio.gather(
*[t for _, t in running if not t.done()],
return_exceptions=True,
),
timeout=5.0,
)
except asyncio.TimeoutError:
logger.warning("[batch_worker] stop() 等待任务落地超时 5s强制返回")
# 标记未结束任务为 cancelled
async with self._lock:
for batch_id, _ in running:
state = self._tasks.get(batch_id)
if state and state["status"] in ("pending", "running"):
state["status"] = "cancelled"
state["abort_reason"] = "bridge 关闭,强制取消"
# ------------------------------------------------------------------
# 风控与统计
# ------------------------------------------------------------------
def _check_risk_control(self, req: BatchSendRequest) -> None:
"""风控校验:单批上限 / 日频次 / 间隔限制。
Raises:
BridgeError: 违反风控规则RATE_LIMITED / INVALID_PARAMS
"""
max_per_batch = getattr(self._config, "batch_max_per_batch", 200)
daily_limit = getattr(self._config, "batch_daily_limit", 3)
min_interval = getattr(self._config, "batch_min_interval_sec", 7200)
# 1. targets 去重后 ≤ batch_max_per_batch
unique_count = len(set(req.targets))
if unique_count > max_per_batch:
raise BridgeError(
code="INVALID_PARAMS",
message=f"群发目标数 {unique_count} 超过单批上限 {max_per_batch}",
http_status=400,
)
# 2. 当日已群发 < batch_daily_limit否则 RATE_LIMITED + retry_after 到次日 0 点
if self._daily_count >= daily_limit:
retry_after = max(1, int(self._daily_reset_time - time.time()))
raise BridgeError(
code="RATE_LIMITED",
message=f"当日群发次数已达上限 {daily_limit},请次日重试",
details={"retry_after": retry_after},
)
# 3. 上一批次结束 ≥ batch_min_interval_sec否则 RATE_LIMITED + retry_after
if self._last_batch_end_time > 0:
elapsed = time.monotonic() - self._last_batch_end_time
if elapsed < min_interval:
retry_after = int(min_interval - elapsed) + 1
raise BridgeError(
code="RATE_LIMITED",
message=f"群发间隔不足,需等待 {retry_after}s 后再发起",
details={"retry_after": retry_after},
)
def _compute_estimated_remaining_sec(self, remaining: int) -> Optional[int]:
"""计算剩余时间remaining * avg_send_secavg 取最近 10 条均值)。
无样本时返回 None
"""
if not self._recent_send_sec or remaining <= 0:
return None
avg = sum(self._recent_send_sec) / len(self._recent_send_sec)
# 加上平均抖动
avg_jitter = (_BATCH_JITTER_MIN_SEC + _BATCH_JITTER_MAX_SEC) / 2
return int(remaining * (avg + avg_jitter))
def _reset_daily_count_if_needed(self) -> None:
"""跨日重置 _daily_count。"""
now = time.time()
if now >= self._daily_reset_time:
self._daily_count = 0
self._daily_reset_time = self._next_midnight_timestamp()
@staticmethod
def _next_midnight_timestamp() -> float:
"""返回次日 0 点(本地时区)的 Unix 时间戳。"""
now = datetime.now()
next_midnight = (now + timedelta(days=1)).replace(
hour=0, minute=0, second=0, microsecond=0
)
return next_midnight.timestamp()
# ------------------------------------------------------------------
# 内部辅助
# ------------------------------------------------------------------
async def _personalize_content(self, content: str, to_wxid: str) -> str:
"""按目标联系人替换 {nickname} 占位符。
{nickname} 占位符时原样返回DB 不可用时回退到 wxid
"""
if "{nickname}" not in content:
return content
nickname = to_wxid # 回退值
if self._db_reader is not None:
try:
contact = await asyncio.to_thread(
self._db_reader.get_contact_detail, to_wxid
)
if contact:
nickname = (
contact.get("remark")
or contact.get("nickname")
or to_wxid
)
except Exception as exc:
logger.debug(
"[batch_worker] 查 nickname 失败 to=%s: %s,回退 wxid",
to_wxid, exc,
)
return content.replace("{nickname}", nickname)
def _build_status(self, batch_id: str, state: dict) -> BatchSendStatus:
"""从任务运行态构造 BatchSendStatus调用方持锁"""
items: list[BatchItemResult] = state["items"]
total = len(items)
success = sum(1 for it in items if it.status == "success")
failed = sum(1 for it in items if it.status == "failed")
skipped = sum(1 for it in items if it.status == "skipped")
processed = success + failed + skipped
progress = (processed / total) if total > 0 else 0.0
remaining = total - processed
return BatchSendStatus(
batch_id=batch_id,
status=state["status"],
total=total,
processed=processed,
success=success,
failed=failed,
skipped=skipped,
progress=round(progress, 4),
estimated_remaining_sec=self._compute_estimated_remaining_sec(remaining),
abort_reason=state.get("abort_reason"),
items=items,
)