From cb5e4940b17e8ed1ac5d5661ce55b449b7c562ec Mon Sep 17 00:00:00 2001 From: Kegan Dougal <7190048+kegsay@users.noreply.github.com> Date: Wed, 7 Jan 2026 10:18:52 +0000 Subject: [PATCH] Bake the public key into the signing entity not the key ID --- synapse/crypto/keyring.py | 24 ++++++++----------- .../storage/databases/main/account_keys.py | 18 +++++++------- synapse/types/__init__.py | 20 +++++++++++++++- tests/crypto/test_keyring.py | 22 ++++++++++------- tests/storage/test_account_keys.py | 7 ++++-- 5 files changed, 56 insertions(+), 35 deletions(-) diff --git a/synapse/crypto/keyring.py b/synapse/crypto/keyring.py index 85f6270fb4..cf298aa0f0 100644 --- a/synapse/crypto/keyring.py +++ b/synapse/crypto/keyring.py @@ -84,7 +84,7 @@ class VerifyJsonRequest: get_json_object: Callable[[], JsonDict] minimum_valid_until_ts: int key_ids: list[str] - key_ids_are_public_keys: bool = False + server_name_is_public_key: bool = False @staticmethod def from_json_object( @@ -120,7 +120,7 @@ class VerifyJsonRequest: lambda: prune_event_dict(event.room_version, event.get_pdu_json()), minimum_valid_until_ms, key_ids=key_ids, - key_ids_are_public_keys=False, + server_name_is_public_key=False, ) @@ -314,21 +314,19 @@ class Keyring: """ assert event.room_version.msc4243_account_keys user_id = UserID.from_string(user_id_str) - key_ids = list(event.signatures.get(user_id.domain, [])) - # only keep the key ID that matches the desired user ID we want to verify as. - # Events can be signed by multiple parties e.g invites, restricted joins - expected_key_id = "ed25519:" + user_id.localpart + key_ids = list(event.signatures.get(user_id.localpart, [])) + expected_key_id = "ed25519:1" key_ids = [key_id for key_id in key_ids if key_id == expected_key_id] assert len(key_ids) == 1 # the user must have signed the event. await self.process_request( VerifyJsonRequest( - user_id.domain, + user_id.localpart, # the server name is the localpart which is a public key # We defer creating the redacted json object, as it uses a lot more # memory than the Event object itself. lambda: prune_event_dict(event.room_version, event.get_pdu_json()), 0, # No validity times key_ids=key_ids, - key_ids_are_public_keys=True, + server_name_is_public_key=True, ) ) @@ -345,18 +343,16 @@ class Keyring: Codes.UNAUTHORIZED, ) - if verify_request.key_ids_are_public_keys: + if verify_request.server_name_is_public_key: # No need to fetch keys as we have them already. assert len(verify_request.key_ids) == 1 - key_id = verify_request.key_ids[0] - key_bytes = decode_base64(key_id.removeprefix("ed25519:")) - verify_key = decode_verify_key_bytes(key_id, key_bytes) - await self._process_json(verify_key, verify_request) + key_bytes = decode_base64(verify_request.server_name) + verify_key = decode_verify_key_bytes(verify_request.key_ids[0], key_bytes) + await self.process_json(verify_key, verify_request) return found_keys: dict[str, FetchKeyResult] = {} - # If we are the originating server, short-circuit the key-fetch for any keys # we already have if self._is_mine_server_name(verify_request.server_name): diff --git a/synapse/storage/databases/main/account_keys.py b/synapse/storage/databases/main/account_keys.py index ada56ef1ee..945f4bab4a 100644 --- a/synapse/storage/databases/main/account_keys.py +++ b/synapse/storage/databases/main/account_keys.py @@ -18,7 +18,7 @@ # # -from typing import TYPE_CHECKING, Collection, Dict, List, Tuple, cast +from typing import TYPE_CHECKING, Collection, cast from signedjson.key import ( decode_signing_key_base64, @@ -36,7 +36,7 @@ from synapse.storage.database import ( LoggingTransaction, make_in_list_sql_clause, ) -from synapse.types import get_domain_from_id, get_localpart_from_id +from synapse.types import get_domain_from_id if TYPE_CHECKING: from synapse.server import HomeServer @@ -53,7 +53,7 @@ class AccountKeysStore(SQLBaseStore): async def get_or_create_account_key_user_id_for_account_name_user_id( self, account_name_user_id: str - ) -> Tuple[str, SigningKey]: + ) -> tuple[str, SigningKey]: """ Get or create an account key for the given account name user ID. The user ID must belong to this server. @@ -89,7 +89,7 @@ class AccountKeysStore(SQLBaseStore): return row[0], decode_account_key(row[1]) # create a new account key for this account inside a txn to ensure we lock correctly. - def create_key_txn(txn: LoggingTransaction) -> Tuple[str, str]: + def create_key_txn(txn: LoggingTransaction) -> tuple[str, str]: key, public_key_str = generate_account_key() account_key_user_id = ( f"@{public_key_str}:{get_domain_from_id(account_name_user_id)}" @@ -111,7 +111,7 @@ class AccountKeysStore(SQLBaseStore): ) sql = "SELECT account_key_user_id, account_key FROM account_keys WHERE account_name_user_id = ?" txn.execute(sql, (account_name_user_id,)) - return cast(Tuple[str, str], txn.fetchone()) + return cast(tuple[str, str], txn.fetchone()) row = await self.db_pool.runInteraction( "get_or_create_account_key_user_id_for_account_name_user_id.create_key_txn", @@ -122,7 +122,7 @@ class AccountKeysStore(SQLBaseStore): async def get_account_name_user_ids_for_account_key_user_ids( self, account_key_user_ids: Collection[str], - ) -> Dict[str, str]: + ) -> dict[str, str]: """ Fetch the verified account name user IDs for the given account key user IDs. Unknown account key user IDs will be omitted from the dict. @@ -140,10 +140,10 @@ class AccountKeysStore(SQLBaseStore): self.database_engine, "account_key_user_id", account_key_user_ids ) - def f(txn: LoggingTransaction) -> List[Tuple[str, str]]: + def f(txn: LoggingTransaction) -> list[tuple[str, str]]: sql = f"SELECT account_key_user_id, account_name_user_id FROM account_keys WHERE {clause} AND account_name_user_id IS NOT NULL" txn.execute(sql, args) - return cast(List[Tuple[str, str]], txn.fetchall()) + return cast(list[tuple[str, str]], txn.fetchall()) rows = await self.db_pool.runInteraction( "get_account_name_user_ids_for_account_key_user_ids", f @@ -151,7 +151,7 @@ class AccountKeysStore(SQLBaseStore): return {row[0]: row[1] for row in rows} -def generate_account_key() -> Tuple[SigningKey, str]: +def generate_account_key() -> tuple[SigningKey, str]: signing_key = generate_signing_key("1") verify_key_str = encode_base64(get_verify_key(signing_key).encode(), urlsafe=True) return signing_key, verify_key_str diff --git a/synapse/types/__init__.py b/synapse/types/__init__.py index 16892b37c0..58d9de9cfa 100644 --- a/synapse/types/__init__.py +++ b/synapse/types/__init__.py @@ -46,7 +46,7 @@ from immutabledict import immutabledict from signedjson.key import decode_verify_key_bytes from signedjson.types import VerifyKey from typing_extensions import Self -from unpaddedbase64 import decode_base64 +from unpaddedbase64 import decode_base64, encode_base64 from zope.interface import Interface from twisted.internet.defer import CancelledError @@ -355,6 +355,24 @@ class UserID(DomainSpecificString): SIGIL = "@" + @classmethod + def from_verify_key( + cls: type["UserID"], domain: str, verify_key: VerifyKey + ) -> "UserID": + """ + Converts an MSC4243 account key into a valid user ID. + + Args: + domain: The unverified domain associated with this verify key + verify_key: The ed25519 public key for this user + Returns: + A valid MSC4243 user ID. + """ + # We cannot use signedjson.encode_verify_key_base64 because that does not do URL-safe base64 + # encoding. + verify_key_str = encode_base64(verify_key.encode(), urlsafe=True) + return UserID.from_string(f"@{verify_key_str}:{domain}") + @attr.s(slots=True, frozen=True, repr=False) class RoomAlias(DomainSpecificString): diff --git a/tests/crypto/test_keyring.py b/tests/crypto/test_keyring.py index 0c6bda914d..beff908c73 100644 --- a/tests/crypto/test_keyring.py +++ b/tests/crypto/test_keyring.py @@ -51,7 +51,7 @@ from synapse.logging.context import ( ) from synapse.server import HomeServer from synapse.storage.keys import FetchKeyResult -from synapse.types import JsonDict +from synapse.types import JsonDict, UserID from synapse.util.clock import Clock from tests import unittest @@ -409,27 +409,31 @@ class KeyringTestCase(unittest.HomeserverTestCase): """ room_version = RoomVersions.MSC4243v12 - # Make a signing key and replace the key ID from '1' to be the base64 public key + # Make a signing key signing_key = signedjson.key.generate_signing_key("1") - verify_key_str = encode_verify_key_base64(get_verify_key(signing_key)) - signing_key.version = verify_key_str domain = "can.be.anything.com" - signing_user_id = f"@{verify_key_str}:{domain}" + # Derive a user ID from the signing key + signing_user_id = UserID.from_verify_key(domain, get_verify_key(signing_key)) event_dict = { "type": "m.room.create", "state_key": "", - "sender": signing_user_id, + "sender": signing_user_id.to_string(), "content": { "room_version": room_version.identifier, }, } event_dict["signatures"] = compute_event_signature( - room_version, event_dict, signature_name=domain, signing_key=signing_key + room_version, + event_dict, + signature_name=signing_user_id.localpart, + signing_key=signing_key, ) event = make_event_from_dict(event_dict, room_version) - kr = keyring.Keyring(self.hs, key_fetchers=None) - self.get_success(kr.verify_event_for_account_key(signing_user_id, event)) + kr = keyring.Keyring(self.hs, test_only_key_fetchers=[]) + self.get_success( + kr.verify_event_for_account_key(signing_user_id.to_string(), event) + ) @logcontext_clean diff --git a/tests/storage/test_account_keys.py b/tests/storage/test_account_keys.py index 7786d146a4..a1abb00115 100644 --- a/tests/storage/test_account_keys.py +++ b/tests/storage/test_account_keys.py @@ -26,7 +26,7 @@ from twisted.internet.testing import MemoryReactor from synapse.server import HomeServer from synapse.types import get_localpart_from_id -from synapse.util import Clock +from synapse.util.clock import Clock from tests import unittest @@ -45,7 +45,10 @@ class AccountKeysTestCase(unittest.HomeserverTestCase): # asserts the localpart is unpadded urlsafe base64 self.assertRegex(key_user_id, r"^@[A-Za-z0-9\-_]{43}:test$") # asserts the public key is the localpart - self.assertEquals(encode_base64(get_verify_key(key).encode(), urlsafe=True), get_localpart_from_id(key_user_id)) + self.assertEquals( + encode_base64(get_verify_key(key).encode(), urlsafe=True), + get_localpart_from_id(key_user_id), + ) # asserts the key ID is 1 self.assertEquals(key.version, "1") # assert that repeated calls return the same key