296 lines
10 KiB
Python
296 lines
10 KiB
Python
"""
|
||
PostgreSQL inbox / outbox 账本。
|
||
|
||
本文件职责:落 agent_inbox / agent_outbox;按条认领 + 超时回收。
|
||
调用:需已执行 migrations/001_runtime_core_v1.sql。
|
||
禁止:在回调线程跑长事务;单次操作短 SQL。
|
||
线程:每个方法短连接或共用连接需自行串行;本实现每操作用独立短连接,避免跨线程共享 cursor。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import json
|
||
import logging
|
||
import threading
|
||
import time
|
||
import uuid
|
||
from datetime import datetime, timezone
|
||
from typing import Any, Optional
|
||
|
||
from agent.channel.queue import InboxItem, OutboxItem, TaskState
|
||
from agent.channel.wecom.models import InboundMessage
|
||
from agent.runtime_db import assert_safe_database, connect_psycopg
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
|
||
def _utcnow() -> datetime:
|
||
return datetime.now(timezone.utc)
|
||
|
||
|
||
class PostgresMessageStore:
|
||
"""
|
||
PG 消息账本。
|
||
|
||
claim_timeout_sec:认领超时后可被其他消费者回收。
|
||
claimed_by:本进程标识,便于排查。
|
||
"""
|
||
|
||
def __init__(
|
||
self,
|
||
dsn: str,
|
||
*,
|
||
claim_timeout_sec: float = 120.0,
|
||
worker_id: str = "",
|
||
) -> None:
|
||
assert_safe_database(dsn)
|
||
self._dsn = dsn
|
||
self._claim_timeout = claim_timeout_sec
|
||
self._worker_id = worker_id or f"pid-{threading.get_ident()}"
|
||
self._inbox_event = threading.Event()
|
||
self._outbox_event = threading.Event()
|
||
|
||
def enqueue_inbound(self, message: InboundMessage) -> tuple[bool, str]:
|
||
"""写入 agent_inbox;message_id 唯一冲突视为重复。"""
|
||
if not message.has_sender():
|
||
return False, "missing_sender_id"
|
||
item_id = str(uuid.uuid4())
|
||
payload = json.dumps(message.to_dict(), ensure_ascii=False)
|
||
try:
|
||
with connect_psycopg(self._dsn) as conn:
|
||
conn.execute(
|
||
"""
|
||
INSERT INTO agent_inbox (
|
||
id, message_id, sender_id, chat_id, chat_type, payload, state
|
||
) VALUES (%s, %s, %s, %s, %s, %s::jsonb, 'pending')
|
||
""",
|
||
(
|
||
item_id,
|
||
message.message_id,
|
||
message.sender_id,
|
||
message.chat_id,
|
||
message.chat_type,
|
||
payload,
|
||
),
|
||
)
|
||
conn.commit()
|
||
self._inbox_event.set()
|
||
logger.info(
|
||
"pg inbox 入队 id=%s msg_id=%s sender=%s",
|
||
item_id,
|
||
message.message_id,
|
||
message.sender_id,
|
||
)
|
||
return True, item_id
|
||
except Exception as exc: # noqa: BLE001
|
||
# unique_violation
|
||
if "unique" in str(exc).lower() or "duplicate" in str(exc).lower():
|
||
return False, "duplicate_message_id"
|
||
logger.exception("pg inbox 入队失败")
|
||
raise
|
||
|
||
def claim_inbound(self) -> Optional[InboxItem]:
|
||
"""认领一条 pending 或超时 claimed。"""
|
||
timeout = int(self._claim_timeout)
|
||
with connect_psycopg(self._dsn) as conn:
|
||
row = conn.execute(
|
||
"""
|
||
WITH cte AS (
|
||
SELECT id FROM agent_inbox
|
||
WHERE state = 'pending'
|
||
OR (
|
||
state = 'claimed'
|
||
AND claimed_at IS NOT NULL
|
||
AND claimed_at < NOW() - (%s || ' seconds')::interval
|
||
)
|
||
ORDER BY created_at
|
||
FOR UPDATE SKIP LOCKED
|
||
LIMIT 1
|
||
)
|
||
UPDATE agent_inbox i
|
||
SET state = 'claimed',
|
||
claimed_at = NOW(),
|
||
claimed_by = %s,
|
||
attempts = attempts + 1,
|
||
updated_at = NOW()
|
||
FROM cte
|
||
WHERE i.id = cte.id
|
||
RETURNING i.id, i.payload, i.attempts, i.last_error, i.claimed_at
|
||
""",
|
||
(str(timeout), self._worker_id),
|
||
).fetchone()
|
||
conn.commit()
|
||
if not row:
|
||
self._inbox_event.clear()
|
||
return None
|
||
payload = row[1]
|
||
if isinstance(payload, str):
|
||
payload = json.loads(payload)
|
||
msg = InboundMessage.from_dict(dict(payload))
|
||
claimed_at = row[4]
|
||
ts = claimed_at.timestamp() if hasattr(claimed_at, "timestamp") else time.time()
|
||
return InboxItem(
|
||
id=str(row[0]),
|
||
message=msg,
|
||
state=TaskState.CLAIMED,
|
||
claimed_at=ts,
|
||
attempts=int(row[2] or 0),
|
||
last_error=str(row[3] or ""),
|
||
)
|
||
|
||
def complete_inbound(self, item_id: str, *, error: str = "") -> None:
|
||
state = "failed" if error else "done"
|
||
with connect_psycopg(self._dsn) as conn:
|
||
conn.execute(
|
||
"""
|
||
UPDATE agent_inbox
|
||
SET state = %s, last_error = %s, updated_at = NOW()
|
||
WHERE id = %s
|
||
""",
|
||
(state, error or "", item_id),
|
||
)
|
||
conn.commit()
|
||
|
||
def enqueue_outbound(
|
||
self,
|
||
*,
|
||
touser: str,
|
||
content: str,
|
||
dedupe_key: str = "",
|
||
) -> tuple[bool, str]:
|
||
if not touser or not content:
|
||
return False, "invalid_outbox"
|
||
item_id = str(uuid.uuid4())
|
||
try:
|
||
with connect_psycopg(self._dsn) as conn:
|
||
conn.execute(
|
||
"""
|
||
INSERT INTO agent_outbox (
|
||
id, dedupe_key, touser, content, payload, state
|
||
) VALUES (%s, %s, %s, %s, '{}'::jsonb, 'pending')
|
||
""",
|
||
(item_id, dedupe_key or "", touser, content),
|
||
)
|
||
conn.commit()
|
||
self._outbox_event.set()
|
||
logger.info("pg outbox 入队 id=%s touser=%s", item_id, touser)
|
||
return True, item_id
|
||
except Exception as exc: # noqa: BLE001
|
||
if dedupe_key and (
|
||
"unique" in str(exc).lower() or "duplicate" in str(exc).lower()
|
||
):
|
||
return False, "duplicate_outbox"
|
||
logger.exception("pg outbox 入队失败")
|
||
raise
|
||
|
||
def claim_outbound(self) -> Optional[OutboxItem]:
|
||
timeout = int(self._claim_timeout)
|
||
with connect_psycopg(self._dsn) as conn:
|
||
row = conn.execute(
|
||
"""
|
||
WITH cte AS (
|
||
SELECT id FROM agent_outbox
|
||
WHERE state = 'pending'
|
||
OR (
|
||
state = 'claimed'
|
||
AND claimed_at IS NOT NULL
|
||
AND claimed_at < NOW() - (%s || ' seconds')::interval
|
||
)
|
||
ORDER BY created_at
|
||
FOR UPDATE SKIP LOCKED
|
||
LIMIT 1
|
||
)
|
||
UPDATE agent_outbox o
|
||
SET state = 'claimed',
|
||
claimed_at = NOW(),
|
||
claimed_by = %s,
|
||
attempts = attempts + 1,
|
||
updated_at = NOW()
|
||
FROM cte
|
||
WHERE o.id = cte.id
|
||
RETURNING o.id, o.touser, o.content, o.attempts, o.last_error,
|
||
o.created_at, o.dedupe_key, o.claimed_at
|
||
""",
|
||
(str(timeout), self._worker_id),
|
||
).fetchone()
|
||
conn.commit()
|
||
if not row:
|
||
self._outbox_event.clear()
|
||
return None
|
||
created = row[5]
|
||
created_ts = created.timestamp() if hasattr(created, "timestamp") else time.time()
|
||
claimed = row[7]
|
||
claimed_ts = claimed.timestamp() if hasattr(claimed, "timestamp") else time.time()
|
||
return OutboxItem(
|
||
id=str(row[0]),
|
||
touser=str(row[1]),
|
||
content=str(row[2]),
|
||
state=TaskState.CLAIMED,
|
||
claimed_at=claimed_ts,
|
||
attempts=int(row[3] or 0),
|
||
last_error=str(row[4] or ""),
|
||
created_at=created_ts,
|
||
dedupe_key=str(row[6] or ""),
|
||
)
|
||
|
||
def complete_outbound(
|
||
self, item_id: str, *, error: str = "", permanent: bool = False
|
||
) -> None:
|
||
with connect_psycopg(self._dsn) as conn:
|
||
if error:
|
||
row = conn.execute(
|
||
"SELECT attempts FROM agent_outbox WHERE id = %s",
|
||
(item_id,),
|
||
).fetchone()
|
||
attempts = int(row[0]) if row else 0
|
||
if permanent or attempts >= 5:
|
||
conn.execute(
|
||
"""
|
||
UPDATE agent_outbox
|
||
SET state = 'failed', last_error = %s, updated_at = NOW()
|
||
WHERE id = %s
|
||
""",
|
||
(error, item_id),
|
||
)
|
||
else:
|
||
conn.execute(
|
||
"""
|
||
UPDATE agent_outbox
|
||
SET state = 'pending', claimed_at = NULL, claimed_by = '',
|
||
last_error = %s, updated_at = NOW()
|
||
WHERE id = %s
|
||
""",
|
||
(error, item_id),
|
||
)
|
||
self._outbox_event.set()
|
||
else:
|
||
conn.execute(
|
||
"""
|
||
UPDATE agent_outbox
|
||
SET state = 'done', last_error = '', updated_at = NOW()
|
||
WHERE id = %s
|
||
""",
|
||
(item_id,),
|
||
)
|
||
conn.commit()
|
||
|
||
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 stats(self) -> dict[str, Any]:
|
||
with connect_psycopg(self._dsn) as conn:
|
||
inbox = conn.execute(
|
||
"SELECT COUNT(*) FROM agent_inbox WHERE state IN ('pending','claimed')"
|
||
).fetchone()
|
||
outbox = conn.execute(
|
||
"SELECT COUNT(*) FROM agent_outbox WHERE state IN ('pending','claimed')"
|
||
).fetchone()
|
||
return {
|
||
"backend": "postgres",
|
||
"inbox": int(inbox[0] if inbox else 0),
|
||
"outbox": int(outbox[0] if outbox else 0),
|
||
}
|