mirror of
https://github.com/element-hq/synapse.git
synced 2026-08-29 01:18:30 +00:00
Add more storage functions for account keys
This commit is contained in:
@@ -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;
|
||||
@@ -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,
|
||||
},
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user