mirror of
https://github.com/hansjone/oclaw.git
synced 2026-10-09 00:40:45 +08:00
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:
parent
f1ed86537c
commit
2446993b11
6 changed files with 801 additions and 79 deletions
155
runtime/application/gateway/channel_turn_gate.py
Normal file
155
runtime/application/gateway/channel_turn_gate.py
Normal 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",
|
||||
]
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)) {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
63
tests/test_channel_turn_gate.py
Normal file
63
tests/test_channel_turn_gate.py
Normal 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
|
||||
168
tests/test_whatsapp_inbound_queue_cancel.py
Normal file
168
tests/test_whatsapp_inbound_queue_cancel.py
Normal 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()
|
||||
Loading…
Add table
Add a link
Reference in a new issue