refactor(whatsapp): 整理WhatsApp适配器代码结构并修复多线程安全问题

主要变更:
1. 重构导入顺序,统一模块导入规范
2. 提取通用方法到session模块,减少代码重复
3. 为缓存类添加线程/异步锁,修复并发安全问题
4. 新增入站处理器和发送管理器模块,拆分业务逻辑
5. 优化凭证队列,改为异步实现
6. 移除废弃的SSE_POLLING能力标识
7. 修复轮询投票解析逻辑
8. 优化Markdown转换规则,避免格式冲突
9. 完善连接控制器的异常处理
10. 新增发送静默消息的API支持
This commit is contained in:
Kris 2026-05-13 16:17:30 +08:00
parent f551745ec2
commit 5c3611ff19
20 changed files with 840 additions and 485 deletions

View File

@ -1,14 +1,11 @@
from yuxi.channels.adapters.whatsapp.accounts import MultiAccountManager, merge_account_config, resolve_default_account
from yuxi.channels.adapters.whatsapp.adapter import WhatsAppAdapter from yuxi.channels.adapters.whatsapp.adapter import WhatsAppAdapter
from yuxi.channels.adapters.whatsapp.security import ( from yuxi.channels.adapters.whatsapp.approve import (
DmPolicy, ApprovalCallbackManager,
GroupPolicy, ApprovalHandler,
WhatsAppSecurityPolicy, ApprovalRequest,
) ApprovalStatus,
from yuxi.channels.adapters.whatsapp.pairing import PairingManager, PairRequest ExecApprovals,
from yuxi.channels.adapters.whatsapp.inbound_pipeline import (
InboundPipeline,
PipelineAction,
PipelineResult,
) )
from yuxi.channels.adapters.whatsapp.channel_meta import ( from yuxi.channels.adapters.whatsapp.channel_meta import (
ChannelCapability, ChannelCapability,
@ -19,27 +16,30 @@ from yuxi.channels.adapters.whatsapp.connection_controller import (
ConnectionController, ConnectionController,
ConnectionState, ConnectionState,
) )
from yuxi.channels.adapters.whatsapp.accounts import MultiAccountManager, merge_account_config, resolve_default_account from yuxi.channels.adapters.whatsapp.directory import GroupEntry, PerGroupConfig
from yuxi.channels.adapters.whatsapp.reactions import ReactionLevel, ReactionLevelController, AckReactionManager
from yuxi.channels.adapters.whatsapp.directory import PerGroupConfig, GroupEntry
from yuxi.channels.adapters.whatsapp.per_dm_config import PerDmConfig
from yuxi.channels.adapters.whatsapp.error_policy import ErrorPolicy, ErrorPolicyConfig from yuxi.channels.adapters.whatsapp.error_policy import ErrorPolicy, ErrorPolicyConfig
from yuxi.channels.adapters.whatsapp.network_errors import ErrorCategory, classify_error, is_recoverable
from yuxi.channels.adapters.whatsapp.sent_message_cache import SentMessageCache
from yuxi.channels.adapters.whatsapp.group_gating import GroupGating from yuxi.channels.adapters.whatsapp.group_gating import GroupGating
from yuxi.channels.adapters.whatsapp.approve import (
ExecApprovals,
ApprovalRequest,
ApprovalStatus,
ApprovalHandler,
ApprovalCallbackManager,
)
from yuxi.channels.adapters.whatsapp.voice import transcoder, VoiceNoteSender
from yuxi.channels.adapters.whatsapp.vision import StickerCache, StickerVision
from yuxi.channels.adapters.whatsapp.health import DoctorDiagnostic, StatusIssueCollector from yuxi.channels.adapters.whatsapp.health import DoctorDiagnostic, StatusIssueCollector
from yuxi.channels.adapters.whatsapp.inbound_pipeline import (
InboundPipeline,
PipelineAction,
PipelineResult,
)
from yuxi.channels.adapters.whatsapp.network_errors import ErrorCategory, classify_error, is_recoverable
from yuxi.channels.adapters.whatsapp.pairing import PairingManager, PairRequest
from yuxi.channels.adapters.whatsapp.per_dm_config import PerDmConfig
from yuxi.channels.adapters.whatsapp.reactions import AckReactionManager, ReactionLevel, ReactionLevelController
from yuxi.channels.adapters.whatsapp.security import (
DmPolicy,
GroupPolicy,
WhatsAppSecurityPolicy,
)
from yuxi.channels.adapters.whatsapp.sent_message_cache import SentMessageCache
from yuxi.channels.adapters.whatsapp.setup.setup_wizard import setup_wizard from yuxi.channels.adapters.whatsapp.setup.setup_wizard import setup_wizard
from yuxi.channels.adapters.whatsapp.ui import InteractiveDispatcher
from yuxi.channels.adapters.whatsapp.targets import TargetResolver from yuxi.channels.adapters.whatsapp.targets import TargetResolver
from yuxi.channels.adapters.whatsapp.ui import InteractiveDispatcher
from yuxi.channels.adapters.whatsapp.vision import StickerCache, StickerVision
from yuxi.channels.adapters.whatsapp.voice import VoiceNoteSender, transcoder
__all__ = [ __all__ = [
"WhatsAppAdapter", "WhatsAppAdapter",

View File

@ -1,6 +1,6 @@
from yuxi.channels.adapters.whatsapp.accounts.accounts import MultiAccountManager
from yuxi.channels.adapters.whatsapp.accounts.account_config import merge_account_config from yuxi.channels.adapters.whatsapp.accounts.account_config import merge_account_config
from yuxi.channels.adapters.whatsapp.accounts.account_selection import resolve_default_account from yuxi.channels.adapters.whatsapp.accounts.account_selection import resolve_default_account
from yuxi.channels.adapters.whatsapp.accounts.accounts import MultiAccountManager
__all__ = [ __all__ = [
"MultiAccountManager", "MultiAccountManager",

View File

@ -1,17 +1,15 @@
from __future__ import annotations from __future__ import annotations
import asyncio import asyncio
import os
import tempfile
import time import time
from collections.abc import Awaitable, Callable from collections.abc import Awaitable, Callable
from typing import Any from typing import Any
import aiohttp
from yuxi.channels.base import BaseChannelAdapter from yuxi.channels.base import BaseChannelAdapter
from yuxi.channels.capabilities import ChannelCapabilities, TTSVoiceCapabilities, TTSCapabilities from yuxi.channels.capabilities import ChannelCapabilities, TTSCapabilities, TTSVoiceCapabilities
from yuxi.channels.exceptions import ChannelNotConnectedError from yuxi.channels.exceptions import ChannelNotConnectedError
from yuxi.channels.infra.circuit_breaker import CircuitBreaker
from yuxi.channels.meta import ChannelMeta as DisplayMeta from yuxi.channels.meta import ChannelMeta as DisplayMeta
from yuxi.channels.models import ( from yuxi.channels.models import (
ChannelMessage, ChannelMessage,
@ -34,22 +32,20 @@ from .dedupe import ButtonDeduplicator, MessageDeduplicator
from .directory import PerGroupConfig from .directory import PerGroupConfig
from .directory_peers import ContactDirectory from .directory_peers import ContactDirectory
from .echo_filter import EchoFilter from .echo_filter import EchoFilter
from .inbound_cache import InboundMessageCache
from .error_policy import ErrorPolicyConfig from .error_policy import ErrorPolicyConfig
from .format import format_outbound as _format_outbound from .format import format_outbound as _format_outbound
from .format import normalize_inbound as _normalize_inbound from .format import normalize_inbound as _normalize_inbound
from .heartbeat import HeartbeatManager from .heartbeat import HeartbeatManager
from .inbound_pipeline import InboundPipeline, PipelineAction, PipelineResult from .inbound_cache import InboundMessageCache
from .markdown import markdown_to_whatsapp, text_sanitizer from .inbound_processor import InboundProcessor
from .media import _MEDIA_SUFFIX_MAP, cleanup_temp_file from .inbound_pipeline import InboundPipeline, PipelineAction
from .media import download_media as _download_media
from .monitor import WhatsAppMonitor from .monitor import WhatsAppMonitor
from .pairing import PairingManager from .pairing import PairingManager
from .per_dm_config import PerDmConfig from .per_dm_config import PerDmConfig
from .reactions.ack_reaction import AckReactionManager from .reactions.ack_reaction import AckReactionManager
from .reactions.reaction_level import ReactionLevelController from .reactions.reaction_level import ReactionLevelController
from .security import DmPolicy, WhatsAppSecurityPolicy from .security import WhatsAppSecurityPolicy
from .send import chunk_message from .send_manager import SendManager
from .sent_message_cache import SentMessageCache from .sent_message_cache import SentMessageCache
from .session import jid_to_thread_key, normalize_phone, resolve_session_scope from .session import jid_to_thread_key, normalize_phone, resolve_session_scope
from .stream.lane_delivery import LaneDelivery from .stream.lane_delivery import LaneDelivery
@ -91,6 +87,7 @@ class WhatsAppAdapter(BaseChannelAdapter):
media=True, media=True,
supports_markdown=True, supports_markdown=True,
supports_streaming=True, supports_streaming=True,
supports_broadcast=False,
streaming_modes=["off", "typing_indicator"], streaming_modes=["off", "typing_indicator"],
text_chunk_limit=4000, text_chunk_limit=4000,
max_media_size_mb=100, max_media_size_mb=100,
@ -146,12 +143,45 @@ class WhatsAppAdapter(BaseChannelAdapter):
max_calls=cfg.get("debounceMaxCalls", 3), max_calls=cfg.get("debounceMaxCalls", 3),
) )
self._creds_queue: CredentialQueue | None = None self._creds_queue: CredentialQueue | None = None
self._inbound_pipeline = self._build_inbound_pipeline()
self._self_jid: str | None = None self._self_jid: str | None = None
self._on_connection_change: Callable[[dict[str, Any]], Awaitable[None]] | None = None self._on_connection_change: Callable[[dict[str, Any]], Awaitable[None]] | None = None
self._recent_messages: dict[str, dict[str, Any]] = {} self._recent_messages: dict[str, dict[str, Any]] = {}
self._last_inbound_at: float | None = None self._last_inbound_at: float | None = None
self._last_message_at: float | None = None self._last_message_at: float | None = None
self._circuit_breaker = CircuitBreaker(failure_threshold=5, recovery_timeout=60, channel_id="whatsapp")
self._stream_buffers: dict[str, str] = {}
self._stream_chunk_counters: dict[str, int] = {}
self._stream_status_msg_ids: dict[str, str] = {}
self._last_message_tracker: list[float | None] = [None]
self._send_mgr = SendManager(
bridge=self._bridge,
circuit_breaker=self._circuit_breaker,
sent_message_cache=self._sent_message_cache,
echo_filter=self._echo_filter,
debouncer=self._debouncer,
reaction_level=self._reaction_level,
error_policy=self._error_policy,
lane_delivery=self._lane_delivery,
config=cfg,
stream_buffers=self._stream_buffers,
stream_chunk_counters=self._stream_chunk_counters,
stream_status_msg_ids=self._stream_status_msg_ids,
recent_messages=self._recent_messages,
last_message_tracker=self._last_message_tracker,
)
self._inbound_processor = InboundProcessor(
echo_filter=self._echo_filter,
deduplicator=self._deduplicator,
button_deduplicator=self._button_deduplicator,
security=self._security,
self_chat_mode=self._self_chat_mode,
per_group_config=self._per_group_config,
per_dm_config=self._per_dm_config,
pairing=self._pairing,
self_jid=None,
)
self._inbound_pipeline = self._inbound_processor.build_pipeline()
async def connect(self) -> None: async def connect(self) -> None:
if self._status in (ChannelStatus.CONNECTED, ChannelStatus.CONNECTING): if self._status in (ChannelStatus.CONNECTED, ChannelStatus.CONNECTING):
@ -219,165 +249,34 @@ class WhatsAppAdapter(BaseChannelAdapter):
return {"status": "error", "error": "QR generation failed"} return {"status": "error", "error": "QR generation failed"}
async def send(self, response: ChannelResponse, silent: bool = False) -> DeliveryResult: async def send(self, response: ChannelResponse, silent: bool = False) -> DeliveryResult:
jid = response.identity.channel_chat_id return await self._send_mgr.send(response, silent)
if not jid and self._default_to:
jid = self._default_to
reply_to = response.reply_to_message_id
if self._reply_to_mode == "off":
reply_to = None
if self._debouncer.should_throttle(jid):
remaining = self._debouncer.window_remaining(jid)
logger.warning(f"[WhatsApp] Throttled message to {jid} (retry in {remaining:.1f}s)")
return DeliveryResult(
success=False,
error=f"Rate limited, retry in {remaining:.1f}s",
)
payload = _format_outbound(response)
content = payload.get("content", response.content)
content = text_sanitizer(content)
if self.supports_markdown and getattr(response, "content_format", "") == "markdown":
try:
content = markdown_to_whatsapp(content)
except Exception as e:
logger.warning(f"[WhatsApp] Markdown parse failed, using plain text: {e}")
prefix = self.config.get("messagePrefix", "")
if prefix:
content = f"{prefix} {content}"
self._echo_filter.record_outbound(jid, content)
self._debouncer.record_send(jid)
self._last_message_at = time.monotonic()
if len(content) > self.text_chunk_limit and self._reply_to_mode != "batched":
chunks = chunk_message(content, self.text_chunk_limit, mode=self._chunk_mode)
results = []
for chunk in chunks:
use_reply = reply_to if self._reply_to_mode == "all" else (reply_to if len(results) == 0 else None)
result = await self._bridge.send_message(
jid=jid,
content=chunk,
reply_to=use_reply,
silent=silent,
)
results.append(result)
return results[0] if results else DeliveryResult(success=False, error="No chunks")
if self._reply_to_mode == "batched" and len(content) > self.text_chunk_limit:
content = content[: self.text_chunk_limit - 3] + "..."
return await self._send_with_retry(jid, content, reply_to, silent)
async def send_media(self, chat_id: str, media_type: str, data: Any) -> DeliveryResult: async def send_media(self, chat_id: str, media_type: str, data: Any) -> DeliveryResult:
suffix = _MEDIA_SUFFIX_MAP.get(media_type, "") return await self._send_mgr.send_media(chat_id, media_type, data)
if isinstance(data, bytes):
with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as f:
f.write(data)
tmp_path = f.name
elif isinstance(data, str) and os.path.exists(data):
tmp_path = data
else:
return DeliveryResult(success=False, error=f"Unsupported media data type: {type(data)}")
try:
return await self._bridge.send_media(
jid=chat_id,
media_type=media_type,
media_path=tmp_path,
)
finally:
if isinstance(data, bytes) and os.path.exists(tmp_path):
cleanup_temp_file(tmp_path)
async def send_reaction(self, chat_id: str, msg_id: str, emoji: str) -> DeliveryResult: async def send_reaction(self, chat_id: str, msg_id: str, emoji: str) -> DeliveryResult:
if not emoji: return await self._send_mgr.send_reaction(chat_id, msg_id, emoji)
return await self._bridge.send_reaction(
jid=chat_id,
message_id=msg_id,
emoji="",
)
if not self._reaction_level.can_send_reaction():
return DeliveryResult(success=False, error="Reactions disabled by reactionLevel config")
return await self._bridge.send_reaction(
jid=chat_id,
message_id=msg_id,
emoji=emoji,
)
async def remove_reaction(self, chat_id: str, msg_id: str) -> DeliveryResult: async def remove_reaction(self, chat_id: str, msg_id: str) -> DeliveryResult:
return await self.send_reaction(chat_id, msg_id, "") return await self._send_mgr.remove_reaction(chat_id, msg_id)
async def send_reaction_from_action(self, action: dict[str, Any]) -> DeliveryResult: async def send_reaction_from_action(self, action: dict[str, Any]) -> DeliveryResult:
chat_jid = action.get("chatJid") or action.get("to") or action.get("chat_id", "") return await self._send_mgr.send_reaction_from_action(action)
msg_id = action.get("messageId") or action.get("msg_id", "")
emoji = action.get("emoji", "")
remove = action.get("remove", False)
participant = action.get("participant", "")
if participant and "@g.us" not in chat_jid:
chat_jid = participant
if remove or not emoji:
return await self.send_reaction(chat_jid, msg_id, "")
return await self.send_reaction(chat_jid, msg_id, emoji)
async def download_media(self, file_id: str) -> bytes: async def download_media(self, file_id: str) -> bytes:
raw = self._recent_messages.get(file_id) return await self._send_mgr.download_media(file_id)
if not raw:
raise ValueError(f"Message context not found for file_id: {file_id}")
remote_jid = raw.get("key", {}).get("remoteJid", "")
message = raw.get("message", {})
return await _download_media(self._bridge, remote_jid, file_id, message)
async def send_stream_chunk(self, chat_id: str, msg_id: str, chunk: str, finished: bool) -> DeliveryResult: async def send_stream_chunk(self, chat_id: str, msg_id: str, chunk: str, finished: bool) -> DeliveryResult:
if finished: return await self._send_mgr.send_stream_chunk(chat_id, msg_id, chunk, finished)
content = text_sanitizer(chunk)
if self.supports_markdown:
content = markdown_to_whatsapp(content)
if self._lane_delivery.reasoning_enabled:
lane_chunks = self._lane_delivery.split_lane_aware(content)
messages = self._lane_delivery.format_for_whatsapp(lane_chunks)
results = []
for msg in messages:
result = await self._bridge.send_message(jid=chat_id, content=msg)
results.append(result)
return results[0] if results else DeliveryResult(success=False, error="No messages")
return await self._bridge.send_message(
jid=chat_id,
content=content,
)
return await self._bridge.send_presence(
jid=chat_id,
presence="composing",
)
async def edit_message(self, chat_id: str, msg_id: str, content: str) -> DeliveryResult: async def edit_message(self, chat_id: str, msg_id: str, content: str) -> DeliveryResult:
return DeliveryResult(success=False, error="WhatsApp does not support editing messages") return await self._send_mgr.edit_message(chat_id, msg_id, content)
async def delete_message(self, chat_id: str, msg_id: str) -> DeliveryResult: async def delete_message(self, chat_id: str, msg_id: str) -> DeliveryResult:
return await self._bridge.delete_message( return await self._send_mgr.delete_message(chat_id, msg_id)
jid=chat_id,
message_id=msg_id,
)
async def send_poll(self, chat_id: str, name: str, options: list[str], selectable_count: int = 1) -> DeliveryResult: async def send_poll(self, chat_id: str, name: str, options: list[str], selectable_count: int = 1) -> DeliveryResult:
if len(options) < 2: return await self._send_mgr.send_poll(chat_id, name, options, selectable_count)
return DeliveryResult(success=False, error="Poll requires at least 2 options")
if len(options) > 12:
return DeliveryResult(success=False, error=f"Poll supports max 12 options, got {len(options)}")
return await self._bridge.create_poll(
jid=chat_id,
name=name,
options=options,
selectable_count=selectable_count,
)
async def send_location( async def send_location(
self, self,
@ -387,23 +286,14 @@ class WhatsAppAdapter(BaseChannelAdapter):
name: str = "", name: str = "",
address: str = "", address: str = "",
) -> DeliveryResult: ) -> DeliveryResult:
return await self._bridge.send_location( return await self._send_mgr.send_location(chat_id, latitude, longitude, name, address)
jid=chat_id,
latitude=latitude,
longitude=longitude,
name=name,
address=address,
)
async def send_contact( async def send_contact(
self, self,
chat_id: str, chat_id: str,
contacts: list[dict[str, str]], contacts: list[dict[str, str]],
) -> DeliveryResult: ) -> DeliveryResult:
return await self._bridge.send_contact( return await self._send_mgr.send_contact(chat_id, contacts)
jid=chat_id,
contacts=contacts,
)
async def send_sticker( async def send_sticker(
self, self,
@ -411,11 +301,7 @@ class WhatsAppAdapter(BaseChannelAdapter):
sticker_path: str, sticker_path: str,
reply_to: str | None = None, reply_to: str | None = None,
) -> DeliveryResult: ) -> DeliveryResult:
return await self._bridge.send_sticker( return await self._send_mgr.send_sticker(chat_id, sticker_path, reply_to)
jid=chat_id,
sticker_path=sticker_path,
reply_to=reply_to,
)
async def send_buttons( async def send_buttons(
self, self,
@ -425,13 +311,7 @@ class WhatsAppAdapter(BaseChannelAdapter):
title: str = "", title: str = "",
footer: str = "", footer: str = "",
) -> DeliveryResult: ) -> DeliveryResult:
return await self._bridge.send_buttons( return await self._send_mgr.send_buttons(chat_id, text, buttons, title, footer)
jid=chat_id,
text=text,
buttons=buttons,
title=title,
footer=footer,
)
async def send_list_message( async def send_list_message(
self, self,
@ -442,14 +322,7 @@ class WhatsAppAdapter(BaseChannelAdapter):
footer: str = "", footer: str = "",
button_text: str = "Select", button_text: str = "Select",
) -> DeliveryResult: ) -> DeliveryResult:
return await self._bridge.send_list_message( return await self._send_mgr.send_list_message(chat_id, text, sections, title, footer, button_text)
jid=chat_id,
text=text,
sections=sections,
title=title,
footer=footer,
button_text=button_text,
)
def normalize_inbound(self, raw: dict[str, Any]) -> ChannelMessage: def normalize_inbound(self, raw: dict[str, Any]) -> ChannelMessage:
return _normalize_inbound(raw, self.channel_id) return _normalize_inbound(raw, self.channel_id)
@ -512,8 +385,8 @@ class WhatsAppAdapter(BaseChannelAdapter):
self._401_retry_pending = False self._401_retry_pending = False
self._connection_ctrl.transition(ConnectionState.LOGGED_OUT) self._connection_ctrl.transition(ConnectionState.LOGGED_OUT)
from .logout_security import perform_logout_cleanup
from .auth_resolve import resolve_auth_dir from .auth_resolve import resolve_auth_dir
from .logout_security import perform_logout_cleanup
auth_dir = resolve_auth_dir(self.channel_id, self.config) auth_dir = resolve_auth_dir(self.channel_id, self.config)
perform_logout_cleanup(auth_dir) perform_logout_cleanup(auth_dir)
@ -547,7 +420,6 @@ class WhatsAppAdapter(BaseChannelAdapter):
result = await self._inbound_pipeline.process(raw_payload) result = await self._inbound_pipeline.process(raw_payload)
if result.action != PipelineAction.ACCEPT: if result.action != PipelineAction.ACCEPT:
self._echo_filter.clear()
return return
self._last_inbound_at = time.monotonic() self._last_inbound_at = time.monotonic()
@ -638,87 +510,8 @@ class WhatsAppAdapter(BaseChannelAdapter):
logger.debug("[WhatsApp] Vision enrichment failed", exc_info=True) logger.debug("[WhatsApp] Vision enrichment failed", exc_info=True)
def _build_inbound_pipeline(self) -> InboundPipeline: def _build_inbound_pipeline(self) -> InboundPipeline:
pipeline = InboundPipeline() self._inbound_processor._self_jid = self._self_jid
pipeline.add_step(self._check_echo, "echo_filter") return self._inbound_processor.build_pipeline()
pipeline.add_step(self._check_dedupe, "dedupe")
pipeline.add_step(self._check_button_dedupe, "button_dedupe")
pipeline.add_step(self._check_security, "security")
return pipeline
def _check_echo(self, payload: dict) -> PipelineResult:
key = payload.get("key", {})
msg = payload.get("message", {})
remote_jid = key.get("remoteJid", "")
text = msg.get("conversation", "")
if not text:
text = (msg.get("extendedTextMessage") or {}).get("text", "")
if self._echo_filter.is_echo(remote_jid, text):
return PipelineResult(action=PipelineAction.DROP, reason="echo")
return PipelineResult(action=PipelineAction.ACCEPT)
def _check_dedupe(self, payload: dict) -> PipelineResult:
msg_id = payload.get("key", {}).get("id", "")
if not msg_id:
return PipelineResult(action=PipelineAction.ACCEPT)
if self._deduplicator.is_duplicate(msg_id):
return PipelineResult(action=PipelineAction.DROP, reason="duplicate")
return PipelineResult(action=PipelineAction.ACCEPT)
def _check_button_dedupe(self, payload: dict) -> PipelineResult:
msg = payload.get("message", {})
button_msg = msg.get("buttonsResponseMessage") or msg.get("templateButtonReplyMessage")
if not button_msg:
return PipelineResult(action=PipelineAction.ACCEPT)
sender = payload.get("key", {}).get("remoteJid", "")
button_id = button_msg.get("selectedButtonId", "")
if not button_id:
button_id = button_msg.get("selectedId", "")
if not button_id:
return PipelineResult(action=PipelineAction.ACCEPT)
if self._button_deduplicator.is_duplicate(sender, button_id):
return PipelineResult(action=PipelineAction.DROP, reason="button_duplicate")
return PipelineResult(action=PipelineAction.ACCEPT)
def _check_security(self, payload: dict) -> PipelineResult:
from_me = payload.get("key", {}).get("fromMe", False)
if from_me:
if self._self_chat_mode:
return PipelineResult(action=PipelineAction.ACCEPT)
return PipelineResult(action=PipelineAction.DROP, reason="self_message")
remote_jid = payload.get("key", {}).get("remoteJid", "")
sender = remote_jid.split("@")[0]
if "@g.us" in remote_jid:
if not self._per_group_config.is_enabled(remote_jid):
return PipelineResult(action=PipelineAction.DROP, reason="group_disabled")
if self._per_group_config.require_mention(remote_jid):
msg = payload.get("message", {})
ext_text = msg.get("extendedTextMessage", {})
context_info = ext_text.get("contextInfo", {})
mentioned_jids = context_info.get("mentionedJid", []) or []
if not any(self._self_jid and mj == self._self_jid for mj in mentioned_jids):
return PipelineResult(action=PipelineAction.DROP, reason="require_mention_not_met")
allowed, reason = self._security.check_group_access(remote_jid)
if not allowed:
logger.debug(f"[WhatsApp] Group access denied for {remote_jid}: {reason}")
return PipelineResult(action=PipelineAction.DROP, reason=reason or "group_blocked")
else:
if not self._per_dm_config.is_enabled(sender):
return PipelineResult(action=PipelineAction.DROP, reason="dm_disabled")
if self._security.dm_policy == DmPolicy.PAIRING and not self._pairing.is_paired(sender):
logger.debug(f"[WhatsApp] DM pairing required for {sender}")
return PipelineResult(action=PipelineAction.DROP, reason="dm_not_paired")
allowed, reason = self._security.check_dm_access(sender)
if not allowed:
logger.debug(f"[WhatsApp] DM access denied for {sender}: {reason}")
return PipelineResult(action=PipelineAction.DROP, reason=reason or "dm_blocked")
return PipelineResult(action=PipelineAction.ACCEPT)
def _setup_monitor_handler(self) -> None: def _setup_monitor_handler(self) -> None:
self._monitor.on_raw_message(self._handle_sse_event) self._monitor.on_raw_message(self._handle_sse_event)
@ -747,55 +540,6 @@ class WhatsAppAdapter(BaseChannelAdapter):
} }
) )
async def _send_with_retry(
self,
jid: str,
content: str,
reply_to: str | None = None,
silent: bool = False,
max_attempts: int = 3,
) -> DeliveryResult:
last_error: str | None = None
for attempt in range(1, max_attempts + 1):
try:
result = await self._bridge.send_message(
jid=jid,
content=content,
reply_to=reply_to,
silent=silent,
)
if result.success and result.message_id:
self._sent_message_cache.put(
result.message_id,
jid,
{"content": content[:200], "reply_to": reply_to},
)
if result.success:
return result
last_error = result.error or "unknown"
except aiohttp.ClientConnectorError as e:
last_error = f"Bridge connection failed: {e}"
logger.warning(
f"[WhatsApp] send attempt {attempt}/{max_attempts} failed - bridge connection error for {jid}: {e}"
)
except TimeoutError as e:
last_error = str(e)
logger.warning(f"[WhatsApp] send attempt {attempt}/{max_attempts} failed - timeout for {jid}: {e}")
except Exception as e:
last_error = str(e)
logger.warning(f"[WhatsApp] send attempt {attempt}/{max_attempts} failed for {jid}: {e}")
if attempt < max_attempts:
delay_ms = 500 * attempt
logger.info(f"[WhatsApp] Retrying send to {jid} in {delay_ms}ms (attempt {attempt})")
await asyncio.sleep(delay_ms / 1000)
logger.error(f"[WhatsApp] send failed after {max_attempts} attempts for {jid}")
if self._error_policy.should_notify(jid, last_error or "send_failed"):
last_error = self._error_policy.format_error_message(last_error or "send failed")
return DeliveryResult(success=False, error=last_error or "send failed")
async def get_user_info(self, channel_user_id: str) -> dict[str, Any]: async def get_user_info(self, channel_user_id: str) -> dict[str, Any]:
phone = normalize_phone(channel_user_id) phone = normalize_phone(channel_user_id)
jid = f"{phone}@s.whatsapp.net" jid = f"{phone}@s.whatsapp.net"
@ -819,8 +563,8 @@ class WhatsAppAdapter(BaseChannelAdapter):
return await self._bridge.wait_scan(timeout) return await self._bridge.wait_scan(timeout)
async def detect_whatsapp_linked(self) -> bool: async def detect_whatsapp_linked(self) -> bool:
from .logout_security import validate_credential_freshness
from .auth_resolve import resolve_auth_dir from .auth_resolve import resolve_auth_dir
from .logout_security import validate_credential_freshness
auth_dir = resolve_auth_dir(self.channel_id, self.config) auth_dir = resolve_auth_dir(self.channel_id, self.config)
if not validate_credential_freshness(auth_dir): if not validate_credential_freshness(auth_dir):
@ -835,8 +579,8 @@ class WhatsAppAdapter(BaseChannelAdapter):
return await self._bridge.get_qr_status() return await self._bridge.get_qr_status()
async def logout(self) -> dict[str, Any]: async def logout(self) -> dict[str, Any]:
from .logout_security import perform_logout_cleanup
from .auth_resolve import resolve_auth_dir from .auth_resolve import resolve_auth_dir
from .logout_security import perform_logout_cleanup
result = await self._bridge.logout() result = await self._bridge.logout()
auth_dir = resolve_auth_dir(self.channel_id, self.config) auth_dir = resolve_auth_dir(self.channel_id, self.config)
@ -887,16 +631,30 @@ class WhatsAppAdapter(BaseChannelAdapter):
return self._last_message_at return self._last_message_at
def resolve_system_prompt(self, chat_id: str) -> str | None: def resolve_system_prompt(self, chat_id: str) -> str | None:
base = None
if "@g.us" in chat_id: if "@g.us" in chat_id:
return self._per_group_config.system_prompt(chat_id) base = self._per_group_config.system_prompt(chat_id)
sender = chat_id.split("@")[0] if "@" in chat_id else chat_id else:
return self._per_dm_config.system_prompt(sender) sender = chat_id.split("@")[0] if "@" in chat_id else chat_id
base = self._per_dm_config.system_prompt(sender)
return self._inject_reaction_guidance(base)
def resolve_system_prompt_for_group(self, group_jid: str) -> str | None: def resolve_system_prompt_for_group(self, group_jid: str) -> str | None:
return self._per_group_config.system_prompt(group_jid) base = self._per_group_config.system_prompt(group_jid)
return self._inject_reaction_guidance(base)
def resolve_system_prompt_for_direct(self, phone: str) -> str | None: def resolve_system_prompt_for_direct(self, phone: str) -> str | None:
return self._per_dm_config.system_prompt(phone) base = self._per_dm_config.system_prompt(phone)
return self._inject_reaction_guidance(base)
def _inject_reaction_guidance(self, base: str | None) -> str | None:
guidance = self.config.get("agentReactionGuidance")
if not guidance:
return base
guidance_text = f"\n\nYou may use WhatsApp reactions to respond with emoji when appropriate. {guidance}"
if base:
return base + guidance_text
return guidance_text.strip()
def resolve_quoted_message_key(self, chat_id: str, quoted_msg_id: str | None = None) -> dict | None: def resolve_quoted_message_key(self, chat_id: str, quoted_msg_id: str | None = None) -> dict | None:
return self._inbound_cache.resolve_quoted_message_key(chat_id, quoted_msg_id) return self._inbound_cache.resolve_quoted_message_key(chat_id, quoted_msg_id)

View File

@ -1,6 +1,6 @@
from yuxi.channels.adapters.whatsapp.approve.exec_approvals import ExecApprovals, ApprovalRequest, ApprovalStatus
from yuxi.channels.adapters.whatsapp.approve.approval_handler import ApprovalHandler
from yuxi.channels.adapters.whatsapp.approve.approval_callbacks import ApprovalCallbackManager from yuxi.channels.adapters.whatsapp.approve.approval_callbacks import ApprovalCallbackManager
from yuxi.channels.adapters.whatsapp.approve.approval_handler import ApprovalHandler
from yuxi.channels.adapters.whatsapp.approve.exec_approvals import ApprovalRequest, ApprovalStatus, ExecApprovals
__all__ = [ __all__ = [
"ExecApprovals", "ExecApprovals",

View File

@ -1,7 +1,7 @@
from __future__ import annotations from __future__ import annotations
from typing import Any
from collections.abc import Callable from collections.abc import Callable
from typing import Any
from yuxi.utils.logging_config import logger from yuxi.utils.logging_config import logger

View File

@ -12,7 +12,7 @@ function buildSendRoutes(app, getSock, notifySSE) {
return res.json({ success: false, error: 'Not connected' }); return res.json({ success: false, error: 'Not connected' });
} }
const { jid, content, reply_to } = req.body; const { jid, content, reply_to, silent } = req.body;
const options = {}; const options = {};
if (reply_to) { if (reply_to) {
options.quoted = { options.quoted = {
@ -21,7 +21,13 @@ function buildSendRoutes(app, getSock, notifySSE) {
}; };
} }
const result = await sock.sendMessage(jid, { text: content }, options); const msgContent = { text: content };
if (silent) {
msgContent.contextInfo = { isSilent: true };
}
const result = await sock.sendMessage(jid, msgContent, options);
res.json({ res.json({
success: true, success: true,
message_id: result?.key?.id || null, message_id: result?.key?.id || null,

View File

@ -117,6 +117,12 @@ class BaileysBridge(ExternalProcessManager):
except TimeoutError: except TimeoutError:
self._process.kill() self._process.kill()
await self._process.wait() await self._process.wait()
if self._process:
if self._process.stdout:
self._process.stdout.close()
if self._process.stderr:
self._process.stderr.close()
self._process = None
logger.info("Baileys bridge stopped") logger.info("Baileys bridge stopped")
async def health_check(self) -> HealthStatus: async def health_check(self) -> HealthStatus:

View File

@ -53,7 +53,6 @@ class WhatsAppChannelMeta:
ChannelCapability.GROUP, ChannelCapability.GROUP,
ChannelCapability.QR_LOGIN, ChannelCapability.QR_LOGIN,
ChannelCapability.PAIRING, ChannelCapability.PAIRING,
ChannelCapability.SSE_POLLING,
] ]
) )
inbound_rate_limit: int = 60 inbound_rate_limit: int = 60

View File

@ -55,7 +55,10 @@ class ConnectionController:
for listener in self._state_listeners: for listener in self._state_listeners:
try: try:
if asyncio.iscoroutinefunction(listener): if asyncio.iscoroutinefunction(listener):
asyncio.ensure_future(listener(old, new_state)) task = asyncio.create_task(listener(old, new_state))
task.add_done_callback(
lambda t: logger.exception("ConnectionController listener failed") if t.exception() else None
)
else: else:
listener(old, new_state) listener(old, new_state)
except Exception: except Exception:

View File

@ -1,7 +1,7 @@
from __future__ import annotations from __future__ import annotations
import asyncio
import json import json
import threading
import time import time
from pathlib import Path from pathlib import Path
from typing import Any from typing import Any
@ -12,105 +12,119 @@ from yuxi.utils.logging_config import logger
class CredentialQueue: class CredentialQueue:
def __init__(self, auth_dir: Path, debounce_ms: int = 500): def __init__(self, auth_dir: Path, debounce_ms: int = 500):
self._auth_dir = auth_dir self._auth_dir = auth_dir
self._lock = threading.Lock() self._lock = asyncio.Lock()
self._write_count = 0 self._write_count = 0
self._debounce_ms = debounce_ms self._debounce_ms = debounce_ms
self._pending_creds: dict[str, Any] | None = None self._pending_creds: dict[str, Any] | None = None
self._last_write_time = 0.0 self._last_write_time = 0.0
self._debounce_timer: threading.Timer | None = None self._debounce_task: asyncio.Task | None = None
def write_auth(self, creds: dict[str, Any]) -> int: async def write_auth(self, creds: dict[str, Any]) -> int:
loop = asyncio.get_running_loop()
now = time.monotonic() now = time.monotonic()
with self._lock: async with self._lock:
self._write_count += 1 self._write_count += 1
count = self._write_count count = self._write_count
if now - self._last_write_time < self._debounce_ms / 1000.0: if now - self._last_write_time < self._debounce_ms / 1000.0:
self._pending_creds = creds self._pending_creds = creds
self._schedule_debounced_write() self._schedule_debounced_write(loop)
return count return count
self._last_write_time = now self._last_write_time = now
self._pending_creds = None self._pending_creds = None
return self._do_write(creds, count) return await self._do_write(creds, count)
def _schedule_debounced_write(self) -> None: def _schedule_debounced_write(self, loop: asyncio.AbstractEventLoop) -> None:
if self._debounce_timer is not None: if self._debounce_task is not None:
self._debounce_timer.cancel() self._debounce_task.cancel()
self._debounce_timer = threading.Timer( async def _delayed_flush():
self._debounce_ms / 1000.0, await asyncio.sleep(self._debounce_ms / 1000.0)
self._flush_debounced, await self._flush_debounced()
)
self._debounce_timer.daemon = True
self._debounce_timer.start()
def _flush_debounced(self) -> None: self._debounce_task = loop.create_task(_delayed_flush())
with self._lock:
async def _flush_debounced(self) -> None:
async with self._lock:
creds = self._pending_creds creds = self._pending_creds
self._pending_creds = None self._pending_creds = None
self._debounce_timer = None self._debounce_task = None
if creds is None: if creds is None:
return return
count = self._write_count count = self._write_count
self._do_write(creds, count) await self._do_write(creds, count)
def _do_write(self, creds: dict[str, Any], count: int) -> int: async def _do_write(self, creds: dict[str, Any], count: int) -> int:
path = self._auth_dir / "creds.json" path = self._auth_dir / "creds.json"
backup = self._auth_dir / "creds.json.bak" backup = self._auth_dir / "creds.json.bak"
tmp_path = self._auth_dir / "creds.json.tmp" tmp_path = self._auth_dir / "creds.json.tmp"
try: try:
data = json.dumps(creds, indent=2, ensure_ascii=False) data = json.dumps(creds, indent=2, ensure_ascii=False)
tmp_path.write_text(data, encoding="utf-8")
if path.exists(): def _sync_write() -> int:
try: tmp_path.write_text(data, encoding="utf-8")
path.replace(backup) if path.exists():
except OSError: try:
pass path.replace(backup)
tmp_path.replace(path) except OSError:
logger.debug(f"CredentialQueue: wrote creds (#{count}) to {path}") pass
tmp_path.replace(path)
return count
result = await asyncio.to_thread(_sync_write)
logger.debug(f"CredentialQueue: wrote creds (#{result}) to {path}")
return result
except (OSError, TypeError) as e: except (OSError, TypeError) as e:
logger.error(f"CredentialQueue: write failed (#{count}): {e}") logger.error(f"CredentialQueue: write failed (#{count}): {e}")
raise raise
return count
def read_auth(self) -> dict[str, Any] | None: async def read_auth(self) -> dict[str, Any] | None:
path = self._auth_dir / "creds.json" path = self._auth_dir / "creds.json"
backup = self._auth_dir / "creds.json.bak" backup = self._auth_dir / "creds.json.bak"
try:
def _sync_read():
if not path.exists(): if not path.exists():
return None return None
return json.loads(path.read_text(encoding="utf-8")) try:
except (OSError, json.JSONDecodeError) as e: return json.loads(path.read_text(encoding="utf-8"))
logger.error(f"CredentialQueue: read failed: {e}") except (OSError, json.JSONDecodeError) as e:
if backup.exists(): logger.error(f"CredentialQueue: read failed: {e}")
try: if backup.exists():
data = json.loads(backup.read_text(encoding="utf-8")) try:
backup.replace(path) data = json.loads(backup.read_text(encoding="utf-8"))
logger.info("CredentialQueue: restored creds from backup") backup.replace(path)
return data logger.info("CredentialQueue: restored creds from backup")
except (OSError, json.JSONDecodeError) as e2: return data
logger.error(f"CredentialQueue: backup recovery failed: {e2}") except (OSError, json.JSONDecodeError) as e2:
return None logger.error(f"CredentialQueue: backup recovery failed: {e2}")
return None
def clear_auth(self) -> bool: return await asyncio.to_thread(_sync_read)
with self._lock:
async def clear_auth(self) -> bool:
async with self._lock:
self._pending_creds = None self._pending_creds = None
if self._debounce_timer is not None: if self._debounce_task is not None:
self._debounce_timer.cancel() self._debounce_task.cancel()
self._debounce_timer = None self._debounce_task = None
path = self._auth_dir / "creds.json" path = self._auth_dir / "creds.json"
backup = self._auth_dir / "creds.json.bak" backup = self._auth_dir / "creds.json.bak"
try:
def _sync_clear():
if path.exists(): if path.exists():
if backup.exists(): if backup.exists():
backup.unlink() backup.unlink()
path.rename(backup) path.rename(backup)
logger.info(f"CredentialQueue: cleared creds, backup at {backup}") logger.info(f"CredentialQueue: cleared creds, backup at {backup}")
return True return True
try:
result = await asyncio.to_thread(_sync_clear)
return result
except OSError as e: except OSError as e:
logger.error(f"CredentialQueue: clear failed: {e}") logger.error(f"CredentialQueue: clear failed: {e}")
return False return False

View File

@ -14,6 +14,7 @@ from yuxi.channels.models import (
MessageType, MessageType,
) )
from .session import _extract_sender_number
from .structured_context import ContextSource, StructuredContextEntry, UntrustedStructuredContext from .structured_context import ContextSource, StructuredContextEntry, UntrustedStructuredContext
_PROTOCOL_REVOKE = 0 _PROTOCOL_REVOKE = 0
@ -153,8 +154,8 @@ def _normalize_poll_vote(poll_update: dict, raw_payload: dict, channel_id: str)
selected = vote_info.get("selectedOptions", []) selected = vote_info.get("selectedOptions", [])
if isinstance(selected, list): if isinstance(selected, list):
selected_options = selected selected_options = selected
elif isinstance(vote_info, list): elif isinstance(vote_info, list):
selected_options = vote_info selected_options = vote_info
option_names = [o.get("name", str(o)) if isinstance(o, dict) else str(o) for o in selected_options] option_names = [o.get("name", str(o)) if isinstance(o, dict) else str(o) for o in selected_options]
vote_text = f"[Poll Vote] {' | '.join(option_names)}" if option_names else "[Poll Vote]" vote_text = f"[Poll Vote] {' | '.join(option_names)}" if option_names else "[Poll Vote]"
@ -370,10 +371,6 @@ def _jid_to_chat_type(jid: str) -> ChatType:
return ChatType.DIRECT return ChatType.DIRECT
def _extract_sender_number(jid: str) -> str:
return jid.split("@")[0]
def _build_structured_context(msg: dict, content_type: str, msg_id: str | None) -> UntrustedStructuredContext: def _build_structured_context(msg: dict, content_type: str, msg_id: str | None) -> UntrustedStructuredContext:
context = UntrustedStructuredContext() context = UntrustedStructuredContext()

View File

@ -2,6 +2,7 @@ from __future__ import annotations
import time import time
from collections import OrderedDict from collections import OrderedDict
from threading import Lock
from typing import Any from typing import Any
from yuxi.utils.logging_config import logger from yuxi.utils.logging_config import logger
@ -20,57 +21,62 @@ class InboundMessageCache:
self._cache: OrderedDict[str, dict[str, Any]] = OrderedDict() self._cache: OrderedDict[str, dict[str, Any]] = OrderedDict()
self._max_size = max_size self._max_size = max_size
self._ttl = ttl_seconds self._ttl = ttl_seconds
self._lock = Lock()
def put(self, msg_id: str, jid: str, content: str, metadata: dict[str, Any] | None = None) -> None: def put(self, msg_id: str, jid: str, content: str, metadata: dict[str, Any] | None = None) -> None:
if not msg_id: if not msg_id:
return return
self._cache[msg_id] = { with self._lock:
"jid": jid, self._cache[msg_id] = {
"content": content[:500], "jid": jid,
"timestamp": time.time(), "content": content[:500],
"metadata": metadata or {}, "timestamp": time.time(),
} "metadata": metadata or {},
self._cache.move_to_end(msg_id) }
self._evict() self._cache.move_to_end(msg_id)
self._evict()
def get(self, msg_id: str) -> dict[str, Any] | None: def get(self, msg_id: str) -> dict[str, Any] | None:
entry = self._cache.get(msg_id) with self._lock:
if entry is None: entry = self._cache.get(msg_id)
return None if entry is None:
if time.time() - entry["timestamp"] > self._ttl: return None
self._cache.pop(msg_id, None) if time.time() - entry["timestamp"] > self._ttl:
return None self._cache.pop(msg_id, None)
return entry return None
return entry
def lookup_inbound_meta(self, target_jid: str, target_msg_id: str | None = None) -> dict[str, Any] | None: def lookup_inbound_meta(self, target_jid: str, target_msg_id: str | None = None) -> dict[str, Any] | None:
if target_msg_id: with self._lock:
exact = self.get(target_msg_id) if target_msg_id:
if exact: exact = self._cache.get(target_msg_id)
return exact if exact and time.time() - exact["timestamp"] <= self._ttl:
return exact
target_is_group = _is_group_jid(target_jid) target_is_group = _is_group_jid(target_jid)
candidates: list[tuple[str, dict[str, Any], float]] = [] candidates: list[tuple[str, dict[str, Any], float]] = []
now = time.time() now = time.time()
for msg_id, entry in self._cache.items(): for msg_id, entry in self._cache.items():
if now - entry["timestamp"] > self._ttl: if now - entry["timestamp"] > self._ttl:
continue continue
entry_jid = entry.get("jid", "") entry_jid = entry.get("jid", "")
if target_is_group != _is_group_jid(entry_jid): if target_is_group != _is_group_jid(entry_jid):
continue continue
candidates.append((msg_id, entry, entry["timestamp"])) candidates.append((msg_id, entry, entry["timestamp"]))
if len(candidates) >= 10: if len(candidates) >= 10:
break break
if not candidates: if not candidates:
return None return None
candidates.sort(key=lambda x: x[2], reverse=True) candidates.sort(key=lambda x: x[2], reverse=True)
best_msg_id, best_entry, _ = candidates[0] best_msg_id, best_entry, _ = candidates[0]
logger.debug( logger.debug(
f"InboundMessageCache: fuzzy match for {target_jid} -> msg_id={best_msg_id}, candidates={len(candidates)}" f"InboundMessageCache: fuzzy match for {target_jid} "
) f"-> msg_id={best_msg_id}, candidates={len(candidates)}"
return best_entry )
return best_entry
def resolve_quoted_message_key(self, target_jid: str, quoted_msg_id: str | None = None) -> dict[str, Any] | None: def resolve_quoted_message_key(self, target_jid: str, quoted_msg_id: str | None = None) -> dict[str, Any] | None:
meta = None meta = None
@ -98,7 +104,9 @@ class InboundMessageCache:
del self._cache[oldest] del self._cache[oldest]
def clear(self) -> None: def clear(self) -> None:
self._cache.clear() with self._lock:
self._cache.clear()
def __len__(self) -> int: def __len__(self) -> int:
return len(self._cache) with self._lock:
return len(self._cache)

View File

@ -0,0 +1,113 @@
from __future__ import annotations
from yuxi.utils.logging_config import logger
from .inbound_pipeline import InboundPipeline, PipelineAction, PipelineResult
from .security import DmPolicy
class InboundProcessor:
def __init__(
self,
echo_filter,
deduplicator,
button_deduplicator,
security,
self_chat_mode: bool,
per_group_config,
per_dm_config,
pairing,
self_jid: str | None = None,
):
self._echo_filter = echo_filter
self._deduplicator = deduplicator
self._button_deduplicator = button_deduplicator
self._security = security
self._self_chat_mode = self_chat_mode
self._per_group_config = per_group_config
self._per_dm_config = per_dm_config
self._pairing = pairing
self._self_jid = self_jid
def build_pipeline(self) -> InboundPipeline:
pipeline = InboundPipeline()
pipeline.add_step(self._check_echo, "echo_filter")
pipeline.add_step(self._check_dedupe, "dedupe")
pipeline.add_step(self._check_button_dedupe, "button_dedupe")
pipeline.add_step(self._check_security, "security")
return pipeline
def _check_echo(self, payload: dict) -> PipelineResult:
key = payload.get("key", {})
msg = payload.get("message", {})
remote_jid = key.get("remoteJid", "")
text = msg.get("conversation", "")
if not text:
text = (msg.get("extendedTextMessage") or {}).get("text", "")
if self._echo_filter.is_echo(remote_jid, text):
return PipelineResult(action=PipelineAction.DROP, reason="echo")
return PipelineResult(action=PipelineAction.ACCEPT)
def _check_dedupe(self, payload: dict) -> PipelineResult:
msg_id = payload.get("key", {}).get("id", "")
if not msg_id:
return PipelineResult(action=PipelineAction.ACCEPT)
if self._deduplicator.is_duplicate(msg_id):
return PipelineResult(action=PipelineAction.DROP, reason="duplicate")
return PipelineResult(action=PipelineAction.ACCEPT)
def _check_button_dedupe(self, payload: dict) -> PipelineResult:
msg = payload.get("message", {})
button_msg = msg.get("buttonsResponseMessage") or msg.get("templateButtonReplyMessage")
if not button_msg:
return PipelineResult(action=PipelineAction.ACCEPT)
sender = payload.get("key", {}).get("remoteJid", "")
button_id = button_msg.get("selectedButtonId", "")
if not button_id:
button_id = button_msg.get("selectedId", "")
if not button_id:
return PipelineResult(action=PipelineAction.ACCEPT)
if self._button_deduplicator.is_duplicate(sender, button_id):
return PipelineResult(action=PipelineAction.DROP, reason="button_duplicate")
return PipelineResult(action=PipelineAction.ACCEPT)
def _check_security(self, payload: dict) -> PipelineResult:
from_me = payload.get("key", {}).get("fromMe", False)
if from_me:
if self._self_chat_mode:
return PipelineResult(action=PipelineAction.ACCEPT)
return PipelineResult(action=PipelineAction.DROP, reason="self_message")
remote_jid = payload.get("key", {}).get("remoteJid", "")
sender = remote_jid.split("@")[0]
if "@g.us" in remote_jid:
if not self._per_group_config.is_enabled(remote_jid):
return PipelineResult(action=PipelineAction.DROP, reason="group_disabled")
if self._per_group_config.require_mention(remote_jid):
msg = payload.get("message", {})
ext_text = msg.get("extendedTextMessage", {})
context_info = ext_text.get("contextInfo", {})
mentioned_jids = context_info.get("mentionedJid", []) or []
if not any(self._self_jid and mj == self._self_jid for mj in mentioned_jids):
return PipelineResult(action=PipelineAction.DROP, reason="require_mention_not_met")
allowed, reason = self._security.check_group_access(remote_jid)
if not allowed:
logger.debug(f"[WhatsApp] Group access denied for {remote_jid}: {reason}")
return PipelineResult(action=PipelineAction.DROP, reason=reason or "group_blocked")
else:
if not self._per_dm_config.is_enabled(sender):
return PipelineResult(action=PipelineAction.DROP, reason="dm_disabled")
if self._security.dm_policy == DmPolicy.PAIRING and not self._pairing.is_paired(sender):
logger.debug(f"[WhatsApp] DM pairing required for {sender}")
return PipelineResult(action=PipelineAction.DROP, reason="dm_not_paired")
allowed, reason = self._security.check_dm_access(sender)
if not allowed:
logger.debug(f"[WhatsApp] DM access denied for {sender}: {reason}")
return PipelineResult(action=PipelineAction.DROP, reason=reason or "dm_blocked")
return PipelineResult(action=PipelineAction.ACCEPT)

View File

@ -51,14 +51,20 @@ def markdown_to_whatsapp(md_text: str) -> str:
text = _convert_tables(text) text = _convert_tables(text)
text = re.sub(r"\*\*\*(.+?)\*\*\*", r"*_\1_*", text) # All formatting -> temp markers to prevent cross-contamination
text = re.sub(r"___(.+?)___", r"*_\1_*", text) text = re.sub(r"\*\*\*(.+?)\*\*\*", r"<BI>\1</BI>", text)
text = re.sub(r"___(.+?)___", r"<BI>\1</BI>", text)
text = re.sub(r"\*\*(.+?)\*\*", r"<B>\1</B>", text)
text = re.sub(r"__(.+?)__", r"<B>\1</B>", text)
text = re.sub(r"(?<!\*)\*(?!\*)(.+?)(?<!\*)\*(?!\*)", r"<I>\1</I>", text)
text = re.sub(r"(?<!_)_(?!_)(.+?)(?<!_)_(?!_)", r"<I>\1</I>", text)
text = re.sub(r"\*\*(.+?)\*\*", r"*\1*", text) text = text.replace("<BI>", "*_")
text = re.sub(r"__(.+?)__", r"*\1*", text) text = text.replace("</BI>", "_*")
text = text.replace("<B>", "*")
text = re.sub(r"(?<!\*)\*(?!\*)(.+?)(?<!\*)\*(?!\*)", r"_\1_", text) text = text.replace("</B>", "*")
text = re.sub(r"(?<!_)_(?!_)(.+?)(?<!_)_(?!_)", r"_\1_", text) text = text.replace("<I>", "_")
text = text.replace("</I>", "_")
text = re.sub(r"~~(.+?)~~", r"~\1~", text) text = re.sub(r"~~(.+?)~~", r"~\1~", text)

View File

@ -13,6 +13,8 @@ from yuxi.channels.models import (
) )
from yuxi.utils.logging_config import logger from yuxi.utils.logging_config import logger
from .session import _extract_sender_number
def normalize_poll_input(poll: dict[str, Any], max_options: int = 12) -> dict[str, Any] | None: def normalize_poll_input(poll: dict[str, Any], max_options: int = 12) -> dict[str, Any] | None:
name = poll.get("name", "").strip() name = poll.get("name", "").strip()
@ -25,6 +27,14 @@ def normalize_poll_input(poll: dict[str, Any], max_options: int = 12) -> dict[st
options = [o.strip() for o in options.split(",") if o.strip()] options = [o.strip() for o in options.split(",") if o.strip()]
options = [o.strip() for o in options if o.strip()] options = [o.strip() for o in options if o.strip()]
seen = set()
unique_options = []
for o in options:
if o not in seen:
seen.add(o)
unique_options.append(o)
options = unique_options
if len(options) < 2: if len(options) < 2:
logger.warning("Poll normalization: less than 2 valid options") logger.warning("Poll normalization: less than 2 valid options")
return None return None
@ -64,8 +74,8 @@ def parse_poll_vote(raw_payload: dict[str, Any], channel_id: str) -> ChannelMess
selected = vote_info.get("selectedOptions", []) selected = vote_info.get("selectedOptions", [])
if isinstance(selected, list): if isinstance(selected, list):
selected_options = selected selected_options = selected
elif isinstance(vote_info, list): elif isinstance(vote_info, list):
selected_options = vote_info selected_options = vote_info
option_names = [o.get("name", str(o)) if isinstance(o, dict) else str(o) for o in selected_options] option_names = [o.get("name", str(o)) if isinstance(o, dict) else str(o) for o in selected_options]
vote_text = f"[Poll Vote] {' | '.join(option_names)}" if option_names else "[Poll Vote]" vote_text = f"[Poll Vote] {' | '.join(option_names)}" if option_names else "[Poll Vote]"
@ -107,7 +117,3 @@ def _jid_to_chat_type(jid: str) -> ChatType:
if "@g.us" in jid: if "@g.us" in jid:
return ChatType.GROUP return ChatType.GROUP
return ChatType.DIRECT return ChatType.DIRECT
def _extract_sender_number(jid: str) -> str:
return jid.split("@")[0]

View File

@ -1,5 +1,5 @@
from yuxi.channels.adapters.whatsapp.reactions.reaction_level import ReactionLevel, ReactionLevelController
from yuxi.channels.adapters.whatsapp.reactions.ack_reaction import AckReactionManager from yuxi.channels.adapters.whatsapp.reactions.ack_reaction import AckReactionManager
from yuxi.channels.adapters.whatsapp.reactions.reaction_level import ReactionLevel, ReactionLevelController
__all__ = [ __all__ = [
"ReactionLevel", "ReactionLevel",

View File

@ -0,0 +1,429 @@
from __future__ import annotations
import asyncio
import os
import tempfile
import time
from typing import Any
import aiohttp
from yuxi.channels.infra.circuit_breaker import CircuitBreaker, CircuitBreakerOpenError
from yuxi.channels.models import DeliveryResult, ChannelResponse
from yuxi.utils.logging_config import logger
from .format import format_outbound as _format_outbound
from .markdown import markdown_to_whatsapp, text_sanitizer
from .media import _MEDIA_SUFFIX_MAP, cleanup_temp_file
from .media import download_media as _download_media
from .network_errors import ErrorCategory, classify_error
from .send import chunk_message
class SendManager:
def __init__(
self,
bridge,
circuit_breaker: CircuitBreaker,
sent_message_cache,
echo_filter,
debouncer,
reaction_level,
error_policy,
lane_delivery,
config: dict[str, Any],
stream_buffers: dict[str, str],
stream_chunk_counters: dict[str, int],
stream_status_msg_ids: dict[str, str],
recent_messages: dict[str, dict[str, Any]],
last_message_tracker: list[float | None],
):
self._bridge = bridge
self._circuit_breaker = circuit_breaker
self._sent_message_cache = sent_message_cache
self._echo_filter = echo_filter
self._debouncer = debouncer
self._reaction_level = reaction_level
self._error_policy = error_policy
self._lane_delivery = lane_delivery
self._config = config
self._stream_buffers = stream_buffers
self._stream_chunk_counters = stream_chunk_counters
self._stream_status_msg_ids = stream_status_msg_ids
self._recent_messages = recent_messages
self._last_message_tracker = last_message_tracker
self._prefix = config.get("messagePrefix", "")
self._chunk_mode = config.get("chunkMode", "length")
self._reply_to_mode = config.get("replyToMode", "first")
self._block_streaming = config.get("blockStreaming", False)
self._default_to = config.get("defaultTo")
self.text_chunk_limit = 4000
self.supports_markdown = True
@property
def last_message_at(self) -> float | None:
return self._last_message_tracker[0] if self._last_message_tracker else None
@last_message_at.setter
def last_message_at(self, value: float | None) -> None:
if self._last_message_tracker:
self._last_message_tracker[0] = value
async def send(self, response: ChannelResponse, silent: bool = False) -> DeliveryResult:
jid = response.identity.channel_chat_id
if not jid and self._default_to:
jid = self._default_to
reply_to = response.reply_to_message_id
if self._reply_to_mode == "off":
reply_to = None
if self._debouncer.should_throttle(jid):
remaining = self._debouncer.window_remaining(jid)
logger.warning(f"[WhatsApp] Throttled message to {jid} (retry in {remaining:.1f}s)")
return DeliveryResult(
success=False,
error=f"Rate limited, retry in {remaining:.1f}s",
)
payload = _format_outbound(response)
content = payload.get("content", response.content)
content = text_sanitizer(content)
if self.supports_markdown and getattr(response, "content_format", "") == "markdown":
try:
content = markdown_to_whatsapp(content)
except Exception as e:
logger.warning(f"[WhatsApp] Markdown parse failed, using plain text: {e}")
if self._prefix:
content = f"{self._prefix} {content}"
self._echo_filter.record_outbound(jid, content)
self._debouncer.record_send(jid)
self.last_message_at = time.monotonic()
async def _do_send() -> DeliveryResult:
send_content = content
if len(send_content) > self.text_chunk_limit and self._reply_to_mode != "batched":
chunks = chunk_message(send_content, self.text_chunk_limit, mode=self._chunk_mode)
results = []
for chunk in chunks:
use_reply = reply_to if self._reply_to_mode == "all" else (reply_to if len(results) == 0 else None)
result = await self._bridge.send_message(
jid=jid,
content=chunk,
reply_to=use_reply,
silent=silent,
)
results.append(result)
return results[0] if results else DeliveryResult(success=False, error="No chunks")
if self._reply_to_mode == "batched" and len(send_content) > self.text_chunk_limit:
send_content = send_content[: self.text_chunk_limit - 3] + "..."
return await self._send_with_retry(jid, send_content, reply_to, silent)
try:
return await self._circuit_breaker.call(_do_send)
except CircuitBreakerOpenError:
return DeliveryResult(success=False, error="Circuit breaker open")
async def send_media(self, chat_id: str, media_type: str, data: Any) -> DeliveryResult:
suffix = _MEDIA_SUFFIX_MAP.get(media_type, "")
if isinstance(data, bytes):
with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as f:
f.write(data)
tmp_path = f.name
elif isinstance(data, str) and os.path.exists(data):
tmp_path = data
else:
return DeliveryResult(success=False, error=f"Unsupported media data type: {type(data)}")
try:
return await self._bridge.send_media(
jid=chat_id,
media_type=media_type,
media_path=tmp_path,
)
finally:
if isinstance(data, bytes) and os.path.exists(tmp_path):
cleanup_temp_file(tmp_path)
async def send_reaction(self, chat_id: str, msg_id: str, emoji: str) -> DeliveryResult:
if not emoji:
return await self._bridge.send_reaction(
jid=chat_id,
message_id=msg_id,
emoji="",
)
if not self._reaction_level.can_send_reaction():
return DeliveryResult(success=False, error="Reactions disabled by reactionLevel config")
return await self._bridge.send_reaction(
jid=chat_id,
message_id=msg_id,
emoji=emoji,
)
async def remove_reaction(self, chat_id: str, msg_id: str) -> DeliveryResult:
return await self.send_reaction(chat_id, msg_id, "")
async def send_reaction_from_action(self, action: dict[str, Any]) -> DeliveryResult:
chat_jid = action.get("chatJid") or action.get("to") or action.get("chat_id", "")
msg_id = action.get("messageId") or action.get("msg_id", "")
emoji = action.get("emoji", "")
remove = action.get("remove", False)
participant = action.get("participant", "")
if participant and "@g.us" not in chat_jid:
chat_jid = participant
if remove or not emoji:
return await self.send_reaction(chat_jid, msg_id, "")
return await self.send_reaction(chat_jid, msg_id, emoji)
async def download_media(self, file_id: str) -> bytes:
raw = self._recent_messages.get(file_id)
if not raw:
raise ValueError(f"Message context not found for file_id: {file_id}")
remote_jid = raw.get("key", {}).get("remoteJid", "")
message = raw.get("message", {})
return await _download_media(self._bridge, remote_jid, file_id, message)
async def send_stream_chunk(self, chat_id: str, msg_id: str, chunk: str, finished: bool) -> DeliveryResult:
stream_cfg = self._config.get("streaming", {})
if not isinstance(stream_cfg, dict):
stream_cfg = {}
if finished:
status_msg_id = self._stream_status_msg_ids.pop(chat_id, None)
if status_msg_id:
try:
await self._bridge.delete_message(jid=chat_id, message_id=status_msg_id)
except Exception:
logger.debug(f"[WhatsApp] Failed to delete stream status message {status_msg_id}", exc_info=True)
self._stream_chunk_counters.pop(chat_id, None)
content = text_sanitizer(chunk)
if self.supports_markdown:
content = markdown_to_whatsapp(content)
if self._lane_delivery.reasoning_enabled:
lane_chunks = self._lane_delivery.split_lane_aware(content)
messages = self._lane_delivery.format_for_whatsapp(lane_chunks)
results = []
for msg in messages:
result = await self._bridge.send_message(jid=chat_id, content=msg)
results.append(result)
return results[0] if results else DeliveryResult(success=False, error="No messages")
return await self._bridge.send_message(
jid=chat_id,
content=content,
)
self._stream_chunk_counters[chat_id] = self._stream_chunk_counters.get(chat_id, 0) + 1
await self._bridge.send_presence(
jid=chat_id,
presence="composing",
)
stream_status_updates = stream_cfg.get("streamStatusUpdates", False)
status_interval = stream_cfg.get("streamStatusInterval", 5)
if stream_status_updates and self._stream_chunk_counters[chat_id] % status_interval == 0:
try:
status_text = stream_cfg.get("streamStatusText", "Processing...")
status_result = await self._bridge.send_message(jid=chat_id, content=status_text)
if status_result.success and status_result.message_id:
old_status = self._stream_status_msg_ids.get(chat_id)
if old_status:
try:
await self._bridge.delete_message(jid=chat_id, message_id=old_status)
except Exception:
pass
self._stream_status_msg_ids[chat_id] = status_result.message_id
except Exception:
logger.debug(f"[WhatsApp] Failed to send stream status message to {chat_id}", exc_info=True)
block_cfg = stream_cfg.get("block", {})
if not isinstance(block_cfg, dict):
block_cfg = {}
block_enabled = block_cfg.get("enabled", True)
block_coalesce = block_cfg.get("coalesce", False)
min_chars = block_cfg.get("coalesce_min_chars", 1500)
max_chars = block_cfg.get("coalesce_min_chars", 4096)
if not block_enabled:
return DeliveryResult(success=True)
if block_coalesce:
self._stream_buffers.setdefault(chat_id, "")
self._stream_buffers[chat_id] += chunk
if len(self._stream_buffers[chat_id]) >= min_chars:
buffered = self._stream_buffers.pop(chat_id, "")
content = text_sanitizer(buffered[:max_chars])
return await self._bridge.send_message(jid=chat_id, content=content)
return DeliveryResult(success=True)
if len(chunk) >= min_chars:
content = text_sanitizer(chunk[:max_chars])
return await self._bridge.send_message(jid=chat_id, content=content)
return DeliveryResult(success=True)
async def edit_message(self, chat_id: str, msg_id: str, content: str) -> DeliveryResult:
return DeliveryResult(success=False, error="WhatsApp does not support editing messages")
async def delete_message(self, chat_id: str, msg_id: str) -> DeliveryResult:
return await self._bridge.delete_message(
jid=chat_id,
message_id=msg_id,
)
async def send_poll(self, chat_id: str, name: str, options: list[str], selectable_count: int = 1) -> DeliveryResult:
if len(options) < 2:
return DeliveryResult(success=False, error="Poll requires at least 2 options")
if len(options) > 12:
return DeliveryResult(success=False, error=f"Poll supports max 12 options, got {len(options)}")
return await self._bridge.create_poll(
jid=chat_id,
name=name,
options=options,
selectable_count=selectable_count,
)
async def send_location(
self,
chat_id: str,
latitude: float,
longitude: float,
name: str = "",
address: str = "",
) -> DeliveryResult:
return await self._bridge.send_location(
jid=chat_id,
latitude=latitude,
longitude=longitude,
name=name,
address=address,
)
async def send_contact(
self,
chat_id: str,
contacts: list[dict[str, str]],
) -> DeliveryResult:
return await self._bridge.send_contact(
jid=chat_id,
contacts=contacts,
)
async def send_sticker(
self,
chat_id: str,
sticker_path: str,
reply_to: str | None = None,
) -> DeliveryResult:
return await self._bridge.send_sticker(
jid=chat_id,
sticker_path=sticker_path,
reply_to=reply_to,
)
async def send_buttons(
self,
chat_id: str,
text: str,
buttons: list[dict[str, str]],
title: str = "",
footer: str = "",
) -> DeliveryResult:
return await self._bridge.send_buttons(
jid=chat_id,
text=text,
buttons=buttons,
title=title,
footer=footer,
)
async def send_list_message(
self,
chat_id: str,
text: str,
sections: list[dict[str, Any]],
title: str = "",
footer: str = "",
button_text: str = "Select",
) -> DeliveryResult:
return await self._bridge.send_list_message(
jid=chat_id,
text=text,
sections=sections,
title=title,
footer=footer,
button_text=button_text,
)
async def _send_with_retry(
self,
jid: str,
content: str,
reply_to: str | None = None,
silent: bool = False,
max_attempts: int = 3,
) -> DeliveryResult:
last_error: str | None = None
for attempt in range(1, max_attempts + 1):
try:
result = await self._bridge.send_message(
jid=jid,
content=content,
reply_to=reply_to,
silent=silent,
)
if result.success and result.message_id:
self._sent_message_cache.put(
result.message_id,
jid,
{"content": content[:200], "reply_to": reply_to},
)
if result.success:
return result
last_error = result.error or "unknown"
except aiohttp.ClientConnectorError as e:
last_error = f"Bridge connection failed: {e}"
logger.warning(
f"[WhatsApp] send attempt {attempt}/{max_attempts} failed - bridge connection error for {jid}: {e}"
)
except TimeoutError as e:
last_error = str(e)
logger.warning(f"[WhatsApp] send attempt {attempt}/{max_attempts} failed - timeout for {jid}: {e}")
except Exception as e:
last_error = str(e)
error_cat = classify_error(str(e))
if error_cat in (ErrorCategory.PERMANENT, ErrorCategory.AUTH):
logger.warning(
f"[WhatsApp] send failed for {jid} - {error_cat.value} error (not retrying): {e}"
)
break
logger.warning(f"[WhatsApp] send attempt {attempt}/{max_attempts} failed for {jid}: {e}")
if attempt < max_attempts:
error_cat = classify_error(last_error or "")
if error_cat == ErrorCategory.RATE_LIMIT:
delay_ms = 2000 * attempt
else:
delay_ms = 500 * attempt
logger.info(f"[WhatsApp] Retrying send to {jid} in {delay_ms}ms (attempt {attempt})")
await asyncio.sleep(delay_ms / 1000)
logger.error(f"[WhatsApp] send failed after {max_attempts} attempts for {jid}")
if self._error_policy.should_notify(jid, last_error or "send_failed"):
last_error = self._error_policy.format_error_message(last_error or "send failed")
return DeliveryResult(success=False, error=last_error or "send failed")

View File

@ -2,6 +2,7 @@ from __future__ import annotations
import time import time
from collections import OrderedDict from collections import OrderedDict
from threading import Lock
from typing import Any from typing import Any
from yuxi.utils.logging_config import logger from yuxi.utils.logging_config import logger
@ -12,26 +13,29 @@ class SentMessageCache:
self._cache: OrderedDict[str, dict[str, Any]] = OrderedDict() self._cache: OrderedDict[str, dict[str, Any]] = OrderedDict()
self._max_size = max_size self._max_size = max_size
self._ttl = ttl_seconds self._ttl = ttl_seconds
self._lock = Lock()
def put(self, msg_id: str, jid: str, metadata: dict[str, Any] | None = None) -> None: def put(self, msg_id: str, jid: str, metadata: dict[str, Any] | None = None) -> None:
if not msg_id: if not msg_id:
return return
self._cache[msg_id] = { with self._lock:
"jid": jid, self._cache[msg_id] = {
"timestamp": time.time(), "jid": jid,
"metadata": metadata or {}, "timestamp": time.time(),
} "metadata": metadata or {},
self._cache.move_to_end(msg_id) }
self._evict() self._cache.move_to_end(msg_id)
self._evict()
def get(self, msg_id: str) -> dict[str, Any] | None: def get(self, msg_id: str) -> dict[str, Any] | None:
entry = self._cache.get(msg_id) with self._lock:
if entry is None: entry = self._cache.get(msg_id)
return None if entry is None:
if time.time() - entry["timestamp"] > self._ttl: return None
self._cache.pop(msg_id, None) if time.time() - entry["timestamp"] > self._ttl:
return None self._cache.pop(msg_id, None)
return entry return None
return entry
def get_jid(self, msg_id: str) -> str | None: def get_jid(self, msg_id: str) -> str | None:
entry = self.get(msg_id) entry = self.get(msg_id)
@ -47,8 +51,10 @@ class SentMessageCache:
del self._cache[oldest] del self._cache[oldest]
def clear(self) -> None: def clear(self) -> None:
self._cache.clear() with self._lock:
self._cache.clear()
logger.info("SentMessageCache: cleared") logger.info("SentMessageCache: cleared")
def __len__(self) -> int: def __len__(self) -> int:
return len(self._cache) with self._lock:
return len(self._cache)

View File

@ -13,6 +13,9 @@ def _extract_sender_number(jid: str) -> str:
return jid.split("@")[0] return jid.split("@")[0]
extract_sender_number = _extract_sender_number
def jid_to_chat_type(jid: str) -> str: def jid_to_chat_type(jid: str) -> str:
if "@g.us" in jid: if "@g.us" in jid:
return "group" return "group"

View File

@ -1,8 +1,9 @@
from __future__ import annotations from __future__ import annotations
import aiohttp
from typing import Any from typing import Any
import aiohttp
from yuxi.utils.logging_config import logger from yuxi.utils.logging_config import logger