diff --git a/backend/package/yuxi/channels/adapters/imessage/adapter.py b/backend/package/yuxi/channels/adapters/imessage/adapter.py index 3f16451a..2386e58e 100644 --- a/backend/package/yuxi/channels/adapters/imessage/adapter.py +++ b/backend/package/yuxi/channels/adapters/imessage/adapter.py @@ -31,16 +31,19 @@ from yuxi.utils.datetime_utils import utc_now_naive from yuxi.utils.logging_config import logger from .accounts import list_enabled_accounts +from .audit import AuditLogger from .bridge_client import BlueBubblesClient from .commands import is_authorized_for_commands, parse_command +from .config_schema import validate_config from .contact_resolver import resolve_contact from .echo_cache import EchoCache -from .envelope import format_inbound_envelope from .formatter import convert_markdown_tables, sanitize_imessage_text +from .media_ai import describe_image, transcribe_audio from .media_security import validate_attachment_path -from .normalize import parse_forwarded from .monitor import IMessageMonitor +from .normalize import parse_forwarded from .pairing import IMessagePairingManager +from .polls import IMessagePollManager from .probe import run_doctor_diagnosis from .rate_limiter import LoopRateLimiter from .reflection_guard import ReflectionGuard @@ -77,7 +80,7 @@ class IMessageAdapter(BaseChannelAdapter): edit=True, unsend=True, media=True, - polls=False, + polls=True, threads=False, pin=False, unpin=False, @@ -102,6 +105,11 @@ class IMessageAdapter(BaseChannelAdapter): self.config.setdefault("password", os.getenv("IMESSAGE_PASSWORD", "")) self._account_id = self.config.get("accountId", self.config.get("account_id", "default")) + try: + validate_config(self.config) + except Exception as e: + raise ValueError(f"iMessage configuration validation failed: {e}") from e + super().__init__(self.config) self._block_streaming = self.config.get("blockStreaming", self.config.get("block_streaming", False)) @@ -151,6 +159,8 @@ class IMessageAdapter(BaseChannelAdapter): self._pairing = IMessagePairingManager( config_writes=self.config.get("configWrites", self.config.get("config_writes", True)) ) + self._audit_logger = AuditLogger(log_dir=self.config.get("auditDir", self.config.get("audit_dir", None))) + self._poll_manager = IMessagePollManager(adapter=self) self._echo_cache = EchoCache() self._sticker_cache = StickerCache() self._last_inbound: dict[str, tuple[float, str]] = {} @@ -164,6 +174,19 @@ class IMessageAdapter(BaseChannelAdapter): self._merge_buffers: dict[str, list[ChannelMessage]] = {} self._merge_tasks: dict[str, asyncio.Task] = {} + media_ai_cfg = self.config.get("mediaAi", self.config.get("media_ai", {})) + self._media_ai_enabled = media_ai_cfg.get("enabled", False) + self._media_ai_config = { + "vision_api_url": media_ai_cfg.get("visionApiUrl", media_ai_cfg.get("vision_api_url", "")), + "vision_api_key": media_ai_cfg.get("visionApiKey", media_ai_cfg.get("vision_api_key", "")), + "transcription_api_url": media_ai_cfg.get( + "transcriptionApiUrl", media_ai_cfg.get("transcription_api_url", "") + ), + "transcription_api_key": media_ai_cfg.get( + "transcriptionApiKey", media_ai_cfg.get("transcription_api_key", "") + ), + } + def _build_account_config(self, acct) -> dict[str, Any]: return { "server_url": acct.server_url, @@ -256,6 +279,7 @@ class IMessageAdapter(BaseChannelAdapter): await monitor.stop() for client in self._clients.values(): await client.close() + await self._audit_logger.flush() self._status = ChannelStatus.DISCONNECTED self._pending_messages.clear() @@ -277,28 +301,67 @@ class IMessageAdapter(BaseChannelAdapter): disable_notification=disable_notification, ) else: - chunks = _chunk_text(content, self.text_chunk_limit) - last_result = DeliveryResult(success=True) - for i, chunk in enumerate(chunks): - chunk_reply_to = reply_to if i == 0 else None + result = await self._send_chunked( + chat_guid=chat_guid, + content=content, + reply_to=reply_to, + ) + + if result.success and result.message_id: + self._pending_messages[result.message_id] = content + self._echo_cache.record_sent(result.message_id, content, chat_guid) + return result + + async def _send_chunked( + self, + chat_guid: str, + content: str, + reply_to: str | None = None, + max_retries: int = 2, + ) -> DeliveryResult: + chunks = _chunk_text(content, self.text_chunk_limit) + results: list[DeliveryResult] = [] + + for i, chunk in enumerate(chunks): + chunk_reply_to = reply_to if i == 0 else None + + chunk_result = None + for attempt in range(max_retries + 1): chunk_result = await self._get_client().send_text_message( chat_guid=chat_guid, content=chunk, reply_to=chunk_reply_to, is_html=self.supports_markdown, ) - if not chunk_result.success: - if i == 0: - return chunk_result - last_result = chunk_result - else: - last_result = chunk_result - result = last_result + if chunk_result.success: + break + if attempt < max_retries: + logger.warning( + f"[iMessage] Chunk {i + 1}/{len(chunks)} send failed " + f"(attempt {attempt + 1}/{max_retries + 1}): {chunk_result.error}" + ) + await asyncio.sleep(1 * (attempt + 1)) - if result.success and result.message_id: - self._pending_messages[result.message_id] = content - self._echo_cache.record_sent(result.message_id, content, chat_guid) - return result + results.append(chunk_result) + + if not chunk_result.success: + if i == 0: + return chunk_result + logger.error( + f"[iMessage] Chunk {i + 1}/{len(chunks)} send failed after " + f"{max_retries + 1} attempts: {chunk_result.error}" + ) + + all_success = all(r.success for r in results) + if all_success: + return results[-1] + + failed_indices = [i + 1 for i, r in enumerate(results) if not r.success] + return DeliveryResult( + success=False, + error=f"Chunks {failed_indices} failed to send", + metadata={"failed_chunks": failed_indices, "total_chunks": len(chunks)}, + ) async def send_media(self, chat_id: str, media_type: str, data: Any) -> DeliveryResult: suffix_map = {"image": ".jpg", "video": ".mp4", "audio": ".caf", "file": ""} @@ -354,6 +417,9 @@ class IMessageAdapter(BaseChannelAdapter): ) async def send_sticker(self, chat_guid: str, sticker_id: str) -> DeliveryResult: + # 注意: BlueBubbles 当前不提供原生 sticker API。 + # Sticker 以 URL 文本消息形式发送,而非使用 iMessage 原生贴纸协议。 + # 待 BlueBubbles 支持原生 sticker API 后可升级此实现。 sticker_entry = self._sticker_cache.get(sticker_id) if not sticker_entry: return DeliveryResult(success=False, error=f"Sticker not found: {sticker_id}") @@ -531,21 +597,13 @@ class IMessageAdapter(BaseChannelAdapter): content, msg_type, attachments = self._extract_content(data, include_attachments=self._include_attachments) mentions = self._parse_mentions(data) - reply_context = self._parse_reply_context(data) - sender_name = await resolve_contact(self._get_client(), sender) - envelope = format_inbound_envelope( - identity=ChannelIdentity( - channel_id=self.channel_id, - channel_type=self.channel_type, - channel_user_id=self._clean_handle(sender), - channel_chat_id=chat_guid, - channel_message_id=msg_guid, - ), - sender_name=sender_name or "", - chat_name=chat_guid, - reply_context=reply_context, - ) + if self._media_ai_enabled and attachments: + ai_descriptions = await self._run_media_ai(attachments) + if ai_descriptions: + content = content + "\n" + "\n".join(ai_descriptions) + + reply_context = self._parse_reply_context(data) if reply_context: reply_tag = format_reply_context( @@ -590,6 +648,11 @@ class IMessageAdapter(BaseChannelAdapter): if chat_type == ChatType.DIRECT: access = self._security.check_dm_access(sender) if not access.allowed: + self._audit_logger.record_dm_access( + sender=sender, + allowed=False, + reason=access.reject_reason or "unknown", + ) logger.info(f"[iMessage] DM blocked: sender={sender}, reason={access.reject_reason}") if access.reply: asyncio.ensure_future( @@ -617,6 +680,8 @@ class IMessageAdapter(BaseChannelAdapter): ) return self._make_passthrough_event(chat_guid, msg_guid, sender, event_type) + self._audit_logger.record_dm_access(sender=sender, allowed=True) + if not self._pairing.is_approved(self._clean_handle(sender)): if self._security.dm_policy == "pairing": pairing_code = self._pairing.request_pairing( @@ -640,6 +705,11 @@ class IMessageAdapter(BaseChannelAdapter): elif chat_type == ChatType.GROUP: access = self._security.check_group_access(chat_guid) if not access.allowed: + self._audit_logger.record_group_access( + chat_guid=chat_guid, + allowed=False, + reason=access.reject_reason or "unknown", + ) logger.info(f"[iMessage] Group blocked: chat={chat_guid}, reason={access.reject_reason}") if access.reply: asyncio.ensure_future( @@ -650,11 +720,16 @@ class IMessageAdapter(BaseChannelAdapter): ) return self._make_passthrough_event(chat_guid, msg_guid, sender, event_type) + self._audit_logger.record_group_access(chat_guid=chat_guid, allowed=True) + mention_check = self._security.check_mention_required(chat_guid, mentions, self._self_handle or "") if not mention_check.allowed: logger.debug(f"[iMessage] Mention required but bot not mentioned in {chat_guid}") return self._make_passthrough_event(chat_guid, msg_guid, sender, event_type) + if content and content.strip().isdigit(): + await self._handle_poll_vote(chat_guid, sender, content.strip()) + return ChannelMessage( identity=ChannelIdentity( channel_id=self.channel_id, @@ -897,6 +972,65 @@ class IMessageAdapter(BaseChannelAdapter): ) return True + if action == "poll": + sub_action = args[0].lower() if args else "" + if sub_action == "create": + remaining = " ".join(args[1:]) if len(args) > 1 else "" + parts = _parse_quoted_args(remaining) + if len(parts) >= 3: + question = parts[0] + option_labels = parts[1:] + poll = await self._poll_manager.create_poll( + chat_guid=chat_guid, + question=question, + option_labels=option_labels, + expires_in_s=300.0, + ) + await self._get_client().send_text_message( + chat_guid=chat_guid, + content=f"Poll created! ID: {poll.poll_id}", + ) + else: + await self._get_client().send_text_message( + chat_guid=chat_guid, + content='Usage: /poll create "Question" "Option1" "Option2" ...', + ) + elif sub_action == "close": + poll_id = args[1] if len(args) > 1 else "" + if poll_id: + result = await self._poll_manager.close_poll(poll_id) + if result: + await self._get_client().send_text_message( + chat_guid=chat_guid, + content="Poll closed.", + ) + else: + await self._get_client().send_text_message( + chat_guid=chat_guid, + content="Usage: /poll close ", + ) + elif sub_action == "list": + active = [p for p in self._poll_manager._polls.values() if p.chat_guid == chat_guid and not p.closed] + if active: + lines = ["Active polls:"] + for p in active: + lines.append(f" {p.poll_id}: {p.question}") + await self._get_client().send_text_message( + chat_guid=chat_guid, + content="\n".join(lines), + ) + else: + await self._get_client().send_text_message( + chat_guid=chat_guid, + content="No active polls in this chat.", + ) + else: + await self._get_client().send_text_message( + chat_guid=chat_guid, + content="/poll create|close|list", + ) + return True + if action == "help": await self._get_client().send_text_message( chat_guid=chat_guid, @@ -906,6 +1040,9 @@ class IMessageAdapter(BaseChannelAdapter): "/pairing — List pending pairing requests\n" "/approve — Approve a pairing\n" "/reject — Reject a pairing\n" + '/poll create "Q" "A" "B" — Create a poll\n' + "/poll close — Close a poll\n" + "/poll list — List active polls\n" "/help — Show this help", ) return True @@ -924,6 +1061,101 @@ class IMessageAdapter(BaseChannelAdapter): def is_paired(self, channel_user_id: str) -> bool: return self._pairing.is_paired(channel_user_id) + async def create_poll( + self, + chat_id: str, + question: str, + options: list[str], + duration_minutes: int = 0, + ) -> dict[str, Any]: + poll = await self._poll_manager.create_poll( + chat_guid=chat_id, + question=question, + option_labels=options, + expires_in_s=duration_minutes * 60 if duration_minutes > 0 else 300.0, + ) + return { + "poll_id": poll.poll_id, + "chat_guid": poll.chat_guid, + "question": poll.question, + "options": [{"index": o.index + 1, "label": o.label} for o in poll.options], + } + + async def close_poll(self, chat_id: str, poll_id: str) -> dict[str, Any]: + result = await self._poll_manager.close_poll(poll_id) + if result is None: + return {"error": f"Poll not found: {poll_id}"} + return result + + async def vote_poll(self, poll_id: str, user_id: str, option_index: int) -> dict[str, Any]: + result = await self._poll_manager.vote(poll_id, user_id, option_index) + if result is None: + return {"error": f"Poll not found: {poll_id}"} + return result + + def get_poll(self, poll_id: str) -> dict[str, Any] | None: + poll = self._poll_manager.get_poll(poll_id) + if poll is None: + return None + return { + "poll_id": poll.poll_id, + "chat_guid": poll.chat_guid, + "question": poll.question, + "closed": poll.closed, + "options": [{"index": o.index + 1, "label": o.label, "vote_count": o.vote_count} for o in poll.options], + } + + async def _handle_poll_vote(self, chat_guid: str, sender: str, digit_str: str) -> None: + active_polls = [p for p in self._poll_manager._polls.values() if p.chat_guid == chat_guid and not p.closed] + if not active_polls: + return + + option_index = int(digit_str) - 1 + poll = active_polls[-1] + + result = await self._poll_manager.vote(poll.poll_id, sender, option_index) + if result and result.get("success"): + logger.info(f"[iMessage/Poll] Vote recorded: user={sender}, poll={poll.poll_id}, option={option_index + 1}") + elif result and result.get("poll_closed"): + await self._get_client().send_text_message( + chat_guid=chat_guid, + content="This poll is already closed.", + ) + + async def _run_media_ai(self, attachments: list[Attachment]) -> list[str]: + descriptions: list[str] = [] + cfg = self._media_ai_config + + for att in attachments: + try: + async with asyncio.timeout(10): + if att.type in (MessageType.IMAGE.value, "image"): + if att.url and cfg["vision_api_url"]: + image_bytes = await self._get_client().download_attachment(att.url) + desc = await describe_image( + image_bytes, + vision_api_url=cfg["vision_api_url"], + vision_api_key=cfg["vision_api_key"], + ) + if desc: + descriptions.append(f"[图片描述: {desc}]") + elif att.type in (MessageType.AUDIO.value, "audio"): + if att.url and cfg["transcription_api_url"]: + audio_bytes = await self._get_client().download_attachment(att.url) + transcript = await transcribe_audio( + audio_bytes, + transcription_api_url=cfg["transcription_api_url"], + transcription_api_key=cfg["transcription_api_key"], + ) + if transcript: + descriptions.append(f"[音频转录: {transcript}]") + except TimeoutError: + logger.warning(f"[iMessage/AI] Media AI timed out for {att.url}") + except Exception as e: + logger.warning(f"[iMessage/AI] Media AI failed for {att.url}: {e}") + + return descriptions + def get_security_warnings(self) -> list[str]: return self._security.collect_warnings() @@ -1046,3 +1278,12 @@ def _msg_type_to_media_str(msg_type) -> str: MessageType.FILE: "file", } return mapping.get(msg_type, "file") + + +def _parse_quoted_args(text: str) -> list[str]: + import shlex + + try: + return shlex.split(text) + except ValueError: + return text.split() diff --git a/backend/package/yuxi/channels/adapters/imessage/agent_tools.py b/backend/package/yuxi/channels/adapters/imessage/agent_tools.py index 625ca47e..9d7ae65d 100644 --- a/backend/package/yuxi/channels/adapters/imessage/agent_tools.py +++ b/backend/package/yuxi/channels/adapters/imessage/agent_tools.py @@ -68,6 +68,68 @@ def describe_imessage_security_status_tool() -> dict: } +def describe_imessage_poll_tool() -> dict: + return { + "type": "function", + "function": { + "name": "imessage_manage_polls", + "description": "管理 iMessage 投票:创建、关闭、查看投票", + "parameters": { + "type": "object", + "properties": { + "action": { + "type": "string", + "enum": ["create", "close", "list", "get"], + "description": "操作类型", + }, + "chat_guid": { + "type": "string", + "description": "聊天 ID(创建/列出时需要)", + }, + "question": { + "type": "string", + "description": "投票问题(创建时需要)", + }, + "options": { + "type": "array", + "items": {"type": "string"}, + "description": "选项列表(创建时需要)", + }, + "poll_id": { + "type": "string", + "description": "投票 ID(关闭/查看时需要)", + }, + "duration_minutes": { + "type": "integer", + "description": "投票持续时间(分钟),默认 5 分钟", + }, + }, + "required": ["action"], + }, + }, + } + + +def describe_imessage_audit_tool() -> dict: + return { + "type": "function", + "function": { + "name": "imessage_audit_log", + "description": "查询 iMessage 渠道审计日志", + "parameters": { + "type": "object", + "properties": { + "limit": { + "type": "integer", + "description": "返回记录数上限,默认 20", + }, + }, + "required": [], + }, + }, + } + + class IMessageAgentToolFactory: def __init__(self, adapter: Any): self._adapter = adapter @@ -77,6 +139,8 @@ class IMessageAgentToolFactory: describe_imessage_pairing_tool(), describe_imessage_allowlist_tool(), describe_imessage_security_status_tool(), + describe_imessage_poll_tool(), + describe_imessage_audit_tool(), ] async def execute_tool(self, tool_name: str, params: dict) -> Any: @@ -86,6 +150,10 @@ class IMessageAgentToolFactory: return await self._manage_allowlist(params) if tool_name == "imessage_security_status": return self._security_status() + if tool_name == "imessage_manage_polls": + return await self._manage_polls(params) + if tool_name == "imessage_audit_log": + return self._audit_log(params) raise ValueError(f"Unknown tool: {tool_name}") async def _manage_pairing(self, params: dict) -> dict: @@ -149,3 +217,57 @@ class IMessageAgentToolFactory: "warnings": warnings, "paired_count": len(self._adapter._pairing.get_paired_handles()), } + + async def _manage_polls(self, params: dict) -> dict: + action = params.get("action", "list") + + if action == "list": + active = [ + { + "poll_id": p.poll_id, + "chat_guid": p.chat_guid, + "question": p.question, + "closed": p.closed, + } + for p in self._adapter._poll_manager._polls.values() + if not p.closed + ] + return {"active_polls": active} + + if action == "create": + chat_guid = params.get("chat_guid", "") + question = params.get("question", "") + options = params.get("options", []) + duration_minutes = params.get("duration_minutes", 5) + if not chat_guid or not question or not options: + return {"success": False, "error": "chat_guid, question, and options are required"} + result = await self._adapter.create_poll( + chat_id=chat_guid, + question=question, + options=options, + duration_minutes=duration_minutes, + ) + return {"success": True, "poll": result} + + if action == "close": + poll_id = params.get("poll_id", "") + if not poll_id: + return {"success": False, "error": "poll_id is required"} + result = await self._adapter.close_poll("", poll_id) + return {"success": True, "result": result} + + if action == "get": + poll_id = params.get("poll_id", "") + if not poll_id: + return {"success": False, "error": "poll_id is required"} + result = self._adapter.get_poll(poll_id) + if result is None: + return {"success": False, "error": f"Poll not found: {poll_id}"} + return {"success": True, "poll": result} + + return {"success": False, "error": f"Unknown action: {action}"} + + def _audit_log(self, params: dict) -> dict: + limit = params.get("limit", 20) + records = self._adapter._audit_logger.get_recent(limit=limit) + return {"records": records, "count": len(records)} diff --git a/backend/package/yuxi/channels/adapters/imessage/commands.py b/backend/package/yuxi/channels/adapters/imessage/commands.py index bade5561..69e1d737 100644 --- a/backend/package/yuxi/channels/adapters/imessage/commands.py +++ b/backend/package/yuxi/channels/adapters/imessage/commands.py @@ -41,6 +41,7 @@ AVAILABLE_COMMANDS = { "approve": {"description": "Approve a pending approval request", "requires_owner": True}, "reject": {"description": "Reject a pending approval request", "requires_owner": True}, "allowlist": {"description": "Manage allowlist entries", "requires_owner": True}, + "poll": {"description": "Manage polls: create, close, list", "requires_owner": False}, "help": {"description": "Show this help message", "requires_owner": False}, } diff --git a/backend/package/yuxi/channels/adapters/imessage/formatter.py b/backend/package/yuxi/channels/adapters/imessage/formatter.py index cef6e8cf..8ff66622 100644 --- a/backend/package/yuxi/channels/adapters/imessage/formatter.py +++ b/backend/package/yuxi/channels/adapters/imessage/formatter.py @@ -3,7 +3,7 @@ from __future__ import annotations import re TABLE_ROW_RE = re.compile(r"^\|(.+)\|$") -TABLE_SEP_RE = re.compile(r"^\|[\s\-:]+\|$") +TABLE_SEP_RE = re.compile(r"^\|[\s\-:\|]+\|$") DIRECTIVE_RE = re.compile(r"<\s*/?directive[^>]*>", re.IGNORECASE) diff --git a/backend/package/yuxi/channels/adapters/imessage/media_security.py b/backend/package/yuxi/channels/adapters/imessage/media_security.py index dd99fe94..8c2fadcf 100644 --- a/backend/package/yuxi/channels/adapters/imessage/media_security.py +++ b/backend/package/yuxi/channels/adapters/imessage/media_security.py @@ -60,12 +60,19 @@ def validate_remote_attachment_url( server_url: str = "", ) -> bool: """校验远程附件 URL 是否安全。""" + from urllib.parse import urlparse + if not url: return False + parsed_url = urlparse(url) + url_host = parsed_url.hostname or "" + if allowed_roots: for root in allowed_roots: - if url.startswith(root): + root_parsed = urlparse(root) + root_host = root_parsed.hostname or "" + if url_host == root_host and url.startswith(root): return True return False diff --git a/backend/package/yuxi/channels/adapters/imessage/monitor.py b/backend/package/yuxi/channels/adapters/imessage/monitor.py index 966bd5e3..aeed41cb 100644 --- a/backend/package/yuxi/channels/adapters/imessage/monitor.py +++ b/backend/package/yuxi/channels/adapters/imessage/monitor.py @@ -10,7 +10,6 @@ import aiohttp from yuxi.utils.logging_config import logger - _MAX_ERROR_MSG_LEN = 200