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

296 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 账本。
本文件职责:落 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),
}