From 2446993b11796e4ac6f1ceec9b3969dcf08d7e5e Mon Sep 17 00:00:00 2001 From: oliver Date: Wed, 22 Jul 2026 10:04:25 +0800 Subject: [PATCH] Fix WhatsApp overlap by serial queue, typing, and outbound delivery. Serialize per-session inbound turns and merge follow-ups after the active run; show composing while busy and enqueue final replies so long agent runs still reach the group. Co-authored-by: Cursor --- .../application/gateway/channel_turn_gate.py | 155 +++++++ .../application/gateway/inbound_service.py | 395 ++++++++++++++---- .../whatsapp_bridge/baileys_runner.ts | 92 +++- runtime/scheduler/whatsapp_mentions.py | 7 + tests/test_channel_turn_gate.py | 63 +++ tests/test_whatsapp_inbound_queue_cancel.py | 168 ++++++++ 6 files changed, 801 insertions(+), 79 deletions(-) create mode 100644 runtime/application/gateway/channel_turn_gate.py create mode 100644 tests/test_channel_turn_gate.py create mode 100644 tests/test_whatsapp_inbound_queue_cancel.py diff --git a/runtime/application/gateway/channel_turn_gate.py b/runtime/application/gateway/channel_turn_gate.py new file mode 100644 index 00000000..7307611d --- /dev/null +++ b/runtime/application/gateway/channel_turn_gate.py @@ -0,0 +1,155 @@ +"""Per-session serial queue: one active turn; overlaps enqueue and merge on drain.""" + +from __future__ import annotations + +import threading +import uuid +from dataclasses import dataclass, field +from typing import Any + + +@dataclass +class ChannelTurnHandle: + session_id: str + run_id: str + + +@dataclass +class _SessionState: + active_run_id: str = "" + pending: list[dict[str, Any]] = field(default_factory=list) + + +def merge_channel_pending_jobs(jobs: list[dict[str, Any]]) -> dict[str, Any]: + """Merge queued follow-ups into one job for the next agent turn.""" + rows = [j for j in (jobs or []) if isinstance(j, dict)] + if not rows: + return {} + if len(rows) == 1: + return dict(rows[0]) + + lang = str(rows[-1].get("lang") or "").strip().lower() + if lang.startswith("zh"): + header = "处理上一问期间又收到多条跟进,请一并回答:" + else: + header = "Several follow-up questions arrived while the previous request was still running. Please answer them together:" + + parts: list[str] = [] + attachments: list[dict[str, Any]] = [] + for idx, job in enumerate(rows, start=1): + text = str(job.get("user_text") or "").strip() + if text: + parts.append(f"{idx}) {text}") + for att in job.get("attachments") or []: + if isinstance(att, dict): + attachments.append(att) + + merged = dict(rows[-1]) + body = "\n\n".join(parts) if parts else "" + merged["user_text"] = f"{header}\n\n{body}".strip() if body else header + merged["attachments"] = attachments + merged["merged_count"] = len(rows) + return merged + + +class ChannelTurnGate: + """Serialize channel agent turns per session_id. + + - Idle ``try_begin``: start a turn immediately. + - Busy ``try_begin``: append job to pending and return None (caller returns quickly). + - ``end_and_take_merged``: clear active; if pending exists, merge all into one job and + start the next turn handle for the same HTTP request to continue. + """ + + def __init__(self) -> None: + self._lock = threading.Lock() + self._sessions: dict[str, _SessionState] = {} + + def try_begin(self, session_id: str, job: dict[str, Any]) -> ChannelTurnHandle | None: + sid = str(session_id or "").strip() + if not sid: + raise ValueError("session_id_required") + payload = dict(job or {}) + with self._lock: + st = self._sessions.get(sid) + if st is None: + st = _SessionState() + self._sessions[sid] = st + if st.active_run_id: + st.pending.append(payload) + return None + run_id = uuid.uuid4().hex + st.active_run_id = run_id + return ChannelTurnHandle(session_id=sid, run_id=run_id) + + def is_current(self, handle: ChannelTurnHandle) -> bool: + sid = str(handle.session_id or "").strip() + rid = str(handle.run_id or "").strip() + if not sid or not rid: + return False + with self._lock: + st = self._sessions.get(sid) + return bool(st and st.active_run_id == rid) + + def pending_count(self, session_id: str) -> int: + sid = str(session_id or "").strip() + with self._lock: + st = self._sessions.get(sid) + return len(st.pending) if st else 0 + + def end_and_take_merged(self, handle: ChannelTurnHandle) -> tuple[dict[str, Any] | None, ChannelTurnHandle | None]: + """Finish current turn. If queue non-empty, merge and return (merged_job, next_handle).""" + sid = str(handle.session_id or "").strip() + rid = str(handle.run_id or "").strip() + if not sid or not rid: + return None, None + with self._lock: + st = self._sessions.get(sid) + if st is None or st.active_run_id != rid: + return None, None + pending = list(st.pending) + st.pending.clear() + if not pending: + st.active_run_id = "" + self._sessions.pop(sid, None) + return None, None + merged = merge_channel_pending_jobs(pending) + next_id = uuid.uuid4().hex + st.active_run_id = next_id + return merged, ChannelTurnHandle(session_id=sid, run_id=next_id) + + def force_end(self, handle: ChannelTurnHandle) -> None: + """Drop active marker if still ours (error paths). Pending jobs are kept.""" + sid = str(handle.session_id or "").strip() + rid = str(handle.run_id or "").strip() + if not sid or not rid: + return + with self._lock: + st = self._sessions.get(sid) + if st is None or st.active_run_id != rid: + return + st.active_run_id = "" + if not st.pending: + self._sessions.pop(sid, None) + + +_GATE = ChannelTurnGate() + + +def get_channel_turn_gate() -> ChannelTurnGate: + return _GATE + + +def reset_channel_turn_gate_for_tests() -> ChannelTurnGate: + global _GATE + _GATE = ChannelTurnGate() + return _GATE + + +__all__ = [ + "ChannelTurnGate", + "ChannelTurnHandle", + "get_channel_turn_gate", + "merge_channel_pending_jobs", + "reset_channel_turn_gate_for_tests", +] diff --git a/runtime/application/gateway/inbound_service.py b/runtime/application/gateway/inbound_service.py index f51b814a..b96b12d9 100644 --- a/runtime/application/gateway/inbound_service.py +++ b/runtime/application/gateway/inbound_service.py @@ -889,6 +889,86 @@ def _should_suppress_channel_reply(*, channel: str, text: str) -> bool: return False +def _env_flag_enabled(name: str, *, default: bool = True) -> bool: + import os + + raw = str(os.getenv(name) or "").strip().lower() + if not raw: + return bool(default) + return raw not in {"0", "false", "no", "off"} + + +def _whatsapp_inbound_queue_delivery_enabled() -> bool: + return _env_flag_enabled("OCLAW_WHATSAPP_INBOUND_QUEUE_DELIVERY", default=True) + + +def _channel_turn_serial_queue_enabled() -> bool: + return _env_flag_enabled("OCLAW_CHANNEL_TURN_SERIAL_QUEUE", default=True) + + +def _enqueue_whatsapp_inbound_reply( + store: Any, + *, + inbound: Any, + account_id: str, + tenant_id: str, + reply_text: str, + reply_attachments: list[dict[str, Any]] | None, + reply_metadata: dict[str, Any] | None, +) -> str: + """Persist final WhatsApp reply on the outbound queue (sidecar poller delivers).""" + from runtime.scheduler.whatsapp_mentions import encode_whatsapp_outbound_source + + text = str(reply_text or "").strip() + reply: dict[str, Any] = { + "channel": "whatsapp", + "chat_id": str(getattr(inbound, "external_chat_id", "") or ""), + "text": text, + "attachments": list(reply_attachments or []), + "metadata": dict(reply_metadata or {}), + } + _maybe_expand_reply_attachments_for_channel(reply) + _maybe_add_media_path_for_wechat_reply(reply) + atts = [a for a in (reply.get("attachments") or []) if isinstance(a, dict)] + media_path = str(reply.get("media_path") or "").strip() or None + meta = dict(reply.get("metadata") or {}) if isinstance(reply.get("metadata"), dict) else {} + mention_jids = meta.get("mention_jids") if isinstance(meta.get("mention_jids"), list) else None + mention_names = meta.get("mention_names") if isinstance(meta.get("mention_names"), list) else None + quote_extra = { + k: meta[k] + for k in ( + "quote_remote_jid", + "quote_stanza_id", + "quote_participant", + "quote_text", + "quote_push_name", + "reply_to_user_id", + "is_group", + ) + if str(meta.get(k) or "").strip() or meta.get(k) is True + } + source = encode_whatsapp_outbound_source( + kind="inbound_reply", + mention_jids=mention_jids, + mention_names=mention_names, + mention_text_ready=bool(mention_jids), + attachments=atts or None, + media_path=media_path, + extra=quote_extra or None, + ) + return str( + store.enqueue_channel_outbound_message( + channel="whatsapp", + chat_id=str(getattr(inbound, "external_chat_id", "") or ""), + text=text, + tenant_id=str(tenant_id or ""), + account_id=str(account_id or "").strip() or "wa-default", + source=source, + ) + or "" + ) + + def process_inbound_payload(payload: dict[str, Any]) -> dict[str, Any]: from runtime.operations.mcp_env import apply_gateway_mcp_env_to_os @@ -951,6 +1031,7 @@ def process_inbound_payload(payload: dict[str, Any]) -> dict[str, Any]: preface = "" channel_session_id = "" channel_turn_uuid = "" + tenant_id = "" if text.lower().startswith("bind "): code = text.split(None, 1)[-1].strip() info = store.consume_bind_code( @@ -1264,79 +1345,220 @@ def process_inbound_payload(payload: dict[str, Any]) -> dict[str, Any]: if not user_text and gw_attachments: user_text = "用户发送了附件,请根据附件内容回复。" if user_text or gw_attachments: - try: - from runtime.gateway import OclawGateway - from runtime.types import StandardMessage - from runtime.lang import resolve_runtime_lang + from runtime.gateway import OclawGateway + from runtime.types import StandardMessage + from runtime.lang import resolve_runtime_lang + from runtime.application.gateway.channel_turn_gate import get_channel_turn_gate - interaction_mode, selected_specialist, dispatch_lang = _resolve_channel_dispatch( - store, channel=inbound.channel, account=account - ) - lang = ( - dispatch_lang - if dispatch_lang in {"zh", "en"} - else resolve_runtime_lang(store=store, user_text=user_text) - ) - gw = OclawGateway(store=store) - msg = StandardMessage( - session_id=str(session_id), - tenant_id=str(tenant_id or ""), - user_id=str(user_id or ""), - role=str(role or "member"), - channel=str(inbound.channel or "inbound"), - text=str(user_text or ""), - attachments=gw_attachments, - metadata={ - "tenant_id": tenant_id, - "user_id": user_id, - "channel": inbound.channel, - "role": role, - "account_id": account_id, - "interaction_mode": interaction_mode, - "selected_specialist": selected_specialist, - "is_group": inbound.is_group, - "external_user_id": inbound.external_user_id, - "external_chat_id": inbound.external_chat_id, - "group_sender_id": inbound.external_user_id, - "mentioned_jids": mention_jids_for_ctx, - "mention_names": mention_names_for_ctx, - "raw_inbound_text": raw_user_text, - "bot_jid": str(meta_for_inbound.get("bot_jid") or bot_jid or ""), - "bot_lid": str(meta_for_inbound.get("bot_lid") or ""), - }, - ) - manager = _build_admin_gateway_executor( - store, - tenant_id=tenant_id, - user_id=user_id, - specialist=selected_specialist, - session_id=str(session_id), - lang=lang, - ) - specialist_factory = lambda sid: _build_admin_gateway_executor( - store, - tenant_id=tenant_id, - user_id=user_id, - specialist=sid, - session_id=str(session_id), - lang=lang, - ) - turn_result = gw.handle_turn( - msg=msg, - lang=lang, - executor=manager, - specialist_executor_factory=specialist_factory, - ) - channel_turn_uuid = str(turn_result.turn_uuid or "").strip() - reply = str(turn_result.reply_text or "").strip() - reply_attachments = _collect_reply_attachments_from_history( - store=store, - session_id=str(session_id), - reply_text=reply, - ) - except Exception as e: - reply = f"抱歉,处理消息时出错:{type(e).__name__}: {e}" - reply_attachments = [] + interaction_mode, selected_specialist, dispatch_lang = _resolve_channel_dispatch( + store, channel=inbound.channel, account=account + ) + lang = ( + dispatch_lang + if dispatch_lang in {"zh", "en"} + else resolve_runtime_lang(store=store, user_text=user_text) + ) + msg_metadata = { + "tenant_id": tenant_id, + "user_id": user_id, + "channel": inbound.channel, + "role": role, + "account_id": account_id, + "interaction_mode": interaction_mode, + "selected_specialist": selected_specialist, + "is_group": inbound.is_group, + "external_user_id": inbound.external_user_id, + "external_chat_id": inbound.external_chat_id, + "group_sender_id": inbound.external_user_id, + "mentioned_jids": mention_jids_for_ctx, + "mention_names": mention_names_for_ctx, + "raw_inbound_text": raw_user_text, + "bot_jid": str(meta_for_inbound.get("bot_jid") or bot_jid or ""), + "bot_lid": str(meta_for_inbound.get("bot_lid") or ""), + } + turn_job: dict[str, Any] = { + "user_text": str(user_text or ""), + "attachments": list(gw_attachments or []), + "msg_metadata": msg_metadata, + "lang": lang, + "interaction_mode": interaction_mode, + "selected_specialist": selected_specialist, + "inbound": inbound, + } + gate = get_channel_turn_gate() + turn_handle = None + if _channel_turn_serial_queue_enabled(): + turn_handle = gate.try_begin(str(session_id), turn_job) + if turn_handle is None: + # Another turn is active; this message will be merged after it. + return { + "ok": True, + "replies": [], + "delivery": "accepted_queued", + } + + gw = OclawGateway(store=store) + manager = _build_admin_gateway_executor( + store, + tenant_id=tenant_id, + user_id=user_id, + specialist=selected_specialist, + session_id=str(session_id), + lang=lang, + ) + specialist_factory = lambda sid: _build_admin_gateway_executor( + store, + tenant_id=tenant_id, + user_id=user_id, + specialist=sid, + session_id=str(session_id), + lang=lang, + ) + + serial_replies: list[dict[str, Any]] = [] + last_outbound_id = "" + wa_queue_delivery = ( + str(inbound.channel or "").strip().lower() == "whatsapp" + and _whatsapp_inbound_queue_delivery_enabled() + ) + try: + while True: + job_text = str(turn_job.get("user_text") or "") + job_atts = [ + a for a in (turn_job.get("attachments") or []) if isinstance(a, dict) + ] + job_meta = ( + dict(turn_job.get("msg_metadata") or {}) + if isinstance(turn_job.get("msg_metadata"), dict) + else dict(msg_metadata) + ) + job_lang = str(turn_job.get("lang") or lang or "zh") + job_inbound = turn_job.get("inbound") or inbound + try: + turn_result = gw.handle_turn( + msg=StandardMessage( + session_id=str(session_id), + tenant_id=str(tenant_id or ""), + user_id=str(user_id or ""), + role=str(role or "member"), + channel=str(inbound.channel or "inbound"), + text=job_text, + attachments=job_atts, + metadata=job_meta, + ), + lang=job_lang, + executor=manager, + specialist_executor_factory=specialist_factory, + ) + channel_turn_uuid = str(turn_result.turn_uuid or "").strip() + turn_reply = str(turn_result.reply_text or "").strip() + turn_atts = _collect_reply_attachments_from_history( + store=store, + session_id=str(session_id), + reply_text=turn_reply, + ) + except Exception as e: + turn_reply = f"抱歉,处理消息时出错:{type(e).__name__}: {e}" + turn_atts = [] + + reply_metadata: dict[str, Any] = {} + if ( + str(inbound.channel or "").strip().lower() == "whatsapp" + and bool(getattr(job_inbound, "is_group", False)) + ): + from runtime.orchestration.group_ingest import ( + build_whatsapp_group_reply_metadata, + ) + + reply_metadata = build_whatsapp_group_reply_metadata( + inbound=job_inbound + ) + + if wa_queue_delivery and (turn_reply or turn_atts): + try: + last_outbound_id = _enqueue_whatsapp_inbound_reply( + store, + inbound=job_inbound, + account_id=account_id, + tenant_id=str(tenant_id or ""), + reply_text=turn_reply, + reply_attachments=list(turn_atts or []), + reply_metadata=reply_metadata, + ) + except Exception: + serial_replies.append( + { + "channel": inbound.channel, + "chat_id": inbound.external_chat_id, + "text": turn_reply, + "attachments": list(turn_atts or []), + "metadata": reply_metadata, + } + ) + elif turn_reply or turn_atts: + serial_replies.append( + { + "channel": inbound.channel, + "chat_id": inbound.external_chat_id, + "text": turn_reply, + "attachments": list(turn_atts or []), + "metadata": reply_metadata, + } + ) + + # Keep last turn as the primary reply fields for non-queue paths. + reply = turn_reply + reply_attachments = list(turn_atts or []) + + if turn_handle is None: + break + next_job, next_handle = gate.end_and_take_merged(turn_handle) + if not next_job or next_handle is None: + turn_handle = None + break + turn_job = next_job + turn_handle = next_handle + finally: + if turn_handle is not None: + try: + gate.force_end(turn_handle) + 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: + # Multiple serial turns: return all sync replies in order. + replies = serial_replies + # Skip the default single-reply assembly below. + 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} else: reply = "收到消息,但内容为空。请直接发送文本。" reply_attachments = [] @@ -1371,11 +1593,38 @@ def process_inbound_payload(payload: dict[str, Any]) -> dict[str, Any]: if not reply_attachments: # If assistant didn't persist attachments, fall back to recent tool-produced media. reply_attachments = _collect_recent_tool_attachments(store=store, session_id=str(session_id)) - reply_metadata: dict[str, Any] = {} + reply_metadata = {} if str(inbound.channel or "").strip().lower() == "whatsapp" and inbound.is_group: from runtime.orchestration.group_ingest import build_whatsapp_group_reply_metadata reply_metadata = build_whatsapp_group_reply_metadata(inbound=inbound) + + # WhatsApp: deliver via outbound queue so long agent runs survive sidecar HTTP drops. + if ( + str(inbound.channel or "").strip().lower() == "whatsapp" + and _whatsapp_inbound_queue_delivery_enabled() + and (reply or reply_attachments) + ): + try: + msg_id = _enqueue_whatsapp_inbound_reply( + store, + inbound=inbound, + account_id=account_id, + tenant_id=str(tenant_id or ""), + reply_text=reply, + reply_attachments=list(reply_attachments or []), + reply_metadata=reply_metadata, + ) + return { + "ok": True, + "replies": [], + "delivery": "queued", + "outbound_message_id": msg_id, + } + except Exception: + # Fall back to sync replies if enqueue fails. + pass + replies = [ { "channel": inbound.channel, diff --git a/runtime/operations/whatsapp_bridge/baileys_runner.ts b/runtime/operations/whatsapp_bridge/baileys_runner.ts index 4dc9e9d9..6bf91430 100644 --- a/runtime/operations/whatsapp_bridge/baileys_runner.ts +++ b/runtime/operations/whatsapp_bridge/baileys_runner.ts @@ -35,6 +35,14 @@ const PROXY_URL = ( ).trim(); const MEDIA_MAX_BYTES = Number(process.env.OCLAW_WHATSAPP_MEDIA_MAX_BYTES || String(20 * 1024 * 1024)) || 20 * 1024 * 1024; const MEDIA_DOWNLOAD_TIMEOUT_MS = Number(process.env.OCLAW_WHATSAPP_MEDIA_TIMEOUT_MS || "90000") || 90_000; +const TYPING_ENABLED = + String(process.env.OCLAW_WHATSAPP_TYPING || "1").trim() !== "0" && + String(process.env.OCLAW_WHATSAPP_TYPING || "1").trim().toLowerCase() !== "false" && + String(process.env.OCLAW_WHATSAPP_TYPING || "1").trim().toLowerCase() !== "off"; +const TYPING_HEARTBEAT_MS = Math.max( + 3000, + Number(process.env.OCLAW_WHATSAPP_TYPING_HEARTBEAT_MS || "8000") || 8000, +); function log(msg: string): void { process.stdout.write(`${new Date().toISOString()} [baileys-whatsapp] ${msg}\n`); @@ -916,11 +924,22 @@ async function pollOutboundQueue(sock: ReturnType): Promise `outbound sent id=${id} chat=${chatId} attachments=${attachments.length} mediaPath=${mediaPath ? "yes" : "no"} text=${text.slice(0, 80)}`, ); } else { - const sendContent = buildOutboundSendContent(text, source, sock, chatId); - const sent = await sock.sendMessage(chatId, sendContent); + const meta = decodeOutboundSource(source); + const content = buildMentionedOutboundText(text, meta, sock, chatId); + const quoted = buildQuotedMessage({ + chatId: String((meta as any).quote_remote_jid || (meta as any).quoteRemoteJid || chatId).trim(), + stanzaId: String((meta as any).quote_stanza_id || (meta as any).quoteStanzaId || "").trim(), + participant: String((meta as any).quote_participant || (meta as any).quoteParticipant || "").trim(), + quoteText: String((meta as any).quote_text || (meta as any).quoteText || "").trim(), + }); + const sent = await sock.sendMessage( + chatId, + content as any, + quoted ? { quoted } : undefined, + ); stanzaId = String((sent as any)?.key?.id || "").trim(); log( - `outbound sent id=${id} chat=${chatId} stanza=${stanzaId || "?"} mentions=${(sendContent.mentions || []).length} mention0=${(sendContent.mentions || [])[0] || ""} text=${sendContent.text.slice(0, 80)}`, + `outbound sent id=${id} chat=${chatId} stanza=${stanzaId || "?"} mentions=${(content.mentions || []).length} mention0=${(content.mentions || [])[0] || ""} text=${content.text.slice(0, 80)}`, ); } try { @@ -953,7 +972,7 @@ async function pollOutboundQueue(sock: ReturnType): Promise } function startOutboundPoller(getSock: () => ReturnType | null): void { - const intervalMs = Number(process.env.OCLAW_WHATSAPP_OUTBOUND_POLL_MS || "2500") || 2500; + const intervalMs = Number(process.env.OCLAW_WHATSAPP_OUTBOUND_POLL_MS || "1000") || 1000; setInterval(() => { const s = getSock(); if (!s) return; @@ -965,6 +984,56 @@ function sleep(ms: number): Promise { return new Promise((r) => setTimeout(r, ms)); } +type TypingEntry = { refs: number; timer: ReturnType | null }; +const typingByChat = new Map(); + +async function sendTypingPresence( + sock: ReturnType | null, + chatId: string, + state: "composing" | "paused", +): Promise { + const s = sock; + const jid = String(chatId || "").trim(); + if (!s || !jid || !TYPING_ENABLED) return; + try { + await s.sendPresenceUpdate(state, jid); + } catch (err) { + if (VERBOSE) log(`presence ${state} failed chat=${jid}: ${String(err).slice(0, 120)}`); + } +} + +function startTypingHeartbeat(sock: ReturnType | null, chatId: string): void { + const jid = String(chatId || "").trim(); + if (!sock || !jid || !TYPING_ENABLED) return; + let entry = typingByChat.get(jid); + if (!entry) { + entry = { refs: 0, timer: null }; + typingByChat.set(jid, entry); + } + entry.refs += 1; + if (entry.refs === 1) { + void sendTypingPresence(sock, jid, "composing"); + entry.timer = setInterval(() => { + void sendTypingPresence(sock, jid, "composing"); + }, TYPING_HEARTBEAT_MS); + } +} + +function stopTypingHeartbeat(sock: ReturnType | null, chatId: string): void { + const jid = String(chatId || "").trim(); + if (!jid || !TYPING_ENABLED) return; + const entry = typingByChat.get(jid); + if (!entry) return; + entry.refs = Math.max(0, entry.refs - 1); + if (entry.refs > 0) return; + if (entry.timer) { + clearInterval(entry.timer); + entry.timer = null; + } + typingByChat.delete(jid); + void sendTypingPresence(sock, jid, "paused"); +} + function _ipLooksHijackedOrUnroutable(ip: string): boolean { const s = String(ip || "").trim(); if (!s) return false; @@ -1199,9 +1268,20 @@ async function main(): Promise { } else if (VERBOSE || attachments.length) { log(`inbound posting chat=${chatId} user=${userId} textLen=${text.length} attachments=${attachments.length}`); } - const out = await postInbound(inbound); + startTypingHeartbeat(s, chatId); + let out: Json = {}; + try { + out = await postInbound(inbound); + } finally { + stopTypingHeartbeat(s, chatId); + } const replies = Array.isArray(out.replies) ? (out.replies as Json[]) : []; - if (isGroup || VERBOSE) log(`inbound ok chat=${chatId} replies=${replies.length}`); + const delivery = String((out as any).delivery || "").trim().toLowerCase(); + if (isGroup || VERBOSE || delivery === "queued" || delivery === "accepted_queued") { + log( + `inbound ok chat=${chatId} replies=${replies.length} delivery=${delivery || "sync"} outboundId=${String((out as any).outbound_message_id || "")}`, + ); + } for (const r of replies) { const outText = String((r as any).text || "").trim(); if (!shouldSendOutboundText(outText) && !(Array.isArray((r as any).attachments) && (r as any).attachments.length)) { diff --git a/runtime/scheduler/whatsapp_mentions.py b/runtime/scheduler/whatsapp_mentions.py index 854d7e22..b67395fa 100644 --- a/runtime/scheduler/whatsapp_mentions.py +++ b/runtime/scheduler/whatsapp_mentions.py @@ -16,8 +16,15 @@ def encode_whatsapp_outbound_source( mention_text_ready: bool = False, attachments: list[dict[str, Any]] | None = None, media_path: str | None = None, + extra: dict[str, Any] | None = None, ) -> str: payload: dict[str, Any] = {"kind": str(kind or "scheduled_job")} + if isinstance(extra, dict): + for key, value in extra.items(): + k = str(key or "").strip() + if not k or k in payload: + continue + payload[k] = value jids = normalize_whatsapp_mention_jids(mention_jids) if jids: payload["mention_jids"] = jids diff --git a/tests/test_channel_turn_gate.py b/tests/test_channel_turn_gate.py new file mode 100644 index 00000000..22024ec5 --- /dev/null +++ b/tests/test_channel_turn_gate.py @@ -0,0 +1,63 @@ +from __future__ import annotations + +from runtime.application.gateway.channel_turn_gate import ( + ChannelTurnGate, + merge_channel_pending_jobs, + reset_channel_turn_gate_for_tests, +) + + +def test_try_begin_enqueues_when_busy() -> None: + gate = ChannelTurnGate() + first = gate.try_begin("sess-a", {"user_text": "A"}) + assert first is not None + second = gate.try_begin("sess-a", {"user_text": "B"}) + assert second is None + assert gate.pending_count("sess-a") == 1 + third = gate.try_begin("sess-a", {"user_text": "C"}) + assert third is None + assert gate.pending_count("sess-a") == 2 + + +def test_end_and_take_merged_combines_pending() -> None: + gate = ChannelTurnGate() + first = gate.try_begin("sess-a", {"user_text": "A", "lang": "en"}) + assert first is not None + assert gate.try_begin("sess-a", {"user_text": "B question", "lang": "en", "attachments": [{"n": 1}]}) is None + assert gate.try_begin("sess-a", {"user_text": "C question", "lang": "en", "attachments": [{"n": 2}]}) is None + + merged, nxt = gate.end_and_take_merged(first) + assert nxt is not None + assert merged is not None + assert int(merged.get("merged_count") or 0) == 2 + text = str(merged.get("user_text") or "") + assert "B question" in text + assert "C question" in text + assert "follow-up" in text.lower() or "Several" in text + assert len(merged.get("attachments") or []) == 2 + assert gate.pending_count("sess-a") == 0 + + done, nxt2 = gate.end_and_take_merged(nxt) + assert done is None and nxt2 is None + + +def test_merge_channel_pending_jobs_zh_header() -> None: + merged = merge_channel_pending_jobs( + [ + {"user_text": "光功率?", "lang": "zh"}, + {"user_text": "还有告警吗?", "lang": "zh"}, + ] + ) + assert "一并回答" in str(merged.get("user_text") or "") + assert "1) 光功率?" in str(merged.get("user_text") or "") + assert "2) 还有告警吗?" in str(merged.get("user_text") or "") + + +def test_isolated_sessions_and_reset() -> None: + g1 = reset_channel_turn_gate_for_tests() + a = g1.try_begin("a", {"user_text": "1"}) + b = g1.try_begin("b", {"user_text": "2"}) + assert a is not None and b is not None + g2 = reset_channel_turn_gate_for_tests() + assert g1 is not g2 + assert g2.pending_count("a") == 0 diff --git a/tests/test_whatsapp_inbound_queue_cancel.py b/tests/test_whatsapp_inbound_queue_cancel.py new file mode 100644 index 00000000..6644216f --- /dev/null +++ b/tests/test_whatsapp_inbound_queue_cancel.py @@ -0,0 +1,168 @@ +from __future__ import annotations + +import json +import tempfile +import threading +import unittest +from pathlib import Path +from unittest import mock + +from runtime.application.gateway import inbound_service as inbound_mod +from runtime.application.gateway.channel_turn_gate import reset_channel_turn_gate_for_tests +from svc.persistence.sqlite_store import SqliteStore + + +class _FakeTurn: + def __init__(self, text: str = "final answer") -> None: + self.reply_text = text + self.turn_uuid = "turn-1" + + +class WhatsappInboundSerialQueueTests(unittest.TestCase): + def setUp(self) -> None: + self._tmp = tempfile.TemporaryDirectory(ignore_cleanup_errors=True) + self.db = Path(self._tmp.name) / "wa_inbound.sqlite" + self.store = SqliteStore(str(self.db)) + tenant = self.store.create_tenant("Team") + self.tenant_id = str(tenant["id"]) + user = self.store.create_user(tenant_id=self.tenant_id, display_name="ops", role="administrator") + self.user_id = str(user["id"]) + self.store.upsert_user_channel_account( + tenant_id=self.tenant_id, + user_id=self.user_id, + channel="whatsapp", + account_id="wa-default", + name="wa-default", + config={}, + is_active=True, + ) + self.store.upsert_channel_identity_v2( + tenant_id=self.tenant_id, + channel="whatsapp", + account_id="wa-default", + external_user_id="628100000@s.whatsapp.net", + user_id=self.user_id, + ) + reset_channel_turn_gate_for_tests() + + def tearDown(self) -> None: + self._tmp.cleanup() + + def _payload(self, text: str = "hello", stanza: str = "stanza1") -> dict: + return { + "channel": "whatsapp", + "account_id": "wa-default", + "user_id": "628100000@s.whatsapp.net", + "chat_id": "120363011111111111@g.us", + "text": f"@bot {text}", + "is_group": True, + "mentions": ["bot@s.whatsapp.net"], + "metadata": { + "bot_jid": "bot@s.whatsapp.net", + "mentions_bot": True, + "group_name": "AI nms", + "raw": { + "id": stanza, + "participant": "628100000@s.whatsapp.net", + "pushName": "Egista", + }, + }, + } + + def _patch_common(self): + return mock.patch.multiple( + inbound_mod, + get_assistant_store=mock.MagicMock(return_value=self.store), + _build_admin_gateway_executor=mock.MagicMock(return_value=object()), + _resolve_channel_dispatch=mock.MagicMock(return_value=("expert", "ops", "en")), + ) + + def test_whatsapp_final_reply_goes_to_outbound_queue(self) -> None: + 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("optical ok") + out = inbound_mod.process_inbound_payload(self._payload("check optics")) + + self.assertTrue(out.get("ok")) + self.assertEqual(out.get("delivery"), "queued") + self.assertEqual(out.get("replies"), []) + self.assertTrue(str(out.get("outbound_message_id") or "").strip()) + pending = self.store.list_pending_channel_outbound_messages( + channel="whatsapp", account_id="wa-default", limit=5 + ) + self.assertEqual(len(pending), 1) + self.assertIn("optical ok", pending[0].get("text") or "") + source = json.loads(str(pending[0].get("source") or "{}")) + self.assertEqual(source.get("kind"), "inbound_reply") + self.assertEqual(source.get("quote_stanza_id"), "stanza1") + + def test_busy_inbound_is_accepted_queued_then_merged(self) -> None: + release = threading.Event() + seen_texts: list[str] = [] + + def _handle_turn(**kwargs): + msg = kwargs.get("msg") + text = str(getattr(msg, "text", "") or "") + seen_texts.append(text) + if len(seen_texts) == 1: + # While first turn runs, enqueue two follow-ups from other threads/requests. + release.wait(timeout=2.0) + return _FakeTurn(f"ans:{len(seen_texts)}") + + 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.side_effect = lambda **kw: _handle_turn(**kw) + + results: dict[str, dict] = {} + + def run_first() -> None: + results["first"] = inbound_mod.process_inbound_payload( + self._payload("first question", stanza="s1") + ) + + t = threading.Thread(target=run_first, daemon=True) + t.start() + # Wait until first turn has started (gate busy). + import time + + for _ in range(100): + if seen_texts: + break + time.sleep(0.02) + + self.assertTrue(seen_texts, "first turn did not start") + out_b = inbound_mod.process_inbound_payload(self._payload("second question", stanza="s2")) + out_c = inbound_mod.process_inbound_payload(self._payload("third question", stanza="s3")) + self.assertEqual(out_b.get("delivery"), "accepted_queued") + self.assertEqual(out_c.get("delivery"), "accepted_queued") + + release.set() + t.join(timeout=5.0) + self.assertFalse(t.is_alive()) + self.assertEqual(results["first"].get("delivery"), "queued") + + self.assertEqual(len(seen_texts), 2) + self.assertIn("first question", seen_texts[0]) + self.assertIn("second question", seen_texts[1]) + self.assertIn("third question", seen_texts[1]) + pending = self.store.list_pending_channel_outbound_messages( + channel="whatsapp", account_id="wa-default", limit=10 + ) + self.assertEqual(len(pending), 2) + texts = " ".join(str(p.get("text") or "") for p in pending) + self.assertIn("ans:1", texts) + self.assertIn("ans:2", texts) + + +if __name__ == "__main__": + unittest.main()