mirror of
https://github.com/hansjone/oclaw.git
synced 2026-10-08 23:33:16 +08:00
Fix duplicate WhatsApp outbound delivery.
Claim pending rows atomically and serialize the sidecar poller so overlapping polls cannot send the same reply twice; also stop inbound queue delivery from falling through to a second enqueue. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
10ae4abc1d
commit
b6e023cd91
6 changed files with 193 additions and 14 deletions
|
|
@ -294,7 +294,11 @@ def create_app() -> FastAPI:
|
|||
|
||||
store = get_assistant_store()
|
||||
aid = str(account_id or os.getenv("AIA_WHATSAPP_ACCOUNT_ID") or "wa-default").strip()
|
||||
items = store.list_pending_channel_outbound_messages(channel="whatsapp", account_id=aid, limit=limit)
|
||||
claimer = getattr(store, "claim_pending_channel_outbound_messages", None)
|
||||
if callable(claimer):
|
||||
items = claimer(channel="whatsapp", account_id=aid, limit=limit)
|
||||
else:
|
||||
items = store.list_pending_channel_outbound_messages(channel="whatsapp", account_id=aid, limit=limit)
|
||||
return {"ok": True, "items": items}
|
||||
|
||||
@app.post("/whatsapp/outbound/ack")
|
||||
|
|
|
|||
|
|
@ -1527,14 +1527,45 @@ def process_inbound_payload(payload: dict[str, Any]) -> dict[str, Any]:
|
|||
except Exception:
|
||||
pass
|
||||
|
||||
if wa_queue_delivery and last_outbound_id and not serial_replies:
|
||||
return {
|
||||
"ok": True,
|
||||
"replies": [],
|
||||
"delivery": "queued",
|
||||
"outbound_message_id": last_outbound_id,
|
||||
}
|
||||
if serial_replies and not wa_queue_delivery:
|
||||
if wa_queue_delivery:
|
||||
# Always leave this branch when queue delivery is on — otherwise
|
||||
# ``reply`` falls through to the bottom enqueue and WhatsApp
|
||||
# gets the same final answer twice.
|
||||
if last_outbound_id and not serial_replies:
|
||||
return {
|
||||
"ok": True,
|
||||
"replies": [],
|
||||
"delivery": "queued",
|
||||
"outbound_message_id": last_outbound_id,
|
||||
}
|
||||
if serial_replies:
|
||||
# Enqueue failed for at least one turn; sync fallback.
|
||||
replies = serial_replies
|
||||
if preface:
|
||||
first = replies[0] if replies else None
|
||||
if isinstance(first, dict) and first.get("text"):
|
||||
first["text"] = f"{preface}\n\n{first.get('text')}"
|
||||
elif preface:
|
||||
replies.insert(
|
||||
0,
|
||||
{
|
||||
"channel": inbound.channel,
|
||||
"chat_id": inbound.external_chat_id,
|
||||
"text": preface,
|
||||
"attachments": [],
|
||||
"metadata": {},
|
||||
},
|
||||
)
|
||||
for r in replies or []:
|
||||
if not isinstance(r, dict):
|
||||
continue
|
||||
ch = str(r.get("channel") or inbound.channel or "").strip().lower()
|
||||
if ch in {"wechat", "weixin", "whatsapp"}:
|
||||
_maybe_expand_reply_attachments_for_channel(r)
|
||||
_maybe_add_media_path_for_wechat_reply(r)
|
||||
return {"ok": True, "replies": replies}
|
||||
return {"ok": True, "replies": [], "delivery": "queued"}
|
||||
if serial_replies:
|
||||
# Multiple serial turns: return all sync replies in order.
|
||||
replies = serial_replies
|
||||
# Skip the default single-reply assembly below.
|
||||
|
|
|
|||
|
|
@ -973,10 +973,17 @@ async function pollOutboundQueue(sock: ReturnType<typeof makeWASocket>): Promise
|
|||
|
||||
function startOutboundPoller(getSock: () => ReturnType<typeof makeWASocket> | null): void {
|
||||
const intervalMs = Number(process.env.OCLAW_WHATSAPP_OUTBOUND_POLL_MS || "1000") || 1000;
|
||||
// Overlapping polls used to claim the same pending row twice (send → duplicate WhatsApp
|
||||
// bubbles) once inbound replies also went through the outbound queue.
|
||||
let pollInFlight = false;
|
||||
setInterval(() => {
|
||||
if (pollInFlight) return;
|
||||
const s = getSock();
|
||||
if (!s) return;
|
||||
void pollOutboundQueue(s);
|
||||
pollInFlight = true;
|
||||
void pollOutboundQueue(s).finally(() => {
|
||||
pollInFlight = false;
|
||||
});
|
||||
}, Math.max(1000, intervalMs));
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -6475,6 +6475,100 @@ class SqliteStore(ScheduledJobStoreMixin):
|
|||
out.append(item)
|
||||
return out
|
||||
|
||||
def claim_pending_channel_outbound_messages(
|
||||
self,
|
||||
*,
|
||||
channel: str,
|
||||
account_id: str,
|
||||
limit: int = 20,
|
||||
lease_seconds: int = 120,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Atomically claim pending outbound rows so overlapping pollers cannot double-send.
|
||||
|
||||
Marks matched rows ``sending``. Stale ``sending`` rows (lease expired) are reclaimed
|
||||
back to ``pending`` first so a crashed sidecar cannot strand messages forever.
|
||||
"""
|
||||
lim = max(1, min(int(limit), 100))
|
||||
lease = max(15, int(lease_seconds or 120))
|
||||
now = datetime.now(timezone.utc)
|
||||
now_iso = now.isoformat()
|
||||
stale_before = (now - timedelta(seconds=lease)).isoformat()
|
||||
ch = str(channel or "").strip()
|
||||
aid = str(account_id or "").strip()
|
||||
with self._connect() as conn:
|
||||
conn.execute(
|
||||
"""
|
||||
UPDATE channel_outbound_message
|
||||
SET status = 'pending'
|
||||
WHERE channel = ? AND account_id = ? AND status = 'sending'
|
||||
AND (last_attempt_at IS NULL OR last_attempt_at = '' OR last_attempt_at < ?)
|
||||
""",
|
||||
(ch, aid, stale_before),
|
||||
)
|
||||
rows = conn.execute(
|
||||
"""
|
||||
SELECT id
|
||||
FROM channel_outbound_message
|
||||
WHERE channel = ? AND account_id = ? AND status = 'pending'
|
||||
AND (next_attempt_at IS NULL OR next_attempt_at = '' OR next_attempt_at <= ?)
|
||||
ORDER BY created_at ASC
|
||||
LIMIT ?
|
||||
""",
|
||||
(ch, aid, now_iso, lim),
|
||||
).fetchall()
|
||||
ids = [str(r["id"] or "").strip() for r in rows if str(r["id"] or "").strip()]
|
||||
claimed: list[Any] = []
|
||||
for mid in ids:
|
||||
cur = conn.execute(
|
||||
"""
|
||||
UPDATE channel_outbound_message
|
||||
SET status = 'sending', last_attempt_at = ?
|
||||
WHERE id = ? AND status = 'pending'
|
||||
""",
|
||||
(now_iso, mid),
|
||||
)
|
||||
if int(cur.rowcount or 0) <= 0:
|
||||
continue
|
||||
row = conn.execute(
|
||||
"""
|
||||
SELECT id, tenant_id, channel, account_id, chat_id, text, source, created_at
|
||||
FROM channel_outbound_message
|
||||
WHERE id = ?
|
||||
""",
|
||||
(mid,),
|
||||
).fetchone()
|
||||
if row is not None:
|
||||
claimed.append(row)
|
||||
out: list[dict[str, Any]] = []
|
||||
for r in claimed:
|
||||
source = str(r["source"] or "")
|
||||
meta: dict[str, Any] = {}
|
||||
if source.startswith("{"):
|
||||
try:
|
||||
parsed = json.loads(source)
|
||||
if isinstance(parsed, dict):
|
||||
meta = parsed
|
||||
except Exception:
|
||||
meta = {}
|
||||
item: dict[str, Any] = {
|
||||
"id": str(r["id"] or ""),
|
||||
"tenant_id": str(r["tenant_id"] or ""),
|
||||
"channel": str(r["channel"] or ""),
|
||||
"account_id": str(r["account_id"] or ""),
|
||||
"chat_id": str(r["chat_id"] or ""),
|
||||
"text": str(r["text"] or ""),
|
||||
"source": source,
|
||||
"created_at": str(r["created_at"] or ""),
|
||||
}
|
||||
atts = [a for a in (meta.get("attachments") or []) if isinstance(a, dict)]
|
||||
if atts:
|
||||
item["attachments"] = atts
|
||||
mp = str(meta.get("media_path") or "").strip()
|
||||
if mp:
|
||||
item["media_path"] = mp
|
||||
out.append(item)
|
||||
return out
|
||||
|
||||
def list_pending_weixin_outbound_messages(
|
||||
self,
|
||||
*,
|
||||
|
|
@ -6546,7 +6640,7 @@ class SqliteStore(ScheduledJobStoreMixin):
|
|||
"""
|
||||
SELECT tenant_id, account_id, chat_id, text, source, send_attempts
|
||||
FROM channel_outbound_message
|
||||
WHERE id = ? AND status = 'pending'
|
||||
WHERE id = ? AND status IN ('pending', 'sending')
|
||||
""",
|
||||
(mid,),
|
||||
).fetchone()
|
||||
|
|
@ -6558,7 +6652,7 @@ class SqliteStore(ScheduledJobStoreMixin):
|
|||
"""
|
||||
UPDATE channel_outbound_message
|
||||
SET status = 'sent', sent_at = ?, error = '', send_attempts = ?, last_attempt_at = ?, next_attempt_at = NULL
|
||||
WHERE id = ? AND status = 'pending'
|
||||
WHERE id = ? AND status IN ('pending', 'sending')
|
||||
""",
|
||||
(ts, attempt_no, ts, mid),
|
||||
)
|
||||
|
|
@ -6568,7 +6662,7 @@ class SqliteStore(ScheduledJobStoreMixin):
|
|||
"""
|
||||
UPDATE channel_outbound_message
|
||||
SET status = 'failed', sent_at = ?, error = ?, send_attempts = ?, last_attempt_at = ?, next_attempt_at = NULL
|
||||
WHERE id = ? AND status = 'pending'
|
||||
WHERE id = ? AND status IN ('pending', 'sending')
|
||||
""",
|
||||
(ts, str(error or ""), attempt_no, ts, mid),
|
||||
)
|
||||
|
|
@ -6579,7 +6673,7 @@ class SqliteStore(ScheduledJobStoreMixin):
|
|||
"""
|
||||
UPDATE channel_outbound_message
|
||||
SET status = 'pending', error = ?, send_attempts = ?, last_attempt_at = ?, next_attempt_at = ?
|
||||
WHERE id = ? AND status = 'pending'
|
||||
WHERE id = ? AND status IN ('pending', 'sending')
|
||||
""",
|
||||
(str(error or ""), attempt_no, ts, next_attempt_at, mid),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -54,3 +54,26 @@ def test_channel_outbound_retry_success_clears_schedule(tmp_path) -> None:
|
|||
assert row["status"] == "sent"
|
||||
assert int(row["send_attempts"] or 0) == 2
|
||||
assert not str(row["next_attempt_at"] or "").strip()
|
||||
|
||||
|
||||
def test_claim_pending_channel_outbound_messages_is_exclusive(tmp_path) -> None:
|
||||
store = SqliteStore(str(tmp_path / "ops.sqlite"))
|
||||
msg_id = store.enqueue_channel_outbound_message(
|
||||
channel="whatsapp",
|
||||
account_id="wa-default",
|
||||
chat_id="120363012345678@g.us",
|
||||
text="once only",
|
||||
)
|
||||
first = store.claim_pending_channel_outbound_messages(
|
||||
channel="whatsapp", account_id="wa-default", limit=5
|
||||
)
|
||||
second = store.claim_pending_channel_outbound_messages(
|
||||
channel="whatsapp", account_id="wa-default", limit=5
|
||||
)
|
||||
assert len(first) == 1
|
||||
assert first[0]["id"] == msg_id
|
||||
assert second == []
|
||||
row = _row(store, msg_id)
|
||||
assert row["status"] == "sending"
|
||||
assert store.ack_channel_outbound_message(message_id=msg_id, ok=True, stanza_id="S1") is True
|
||||
assert _row(store, msg_id)["status"] == "sent"
|
||||
|
|
|
|||
|
|
@ -101,6 +101,26 @@ class WhatsappInboundSerialQueueTests(unittest.TestCase):
|
|||
self.assertEqual(source.get("kind"), "inbound_reply")
|
||||
self.assertEqual(source.get("quote_stanza_id"), "stanza1")
|
||||
|
||||
def test_whatsapp_queue_delivery_does_not_double_enqueue(self) -> None:
|
||||
"""Regression: agent-path enqueue must not fall through to bottom enqueue."""
|
||||
with self._patch_common(), mock.patch("runtime.gateway.OclawGateway") as gw_cls, mock.patch(
|
||||
"runtime.orchestration.group_ingest.should_process_group_inbound", return_value=True
|
||||
), mock.patch(
|
||||
"runtime.application.gateway.whatsapp_inbound_access.handle_whatsapp_access",
|
||||
return_value=None,
|
||||
):
|
||||
gw = gw_cls.return_value
|
||||
gw.handle_turn.return_value = _FakeTurn("only once")
|
||||
out = inbound_mod.process_inbound_payload(self._payload("ping"))
|
||||
|
||||
self.assertEqual(out.get("delivery"), "queued")
|
||||
self.assertEqual(out.get("replies"), [])
|
||||
pending = self.store.list_pending_channel_outbound_messages(
|
||||
channel="whatsapp", account_id="wa-default", limit=10
|
||||
)
|
||||
self.assertEqual(len(pending), 1)
|
||||
self.assertEqual(pending[0].get("text"), "only once")
|
||||
|
||||
def test_busy_inbound_is_accepted_queued_then_merged(self) -> None:
|
||||
release = threading.Event()
|
||||
seen_texts: list[str] = []
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue