fix(mattermost): harden channel lifecycle and streaming

Maintainer edit: keep the Mattermost adapter running under the gateway, fail closed when team filtering cannot verify the team, isolate new thread sessions immediately, and make buffered stream finalization retry-safe.
This commit is contained in:
chengyongru
2026-07-06 12:14:57 +08:00
committed by Xubin Ren
parent ef9780719d
commit cf35238834
2 changed files with 254 additions and 35 deletions
+97 -33
View File
@@ -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 ""
+157 -2
View File
@@ -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)
# ---------------------------------------------------------------------------