diff --git a/docs/chat-apps.md b/docs/chat-apps.md index 84522f53be0..eaf04cdb20c 100644 --- a/docs/chat-apps.md +++ b/docs/chat-apps.md @@ -382,7 +382,7 @@ nanobot plugins enable matrix | `groupAllowFrom` | Room allowlist (used when policy is `allowlist`). | | `allowRoomMentions` | Accept `@room` mentions in mention mode. | | `e2eeEnabled` | E2EE support (default `true`). Set `false` for plaintext-only. | -| `sasVerification` | Auto-complete SAS device verification requests from allowed users (default `false`). Useful for Element X, which does not expose manual trust for third-party devices. | +| `sasVerification` | Complete Element-initiated SAS device verification for allowed users (default `false`). This does not add cross-signing, clear Element's cross-signing trust warning, or let the bot initiate verification. | | `maxMediaBytes` | Max attachment size (default `20MB`). Set `0` to block all media. | diff --git a/nanobot/channels/matrix/runtime.py b/nanobot/channels/matrix/runtime.py index f0d9187ef60..8a0a3560ea8 100644 --- a/nanobot/channels/matrix/runtime.py +++ b/nanobot/channels/matrix/runtime.py @@ -47,6 +47,8 @@ SyncError, SyncResponse, ToDeviceError, + ToDeviceMessage, + UnknownToDeviceEvent, UploadError, ) from nio.crypto.attachments import decrypt_attachment @@ -75,6 +77,9 @@ _ATTACH_UPLOAD_FAILED = "[attachment: {} - upload failed]" _DEFAULT_ATTACH_NAME = "attachment" _MSGTYPE_MAP = {"m.image": "image", "m.audio": "audio", "m.video": "video", "m.file": "file"} +_SAS_METHOD = "m.sas.v1" +_SAS_REQUEST_MAX_AGE_MS = 10 * 60 * 1000 +_SAS_REQUEST_MAX_FUTURE_MS = 5 * 60 * 1000 MATRIX_MEDIA_EVENT_FILTER = (RoomMessageMedia, RoomEncryptedMedia) MatrixMediaEvent: TypeAlias = RoomMessageMedia | RoomEncryptedMedia @@ -200,6 +205,15 @@ class _StreamBuf: event_id: str | None = None last_edit: float = 0.0 + +@dataclass(frozen=True) +class _SasVerificationRequest: + """An allowed Element verification request awaiting SAS completion.""" + + sender: str + device_id: str + timestamp_ms: int + def _render_markdown_html(text: str) -> str | None: """Render markdown to sanitized HTML; returns None for plain text.""" try: @@ -321,6 +335,7 @@ def __init__( self._server_upload_limit_bytes: int | None = None self._server_upload_limit_checked = False self._stream_bufs: dict[str, _StreamBuf] = {} + self._sas_verification_requests: dict[str, _SasVerificationRequest] = {} self._started_at_ms: int = 0 self._media_download_semaphore = asyncio.Semaphore( max(1, int(self.config.max_concurrent_media_downloads)) @@ -696,7 +711,7 @@ def _register_to_device_callbacks(self) -> None: client = self._callback_registrar() client.add_to_device_callback( self._on_key_verification_event, - (KeyVerificationEvent,), + (KeyVerificationEvent, UnknownToDeviceEvent), ) def _register_response_callbacks(self) -> None: @@ -709,7 +724,10 @@ def _register_response_callbacks(self) -> None: def _is_sas_sender_allowed(self, sender: str) -> bool: return bool(sender and self.is_allowed(sender)) - async def _on_key_verification_event(self, event: KeyVerificationEvent) -> None: + async def _on_key_verification_event( + self, + event: KeyVerificationEvent | UnknownToDeviceEvent, + ) -> None: try: await self._handle_key_verification_event(event) except asyncio.CancelledError: @@ -717,15 +735,150 @@ async def _on_key_verification_event(self, event: KeyVerificationEvent) -> None: except Exception: self.logger.exception("Matrix SAS verification handling failed") - async def _handle_key_verification_event(self, event: KeyVerificationEvent) -> None: + @staticmethod + def _unknown_verification_content( + event: UnknownToDeviceEvent, + ) -> tuple[str, dict[str, object]] | None: + event_type = event.type + source = event.source + content = source.get("content") + if not isinstance(content, dict): + return None + return event_type, cast(dict[str, object], content) + + @staticmethod + def _content_string(content: dict[str, object], key: str) -> str: + value = content.get(key) + return value if isinstance(value, str) else "" + + def _prune_sas_verification_requests(self, now_ms: int) -> None: + oldest_allowed = now_ms - _SAS_REQUEST_MAX_AGE_MS + self._sas_verification_requests = { + transaction_id: request + for transaction_id, request in self._sas_verification_requests.items() + if request.timestamp_ms >= oldest_allowed + } + + async def _send_sas_control_message( + self, + *, + event_type: str, + sender: str, + device_id: str, + content: dict[str, object], + ) -> bool: + if not self.client: + return False + response = await self.client.to_device( + ToDeviceMessage( + type=event_type, + recipient=sender, + recipient_device=device_id, + content=content, + ) + ) + if isinstance(response, ToDeviceError): + self.logger.warning("Matrix SAS {} failed for {}: {}", event_type, sender, response) + return False + return True + + async def _handle_unknown_verification_event( + self, + event: UnknownToDeviceEvent, + sender: str, + ) -> None: + parsed = self._unknown_verification_content(event) + if parsed is None: + return + event_type, content = parsed + if event_type not in { + "m.key.verification.request", + "m.key.verification.ready", + "m.key.verification.done", + }: + return + + transaction_id = self._content_string(content, "transaction_id") + if not transaction_id: + return + + if event_type == "m.key.verification.request": + from_device = self._content_string(content, "from_device") + methods = content.get("methods") + timestamp = content.get("timestamp") + if ( + not from_device + or not isinstance(methods, list) + or _SAS_METHOD not in methods + or isinstance(timestamp, bool) + or not isinstance(timestamp, int) + ): + return + + now_ms = int(time.time() * 1000) + if not ( + now_ms - _SAS_REQUEST_MAX_AGE_MS + <= timestamp + <= now_ms + _SAS_REQUEST_MAX_FUTURE_MS + ): + self.logger.info("Ignoring expired Matrix SAS request from {}", sender) + return + + self._prune_sas_verification_requests(now_ms) + request = _SasVerificationRequest(sender, from_device, timestamp) + existing = self._sas_verification_requests.get(transaction_id) + if existing is not None and existing != request: + self.logger.warning( + "Ignoring conflicting Matrix SAS transaction {} from {}", + transaction_id, + sender, + ) + return + + own_device = str(self.client.device_id or "") if self.client else "" + if not own_device: + return + sent = await self._send_sas_control_message( + event_type="m.key.verification.ready", + sender=sender, + device_id=from_device, + content={ + "from_device": own_device, + "methods": [_SAS_METHOD], + "transaction_id": transaction_id, + }, + ) + if sent: + self._sas_verification_requests[transaction_id] = request + return + + if event_type == "m.key.verification.done": + request = self._sas_verification_requests.get(transaction_id) + if request is not None and request.sender == sender: + self._sas_verification_requests.pop(transaction_id, None) + self.logger.info("Matrix SAS verification finished with {}", sender) + + # Ready is deliberately ignored: this channel does not initiate verification. + + async def _handle_key_verification_event( + self, + event: KeyVerificationEvent | UnknownToDeviceEvent, + ) -> None: if not (self.config.e2ee_enabled and self.config.sas_verification): return if not self.client: return sender = str(getattr(event, "sender", "") or "") + if not self._is_sas_sender_allowed(sender): + return + + if isinstance(event, UnknownToDeviceEvent): + await self._handle_unknown_verification_event(event, sender) + return + transaction_id = str(getattr(event, "transaction_id", "") or "") - if not transaction_id or not self._is_sas_sender_allowed(sender): + if not transaction_id: return if isinstance(event, KeyVerificationStart): @@ -756,9 +909,27 @@ async def _handle_key_verification_event(self, event: KeyVerificationEvent) -> N sas = getattr(self.client, "key_verifications", {}).get(transaction_id) if sas is not None and getattr(sas, "verified", False): self.logger.info("Matrix SAS verification completed for {}", sender) + request = self._sas_verification_requests.get(transaction_id) + other_device = str(getattr(getattr(sas, "other_olm_device", None), "id", "")) + if ( + request is not None + and request.sender == sender + and request.device_id == other_device + ): + sent = await self._send_sas_control_message( + event_type="m.key.verification.done", + sender=sender, + device_id=request.device_id, + content={"transaction_id": transaction_id}, + ) + if sent: + self._sas_verification_requests.pop(transaction_id, None) return if isinstance(event, KeyVerificationCancel): + request = self._sas_verification_requests.get(transaction_id) + if request is not None and request.sender == sender: + self._sas_verification_requests.pop(transaction_id, None) self.logger.info( "Matrix SAS verification cancelled by {}: {}", sender, diff --git a/nanobot/channels/matrix/tests/test_matrix_channel.py b/nanobot/channels/matrix/tests/test_matrix_channel.py index ae5729a5203..4774f7331a0 100644 --- a/nanobot/channels/matrix/tests/test_matrix_channel.py +++ b/nanobot/channels/matrix/tests/test_matrix_channel.py @@ -219,10 +219,11 @@ async def close(self) -> None: class _FakeSas: - def __init__(self, *, verified: bool = False) -> None: + def __init__(self, *, verified: bool = False, device_id: str = "ALICEDEVICE") -> None: self.share_key_called = False self.get_mac_called = False self.verified = verified + self.other_olm_device = SimpleNamespace(id=device_id) def share_key(self): self.share_key_called = True @@ -274,6 +275,18 @@ def _patch_key_verification_events(monkeypatch) -> None: monkeypatch.setattr(matrix_module, "KeyVerificationMac", _FakeKeyVerificationMac) +def _unknown_verification_event( + event_type: str, + *, + sender: str = "@alice:matrix.org", + transaction_id: str = "tx1", + **content: object, +): + event_content = {"transaction_id": transaction_id, **content} + source = {"type": event_type, "sender": sender, "content": event_content} + return matrix_module.UnknownToDeviceEvent(source, sender, event_type) + + def _make_config(**kwargs) -> MatrixConfig: kwargs.setdefault("allow_from", ["*"]) return MatrixConfig( @@ -344,7 +357,10 @@ def test_register_to_device_callbacks_when_sas_verification_enabled() -> None: channel._register_to_device_callbacks() assert client.to_device_callbacks == [ - (channel._on_key_verification_event, (matrix_module.KeyVerificationEvent,)) + ( + channel._on_key_verification_event, + (matrix_module.KeyVerificationEvent, matrix_module.UnknownToDeviceEvent), + ) ] @@ -424,6 +440,104 @@ async def test_sas_verification_ignores_when_disabled(monkeypatch) -> None: assert client.to_device_calls == [] +@pytest.mark.asyncio +async def test_sas_verification_request_sends_ready_to_allowed_device(monkeypatch) -> None: + monkeypatch.setattr(matrix_module.time, "time", lambda: 1_000.0) + channel = MatrixChannel( + _make_config( + allow_from=["@alice:matrix.org"], + e2ee_enabled=True, + sas_verification=True, + ), + MessageBus(), + ) + client = _FakeAsyncClient("", "", "", None) + client.device_id = "BOTDEVICE" + channel.client = client + + event = _unknown_verification_event( + "m.key.verification.request", + from_device="ALICEDEVICE", + methods=["m.sas.v1"], + timestamp=1_000_000, + ) + await channel._handle_key_verification_event(event) + + assert len(client.to_device_calls) == 1 + ready = client.to_device_calls[0] + assert ready.type == "m.key.verification.ready" + assert ready.recipient == "@alice:matrix.org" + assert ready.recipient_device == "ALICEDEVICE" + assert ready.content == { + "from_device": "BOTDEVICE", + "methods": ["m.sas.v1"], + "transaction_id": "tx1", + } + assert channel._sas_verification_requests["tx1"].device_id == "ALICEDEVICE" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("sender", "methods", "timestamp"), + [ + ("@mallory:matrix.org", ["m.sas.v1"], 1_000_000), + ("@alice:matrix.org", ["m.qr_code.scan.v1"], 1_000_000), + ("@alice:matrix.org", ["m.sas.v1"], 1), + ("@alice:matrix.org", ["m.sas.v1"], 2_000_000), + ], +) +async def test_sas_verification_request_rejects_untrusted_or_invalid_input( + monkeypatch, + sender: str, + methods: list[str], + timestamp: int, +) -> None: + monkeypatch.setattr(matrix_module.time, "time", lambda: 1_000.0) + channel = MatrixChannel( + _make_config( + allow_from=["@alice:matrix.org"], + e2ee_enabled=True, + sas_verification=True, + ), + MessageBus(), + ) + client = _FakeAsyncClient("", "", "", None) + client.device_id = "BOTDEVICE" + channel.client = client + + event = _unknown_verification_event( + "m.key.verification.request", + sender=sender, + from_device="ALICEDEVICE", + methods=methods, + timestamp=timestamp, + ) + await channel._handle_key_verification_event(event) + + assert client.to_device_calls == [] + assert channel._sas_verification_requests == {} + + +@pytest.mark.asyncio +async def test_sas_verification_ready_is_ignored_without_bot_initiated_flow() -> None: + channel = MatrixChannel( + _make_config(allow_from=["@alice:matrix.org"], sas_verification=True), + MessageBus(), + ) + client = _FakeAsyncClient("", "", "", None) + channel.client = client + + event = _unknown_verification_event( + "m.key.verification.ready", + from_device="ALICEDEVICE", + methods=["m.sas.v1"], + ) + await channel._handle_key_verification_event(event) + + assert client.to_device_calls == [] + assert channel._sas_verification_requests == {} + + @pytest.mark.asyncio async def test_sas_verification_key_confirms_allowed_sender(monkeypatch) -> None: _patch_key_verification_events(monkeypatch) @@ -463,6 +577,74 @@ async def test_sas_verification_mac_does_not_resend_mac(monkeypatch) -> None: assert client.to_device_calls == [] +@pytest.mark.asyncio +async def test_sas_verification_mac_sends_done_for_element_request(monkeypatch) -> None: + _patch_key_verification_events(monkeypatch) + monkeypatch.setattr(matrix_module.time, "time", lambda: 1_000.0) + channel = MatrixChannel( + _make_config( + allow_from=["@alice:matrix.org"], + e2ee_enabled=True, + sas_verification=True, + ), + MessageBus(), + ) + client = _FakeAsyncClient("", "", "", None) + client.device_id = "BOTDEVICE" + channel.client = client + + request = _unknown_verification_event( + "m.key.verification.request", + from_device="ALICEDEVICE", + methods=["m.sas.v1"], + timestamp=1_000_000, + ) + await channel._handle_key_verification_event(request) + client.key_verifications["tx1"] = _FakeSas(verified=True) + + await channel._handle_key_verification_event(_FakeKeyVerificationMac()) + + assert [message.type for message in client.to_device_calls] == [ + "m.key.verification.ready", + "m.key.verification.done", + ] + done = client.to_device_calls[1] + assert done.recipient == "@alice:matrix.org" + assert done.recipient_device == "ALICEDEVICE" + assert done.content == {"transaction_id": "tx1"} + assert channel._sas_verification_requests == {} + + +@pytest.mark.asyncio +async def test_sas_verification_done_clears_matching_request(monkeypatch) -> None: + monkeypatch.setattr(matrix_module.time, "time", lambda: 1_000.0) + channel = MatrixChannel( + _make_config( + allow_from=["@alice:matrix.org"], + e2ee_enabled=True, + sas_verification=True, + ), + MessageBus(), + ) + client = _FakeAsyncClient("", "", "", None) + client.device_id = "BOTDEVICE" + channel.client = client + + request = _unknown_verification_event( + "m.key.verification.request", + from_device="ALICEDEVICE", + methods=["m.sas.v1"], + timestamp=1_000_000, + ) + await channel._handle_key_verification_event(request) + + done = _unknown_verification_event("m.key.verification.done") + await channel._handle_key_verification_event(done) + + assert channel._sas_verification_requests == {} + assert len(client.to_device_calls) == 1 + + def test_media_event_filter_does_not_match_text_events() -> None: assert not issubclass(matrix_module.RoomMessageText, matrix_module.MATRIX_MEDIA_EVENT_FILTER)