mirror of
https://github.com/element-hq/synapse.git
synced 2026-08-28 23:08:22 +00:00
Bake the public key into the signing entity not the key ID
This commit is contained in:
+10
-14
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user