diff --git a/nanobot/channels/mattermost.py b/nanobot/channels/mattermost.py index 85a03510..1a229731 100644 --- a/nanobot/channels/mattermost.py +++ b/nanobot/channels/mattermost.py @@ -16,7 +16,7 @@ from nanobot.bus.queue import MessageBus from nanobot.channels.base import BaseChannel from nanobot.config.paths import get_media_dir from nanobot.config_base import Base -from nanobot.pairing import PAIRING_CODE_META_KEY, format_pairing_reply, generate_code +from nanobot.pairing import PAIRING_CODE_META_KEY, format_pairing_reply, generate_code, is_approved from nanobot.utils.helpers import safe_filename, split_message MATTERMOST_MAX_MESSAGE_LEN = 16383 @@ -93,6 +93,7 @@ class MattermostChannel(BaseChannel): self._usernames: dict[str, str] = {} self._user_emails: dict[str, str] = {} self._channel_types: dict[str, str] = {} + self._channel_team_ids: dict[str, str] = {} self._stream_posts: dict[str, str] = {} self._stream_buffers: dict[str, str] = {} self._stream_last_content: dict[str, str] = {} @@ -129,13 +130,18 @@ class MattermostChannel(BaseChannel): self._running = True self._ws_task = asyncio.create_task(self._ws_listen_loop()) + try: + await self._ws_task + finally: + self._ws_task = None async def stop(self) -> None: self._running = False - if self._ws_task: - self._ws_task.cancel() + task = self._ws_task + if task and task is not asyncio.current_task(): + task.cancel() try: - await self._ws_task + await task except asyncio.CancelledError: pass self._ws_task = None @@ -212,8 +218,10 @@ class MattermostChannel(BaseChannel): is_dm = channel_type == "dm" team_id = broadcast.get("team_id", "") - if self.config.team_id and team_id and team_id != self.config.team_id: - if not is_dm: + if self.config.team_id and not is_dm: + if not team_id: + team_id = await self.resolve_channel_team_id(channel_id) + if team_id != self.config.team_id: return if not await self._is_allowed(sender_id, channel_id, channel_type): @@ -241,9 +249,7 @@ class MattermostChannel(BaseChannel): thread_ts = root_id if root_id else None if self.config.reply_in_thread and not thread_ts and not is_dm: thread_ts = post_id - session_key = ( - f"mattermost:{channel_id}:{thread_ts}" if thread_ts and root_id else None - ) + session_key = f"mattermost:{channel_id}:{thread_ts}" if thread_ts else None try: await self._add_reaction(channel_id, post_id, self.config.react_emoji) @@ -296,6 +302,10 @@ class MattermostChannel(BaseChannel): return channel_type = await self.resolve_channel_type(channel_id) + if self.config.team_id and channel_type != "dm": + team_id = await self.resolve_channel_team_id(channel_id) + if team_id != self.config.team_id: + return if not await self._is_allowed(sender_id, channel_id, channel_type): return @@ -334,6 +344,8 @@ class MattermostChannel(BaseChannel): if channel_type == "dm": if not self.config.dm.enabled: return False + if is_approved(self.name, str(sender_id)): + return True if self.config.dm.policy == "allowlist": return await self._match_sender(sender_id, self.config.dm.allow_from) return True @@ -494,54 +506,91 @@ class MattermostChannel(BaseChannel): # Streaming ----------------------------------------------------------------- - async def send_delta(self, chat_id: str, delta: str, metadata: dict[str, Any] | None = None) -> None: + async def send_delta( + self, + chat_id: str, + delta: str, + metadata: dict[str, Any] | None = None, + *, + stream_id: str | None = None, + stream_end: bool = False, + resuming: bool = False, + ) -> None: if not self._http_client: return meta = metadata or {} - stream_id = meta.get("_stream_id", chat_id) + stream_id = stream_id or meta.get("_stream_id") or chat_id + stream_end = stream_end or bool(meta.get("_stream_end")) + resuming = resuming or bool(meta.get("_resuming")) - if meta.get("_stream_end"): - stream_root = self._stream_root_ids.pop(stream_id, None) - self._stream_posts.pop(stream_id, None) - self._stream_last_content.pop(stream_id, None) - buf = self._stream_buffers.pop(stream_id, None) - final = self._stream_committed.pop(stream_id, None) or buf + if stream_end: + committed = self._stream_committed.get(stream_id, "") + buf = self._stream_buffers.get(stream_id, "") + final = committed or buf + if delta: + final += delta - if not meta.get("_progress") and meta.get("message_id"): + if resuming: + self._clear_stream_state(stream_id) + return + + if final and not meta.get("_progress"): + mm_meta = (meta.get("mattermost", {}) or {}) if isinstance(meta.get("mattermost"), dict) else {} + root_id = ( + mm_meta.get("root_id") + or mm_meta.get("thread_ts") + or meta.get("root_id") + or self._stream_root_ids.get(stream_id) + ) + chunks = split_message(final, MATTERMOST_MAX_MESSAGE_LEN) + first_post_id: str | None = None try: - await self._remove_reaction(meta["message_id"], self.config.react_emoji) - except Exception: - self.logger.debug("remove reaction failed") - - if final and not meta.get("_progress") and not meta.get("_resuming"): - try: - mm_meta = (meta.get("mattermost", {}) or {}) if isinstance(meta.get("mattermost"), dict) else {} - root_id = mm_meta.get("root_id") or mm_meta.get("thread_ts") or meta.get("root_id") or stream_root - chunks = split_message(final, MATTERMOST_MAX_MESSAGE_LEN) - for i, chunk in enumerate(chunks): + for chunk in chunks: post = await self._create_post( chat_id, chunk, root_id=root_id if self.config.reply_in_thread else None, ) - if i == 0 and self.config.done_emoji: - try: - await self._add_reaction(chat_id, post["id"], self.config.done_emoji) - except Exception: - self.logger.debug("done reaction failed") + if first_post_id is None: + first_post_id = post.get("id") except Exception: self.logger.exception("stream final post failed") + raise + + if meta.get("message_id"): + try: + await self._remove_reaction(meta["message_id"], self.config.react_emoji) + except Exception: + self.logger.debug("remove reaction failed") + if first_post_id and self.config.done_emoji: + try: + await self._add_reaction(chat_id, first_post_id, self.config.done_emoji) + except Exception: + self.logger.debug("done reaction failed") + + self._clear_stream_state(stream_id) return if not delta.strip(): return + mm_meta = (meta.get("mattermost", {}) or {}) if isinstance(meta.get("mattermost"), dict) else {} + root_id = mm_meta.get("root_id") or mm_meta.get("thread_ts") or meta.get("root_id") + if root_id: + self._stream_root_ids[stream_id] = root_id committed = self._stream_committed.get(stream_id, "") buf = committed + delta self._stream_buffers[stream_id] = buf self._stream_committed[stream_id] = buf return + def _clear_stream_state(self, stream_id: str) -> None: + self._stream_root_ids.pop(stream_id, None) + self._stream_posts.pop(stream_id, None) + self._stream_buffers.pop(stream_id, None) + self._stream_last_content.pop(stream_id, None) + self._stream_committed.pop(stream_id, None) + # API helpers --------------------------------------------------------------- async def _api_get(self, path: str) -> dict[str, Any]: @@ -644,6 +693,21 @@ class MattermostChannel(BaseChannel): data = await self._api_get(f"/api/v4/channels/{channel_id}") ctype = _CHANNEL_TYPES.get(data.get("type", ""), "public") self._channel_types[channel_id] = ctype + if "team_id" in data: + self._channel_team_ids[channel_id] = data.get("team_id", "") or "" return ctype except Exception: return "public" + + async def resolve_channel_team_id(self, channel_id: str) -> str: + if channel_id in self._channel_team_ids: + return self._channel_team_ids[channel_id] + try: + data = await self._api_get(f"/api/v4/channels/{channel_id}") + team_id = data.get("team_id", "") or "" + self._channel_team_ids[channel_id] = team_id + if "type" in data: + self._channel_types[channel_id] = _CHANNEL_TYPES.get(data.get("type", ""), "public") + return team_id + except Exception: + return "" diff --git a/tests/channels/test_mattermost_channel.py b/tests/channels/test_mattermost_channel.py index 59566a84..b3d7671a 100644 --- a/tests/channels/test_mattermost_channel.py +++ b/tests/channels/test_mattermost_channel.py @@ -2,6 +2,7 @@ from __future__ import annotations +import asyncio import json from typing import Any from unittest.mock import AsyncMock, patch @@ -155,14 +156,29 @@ def test_config_default_config_classmethod(): async def test_start_identifies_bot(): channel, fake = _make_channel({"serverUrl": "https://chat.example.com", "token": "tok"}) calls_before = len(fake.get_calls) - with patch("websockets.connect", AsyncMock(side_effect=Exception("no-op"))): - await channel.start() + + async def fake_listen_loop(): + while channel._running: + await asyncio.sleep(0.01) + + with patch.object(channel, "_ws_listen_loop", fake_listen_loop): + start_task = asyncio.create_task(channel.start()) + for _ in range(50): + if channel._self_id: + break + await asyncio.sleep(0.01) + assert channel._self_id == "botuserid123" assert channel._self_username == "nanobot" assert channel._self_email == "bot@example.com" + assert not start_task.done() user_me_calls = [c for c in fake.get_calls[calls_before:] if "/api/v4/users/me" in c["path"]] assert len(user_me_calls) == 1 await channel.stop() + try: + await start_task + except asyncio.CancelledError: + pass @pytest.mark.asyncio @@ -540,6 +556,65 @@ async def test_stream_chunk_boundary_finalizes_and_creates_new(): assert channel._stream_buffers["s1"] == "Hello world" +@pytest.mark.asyncio +async def test_stream_end_keyword_resuming_does_not_post_or_mark_done(): + channel, fake = _make_channel() + channel._self_id = "bot_id" + await channel.send_delta("chan_1", "Working", stream_id="s1") + + await channel.send_delta( + "chan_1", + "", + {"message_id": "orig_post_1"}, + stream_id="s1", + stream_end=True, + resuming=True, + ) + + posts = [c for c in fake.post_calls if c["path"] == "/api/v4/posts"] + reactions = [c for c in fake.post_calls if c["path"] == "/api/v4/reactions"] + assert posts == [] + assert reactions == [] + assert "s1" not in channel._stream_buffers + + +@pytest.mark.asyncio +async def test_stream_end_failure_keeps_buffer_for_retry(): + channel, fake = _make_channel() + channel._self_id = "bot_id" + await channel.send_delta("chan_1", "final answer", stream_id="s1") + + async def fail_create_post(*args, **kwargs): + raise RuntimeError("network down") + + channel._create_post = fail_create_post + with pytest.raises(RuntimeError): + await channel.send_delta("chan_1", "", stream_id="s1", stream_end=True) + + assert channel._stream_buffers["s1"] == "final answer" + assert channel._stream_committed["s1"] == "final answer" + + +@pytest.mark.asyncio +async def test_coalesced_stream_end_posts_inline_content(): + channel, fake = _make_channel() + channel._self_id = "bot_id" + fake.set_post_response("/api/v4/posts", {"id": "stream_post_1"}) + + await channel.send_delta( + "chan_1", + "coalesced final", + {"mattermost": {"root_id": "root_1"}}, + stream_id="s1", + stream_end=True, + ) + + posts = [c for c in fake.post_calls if c["path"] == "/api/v4/posts"] + assert len(posts) == 1 + assert posts[0]["json"]["message"] == "coalesced final" + assert posts[0]["json"]["root_id"] == "root_1" + + # --------------------------------------------------------------------------- # Reactions # --------------------------------------------------------------------------- @@ -632,6 +707,31 @@ async def test_team_filtering_dm_bypass(): mock_handle.assert_awaited_once() +@pytest.mark.asyncio +async def test_team_filtering_resolves_missing_broadcast_team_and_rejects_wrong_team(): + channel, fake = _make_channel({"teamId": "team_a"}) + channel._self_id = "bot_id" + fake.set_get_response("/api/v4/channels/c1", { + "id": "c1", + "type": "O", + "team_id": "team_b", + }) + with patch.object(channel, "_handle_message", AsyncMock()) as mock_handle: + ws_msg = { + "event": "posted", + "data": { + "channel_type": "O", + "post": json.dumps({ + "id": "p1", "user_id": "u1", + "channel_id": "c1", "message": "hi", "root_id": "", + }), + }, + "broadcast": {"channel_id": "c1"}, + } + await channel._handle_ws_message(ws_msg) + mock_handle.assert_not_awaited() + + # --------------------------------------------------------------------------- # Thread session key # --------------------------------------------------------------------------- @@ -661,6 +761,31 @@ async def test_thread_session_key(): assert kwargs["session_key"] == "mattermost:c1:root_99" +@pytest.mark.asyncio +async def test_top_level_mention_uses_thread_session_key(): + channel, fake = _make_channel({"replyInThread": True}) + channel._self_id = "bot_id" + channel._self_username = "nanobot" + with patch.object(channel, "_handle_message", AsyncMock()) as mock_handle: + with patch.object(channel, "_is_allowed", AsyncMock(return_value=True)): + ws_msg = { + "event": "posted", + "data": { + "channel_type": "O", + "post": json.dumps({ + "id": "post_1", "user_id": "u1", + "channel_id": "c1", "message": "@nanobot start thread", + "root_id": "", + }), + }, + "broadcast": {}, + } + await channel._handle_ws_message(ws_msg) + kwargs = mock_handle.call_args[1] + assert kwargs["session_key"] == "mattermost:c1:post_1" + assert kwargs["metadata"]["mattermost"]["thread_ts"] == "post_1" + + # --------------------------------------------------------------------------- # Action event (interactive buttons) # --------------------------------------------------------------------------- @@ -709,6 +834,29 @@ async def test_action_event_denied_dm(): mock_handle.assert_not_awaited() +@pytest.mark.asyncio +async def test_action_event_rejects_wrong_team(): + channel, fake = _make_channel({"teamId": "team_a"}) + channel._self_id = "bot_id" + fake.set_get_response("/api/v4/channels/c1", { + "id": "c1", + "type": "O", + "team_id": "team_b", + }) + with patch.object(channel, "_handle_message", AsyncMock()) as mock_handle: + ws_msg = { + "event": "action", + "data": { + "user_id": "u1", + "channel_id": "c1", + "context": {"selected_option": "Approve"}, + }, + "broadcast": {}, + } + await channel._handle_ws_message(ws_msg) + mock_handle.assert_not_awaited() + + # --------------------------------------------------------------------------- # Post deleted event # --------------------------------------------------------------------------- @@ -763,6 +911,13 @@ async def test_dm_allowlist_with_username_match(): assert await channel._is_allowed("u2", "dm_chan", "dm") is False +@pytest.mark.asyncio +async def test_dm_allowlist_accepts_pairing_approval(): + channel, fake = _make_channel({"dm": {"policy": "allowlist", "allowFrom": ["u_allowed"]}}) + with patch("nanobot.channels.mattermost.is_approved", return_value=True): + assert await channel._is_allowed("u_paired", "dm_chan", "dm") is True + + # --------------------------------------------------------------------------- # Denied DM sends pairing code (not empty message) # ---------------------------------------------------------------------------