mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-09-04 10:11:46 +03:00
fix(matrix): complete Element SAS request flow
This commit is contained in:
+1
-1
@@ -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. |
|
||||
|
||||
|
||||
|
||||
@@ -47,6 +47,8 @@ try:
|
||||
SyncError,
|
||||
SyncResponse,
|
||||
ToDeviceError,
|
||||
ToDeviceMessage,
|
||||
UnknownToDeviceEvent,
|
||||
UploadError,
|
||||
)
|
||||
from nio.crypto.attachments import decrypt_attachment
|
||||
@@ -75,6 +77,9 @@ _ATTACH_FAILED = "[attachment: {} - download failed]"
|
||||
_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 @@ class MatrixChannel(BaseChannel):
|
||||
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 @@ class MatrixChannel(BaseChannel):
|
||||
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 @@ class MatrixChannel(BaseChannel):
|
||||
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 @@ class MatrixChannel(BaseChannel):
|
||||
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 @@ class MatrixChannel(BaseChannel):
|
||||
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,
|
||||
|
||||
@@ -220,10 +220,11 @@ class _FakeAsyncClient:
|
||||
|
||||
|
||||
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
|
||||
@@ -275,6 +276,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(
|
||||
@@ -345,7 +358,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),
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
@@ -425,6 +441,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)
|
||||
@@ -464,6 +578,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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user