Bake the public key into the signing entity not the key ID

This commit is contained in:
Kegan Dougal
2026-01-07 10:18:52 +00:00
parent 5233b1e626
commit cb5e4940b1
5 changed files with 56 additions and 35 deletions
+10 -14
View File
@@ -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):
@@ -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
+19 -1
View File
@@ -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):
+13 -9
View File
@@ -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
+5 -2
View File
@@ -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