From f7d05558b6fa7d5c8c387613ee3650ebda400940 Mon Sep 17 00:00:00 2001 From: Jason Robinson Date: Tue, 7 Jul 2026 18:50:59 +0300 Subject: [PATCH] Refactor profile updates behind the profile updates stream writer Setting a profile field (including displayname / avatar_uri) now needs to happen on a profile updates stream writer. There are new "dispatch" methods to do the profile field set/delete, that handle pushing the update over replication, if required. Profile field updates and profile updates stream writes are now in the same database transaction. --- synapse/handlers/profile.py | 214 ++++++++++++++++--- synapse/handlers/room_member.py | 76 ++++--- synapse/handlers/sso.py | 18 +- synapse/module_api/__init__.py | 5 +- synapse/replication/http/profile.py | 104 ++++++++-- synapse/rest/admin/users.py | 26 ++- synapse/rest/client/profile.py | 6 +- synapse/rest/synapse/mas/users.py | 43 ++-- synapse/storage/databases/main/profile.py | 237 ++++++++++++++-------- tests/handlers/test_profile.py | 170 +++++++++++----- tests/handlers/test_register.py | 10 +- tests/handlers/test_sync.py | 1 + tests/rest/admin/test_user.py | 43 +++- tests/rest/client/test_account.py | 32 ++- tests/rest/synapse/mas/test_users.py | 16 +- tests/storage/test_main.py | 10 +- tests/storage/test_profile.py | 35 +++- 17 files changed, 771 insertions(+), 275 deletions(-) diff --git a/synapse/handlers/profile.py b/synapse/handlers/profile.py index 3b857d6b12..59773df060 100644 --- a/synapse/handlers/profile.py +++ b/synapse/handlers/profile.py @@ -34,7 +34,10 @@ from synapse.api.errors import ( StoreError, SynapseError, ) -from synapse.replication.http.profile import ReplicationProfileRecordFieldUpdates +from synapse.replication.http.profile import ( + ReplicationProfileDeleteField, + ReplicationProfileSetField, +) from synapse.storage.databases.main.media_repository import LocalMedia, RemoteMedia from synapse.storage.roommember import ProfileInfo from synapse.types import ( @@ -105,13 +108,50 @@ class ProfileHandler: self._is_profile_worker = ( hs.get_instance_name() in hs.config.worker.writers.profile_updates ) - self._record_profile_updates_client = ( - ReplicationProfileRecordFieldUpdates.make_client(self.hs) + self._delete_profile_field_client = ReplicationProfileDeleteField.make_client( + self.hs ) + self._set_profile_field_client = ReplicationProfileSetField.make_client(self.hs) self._profile_updates_writer_instance = ( self.hs.config.worker.writers.profile_updates[0] ) + async def get_targets_for_profile_updates( + self, user_id: UserID + ) -> dict[str, set[str]]: + """ + Get targets for the profile updates stream. + + Collects the rooms and users that are in the same rooms, and thus would be + interested in profile changes of our user. This will always include the user + themselves, even if the user is not in any rooms. + + Args: + user_id: The user ID that is updating a profile + + Returns: + Dictionary containing two sets, one for `rooms`, one for `users` + """ + if not self._msc4429_enabled: + return { + "rooms": set(), + "users": set(), + } + + room_ids = await self.store.get_rooms_for_user(user_id.to_string()) + if not room_ids: + return { + "rooms": set(), + "users": {user_id.to_string()}, + } + + return { + "rooms": set(room_ids), + "users": await self.store.get_local_users_who_share_room_with_user( + user_id.to_string() + ), + } + async def record_profile_updates( self, user_id: UserID, updated_fields: set[str] ) -> None: @@ -243,13 +283,14 @@ class ProfileHandler: async def set_displayname( self, + *, target_user: UserID, requester: Requester, new_displayname: str, - *, + profile_update_target_user_ids: set[str], by_admin: bool = False, propagate: bool = True, - ) -> None: + ) -> int | None: """Set the displayname of a user Preconditions: @@ -263,6 +304,8 @@ class ProfileHandler: target_user: the user whose displayname is to be changed. requester: The user attempting to make this change. new_displayname: The displayname to give this user. + profile_update_target_user_ids: User id's to trigger profile updates stream + updates for. by_admin: Whether this change was made by an administrator. propagate: Whether this change also applies to the user's membership events. """ @@ -304,10 +347,11 @@ class ProfileHandler: authenticated_entity=requester.authenticated_entity, ) - await self.store.set_profile_displayname(target_user, displayname_to_set) - await self._dispatch_record_profile_updates( + stream_id = await self.store.set_profile_field( target_user, - {ProfileFields.DISPLAYNAME}, + ProfileFields.DISPLAYNAME, + displayname_to_set, + profile_update_target_user_ids, ) profile = await self.store.get_profileinfo(target_user) @@ -323,6 +367,8 @@ class ProfileHandler: if propagate: await self._update_join_states(requester, target_user) + return stream_id + async def get_avatar_url(self, target_user: UserID) -> str | None: """ Fetch a user's avatar URL from their profile. @@ -358,13 +404,14 @@ class ProfileHandler: async def set_avatar_url( self, + *, target_user: UserID, requester: Requester, new_avatar_url: str, - *, + profile_update_target_user_ids: set[str], by_admin: bool = False, propagate: bool = True, - ) -> None: + ) -> int | None: """Set a new avatar URL for a user. Preconditions: @@ -378,6 +425,8 @@ class ProfileHandler: target_user: the user whose avatar URL is to be changed. requester: The user attempting to make this change. new_avatar_url: The avatar URL to give this user. + profile_update_target_user_ids: User ID's to trigger profile update stream + updates for. by_admin: Whether this change was made by an administrator. propagate: Whether this change also applies to the user's membership events. """ @@ -417,10 +466,11 @@ class ProfileHandler: target_user, authenticated_entity=requester.authenticated_entity ) - await self.store.set_profile_avatar_url(target_user, avatar_url_to_set) - await self._dispatch_record_profile_updates( + stream_id = await self.store.set_profile_field( target_user, - {ProfileFields.AVATAR_URL}, + ProfileFields.AVATAR_URL, + avatar_url_to_set, + profile_update_target_user_ids, ) profile = await self.store.get_profileinfo(target_user) @@ -436,6 +486,8 @@ class ProfileHandler: if propagate: await self._update_join_states(requester, target_user) + return stream_id + async def user_left_room(self, user_id: UserID, room_id: str) -> None: """ A user left a room. We now: @@ -578,6 +630,50 @@ class ProfileHandler: deactivation=True, ) + async def dispatch_set_profile_field( + self, + *, + target_user: UserID, + requester: Requester, + field_name: str, + new_value: JsonValue | dict[str, JsonValue], + by_admin: bool = False, + propagate: bool = False, + ) -> None: + """ + Dispatch setting a profile field value. This either happens in the same + instance, if configured for profile updates, or via replication in the + right instance. + + Args: + target_user: the user whose profile field is to be changed. + requester: The user attempting to make this change. + field_name: The field name to update. + new_value: New value for the profile field. + by_admin: Whether this change was made by an administrator. + propagate: Whether this change also applies to the user's membership events. + """ + if self._is_profile_worker: + await self.set_field( + target_user=target_user, + requester=requester, + field_name=field_name, + new_value=new_value, + by_admin=by_admin, + propagate=propagate, + ) + else: + # Offload to the right worker via http replication + await self._set_profile_field_client( + instance_name=self._profile_updates_writer_instance, + user_id=target_user.to_string(), + requester=requester, + field_name=field_name, + new_value=new_value, + by_admin=by_admin, + propagate=propagate, + ) + async def _dispatch_record_profile_updates( self, user_id: UserID, updated_fields: set[str] ) -> None: @@ -602,7 +698,7 @@ class ProfileHandler: ) else: # Offload to the right worker via http replication - await self._record_profile_updates_client( + await self._set_profile_field_client( instance_name=self._profile_updates_writer_instance, user_id=user_id.to_string(), updated_fields=updated_fields, @@ -731,51 +827,64 @@ class ProfileHandler: field_name: str, new_value: JsonValue | dict[str, JsonValue], by_admin: bool = False, - propagate: bool = False, + propagate: bool = True, ) -> None: """Wrapper function for setting any profile field for a user.""" + profile_update_targets = await self.get_targets_for_profile_updates(target_user) + if field_name == ProfileFields.DISPLAYNAME: if not isinstance(new_value, str): raise SynapseError( 400, "'displayname' must be a string", errcode=Codes.INVALID_PARAM ) - await self.set_displayname( + stream_id = await self.set_displayname( target_user=target_user, requester=requester, new_displayname=new_value, by_admin=by_admin, propagate=propagate, + profile_update_target_user_ids=profile_update_targets["users"], ) elif field_name == ProfileFields.AVATAR_URL: if not isinstance(new_value, str): raise SynapseError( 400, "'avatar_url' must be a string", errcode=Codes.INVALID_PARAM ) - await self.set_avatar_url( + stream_id = await self.set_avatar_url( target_user=target_user, requester=requester, new_avatar_url=new_value, by_admin=by_admin, propagate=propagate, + profile_update_target_user_ids=profile_update_targets["users"], ) else: - await self.set_profile_field( + stream_id = await self.set_profile_field( target_user=target_user, requester=requester, field_name=field_name, new_value=new_value, by_admin=by_admin, + profile_update_target_user_ids=profile_update_targets["users"], + ) + + if stream_id and profile_update_targets["rooms"]: + self._notifier.on_new_event( + StreamKeyType.PROFILE_UPDATES, + stream_id, + rooms=profile_update_targets["rooms"], ) async def set_profile_field( self, + *, target_user: UserID, requester: Requester, field_name: str, new_value: JsonValue | dict[str, JsonValue], - *, + profile_update_target_user_ids: set[str], by_admin: bool = False, - ) -> None: + ) -> int | None: """Set a new profile field for a user. Preconditions: @@ -788,6 +897,8 @@ class ProfileHandler: requester: The user attempting to make this change. field_name: The name of the profile field to update. new_value: The new field value for this user. + profile_update_target_user_ids: User ID's to trigger profile update stream + updates for. by_admin: Whether this change was made by an administrator. """ if not self.hs.is_mine(target_user): @@ -796,8 +907,12 @@ class ProfileHandler: if not by_admin and target_user != requester.user: raise AuthError(403, "Cannot set another user's profile") - await self.store.set_profile_field(target_user, field_name, new_value) - await self._dispatch_record_profile_updates(target_user, {field_name}) + stream_id = await self.store.set_profile_field( + target_user, + field_name, + new_value, + profile_update_target_user_ids, + ) # Custom fields do not propagate into the user directory *or* rooms. profile = await self.store.get_profileinfo(target_user) @@ -805,6 +920,48 @@ class ProfileHandler: target_user.to_string(), profile, by_admin, deactivation=False ) + return stream_id + + async def dispatch_delete_profile_field( + self, + *, + target_user: UserID, + requester: Requester, + field_name: str, + by_admin: bool = False, + ) -> None: + """ + Dispatch deleting a profile field value. This either happens in the same + instance, if configured for profile updates, or via replication in the + right instance. + + To delete a displayname / avatar_uri, use the `dispatch_set_profile_field` + method, using an empty string as the value. + + Args: + target_user: the user whose profile field is to be changed. + requester: The user attempting to make this change. + field_name: The field name to update. + by_admin: Whether this change was made by an administrator. + """ + assert field_name not in (ProfileFields.DISPLAYNAME, ProfileFields.AVATAR_URL) + if self._is_profile_worker: + await self.delete_profile_field( + target_user=target_user, + requester=requester, + field_name=field_name, + by_admin=by_admin, + ) + else: + # Offload to the right worker via http replication + await self._delete_profile_field_client( + instance_name=self._profile_updates_writer_instance, + user_id=target_user.to_string(), + requester=requester, + field_name=field_name, + by_admin=by_admin, + ) + async def delete_profile_field( self, target_user: UserID, @@ -832,8 +989,10 @@ class ProfileHandler: if not by_admin and target_user != requester.user: raise AuthError(400, "Cannot set another user's profile") - await self.store.delete_profile_field(target_user, field_name) - await self._dispatch_record_profile_updates(target_user, {field_name}) + profile_update_targets = await self.get_targets_for_profile_updates(target_user) + stream_id = await self.store.delete_profile_field( + target_user, field_name, profile_update_targets["users"] + ) # Custom fields do not propagate into the user directory *or* rooms. profile = await self.store.get_profileinfo(target_user) @@ -841,6 +1000,13 @@ class ProfileHandler: target_user.to_string(), profile, by_admin, deactivation=False ) + if stream_id: + self._notifier.on_new_event( + StreamKeyType.PROFILE_UPDATES, + stream_id, + rooms=profile_update_targets["rooms"], + ) + async def on_profile_query(self, args: JsonDict) -> JsonDict: """Handles federation profile query requests.""" diff --git a/synapse/handlers/room_member.py b/synapse/handlers/room_member.py index e34fc79d4a..61279d88bb 100644 --- a/synapse/handlers/room_member.py +++ b/synapse/handlers/room_member.py @@ -555,18 +555,27 @@ class RoomMemberHandler(metaclass=abc.ABCMeta): ) elif self._msc4429_enabled and event.membership == Membership.JOIN: - # Notify the profile handler. We only want to do this once - # in a multi-worker setup, so we can't dispatch a hook to all workers. - if self._is_profile_worker: - await self.profile_handler.user_joined_room(target, room_id) - else: - # Offload to the right worker via http replication - await self._profile_user_room_membership_change_client( - instance_name=self._profile_updates_writer_instance, - user_id=target.to_string(), - room_id=room_id, - membership=Membership.JOIN, + prev_member_event = None + if prev_member_event_id: + prev_member_event = await self.store.get_event( + prev_member_event_id ) + if ( + not prev_member_event + or prev_member_event.membership == Membership.LEAVE + ): + # Notify the profile handler. We only want to do this once + # in a multi-worker setup, so we can't dispatch a hook to all workers. + if self._is_profile_worker: + await self.profile_handler.user_joined_room(target, room_id) + else: + # Offload to the right worker via http replication + await self._profile_user_room_membership_change_client( + instance_name=self._profile_updates_writer_instance, + user_id=target.to_string(), + room_id=room_id, + membership=Membership.JOIN, + ) break except PartialStateConflictError as e: @@ -1570,14 +1579,14 @@ class RoomMemberHandler(metaclass=abc.ABCMeta): ratelimit=ratelimit, ) - if event.membership == Membership.LEAVE: - prev_state_ids = await context.get_prev_state_ids( - StateFilter.from_types([(EventTypes.Member, event.state_key)]) - ) - prev_member_event_id = prev_state_ids.get( - (EventTypes.Member, event.state_key), None - ) + prev_state_ids = await context.get_prev_state_ids( + StateFilter.from_types([(EventTypes.Member, event.state_key)]) + ) + prev_member_event_id = prev_state_ids.get( + (EventTypes.Member, event.state_key), None + ) + if event.membership == Membership.LEAVE: if prev_member_event_id: prev_member_event = await self.store.get_event(prev_member_event_id) if prev_member_event.membership == Membership.JOIN: @@ -1599,18 +1608,25 @@ class RoomMemberHandler(metaclass=abc.ABCMeta): membership=Membership.LEAVE, ) elif self._msc4429_enabled and event.membership == Membership.JOIN: - # Notify the profile handler. We only want to do this once - # in a multi-worker setup, so we can't dispatch a hook to all workers. - if self._is_profile_worker: - await self.profile_handler.user_joined_room(target_user, room_id) - else: - # Offload to the right worker via http replication - await self._profile_user_room_membership_change_client( - instance_name=self._profile_updates_writer_instance, - user_id=target_user.to_string(), - room_id=room_id, - membership=Membership.JOIN, - ) + prev_member_event = None + if prev_member_event_id: + prev_member_event = await self.store.get_event(prev_member_event_id) + if ( + not prev_member_event + or prev_member_event.membership == Membership.LEAVE + ): + # Notify the profile handler. We only want to do this once + # in a multi-worker setup, so we can't dispatch a hook to all workers. + if self._is_profile_worker: + await self.profile_handler.user_joined_room(target_user, room_id) + else: + # Offload to the right worker via http replication + await self._profile_user_room_membership_change_client( + instance_name=self._profile_updates_writer_instance, + user_id=target_user.to_string(), + room_id=room_id, + membership=Membership.JOIN, + ) async def _can_guest_join(self, partial_current_state_ids: StateMap[str]) -> bool: """ diff --git a/synapse/handlers/sso.py b/synapse/handlers/sso.py index bb5ca329e0..f9d9475711 100644 --- a/synapse/handlers/sso.py +++ b/synapse/handlers/sso.py @@ -530,10 +530,11 @@ class SsoHandler: user_id, authenticated_entity=user_id, ) - await self._profile_handler.set_displayname( - user_id_obj, - requester, - attributes.display_name, + await self._profile_handler.dispatch_set_profile_field( + target_user=user_id_obj, + requester=requester, + field_name=ProfileFields.DISPLAYNAME, + new_value=attributes.display_name, by_admin=True, ) if attributes.picture: @@ -842,10 +843,11 @@ class SsoHandler: ) # save it as user avatar - await self._profile_handler.set_avatar_url( - uid, - create_requester(uid), - str(avatar_mxc_url), + await self._profile_handler.dispatch_set_profile_field( + target_user=uid, + requester=create_requester(uid), + field_name=ProfileFields.AVATAR_URL, + new_value=str(avatar_mxc_url), ) logger.info("successfully saved the user avatar") diff --git a/synapse/module_api/__init__.py b/synapse/module_api/__init__.py index 3341e49b85..76ee404a90 100644 --- a/synapse/module_api/__init__.py +++ b/synapse/module_api/__init__.py @@ -2027,10 +2027,11 @@ class ModuleApi: deactivation, ) - await self._hs.get_profile_handler().set_displayname( + await self._hs.get_profile_handler().dispatch_set_profile_field( target_user=user_id, requester=requester, - new_displayname=new_displayname, + field_name=ProfileFields.DISPLAYNAME, + new_value=new_displayname, by_admin=True, ) diff --git a/synapse/replication/http/profile.py b/synapse/replication/http/profile.py index 2ebae2c007..d913f5218c 100644 --- a/synapse/replication/http/profile.py +++ b/synapse/replication/http/profile.py @@ -21,7 +21,8 @@ from twisted.web.server import Request from synapse.api.constants import Membership from synapse.http.server import HttpServer from synapse.replication.http._base import ReplicationEndpoint -from synapse.types import JsonDict, UserID +from synapse.synapse_rust.types import Requester +from synapse.types import JsonDict, JsonValue, UserID, create_requester if TYPE_CHECKING: from synapse.server import HomeServer @@ -86,15 +87,20 @@ class ReplicationProfileUserRoomMembershipChange(ReplicationEndpoint): return (200, {}) -class ReplicationProfileRecordFieldUpdates(ReplicationEndpoint): - """Record user profile field updates for the profile updates stream. +class ReplicationProfileSetField(ReplicationEndpoint): + """Update a profile field for a user. The POST looks like: - POST /_synapse/replication/profile_record_field_updates/ + POST /_synapse/replication/profile_set_field/ { - "updated_fields": ["list", "of", "fields"] + "target_user": "@user:hs", + "requester": "@admin:hs", + "field_name": "displayname", + "new_value": "Alice", + "by_admin": true, + "propagate": false } 200 OK @@ -102,7 +108,7 @@ class ReplicationProfileRecordFieldUpdates(ReplicationEndpoint): {} """ - NAME = "profile_record_field_updates" + NAME = "profile_set_field" PATH_ARGS = ("user_id",) METHOD = "POST" CACHE = False @@ -114,21 +120,88 @@ class ReplicationProfileRecordFieldUpdates(ReplicationEndpoint): @staticmethod async def _serialize_payload( # type: ignore[override] - user_id: str, - updated_fields: set[str], + user_id: UserID, + requester: Requester, + field_name: str, + new_value: JsonValue | dict[str, JsonValue], + by_admin: bool, + propagate: bool, ) -> JsonDict: - assert len(updated_fields) > 0 return { - "updated_fields": list(updated_fields), + "target_user": user_id.to_string(), + "requester": requester.user.to_string(), + "field_name": field_name, + "new_value": new_value, + "by_admin": by_admin, + "propagate": propagate, } async def _handle_request( # type: ignore[override] self, request: Request, content: JsonDict, user_id: str ) -> tuple[int, JsonDict]: - assert len(content["updated_fields"]) > 0 - await self._profile_handler.record_profile_updates( - user_id=UserID.from_string(user_id), - updated_fields=set(content["updated_fields"]), + await self._profile_handler.set_field( + target_user=UserID.from_string(user_id), + requester=create_requester(content["requester"]), + field_name=content["field_name"], + new_value=content["new_value"], + by_admin=content["by_admin"], + propagate=content["propagate"], + ) + + return (200, {}) + + +class ReplicationProfileDeleteField(ReplicationEndpoint): + """Delete a profile field for a user. + + The POST looks like: + + POST /_synapse/replication/profile_delete_field/ + + { + "target_user": "@user:hs", + "requester": "@admin:hs", + "field_name": "displayname", + "by_admin": true + } + + 200 OK + + {} + """ + + NAME = "profile_delete_field" + PATH_ARGS = ("user_id",) + METHOD = "POST" + CACHE = False + + def __init__(self, hs: "HomeServer"): + super().__init__(hs) + + self._profile_handler = hs.get_profile_handler() + + @staticmethod + async def _serialize_payload( # type: ignore[override] + user_id: UserID, + requester: Requester, + field_name: str, + by_admin: bool, + ) -> JsonDict: + return { + "target_user": user_id.to_string(), + "requester": requester.user.to_string(), + "field_name": field_name, + "by_admin": by_admin, + } + + async def _handle_request( # type: ignore[override] + self, request: Request, content: JsonDict, user_id: str + ) -> tuple[int, JsonDict]: + await self._profile_handler.delete_profile_field( + target_user=UserID.from_string(user_id), + requester=create_requester(content["requester"]), + field_name=content["field_name"], + by_admin=content["by_admin"], ) return (200, {}) @@ -137,4 +210,5 @@ class ReplicationProfileRecordFieldUpdates(ReplicationEndpoint): def register_servlets(hs: "HomeServer", http_server: HttpServer) -> None: if hs.config.server.include_profile_updates_in_sync: ReplicationProfileUserRoomMembershipChange(hs).register(http_server) - ReplicationProfileRecordFieldUpdates(hs).register(http_server) + ReplicationProfileSetField(hs).register(http_server) + ReplicationProfileDeleteField(hs).register(http_server) diff --git a/synapse/rest/admin/users.py b/synapse/rest/admin/users.py index 8265c2d789..eeb844d510 100644 --- a/synapse/rest/admin/users.py +++ b/synapse/rest/admin/users.py @@ -28,7 +28,7 @@ from typing import TYPE_CHECKING import attr from pydantic import StrictBool, StrictInt, StrictStr -from synapse.api.constants import Direction +from synapse.api.constants import Direction, ProfileFields from synapse.api.errors import Codes, NotFoundError, SynapseError from synapse.http.servlet import ( RestServlet, @@ -368,8 +368,12 @@ class UserRestServletV2(UserRestServletV2Get): if user: # modify user if "displayname" in body: - await self.profile_handler.set_displayname( - target_user, requester, body["displayname"], by_admin=True + await self.profile_handler.dispatch_set_profile_field( + target_user=target_user, + requester=requester, + field_name=ProfileFields.DISPLAYNAME, + new_value=body["displayname"], + by_admin=True, ) if threepids is not None: @@ -417,8 +421,12 @@ class UserRestServletV2(UserRestServletV2Get): ) if "avatar_url" in body: - await self.profile_handler.set_avatar_url( - target_user, requester, body["avatar_url"], by_admin=True + await self.profile_handler.dispatch_set_profile_field( + target_user=target_user, + requester=requester, + field_name=ProfileFields.AVATAR_URL, + new_value=body["avatar_url"], + by_admin=True, ) if "admin" in body: @@ -525,8 +533,12 @@ class UserRestServletV2(UserRestServletV2Get): ) if "avatar_url" in body and isinstance(body["avatar_url"], str): - await self.profile_handler.set_avatar_url( - target_user, requester, body["avatar_url"], by_admin=True + await self.profile_handler.dispatch_set_profile_field( + target_user=target_user, + requester=requester, + field_name=ProfileFields.AVATAR_URL, + new_value=body["avatar_url"], + by_admin=True, ) user_info_dict = await self.admin_handler.get_user(target_user) diff --git a/synapse/rest/client/profile.py b/synapse/rest/client/profile.py index 93a5b102d1..dc632822c1 100644 --- a/synapse/rest/client/profile.py +++ b/synapse/rest/client/profile.py @@ -209,7 +209,7 @@ class ProfileFieldRestServlet(RestServlet): Codes.USER_ACCOUNT_SUSPENDED, ) - await self.profile_handler.set_field( + await self.profile_handler.dispatch_set_profile_field( target_user=user, requester=requester, field_name=field_name, @@ -263,7 +263,7 @@ class ProfileFieldRestServlet(RestServlet): ) if field_name in (ProfileFields.DISPLAYNAME, ProfileFields.AVATAR_URL): - await self.profile_handler.set_field( + await self.profile_handler.dispatch_set_profile_field( target_user=user, requester=requester, field_name=field_name, @@ -272,7 +272,7 @@ class ProfileFieldRestServlet(RestServlet): propagate=propagate, ) else: - await self.profile_handler.delete_profile_field( + await self.profile_handler.dispatch_delete_profile_field( target_user=user, requester=requester, field_name=field_name, diff --git a/synapse/rest/synapse/mas/users.py b/synapse/rest/synapse/mas/users.py index 01db41bcfa..cc5717b103 100644 --- a/synapse/rest/synapse/mas/users.py +++ b/synapse/rest/synapse/mas/users.py @@ -19,6 +19,7 @@ from typing import TYPE_CHECKING, Any, TypedDict from pydantic import StrictBool, StrictStr, model_validator +from synapse.api.constants import ProfileFields from synapse.api.errors import NotFoundError, SynapseError from synapse.http.servlet import ( parse_and_validate_json_object_from_request, @@ -162,18 +163,18 @@ class MasProvisionUserResource(MasBaseResource): ) else: created = False + new_displayname = None if body.unset_displayname: - await self.profile_handler.set_displayname( - target_user=user_id, - requester=requester, - new_displayname="", - by_admin=True, - ) + new_displayname = "" elif body.set_displayname is not None: - await self.profile_handler.set_displayname( + new_displayname = body.set_displayname + + if new_displayname is not None: + await self.profile_handler.dispatch_set_profile_field( target_user=user_id, requester=requester, - new_displayname=body.set_displayname, + field_name=ProfileFields.DISPLAYNAME, + new_value=new_displayname, by_admin=True, ) @@ -221,18 +222,18 @@ class MasProvisionUserResource(MasBaseResource): if body.locked is not None: await self.store.set_user_locked_status(user_id.to_string(), body.locked) + new_avatar_url_value = None if body.unset_avatar_url: - await self.profile_handler.set_avatar_url( - target_user=user_id, - requester=requester, - new_avatar_url="", - by_admin=True, - ) + new_avatar_url_value = "" elif body.set_avatar_url is not None: - await self.profile_handler.set_avatar_url( + new_avatar_url_value = body.set_avatar_url + + if new_avatar_url_value is not None: + await self.profile_handler.dispatch_set_profile_field( target_user=user_id, requester=requester, - new_avatar_url=body.set_avatar_url, + field_name=ProfileFields.AVATAR_URL, + new_value=new_avatar_url_value, by_admin=True, ) @@ -380,10 +381,11 @@ class MasSetDisplayNameResource(MasBaseResource): requester = create_requester(user_id=user_id) - await self.profile_handler.set_displayname( + await self.profile_handler.dispatch_set_profile_field( target_user=requester.user, requester=requester, - new_displayname=body.displayname, + field_name=ProfileFields.DISPLAYNAME, + new_value=body.displayname, by_admin=True, ) @@ -424,10 +426,11 @@ class MasUnsetDisplayNameResource(MasBaseResource): requester = create_requester(user_id=user_id) - await self.profile_handler.set_displayname( + await self.profile_handler.dispatch_set_profile_field( target_user=requester.user, requester=requester, - new_displayname="", + field_name=ProfileFields.DISPLAYNAME, + new_value="", by_admin=True, ) diff --git a/synapse/storage/databases/main/profile.py b/synapse/storage/databases/main/profile.py index 63e51b0515..4876476b99 100644 --- a/synapse/storage/databases/main/profile.py +++ b/synapse/storage/databases/main/profile.py @@ -79,6 +79,7 @@ class ProfileWorkerStore(SQLBaseStore): "populate_full_user_id_profiles", self.populate_full_user_id_profiles ) + self._msc4429_enabled = hs.config.server.include_profile_updates_in_sync self._can_write_to_profile_updates = ( self._instance_name in hs.config.worker.writers.profile_updates ) @@ -758,92 +759,46 @@ class ProfileWorkerStore(SQLBaseStore): if total_bytes > MAX_PROFILE_SIZE: raise StoreError(400, "Profile too large", Codes.PROFILE_TOO_LARGE) - async def set_profile_displayname( - self, user_id: UserID, new_displayname: str | None - ) -> None: - """ - Set the display name of a user. - - Args: - user_id: The user's ID. - new_displayname: The new display name. If this is None, the user's display - name is removed. - """ - user_localpart = user_id.localpart - - def set_profile_displayname(txn: LoggingTransaction) -> None: - if new_displayname is not None: - self._check_profile_size( - txn, user_id, ProfileFields.DISPLAYNAME, new_displayname - ) - - self.db_pool.simple_upsert_txn( - txn, - table="profiles", - keyvalues={"user_id": user_localpart}, - values={ - "displayname": new_displayname, - "full_user_id": user_id.to_string(), - }, - ) - - await self.db_pool.runInteraction( - "set_profile_displayname", set_profile_displayname - ) - - async def set_profile_avatar_url( - self, user_id: UserID, new_avatar_url: str | None - ) -> None: - """ - Set the avatar of a user. - - Args: - user_id: The user's ID. - new_avatar_url: The new avatar URL. If this is None, the user's avatar is - removed. - """ - user_localpart = user_id.localpart - - def set_profile_avatar_url(txn: LoggingTransaction) -> None: - if new_avatar_url is not None: - self._check_profile_size( - txn, user_id, ProfileFields.AVATAR_URL, new_avatar_url - ) - - self.db_pool.simple_upsert_txn( - txn, - table="profiles", - keyvalues={"user_id": user_localpart}, - values={ - "avatar_url": new_avatar_url, - "full_user_id": user_id.to_string(), - }, - ) - - await self.db_pool.runInteraction( - "set_profile_avatar_url", set_profile_avatar_url - ) - - async def set_profile_field( + def _set_profile_field_txn( self, + txn: LoggingTransaction, user_id: UserID, field_name: str, new_value: JsonValue | dict[str, JsonValue], - ) -> None: + target_users: set[str], + ) -> int | None: """ - Set a custom profile field for a user. + Wrapper function to set a profile field value and write to the profile + update stream tables in one transaction. Args: - user_id: The user's ID. - field_name: The name of the custom profile field. - new_value: The value of the custom profile field. + txn: The transaction to use + user_id: The user to set the profile field for + field_name: The field to set the value for + new_value: New value for the profile field + target_users: Users to trigger a profile update stream row for + + Returns: + The profile updates stream ID that was created in this transaction """ + if self._msc4429_enabled: + assert self._can_write_to_profile_updates - # Encode to canonical JSON. - canonical_value = encode_canonical_json(new_value) + self._check_profile_size(txn, user_id, field_name, new_value) - def set_profile_field(txn: LoggingTransaction) -> None: - self._check_profile_size(txn, user_id, field_name, new_value) + if field_name in (ProfileFields.DISPLAYNAME, ProfileFields.AVATAR_URL): + self.db_pool.simple_upsert_txn( + txn, + table="profiles", + keyvalues={"user_id": user_id.localpart}, + values={ + field_name: new_value, + "full_user_id": user_id.to_string(), + }, + ) + else: + # Encode to canonical JSON. + canonical_value = encode_canonical_json(new_value) if isinstance(self.database_engine, PostgresEngine): from psycopg2.extras import Json @@ -851,10 +806,10 @@ class ProfileWorkerStore(SQLBaseStore): # Note that the || jsonb operator is not recursive, any duplicate # keys will be taken from the second value. sql = """ - INSERT INTO profiles (user_id, full_user_id, fields) VALUES (?, ?, JSON_BUILD_OBJECT(?, ?::jsonb)) - ON CONFLICT (user_id) - DO UPDATE SET full_user_id = EXCLUDED.full_user_id, fields = COALESCE(profiles.fields, '{}'::jsonb) || EXCLUDED.fields - """ + INSERT INTO profiles (user_id, full_user_id, fields) VALUES (?, ?, JSON_BUILD_OBJECT(?, ?::jsonb)) + ON CONFLICT (user_id) + DO UPDATE SET full_user_id = EXCLUDED.full_user_id, fields = COALESCE(profiles.fields, '{}'::jsonb) || EXCLUDED.fields \ + """ txn.execute( sql, @@ -871,10 +826,10 @@ class ProfileWorkerStore(SQLBaseStore): # You may be tempted to use json_patch instead of providing the parameters # twice, but that recursively merges objects instead of replacing. sql = """ - INSERT INTO profiles (user_id, full_user_id, fields) VALUES (?, ?, JSON_OBJECT(?, JSON(?))) - ON CONFLICT (user_id) - DO UPDATE SET full_user_id = EXCLUDED.full_user_id, fields = JSON_SET(COALESCE(profiles.fields, '{}'), ?, JSON(?)) - """ + INSERT INTO profiles (user_id, full_user_id, fields) VALUES (?, ?, JSON_OBJECT(?, JSON(?))) + ON CONFLICT (user_id) + DO UPDATE SET full_user_id = EXCLUDED.full_user_id, fields = JSON_SET(COALESCE(profiles.fields, '{}'), ?, JSON(?)) \ + """ # This will error if field_name has double quotes in it, but that's not # possible due to the grammar. json_field_name = f'$."{field_name}"' @@ -891,9 +846,99 @@ class ProfileWorkerStore(SQLBaseStore): ), ) - await self.db_pool.runInteraction("set_profile_field", set_profile_field) + if not self._msc4429_enabled: + return None - async def delete_profile_field(self, user_id: UserID, field_name: str) -> None: + # Record updates in the profile updates stream + stream_id = self._record_profile_updates_txn( + txn, + user_id, + field_name, + target_users, + ) + + return stream_id + + def _record_profile_updates_txn( + self, + txn: LoggingTransaction, + user_id: UserID, + field_name: str, + target_users: set[str], + ) -> int | None: + """ + Record updates into the profile updates stream tables. + + Args: + txn: Transaction to use + user_id: User ID that made the profile update + field_name: The field to set the value for + target_users: Set of users to create profile update stream rows for + + Returns: + The stream ID created in this transaction + """ + if not self._msc4429_enabled: + return None + + # Record the profile update + stream_id = self._profile_updates_id_gen.get_next_txn(txn) + self.db_pool.simple_insert_txn( + txn, + table="profile_updates", + values={ + "stream_id": stream_id, + "instance_name": self._instance_name, + "user_id": user_id.to_string(), + "action": ProfileUpdateAction.UPDATE.value, + "field_name": field_name, + "inserted_ts": self.clock.time_msec(), + }, + ) + + # Add per user tracking rows + inserted_ts = self.clock.time_msec() + values = [(stream_id, user_id, inserted_ts) for user_id in target_users] + self.db_pool.simple_insert_many_txn( + txn, + table="profile_updates_per_user", + keys=[ + "stream_id", + "user_id", + "inserted_ts", + ], + values=values, + ) + return stream_id + + async def set_profile_field( + self, + user_id: UserID, + field_name: str, + new_value: JsonValue | dict[str, JsonValue], + target_users: set[str], + ) -> int | None: + """ + Set a custom profile field for a user. + + Args: + user_id: The user's ID. + field_name: The name of the custom profile field. + new_value: The value of the custom profile field. + target_users: Set of users to trigger profile updates for. + """ + return await self.db_pool.runInteraction( + "set_profile_field", + self._set_profile_field_txn, + user_id, + field_name, + new_value, + target_users, + ) + + async def delete_profile_field( + self, user_id: UserID, field_name: str, target_users: set[str] + ) -> int | None: """ Remove a custom profile field for a user. @@ -902,7 +947,10 @@ class ProfileWorkerStore(SQLBaseStore): field_name: The name of the custom profile field. """ - def delete_profile_field(txn: LoggingTransaction) -> None: + if self._msc4429_enabled: + assert self._can_write_to_profile_updates + + def delete_profile_field(txn: LoggingTransaction) -> int | None: if isinstance(self.database_engine, PostgresEngine): sql = """ UPDATE profiles SET fields = fields - ? @@ -923,7 +971,20 @@ class ProfileWorkerStore(SQLBaseStore): (f'$."{field_name}"', user_id.localpart), ) - await self.db_pool.runInteraction("delete_profile_field", delete_profile_field) + if not self._msc4429_enabled: + return None + + stream_id = self._record_profile_updates_txn( + txn, + user_id, + field_name, + target_users, + ) + return stream_id + + return await self.db_pool.runInteraction( + "delete_profile_field", delete_profile_field + ) async def delete_profile(self, user_id: UserID) -> None: """ diff --git a/tests/handlers/test_profile.py b/tests/handlers/test_profile.py index 0367e950d4..68489365b7 100644 --- a/tests/handlers/test_profile.py +++ b/tests/handlers/test_profile.py @@ -26,7 +26,7 @@ from parameterized import parameterized from twisted.internet.testing import MemoryReactor import synapse.types -from synapse.api.constants import EventTypes, ProfileUpdateAction +from synapse.api.constants import EventTypes, ProfileFields, ProfileUpdateAction from synapse.api.errors import AuthError, SynapseError from synapse.rest import admin from synapse.rest.client import login, room @@ -89,7 +89,14 @@ class ProfileTestCase(unittest.HomeserverTestCase): self.on_new_event = self.mock_hs_notifier.on_new_event def test_get_my_name(self) -> None: - self.get_success(self.store.set_profile_displayname(self.frank, "Frank")) + self.get_success( + self.handler.set_field( + target_user=self.frank, + requester=synapse.types.create_requester(self.frank), + field_name=ProfileFields.DISPLAYNAME, + new_value="Frank", + ) + ) displayname = self.get_success(self.handler.get_displayname(self.frank)) @@ -97,8 +104,11 @@ class ProfileTestCase(unittest.HomeserverTestCase): def test_set_my_name(self) -> None: self.get_success( - self.handler.set_displayname( - self.frank, synapse.types.create_requester(self.frank), "Frank Jr." + self.handler.set_field( + target_user=self.frank, + requester=synapse.types.create_requester(self.frank), + field_name=ProfileFields.DISPLAYNAME, + new_value="Frank Jr.", ) ) @@ -109,8 +119,11 @@ class ProfileTestCase(unittest.HomeserverTestCase): # Set displayname again self.get_success( - self.handler.set_displayname( - self.frank, synapse.types.create_requester(self.frank), "Frank" + self.handler.set_field( + target_user=self.frank, + requester=synapse.types.create_requester(self.frank), + field_name=ProfileFields.DISPLAYNAME, + new_value="Frank", ) ) @@ -121,8 +134,11 @@ class ProfileTestCase(unittest.HomeserverTestCase): # Set displayname to an empty string self.get_success( - self.handler.set_displayname( - self.frank, synapse.types.create_requester(self.frank), "" + self.handler.set_field( + target_user=self.frank, + requester=synapse.types.create_requester(self.frank), + field_name=ProfileFields.DISPLAYNAME, + new_value="", ) ) @@ -134,8 +150,11 @@ class ProfileTestCase(unittest.HomeserverTestCase): """Test that `set_displayname` updates membership events in rooms.""" self.get_success( - self.handler.set_displayname( - self.frank, synapse.types.create_requester(self.frank), "Frank" + self.handler.set_field( + target_user=self.frank, + requester=synapse.types.create_requester(self.frank), + field_name=ProfileFields.DISPLAYNAME, + new_value="Frank", ) ) @@ -153,8 +172,11 @@ class ProfileTestCase(unittest.HomeserverTestCase): self.assertEqual(membership[state_tuple].content["displayname"], "Frank") self.get_success( - self.handler.set_displayname( - self.frank, synapse.types.create_requester(self.frank), "Frank Jr." + self.handler.set_field( + target_user=self.frank, + requester=synapse.types.create_requester(self.frank), + field_name=ProfileFields.DISPLAYNAME, + new_value="Frank Jr.", ) ) @@ -694,8 +716,11 @@ class ProfileTestCase(unittest.HomeserverTestCase): """Test that `set_displayname` returns immediately and that room membership updates are still done in background.""" self.get_success( - self.handler.set_displayname( - self.frank, synapse.types.create_requester(self.frank), "Frank" + self.handler.set_field( + target_user=self.frank, + requester=synapse.types.create_requester(self.frank), + field_name=ProfileFields.DISPLAYNAME, + new_value="Frank", ) ) @@ -716,8 +741,11 @@ class ProfileTestCase(unittest.HomeserverTestCase): ): state_tuple = (EventTypes.Member, self.frank.to_string()) self.get_success( - self.handler.set_displayname( - self.frank, synapse.types.create_requester(self.frank), "Frank Jr." + self.handler.set_field( + target_user=self.frank, + requester=synapse.types.create_requester(self.frank), + field_name=ProfileFields.DISPLAYNAME, + new_value="Frank Jr.", ) ) @@ -744,8 +772,11 @@ class ProfileTestCase(unittest.HomeserverTestCase): """Test that room membership updates triggered by changing the avatar or the display name are resumed after a restart.""" self.get_success( - self.handler.set_displayname( - self.frank, synapse.types.create_requester(self.frank), "Frank" + self.handler.set_field( + target_user=self.frank, + requester=synapse.types.create_requester(self.frank), + field_name=ProfileFields.DISPLAYNAME, + new_value="Frank", ) ) @@ -782,8 +813,11 @@ class ProfileTestCase(unittest.HomeserverTestCase): ): state_tuple = (EventTypes.Member, self.frank.to_string()) self.get_success( - self.handler.set_displayname( - self.frank, synapse.types.create_requester(self.frank), "Frank Jr." + self.handler.set_field( + target_user=self.frank, + requester=synapse.types.create_requester(self.frank), + field_name=ProfileFields.DISPLAYNAME, + new_value="Frank Jr.", ) ) @@ -847,7 +881,14 @@ class ProfileTestCase(unittest.HomeserverTestCase): @override_config({"enable_set_displayname": False}) def test_set_my_name_if_disabled(self) -> None: # Setting displayname for the first time is allowed - self.get_success(self.store.set_profile_displayname(self.frank, "Frank")) + self.get_success( + self.store.set_profile_field( + user_id=self.frank, + field_name=ProfileFields.DISPLAYNAME, + new_value="Frank", + target_users=set(), + ) + ) self.assertEqual( (self.get_success(self.store.get_profile_displayname(self.frank))), @@ -856,16 +897,22 @@ class ProfileTestCase(unittest.HomeserverTestCase): # Setting displayname a second time is forbidden self.get_failure( - self.handler.set_displayname( - self.frank, synapse.types.create_requester(self.frank), "Frank Jr." + self.handler.set_field( + target_user=self.frank, + requester=synapse.types.create_requester(self.frank), + field_name=ProfileFields.DISPLAYNAME, + new_value="Frank Jr.", ), SynapseError, ) def test_set_my_name_noauth(self) -> None: self.get_failure( - self.handler.set_displayname( - self.frank, synapse.types.create_requester(self.bob), "Frank Jr." + self.handler.set_field( + target_user=self.frank, + requester=synapse.types.create_requester(self.bob), + field_name=ProfileFields.DISPLAYNAME, + new_value="Frank Jr.", ), AuthError, ) @@ -888,8 +935,11 @@ class ProfileTestCase(unittest.HomeserverTestCase): self.store.create_profile(UserID.from_string("@caroline:test")) ) self.get_success( - self.store.set_profile_displayname( - UserID.from_string("@caroline:test"), "Caroline" + self.handler.set_field( + target_user=UserID.from_string("@caroline:test"), + requester=synapse.types.create_requester("@caroline:test"), + field_name=ProfileFields.DISPLAYNAME, + new_value="Caroline", ) ) @@ -907,16 +957,33 @@ class ProfileTestCase(unittest.HomeserverTestCase): def test_get_my_avatar(self) -> None: self.get_success( - self.store.set_profile_avatar_url(self.frank, "http://my.server/me.png") + self.handler.set_field( + target_user=self.frank, + requester=synapse.types.create_requester(self.frank), + field_name=ProfileFields.AVATAR_URL, + new_value="http://my.server/me.png", + ) ) avatar_url = self.get_success(self.handler.get_avatar_url(self.frank)) self.assertEqual("http://my.server/me.png", avatar_url) def test_get_profile_empty_displayname(self) -> None: - self.get_success(self.store.set_profile_displayname(self.frank, None)) self.get_success( - self.store.set_profile_avatar_url(self.frank, "http://my.server/me.png") + self.store.set_profile_field( + user_id=self.frank, + field_name=ProfileFields.DISPLAYNAME, + new_value=None, + target_users=set(), + ) + ) + self.get_success( + self.store.set_profile_field( + user_id=self.frank, + field_name=ProfileFields.AVATAR_URL, + new_value="http://my.server/me.png", + target_users=set(), + ) ) profile = self.get_success(self.handler.get_profile(self.frank.to_string())) @@ -925,10 +992,11 @@ class ProfileTestCase(unittest.HomeserverTestCase): def test_set_my_avatar(self) -> None: self.get_success( - self.handler.set_avatar_url( - self.frank, - synapse.types.create_requester(self.frank), - "http://my.server/pic.gif", + self.handler.set_field( + target_user=self.frank, + requester=synapse.types.create_requester(self.frank), + field_name=ProfileFields.AVATAR_URL, + new_value="http://my.server/pic.gif", ) ) @@ -939,10 +1007,11 @@ class ProfileTestCase(unittest.HomeserverTestCase): # Set avatar again self.get_success( - self.handler.set_avatar_url( - self.frank, - synapse.types.create_requester(self.frank), - "http://my.server/me.png", + self.handler.set_field( + target_user=self.frank, + requester=synapse.types.create_requester(self.frank), + field_name=ProfileFields.AVATAR_URL, + new_value="http://my.server/me.png", ) ) @@ -953,10 +1022,11 @@ class ProfileTestCase(unittest.HomeserverTestCase): # Set avatar to an empty string self.get_success( - self.handler.set_avatar_url( - self.frank, - synapse.types.create_requester(self.frank), - "", + self.handler.set_field( + target_user=self.frank, + requester=synapse.types.create_requester(self.frank), + field_name=ProfileFields.AVATAR_URL, + new_value="", ) ) @@ -968,7 +1038,12 @@ class ProfileTestCase(unittest.HomeserverTestCase): def test_set_my_avatar_if_disabled(self) -> None: # Setting displayname for the first time is allowed self.get_success( - self.store.set_profile_avatar_url(self.frank, "http://my.server/me.png") + self.handler.set_field( + target_user=self.frank, + requester=synapse.types.create_requester(self.frank), + field_name=ProfileFields.AVATAR_URL, + new_value="http://my.server/me.png", + ) ) self.assertEqual( @@ -978,10 +1053,11 @@ class ProfileTestCase(unittest.HomeserverTestCase): # Set avatar a second time is forbidden self.get_failure( - self.handler.set_avatar_url( - self.frank, - synapse.types.create_requester(self.frank), - "http://my.server/pic.gif", + self.handler.set_field( + target_user=self.frank, + requester=synapse.types.create_requester(self.frank), + field_name=ProfileFields.AVATAR_URL, + new_value="http://my.server/pic.gif", ), SynapseError, ) diff --git a/tests/handlers/test_register.py b/tests/handlers/test_register.py index 0db7f30b1f..182ff7a8fc 100644 --- a/tests/handlers/test_register.py +++ b/tests/handlers/test_register.py @@ -25,7 +25,7 @@ from unittest.mock import AsyncMock, Mock from twisted.internet.testing import MemoryReactor from synapse.api.auth.internal import InternalAuth -from synapse.api.constants import UserTypes +from synapse.api.constants import ProfileFields, UserTypes from synapse.api.errors import ( CodeMessageException, Codes, @@ -824,8 +824,12 @@ class RegistrationTestCase(unittest.HomeserverTestCase): if displayname is not None: # logger.info("setting user display name: %s -> %s", user_id, displayname) - await self.hs.get_profile_handler().set_displayname( - user, requester, displayname, by_admin=True + await self.hs.get_profile_handler().set_field( + target_user=user, + requester=requester, + field_name=ProfileFields.DISPLAYNAME, + new_value=displayname, + by_admin=True, ) return user_id, token diff --git a/tests/handlers/test_sync.py b/tests/handlers/test_sync.py index f9062236a4..a44de05cd1 100644 --- a/tests/handlers/test_sync.py +++ b/tests/handlers/test_sync.py @@ -1180,6 +1180,7 @@ class SyncProfileUpdatesTestCase(tests.unittest.HomeserverTestCase): user_id=UserID.from_string(self.user), field_name="m.status", new_value={"text": "Swimming in the Great Lakes!", "emoji": "🏊"}, + target_users=set(), ) ) self.helper.join( diff --git a/tests/rest/admin/test_user.py b/tests/rest/admin/test_user.py index 72937df9a6..72188dd4c6 100644 --- a/tests/rest/admin/test_user.py +++ b/tests/rest/admin/test_user.py @@ -40,6 +40,7 @@ from synapse.api.constants import ( EventContentFields, EventTypes, LoginType, + ProfileFields, UserTypes, ) from synapse.api.errors import Codes, HttpResponseException, ResourceLimitError @@ -944,18 +945,27 @@ class UsersListTestCase(unittest.HomeserverTestCase): # Set avatar URL to all users, that no user has a NULL value to avoid # different sort order between SQlite and PostreSQL self.get_success( - self.store.set_profile_avatar_url( - UserID.from_string("@user1:test"), "mxc://url3" + self.store.set_profile_field( + user_id=UserID.from_string("@user1:test"), + field_name=ProfileFields.AVATAR_URL, + new_value="mxc://url3", + target_users=set(), ) ) self.get_success( - self.store.set_profile_avatar_url( - UserID.from_string("@user2:test"), "mxc://url2" + self.store.set_profile_field( + user_id=UserID.from_string("@user2:test"), + field_name=ProfileFields.AVATAR_URL, + new_value="mxc://url2", + target_users=set(), ) ) self.get_success( - self.store.set_profile_avatar_url( - UserID.from_string("@admin:test"), "mxc://url1" + self.store.set_profile_field( + user_id=UserID.from_string("@admin:test"), + field_name=ProfileFields.AVATAR_URL, + new_value="mxc://url1", + target_users=set(), ) ) @@ -1547,8 +1557,11 @@ class DeactivateAccountTestCase(unittest.HomeserverTestCase): # set attributes for user self.get_success( - self.store.set_profile_avatar_url( - UserID.from_string("@user:test"), "mxc://servername/mediaid" + self.store.set_profile_field( + user_id=UserID.from_string("@user:test"), + field_name=ProfileFields.AVATAR_URL, + new_value="mxc://servername/mediaid", + target_users=set(), ) ) self.get_success( @@ -1680,7 +1693,12 @@ class DeactivateAccountTestCase(unittest.HomeserverTestCase): """ # Patch `self.other_user` to have an empty string as their avatar. self.get_success( - self.store.set_profile_avatar_url(UserID.from_string("@user:test"), "") + self.store.set_profile_field( + user_id=UserID.from_string("@user:test"), + field_name=ProfileFields.AVATAR_URL, + new_value="", + target_users=set(), + ) ) # Check we can still erase them. @@ -2759,8 +2777,11 @@ class UserRestTestCase(unittest.HomeserverTestCase): # set attributes for user self.get_success( - self.store.set_profile_avatar_url( - UserID.from_string("@user:test"), "mxc://servername/mediaid" + self.store.set_profile_field( + user_id=UserID.from_string("@user:test"), + field_name=ProfileFields.AVATAR_URL, + new_value="mxc://servername/mediaid", + target_users=set(), ) ) self.get_success( diff --git a/tests/rest/client/test_account.py b/tests/rest/client/test_account.py index 42102230f0..30dc916a18 100644 --- a/tests/rest/client/test_account.py +++ b/tests/rest/client/test_account.py @@ -30,7 +30,7 @@ from twisted.internet.interfaces import IReactorTCP from twisted.internet.testing import MemoryReactor import synapse.rest.admin -from synapse.api.constants import LoginType, Membership +from synapse.api.constants import LoginType, Membership, ProfileFields from synapse.api.errors import Codes, HttpResponseException, SynapseError from synapse.appservice import ApplicationService from synapse.rest import admin @@ -514,13 +514,19 @@ class DeactivateTestCase(unittest.HomeserverTestCase): # Set some profile data that can be checked for after the user is erased self.get_success( - profile_handler.set_displayname( - user_id, create_requester(user_id), "Kermit the Frog" + profile_handler.set_field( + target_user=user_id, + requester=create_requester(user_id), + field_name=ProfileFields.DISPLAYNAME, + new_value="Kermit the Frog", ) ) self.get_success( - profile_handler.set_avatar_url( - user_id, create_requester(user_id), "http://test/Kermit.jpg" + profile_handler.set_field( + target_user=user_id, + requester=create_requester(user_id), + field_name=ProfileFields.AVATAR_URL, + new_value="http://test/Kermit.jpg", ) ) # Verify it is set @@ -572,9 +578,21 @@ class DeactivateTestCase(unittest.HomeserverTestCase): # Can not use the profile handler to set a display name when it is disabled. Use # the database directly store = self.hs.get_datastores().main - self.get_success(store.set_profile_displayname(user_id, "Kermit the Frog")) self.get_success( - store.set_profile_avatar_url(user_id, "http://test/Kermit.jpg") + store.set_profile_field( + user_id=user_id, + field_name=ProfileFields.DISPLAYNAME, + new_value="Kermit the Frog", + target_users=set(), + ) + ) + self.get_success( + store.set_profile_field( + user_id=user_id, + field_name=ProfileFields.AVATAR_URL, + new_value="http://test/Kermit.jpg", + target_users=set(), + ) ) # Verify it is set diff --git a/tests/rest/synapse/mas/test_users.py b/tests/rest/synapse/mas/test_users.py index 6f44761bb8..be6ff52fcc 100644 --- a/tests/rest/synapse/mas/test_users.py +++ b/tests/rest/synapse/mas/test_users.py @@ -17,6 +17,7 @@ from parameterized import parameterized from twisted.internet.testing import MemoryReactor +from synapse.api.constants import ProfileFields from synapse.api.errors import StoreError from synapse.appservice import ApplicationService from synapse.server import HomeServer @@ -54,9 +55,11 @@ class MasQueryUserResource(BaseTestCase): ) ) self.get_success( - store.set_profile_avatar_url( + store.set_profile_field( user_id=alice, - new_avatar_url="mxc://example.com/avatar", + field_name=ProfileFields.AVATAR_URL, + new_value="mxc://example.com/avatar", + target_users=set(), ) ) @@ -729,7 +732,14 @@ class MasDeleteUserResource(BaseTestCase): store = self.hs.get_datastores().main # Add custom profile field - self.get_success(store.set_profile_field(alice, "io.element.example", "hello")) + self.get_success( + store.set_profile_field( + user_id=alice, + field_name="io.element.example", + new_value="hello", + target_users=set(), + ) + ) # Ensure we're testing what we think we are: # check the user has profile data at the start of the test diff --git a/tests/storage/test_main.py b/tests/storage/test_main.py index 7b5774b8c1..235e4ed678 100644 --- a/tests/storage/test_main.py +++ b/tests/storage/test_main.py @@ -18,8 +18,7 @@ # [This file includes modifications made by New Vector Limited] # # - - +from synapse.api.constants import ProfileFields from synapse.types import UserID from tests import unittest @@ -38,7 +37,12 @@ class DataStoreTestCase(unittest.HomeserverTestCase): self.get_success(self.store.register_user(self.user.to_string(), "pass")) self.get_success(self.store.create_profile(self.user)) self.get_success( - self.store.set_profile_displayname(self.user, self.displayname) + self.store.set_profile_field( + user_id=self.user, + field_name=ProfileFields.DISPLAYNAME, + new_value=self.displayname, + target_users=set(), + ) ) users, total = self.get_success( diff --git a/tests/storage/test_profile.py b/tests/storage/test_profile.py index dbaf298697..c099841267 100644 --- a/tests/storage/test_profile.py +++ b/tests/storage/test_profile.py @@ -21,6 +21,7 @@ from twisted.internet.testing import MemoryReactor +from synapse.api.constants import ProfileFields from synapse.server import HomeServer from synapse.storage.database import LoggingTransaction from synapse.storage.engines import PostgresEngine @@ -39,7 +40,14 @@ class ProfileStoreTestCase(unittest.HomeserverTestCase): def test_displayname(self) -> None: self.get_success(self.store.create_profile(self.u_frank)) - self.get_success(self.store.set_profile_displayname(self.u_frank, "Frank")) + self.get_success( + self.store.set_profile_field( + user_id=self.u_frank, + field_name=ProfileFields.DISPLAYNAME, + new_value="Frank", + target_users=set(), + ) + ) self.assertEqual( "Frank", @@ -47,7 +55,14 @@ class ProfileStoreTestCase(unittest.HomeserverTestCase): ) # test set to None - self.get_success(self.store.set_profile_displayname(self.u_frank, None)) + self.get_success( + self.store.set_profile_field( + user_id=self.u_frank, + field_name=ProfileFields.DISPLAYNAME, + new_value=None, + target_users=set(), + ) + ) self.assertIsNone( self.get_success(self.store.get_profile_displayname(self.u_frank)) @@ -57,7 +72,12 @@ class ProfileStoreTestCase(unittest.HomeserverTestCase): self.get_success(self.store.create_profile(self.u_frank)) self.get_success( - self.store.set_profile_avatar_url(self.u_frank, "http://my.site/here") + self.store.set_profile_field( + user_id=self.u_frank, + field_name=ProfileFields.AVATAR_URL, + new_value="http://my.site/here", + target_users=set(), + ) ) self.assertEqual( @@ -66,7 +86,14 @@ class ProfileStoreTestCase(unittest.HomeserverTestCase): ) # test set to None - self.get_success(self.store.set_profile_avatar_url(self.u_frank, None)) + self.get_success( + self.store.set_profile_field( + user_id=self.u_frank, + field_name=ProfileFields.AVATAR_URL, + new_value=None, + target_users=set(), + ) + ) self.assertIsNone( self.get_success(self.store.get_profile_avatar_url(self.u_frank))