diff --git a/synapse/storage/databases/main/account_keys.py b/synapse/storage/databases/main/account_keys.py index 945f4bab4a..d0cd33040f 100644 --- a/synapse/storage/databases/main/account_keys.py +++ b/synapse/storage/databases/main/account_keys.py @@ -22,11 +22,11 @@ from typing import TYPE_CHECKING, Collection, cast from signedjson.key import ( decode_signing_key_base64, + encode_signing_key_base64, generate_signing_key, get_verify_key, ) from signedjson.types import SigningKey -from unpaddedbase64 import encode_base64 from synapse.api.errors import SynapseError from synapse.storage._base import SQLBaseStore @@ -36,7 +36,7 @@ from synapse.storage.database import ( LoggingTransaction, make_in_list_sql_clause, ) -from synapse.types import get_domain_from_id +from synapse.types import UserID, get_domain_from_id if TYPE_CHECKING: from synapse.server import HomeServer @@ -51,7 +51,7 @@ class AccountKeysStore(SQLBaseStore): ): super().__init__(database, db_conn, hs) - async def get_or_create_account_key_user_id_for_account_name_user_id( + async def get_or_create_local_account_key_user_id( self, account_name_user_id: str ) -> tuple[str, SigningKey]: """ @@ -71,7 +71,7 @@ class AccountKeysStore(SQLBaseStore): raise SynapseError( 500, ( - "get_or_create_account_key_user_id_for_account_name_user_id: this server cannot" + "get_or_create_local_account_key_user_id: this server cannot" f" create an account key for other servers: {account_name_user_id}" ), ) @@ -81,40 +81,42 @@ class AccountKeysStore(SQLBaseStore): keyvalues={ "account_name_user_id": account_name_user_id, }, - retcols=["account_key_user_id", "account_key"], + retcols=["account_key_user_id", "signing_key"], allow_none=True, - desc="get_or_create_account_key_user_id_for_account_name_user_id.get_key_txn", + desc="get_or_create_local_account_key_user_id.get_key_txn", ) if row is not None: 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]: - key, public_key_str = generate_account_key() - account_key_user_id = ( - f"@{public_key_str}:{get_domain_from_id(account_name_user_id)}" + signing_key = generate_signing_key("1") + account_key_user_id = UserID.from_verify_key( + get_domain_from_id(account_name_user_id), + get_verify_key(signing_key), ) # Race to insert the key. The first one to make it will be returned here as we don't clobber sql = ( - "INSERT INTO account_keys(account_name_user_id, account_key_user_id, account_key)" - " VALUES(?, ?, ?)" + "INSERT INTO account_keys(account_name_user_id, account_key_user_id, account_domain, signing_key)" + " VALUES(?, ?, ?, ?)" " ON CONFLICT DO NOTHING" ) txn.execute( sql, ( account_name_user_id, - account_key_user_id, - encode_base64(key.encode(), urlsafe=True), + account_key_user_id.to_string(), + account_key_user_id.domain, + encode_signing_key_base64(signing_key), ), ) - sql = "SELECT account_key_user_id, account_key FROM account_keys WHERE account_name_user_id = ?" + sql = "SELECT account_key_user_id, signing_key FROM account_keys WHERE account_name_user_id = ?" txn.execute(sql, (account_name_user_id,)) 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", + "get_or_create_local_account_key_user_id.create_key_txn", create_key_txn, ) return row[0], decode_account_key(row[1]) @@ -129,11 +131,11 @@ class AccountKeysStore(SQLBaseStore): Args: account_key_user_ids: A list of user IDs in account key format e.g - ["@l8Hft5qXKn1vfHrg3p4+W8gELQVo8N13JkluMfmn2sQ:example.com"] + ["@l8Hft5qXKn1vfHrg3p4-W8gELQVo8N13JkluMfmn2sQ:example.com"] Returns: A map of account key user IDs to account name user IDs e.g. - {"@l8Hft5qXKn1vfHrg3p4+W8gELQVo8N13JkluMfmn2sQ:example.com":"@alice:example.com"} + {"@l8Hft5qXKn1vfHrg3p4-W8gELQVo8N13JkluMfmn2sQ:example.com":"@alice:example.com"} """ clause, args = make_in_list_sql_clause( @@ -150,11 +152,66 @@ class AccountKeysStore(SQLBaseStore): ) return {row[0]: row[1] for row in rows} + async def store_verified_account_name_user_ids( + self, key_to_name: dict[str, str], timestamp: int + ) -> None: + """ + Store the verified account names for the given account key user IDs. -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 + Args: + key_to_name: A map from account key user ID to verified account name user ID. + timestamp: The current time, used for marking when this account was verified. + """ + await self.db_pool.simple_upsert_many( + table="account_keys", + key_names=["account_key_user_id"], + key_values=[(k,) for k in key_to_name.keys()], + value_names=["account_name_user_id", "account_domain", "verified_at_ms"], + value_values=[ + (n, get_domain_from_id(n), timestamp) for n in key_to_name.values() + ], + desc="store_verified_account_name_user_ids", + ) + + async def store_unverified_account_key_user_ids( + self, + user_ids: list[str], + ) -> None: + """ + Store unverified account key user IDs. Does nothing if the account key user ID is already + verified. + + Args: + user_ids: A list of account key user IDs. + """ + await self.db_pool.simple_upsert_many( + table="account_keys", + key_names=["account_key_user_id"], + key_values=[(k,) for k in user_ids], + desc="store_unverified_account_key_user_ids", + value_names=["account_domain"], + value_values=[(get_domain_from_id(k),) for k in user_ids], + ) + + async def get_unverified_account_key_user_ids( + self, + domain: str, + ) -> list[str]: + """ + Get a list of unverified user IDs for the given domain. + + Args: + domain: The domain to query + Returns: + A list of unverified account key user IDs. + """ + # simple_select_onecol does not support IS NULL + result = await self.db_pool.execute( + "get_unverified_account_key_user_ids", + "SELECT account_key_user_id FROM account_keys WHERE account_domain = ? AND account_name_user_id IS NULL", + domain, + ) + return [r[0] for r in result] def decode_account_key(signing_key: str) -> SigningKey: diff --git a/synapse/storage/schema/main/delta/93/01_account_keys.sql b/synapse/storage/schema/main/delta/93/01_account_keys.sql deleted file mode 100644 index ccebf713f5..0000000000 --- a/synapse/storage/schema/main/delta/93/01_account_keys.sql +++ /dev/null @@ -1,25 +0,0 @@ --- --- This file is licensed under the Affero General Public License (AGPL) version 3. --- --- Copyright (C) 2025 New Vector, Ltd --- --- This program is free software: you can redistribute it and/or modify --- it under the terms of the GNU Affero General Public License as --- published by the Free Software Foundation, either version 3 of the --- License, or (at your option) any later version. --- --- See the GNU Affero General Public License for more details: --- . - --- Keeps a record of MSC4243 account key <--> account name mappings for all servers. --- This mapping is permanent. -CREATE TABLE account_keys ( - account_key_user_id TEXT PRIMARY KEY NOT NULL, - -- nullable if we cannot talk to the remote server. - account_name_user_id TEXT, - -- the private key as urlsafe base64, only for local accounts - account_key TEXT, - UNIQUE(account_key_user_id, account_name_user_id) -); - -CREATE INDEX account_keys_key_for_name ON account_keys (account_name_user_id) WHERE account_name_user_id IS NOT NULL; diff --git a/synapse/storage/schema/main/delta/93/04_account_keys.sql b/synapse/storage/schema/main/delta/93/04_account_keys.sql new file mode 100644 index 0000000000..bc7525a9ad --- /dev/null +++ b/synapse/storage/schema/main/delta/93/04_account_keys.sql @@ -0,0 +1,36 @@ +-- +-- This file is licensed under the Affero General Public License (AGPL) version 3. +-- +-- Copyright (C) 2026 New Vector, Ltd +-- +-- This program is free software: you can redistribute it and/or modify +-- it under the terms of the GNU Affero General Public License as +-- published by the Free Software Foundation, either version 3 of the +-- License, or (at your option) any later version. +-- +-- See the GNU Affero General Public License for more details: +-- . + +-- Keeps a record of MSC4243 account key <--> account name mappings for all servers. +-- This mapping is permanent. +CREATE TABLE account_keys ( + -- A user ID with the localpart as the ed25519 verify key. + account_key_user_id TEXT PRIMARY KEY NOT NULL, + -- A user ID with the localpart as a human-readable account name. Nullable if we cannot talk to the remote server. + account_name_user_id TEXT, + -- the claimed domain of the user ID for this account. Used for selecting all unverified keys for a domain. + account_domain TEXT NOT NULL, + -- the private key as urlsafe base64, only for local accounts + signing_key TEXT, + + -- The following timestamps are used for auditing purposes, as the server merely needs a boolean + -- the unix millis timestamp when the domain was verified + verified_at_ms BIGINT, + -- the unix millis timestamp when the account was erased. + erased_at_ms BIGINT +); + +-- Make account key user ID to name lookups faster. Federation uses this to quickly map PDU senders to account names. +CREATE INDEX account_keys_key_to_name ON account_keys (account_key_user_id); +-- Make account name user ID to key lookups faster. Local users use this to quickly map to the account key. +CREATE INDEX account_keys_name_to_key ON account_keys (account_name_user_id) WHERE account_name_user_id IS NOT NULL; diff --git a/tests/storage/test_account_keys.py b/tests/storage/test_account_keys.py index a1abb00115..a061e4a2fd 100644 --- a/tests/storage/test_account_keys.py +++ b/tests/storage/test_account_keys.py @@ -35,34 +35,38 @@ class AccountKeysTestCase(unittest.HomeserverTestCase): def prepare(self, reactor: MemoryReactor, clock: Clock, hs: HomeServer) -> None: self.store = self.hs.get_datastores().main self.user = "@user:test" + self.user2 = "@user2:test" - def test_get_or_create_account_key_user_id_for_account_name_user_id(self) -> None: - key_user_id, key = self.get_success( - self.store.get_or_create_account_key_user_id_for_account_name_user_id( - self.user - ) + def test_get_or_create_local_account_key_user_id(self) -> None: + key_user_id, signing_key = self.get_success( + self.store.get_or_create_local_account_key_user_id(self.user) ) # 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), + encode_base64(get_verify_key(signing_key).encode(), urlsafe=True), get_localpart_from_id(key_user_id), ) # asserts the key ID is 1 - self.assertEquals(key.version, "1") + self.assertEquals(signing_key.version, "1") # assert that repeated calls return the same key key_user_id2, key2 = self.get_success( - self.store.get_or_create_account_key_user_id_for_account_name_user_id( - self.user - ) + self.store.get_or_create_local_account_key_user_id(self.user) ) self.assertEquals(key_user_id, key_user_id2) - self.assertEquals(key.encode(), key2.encode()) + self.assertEquals(signing_key.encode(), key2.encode()) + + # assert that calling for a different user makes a different key + key_user_id_2, signing_key_2 = self.get_success( + self.store.get_or_create_local_account_key_user_id(self.user2) + ) + self.assertNotEquals(signing_key.encode(), signing_key_2.encode()) + self.assertNotEquals(key_user_id, key_user_id_2) def test_get_account_name_user_ids_for_account_key_user_ids(self) -> None: key_user_id, _ = self.get_success( - self.store.get_or_create_account_key_user_id_for_account_name_user_id( + self.store.get_or_create_local_account_key_user_id( self.user, ) ) @@ -75,12 +79,12 @@ class AccountKeysTestCase(unittest.HomeserverTestCase): def test_get_account_name_user_ids_for_account_key_user_ids_multiple(self) -> None: key_user_id_alice, _ = self.get_success( - self.store.get_or_create_account_key_user_id_for_account_name_user_id( + self.store.get_or_create_local_account_key_user_id( "@alice:test", ) ) key_user_id_bob, _ = self.get_success( - self.store.get_or_create_account_key_user_id_for_account_name_user_id( + self.store.get_or_create_local_account_key_user_id( "@bob:test", ) ) @@ -93,3 +97,108 @@ class AccountKeysTestCase(unittest.HomeserverTestCase): self.assertEquals(result[key_user_id_alice], "@alice:test") self.assertEquals(result[key_user_id_bob], "@bob:test") self.assertEquals(result.get(key_user_id_unknown, None), None) + + def test_store_verified_account_name_user_ids(self) -> None: + local_key_user_id = "@6fey6W1wS3-vbvUmHZnTd6Gi3o-TIxvIcwtEQP4nrW0:test" + local_name_user_id = "@alice:test" + remote_key_user_id = "@fjqYXanwu0q4AvejqOQVMQ8pxwG3q3ZZQTOCD0ncR30:remote" + remote_name_user_id = "@bob:remote" + self.get_success( + self.store.store_verified_account_name_user_ids( + { + local_key_user_id: local_name_user_id, + remote_key_user_id: remote_name_user_id, + }, + self.clock.time_msec(), + ) + ) + result = self.get_success( + self.store.get_account_name_user_ids_for_account_key_user_ids( + [local_key_user_id, remote_key_user_id] + ), + ) + self.assertEquals(result[local_key_user_id], local_name_user_id) + self.assertEquals(result[remote_key_user_id], remote_name_user_id) + + def test_get_unverified_account_key_user_ids(self) -> None: + user1 = "@fjqYXanwu0q4AvejqOQVMQ8pxwG3q3ZZQTOCD0ncR30:remote" + user2 = "@6fey6W1wS3-vbvUmHZnTd6Gi3o-TIxvIcwtEQP4nrW0:remote" + user3 = "@9GgRdarGTiGMNxBoVf3C8d00UUF9kpOO0JdWsoyZ6eY:somewhere" + self.get_success( + self.store.store_unverified_account_key_user_ids([user1, user2, user3]) + ) + got_user_ids = self.get_success( + self.store.get_unverified_account_key_user_ids("remote") + ) + got_user_ids.sort() + self.assertEquals(got_user_ids, [user2, user1]) + + def test_unverified_to_verified_account(self) -> None: + user1 = "@fjqYXanwu0q4AvejqOQVMQ8pxwG3q3ZZQTOCD0ncR30:remote" + user1_name = "@alice:remote" + user2 = "@6fey6W1wS3-vbvUmHZnTd6Gi3o-TIxvIcwtEQP4nrW0:remote" + # Store both users as unverified + self.get_success( + self.store.store_unverified_account_key_user_ids([user1, user2]) + ) + # Verify user 1 + self.get_success( + self.store.store_verified_account_name_user_ids( + { + user1: user1_name, + }, + self.clock.time_msec(), + ), + ) + # User 1 must not be unverified + got_user_ids = self.get_success( + self.store.get_unverified_account_key_user_ids("remote") + ) + self.assertEquals(got_user_ids, [user2]) + # User 1 must be verified + key_to_name = self.get_success( + self.store.get_account_name_user_ids_for_account_key_user_ids( + [user1], + ) + ) + self.assertEquals( + key_to_name, + { + user1: user1_name, + }, + ) + + def test_verified_to_unverified_account_noops(self) -> None: + user1 = "@fjqYXanwu0q4AvejqOQVMQ8pxwG3q3ZZQTOCD0ncR30:remote" + user1_name = "@alice:remote" + user2 = "@6fey6W1wS3-vbvUmHZnTd6Gi3o-TIxvIcwtEQP4nrW0:remote" + # Verify user 1 + self.get_success( + self.store.store_verified_account_name_user_ids( + { + user1: user1_name, + }, + self.clock.time_msec(), + ), + ) + # Store both users as unverified + self.get_success( + self.store.store_unverified_account_key_user_ids([user1, user2]) + ) + # User 1 must not be unverified + got_user_ids = self.get_success( + self.store.get_unverified_account_key_user_ids("remote") + ) + self.assertEquals(got_user_ids, [user2]) + # User 1 must be verified + key_to_name = self.get_success( + self.store.get_account_name_user_ids_for_account_key_user_ids( + [user1], + ) + ) + self.assertEquals( + key_to_name, + { + user1: user1_name, + }, + )