Add more storage functions for account keys

This commit is contained in:
Kegan Dougal
2026-01-07 16:26:49 +00:00
parent cb5e4940b1
commit ce9eae7137
4 changed files with 237 additions and 60 deletions
+78 -21
View File
@@ -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:
@@ -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:
-- <https://www.gnu.org/licenses/agpl-3.0.html>.
-- 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;
@@ -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:
-- <https://www.gnu.org/licenses/agpl-3.0.html>.
-- 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;
+123 -14
View File
@@ -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,
},
)