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 <cursoragent@cursor.com>
This commit is contained in:
oliver 2026-07-22 10:04:25 +08:00
parent f1ed86537c
commit 2446993b11
6 changed files with 801 additions and 79 deletions

View file

@ -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",
]

View file

@ -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,

View file

@ -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<typeof makeWASocket>): 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<typeof makeWASocket>): Promise
}
function startOutboundPoller(getSock: () => ReturnType<typeof makeWASocket> | 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<void> {
return new Promise((r) => setTimeout(r, ms));
}
type TypingEntry = { refs: number; timer: ReturnType<typeof setInterval> | null };
const typingByChat = new Map<string, TypingEntry>();
async function sendTypingPresence(
sock: ReturnType<typeof makeWASocket> | null,
chatId: string,
state: "composing" | "paused",
): Promise<void> {
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<typeof makeWASocket> | 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<typeof makeWASocket> | 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<void> {
} 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)) {

View file

@ -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

View file

@ -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

View file

@ -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()