Files
inquiry_robot/inquiry-agent/agent/channel/queue.py
T

316 lines
10 KiB
Python
Raw 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.
"""
进程内 / 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
# 单测没有出站线程:一直 PENDING 就不要空等 12 秒。测服有出站线程会立刻认领。
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()