316 lines
10 KiB
Python
316 lines
10 KiB
Python
"""
|
||
进程内 / PostgreSQL inbox·outbox 账本。
|
||
|
||
本文件职责:定义任务结构、Memory 实现、按配置选择后端工厂。
|
||
禁止:Redis 全局 Worker 单例锁;禁止在回调线程里等 LLM。
|
||
有 PG_DSN 且 MESSAGE_STORE_BACKEND=postgres|auto 时用表;否则内存(本机联调)。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import logging
|
||
import threading
|
||
import time
|
||
import uuid
|
||
from collections import OrderedDict
|
||
from dataclasses import dataclass, field
|
||
from enum import Enum
|
||
from typing import Any, Optional, Union
|
||
|
||
from agent.channel.wecom.models import InboundMessage
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
|
||
class TaskState(str, Enum):
|
||
PENDING = "pending"
|
||
CLAIMED = "claimed"
|
||
DONE = "done"
|
||
FAILED = "failed"
|
||
|
||
|
||
@dataclass
|
||
class InboxItem:
|
||
"""inbox 一条入站任务。"""
|
||
|
||
id: str
|
||
message: InboundMessage
|
||
state: TaskState = TaskState.PENDING
|
||
claimed_at: float = 0.0
|
||
attempts: int = 0
|
||
last_error: str = ""
|
||
|
||
|
||
@dataclass
|
||
class OutboxItem:
|
||
"""
|
||
outbox 一条出站任务。
|
||
|
||
touser:必须是入站 sender_id,禁止另换 ID。
|
||
发送成功 ≠ 业务成功;此处仅通道送达。
|
||
"""
|
||
|
||
id: str
|
||
touser: str
|
||
content: str
|
||
state: TaskState = TaskState.PENDING
|
||
claimed_at: float = 0.0
|
||
attempts: int = 0
|
||
last_error: str = ""
|
||
created_at: float = field(default_factory=time.time)
|
||
dedupe_key: str = ""
|
||
payload: dict[str, Any] = field(default_factory=dict)
|
||
|
||
|
||
class MemoryMessageStore:
|
||
"""
|
||
线程安全的内存 inbox/outbox。
|
||
|
||
认领:同一条任务同一时刻只被一个消费者持有;超时可回收(默认 120s)。
|
||
去重:相同 message_id / dedupe_key 不重复入队。
|
||
"""
|
||
|
||
def __init__(self, *, claim_timeout_sec: float = 120.0, dedupe_maxlen: int = 5000) -> None:
|
||
self._lock = threading.RLock()
|
||
self._inbox: OrderedDict[str, InboxItem] = OrderedDict()
|
||
self._outbox: OrderedDict[str, OutboxItem] = OrderedDict()
|
||
self._seen_msg: OrderedDict[str, float] = OrderedDict()
|
||
self._seen_out: OrderedDict[str, float] = OrderedDict()
|
||
self._claim_timeout = claim_timeout_sec
|
||
self._dedupe_maxlen = dedupe_maxlen
|
||
self._inbox_event = threading.Event()
|
||
self._outbox_event = threading.Event()
|
||
|
||
def enqueue_inbound(self, message: InboundMessage) -> tuple[bool, str]:
|
||
"""
|
||
写入 inbox。
|
||
|
||
返回:(是否新入队, 说明)。无 sender_id 拒绝;message_id 重复则跳过。
|
||
副作用:唤醒 inbox 等待线程。
|
||
"""
|
||
if not message.has_sender():
|
||
return False, "missing_sender_id"
|
||
with self._lock:
|
||
if message.message_id in self._seen_msg:
|
||
return False, "duplicate_message_id"
|
||
self._remember(self._seen_msg, message.message_id)
|
||
item = InboxItem(id=str(uuid.uuid4()), message=message)
|
||
self._inbox[item.id] = item
|
||
self._inbox_event.set()
|
||
logger.info(
|
||
"inbox 入队 id=%s msg_id=%s sender=%s type=%s",
|
||
item.id,
|
||
message.message_id,
|
||
message.sender_id,
|
||
message.msg_type,
|
||
)
|
||
return True, item.id
|
||
|
||
def claim_inbound(self) -> Optional[InboxItem]:
|
||
"""认领一条 pending/超时 inbox;无任务返回 None。"""
|
||
now = time.time()
|
||
with self._lock:
|
||
for item in self._inbox.values():
|
||
if item.state == TaskState.PENDING or (
|
||
item.state == TaskState.CLAIMED
|
||
and item.claimed_at
|
||
and now - item.claimed_at > self._claim_timeout
|
||
):
|
||
item.state = TaskState.CLAIMED
|
||
item.claimed_at = now
|
||
item.attempts += 1
|
||
return item
|
||
self._inbox_event.clear()
|
||
return None
|
||
|
||
def complete_inbound(self, item_id: str, *, error: str = "") -> None:
|
||
"""标记 inbox 完成或失败。"""
|
||
with self._lock:
|
||
item = self._inbox.get(item_id)
|
||
if not item:
|
||
return
|
||
if error:
|
||
item.state = TaskState.FAILED
|
||
item.last_error = error
|
||
else:
|
||
item.state = TaskState.DONE
|
||
|
||
def enqueue_outbound(
|
||
self,
|
||
*,
|
||
touser: str,
|
||
content: str,
|
||
dedupe_key: str = "",
|
||
payload: Optional[dict[str, Any]] = None,
|
||
) -> tuple[bool, str]:
|
||
"""写入 outbox;dedupe_key 重复则跳过。payload 可带企微模板卡。"""
|
||
if not touser or not content:
|
||
return False, "invalid_outbox"
|
||
with self._lock:
|
||
if dedupe_key and dedupe_key in self._seen_out:
|
||
return False, "duplicate_outbox"
|
||
if dedupe_key:
|
||
self._remember(self._seen_out, dedupe_key)
|
||
item = OutboxItem(
|
||
id=str(uuid.uuid4()),
|
||
touser=touser,
|
||
content=content,
|
||
dedupe_key=dedupe_key,
|
||
payload=dict(payload or {}),
|
||
)
|
||
self._outbox[item.id] = item
|
||
self._outbox_event.set()
|
||
logger.info("outbox 入队 id=%s touser=%s", item.id, touser)
|
||
return True, item.id
|
||
|
||
def has_outbound_dedupe(self, prefix: str) -> bool:
|
||
"""去重键是否已出过站(前缀匹配)。给 wait_file 判断文件/成交卡是否已发。"""
|
||
head = (prefix or "").strip()
|
||
if not head:
|
||
return False
|
||
with self._lock:
|
||
return any(str(k).startswith(head) for k in self._seen_out)
|
||
|
||
def claim_outbound(self) -> Optional[OutboxItem]:
|
||
"""认领一条 pending/超时 outbox。"""
|
||
now = time.time()
|
||
with self._lock:
|
||
for item in self._outbox.values():
|
||
if item.state == TaskState.PENDING or (
|
||
item.state == TaskState.CLAIMED
|
||
and item.claimed_at
|
||
and now - item.claimed_at > self._claim_timeout
|
||
):
|
||
item.state = TaskState.CLAIMED
|
||
item.claimed_at = now
|
||
item.attempts += 1
|
||
return item
|
||
self._outbox_event.clear()
|
||
return None
|
||
|
||
def complete_outbound(self, item_id: str, *, error: str = "", permanent: bool = False) -> None:
|
||
"""
|
||
标记 outbox 完成或失败。
|
||
|
||
permanent=True:不再重试(如 IP 未加白 60020)。
|
||
"""
|
||
with self._lock:
|
||
item = self._outbox.get(item_id)
|
||
if not item:
|
||
return
|
||
if error:
|
||
if permanent or item.attempts >= 5:
|
||
item.state = TaskState.FAILED
|
||
else:
|
||
item.state = TaskState.PENDING
|
||
item.claimed_at = 0.0
|
||
self._outbox_event.set()
|
||
item.last_error = error
|
||
else:
|
||
item.state = TaskState.DONE
|
||
|
||
def wait_inbox(self, timeout: float = 1.0) -> None:
|
||
self._inbox_event.wait(timeout=timeout)
|
||
|
||
def wait_outbox(self, timeout: float = 1.0) -> None:
|
||
self._outbox_event.wait(timeout=timeout)
|
||
|
||
def wait_outbound_done(self, item_id: str, timeout_sec: float = 12.0) -> bool:
|
||
"""
|
||
等到指定出站已发送完成或失败。
|
||
|
||
查价前必须先等询价确认发出,避免暂无报价卡抢先到销售手机。
|
||
只等这一条,不锁其它工单。超时返回 False。
|
||
"""
|
||
raw = (item_id or "").strip()
|
||
if not raw:
|
||
return False
|
||
started = time.time()
|
||
deadline = started + max(0.1, float(timeout_sec))
|
||
saw_claim = False
|
||
while time.time() < deadline:
|
||
state = ""
|
||
with self._lock:
|
||
item = self._outbox.get(raw)
|
||
if item:
|
||
state = item.state
|
||
if state in {TaskState.DONE, TaskState.FAILED}:
|
||
return state == TaskState.DONE
|
||
if state == TaskState.CLAIMED:
|
||
saw_claim = True
|
||
# 没有出站线程时(单测)不要空等;已被认领则等到发完。
|
||
if not saw_claim and state == TaskState.PENDING and time.time() - started >= 0.4:
|
||
return True
|
||
if not state and time.time() - started >= 0.3:
|
||
return True
|
||
self._outbox_event.wait(0.05)
|
||
return False
|
||
|
||
def stats(self) -> dict[str, Any]:
|
||
with self._lock:
|
||
return {
|
||
"backend": "memory",
|
||
"inbox": len(self._inbox),
|
||
"outbox": len(self._outbox),
|
||
"seen_msg": len(self._seen_msg),
|
||
}
|
||
|
||
def _remember(self, store: OrderedDict[str, float], key: str) -> None:
|
||
store[key] = time.time()
|
||
while len(store) > self._dedupe_maxlen:
|
||
store.popitem(last=False)
|
||
|
||
|
||
MessageStoreImpl = Union[MemoryMessageStore, Any]
|
||
|
||
_STORE: Optional[MessageStoreImpl] = None
|
||
_STORE_LOCK = threading.Lock()
|
||
|
||
|
||
def get_message_store() -> MessageStoreImpl:
|
||
"""
|
||
获取消息账本单例。
|
||
|
||
MESSAGE_STORE_BACKEND=postgres 且 PG 可用 → PostgresMessageStore;
|
||
auto 有 DSN 则尝试 PG,失败回退内存;memory 强制内存。
|
||
"""
|
||
global _STORE
|
||
with _STORE_LOCK:
|
||
if _STORE is None:
|
||
_STORE = _create_store()
|
||
return _STORE
|
||
|
||
|
||
def reset_message_store_for_tests() -> None:
|
||
"""仅自检用:清空单例。"""
|
||
global _STORE
|
||
with _STORE_LOCK:
|
||
_STORE = None
|
||
|
||
|
||
def _create_store() -> MessageStoreImpl:
|
||
from agent.config import get_settings
|
||
|
||
settings = get_settings()
|
||
backend = (getattr(settings, "message_store_backend", None) or "auto").strip().lower()
|
||
dsn = settings.resolve_pg_dsn()
|
||
|
||
use_pg = backend == "postgres" or (backend == "auto" and bool(dsn))
|
||
if use_pg and dsn:
|
||
try:
|
||
from agent.channel.pg_store import PostgresMessageStore
|
||
from agent.runtime_db import assert_safe_database
|
||
|
||
assert_safe_database(dsn)
|
||
store = PostgresMessageStore(dsn)
|
||
store.stats()
|
||
logger.info("消息账本后端=postgres")
|
||
return store
|
||
except Exception as exc: # noqa: BLE001
|
||
if backend == "postgres":
|
||
raise
|
||
logger.warning("PG 消息账本不可用,回退内存:%s", exc)
|
||
|
||
logger.info("消息账本后端=memory")
|
||
return MemoryMessageStore()
|