# # This file is licensed under the Affero General Public License (AGPL) version 3. # # Copyright 2014-2016 OpenMarket Ltd # Copyright (C) 2023 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: # . # # Originally licensed under the Apache License, Version 2.0: # . # # [This file includes modifications made by New Vector Limited] # # from typing import Any, Awaitable, Callable from unittest.mock import AsyncMock, Mock, patch from parameterized import parameterized from twisted.internet.testing import MemoryReactor import synapse.types from synapse.api.constants import ( EventTypes, ProfileFields, ProfileUpdateAction, ReceiptTypes, ) from synapse.api.errors import AuthError, SynapseError from synapse.rest import admin from synapse.rest.client import knock, login, room from synapse.server import HomeServer from synapse.storage.databases.main.profile import ProfileUpdate from synapse.types import JsonDict, StreamKeyType, UserID from synapse.types.state import StateFilter from synapse.util.clock import Clock from synapse.util.duration import Duration from synapse.util.task_scheduler import TaskStatus from tests import unittest from tests.unittest import override_config class ProfileTestCase(unittest.HomeserverTestCase): """Tests profile management.""" servlets = [ admin.register_servlets, login.register_servlets, room.register_servlets, knock.register_servlets, ] def make_homeserver(self, reactor: MemoryReactor, clock: Clock) -> HomeServer: self.mock_federation = AsyncMock() self.mock_registry = Mock() self.query_handlers: dict[str, Callable[[dict], Awaitable[JsonDict]]] = {} def register_query_handler( query_type: str, handler: Callable[[dict], Awaitable[JsonDict]] ) -> None: self.query_handlers[query_type] = handler self.mock_registry.register_query_handler = register_query_handler self.mock_hs_notifier = Mock() hs = self.setup_test_homeserver( notifier=self.mock_hs_notifier, federation_client=self.mock_federation, federation_server=Mock(), federation_registry=self.mock_registry, ) return hs def prepare(self, reactor: MemoryReactor, clock: Clock, hs: HomeServer) -> None: self.store = hs.get_datastores().main self.storage_controllers = self.hs.get_storage_controllers() self.task_scheduler = hs.get_task_scheduler() self.frank = UserID.from_string("@1234abcd:test") self.bob = UserID.from_string("@4567:test") self.alice = UserID.from_string("@alice:remote") self.register_user(self.frank.localpart, "frankpassword") self.frank_token = self.login(self.frank.localpart, "frankpassword") self.handler = hs.get_profile_handler() self.on_new_event = self.mock_hs_notifier.on_new_event def test_get_my_name(self) -> None: 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)) self.assertEqual("Frank", displayname) def test_set_my_name(self) -> None: 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 Jr.", ) ) self.assertEqual( (self.get_success(self.store.get_profile_displayname(self.frank))), "Frank Jr.", ) # Set displayname again 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", ) ) self.assertEqual( (self.get_success(self.store.get_profile_displayname(self.frank))), "Frank", ) # Set displayname to an empty string self.get_success( self.handler.set_field( target_user=self.frank, requester=synapse.types.create_requester(self.frank), field_name=ProfileFields.DISPLAYNAME, new_value="", ) ) self.assertIsNone( self.get_success(self.store.get_profile_displayname(self.frank)) ) def test_update_room_membership_on_set_displayname(self) -> None: """Test that `set_displayname` updates membership events in rooms.""" 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", ) ) room_id = self.helper.create_room_as( self.frank.to_string(), tok=self.frank_token ) state_tuple = (EventTypes.Member, self.frank.to_string()) membership = self.get_success( self.storage_controllers.state.get_current_state( room_id, StateFilter.from_types([state_tuple]) ) ) self.assertEqual(membership[state_tuple].content["displayname"], "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 Jr.", ) ) membership = self.get_success( self.storage_controllers.state.get_current_state( room_id, StateFilter.from_types([state_tuple]) ) ) self.assertEqual(membership[state_tuple].content["displayname"], "Frank Jr.") @parameterized.expand( [ ["displayname", "Frank"], ["avatar_url", "mxc://foobar"], ["m.status", '{"text": "Holiday", "emoji": "🏖"}'], ] ) def test_update_profile_does_not_update_stream_on_set_field_if_msc4429_not_enabled( self, field_name: str, new_value: str, ) -> None: """Test that profile updates don't get recorded in the profile updates stream if MSC4429 is not enabled.""" self.get_success( self.handler.set_field( target_user=self.frank, requester=synapse.types.create_requester(self.frank), field_name=field_name, new_value=new_value, ) ) updates = self.get_success( self.store.get_updated_profile_updates( from_id=1, to_id=2, limit=1, ) ) self.assertEqual(len(updates), 0) @parameterized.expand( [ ["displayname", "Frank"], ["avatar_url", "mxc://foobar"], ["m.status", '{"text": "Holiday", "emoji": "🏖"}'], ] ) def test_update_profile_does_not_notify_notifier_on_set_field_if_msc4429_not_enabled( self, field_name: str, new_value: str, ) -> None: """Test that profile updates do not cause the profile updates stream notifier to wake up if MSC4429 is not enabled.""" self.get_success( self.handler.set_field( target_user=self.frank, requester=synapse.types.create_requester(self.frank), field_name=field_name, new_value=new_value, ) ) calls_found = [ call for call in self.on_new_event.mock_calls if call.args[0] == StreamKeyType.PROFILE_UPDATES ] self.assertEqual(len(calls_found), 0) @parameterized.expand( [ ["displayname", "Frank"], ["avatar_url", "mxc://foobar"], ["m.status", '{"text": "Holiday", "emoji": "🏖"}'], ] ) @override_config({"include_profile_updates_in_sync": True}) def test_update_profile_does_not_notify_notifier_on_set_field_if_user_not_in_rooms( self, field_name: str, new_value: str ) -> None: """Test that profile updates do not cause the profile updates stream notifier to wake up if the user is not in any rooms, if MSC4429 is enabled.""" self.get_success( self.handler.set_field( target_user=self.frank, requester=synapse.types.create_requester(self.frank), field_name=field_name, new_value=new_value, ) ) calls_found = [ call for call in self.on_new_event.mock_calls if call.args[0] == StreamKeyType.PROFILE_UPDATES ] self.assertEqual(len(calls_found), 0) @parameterized.expand( [ ["displayname", "Frank"], ["avatar_url", "mxc://foobar"], ["m.status", '{"text": "Holiday", "emoji": "🏖"}'], ] ) @override_config({"include_profile_updates_in_sync": True}) def test_update_profile_updates_stream_on_set_field( self, field_name: str, new_value: str ) -> None: """Test that profile updates get recorded in the profile updates stream if MSC4429 is enabled.""" self.get_success( self.handler.set_field( target_user=self.frank, requester=synapse.types.create_requester(self.frank), field_name=field_name, new_value=new_value, ) ) updates = self.get_success( self.store.get_updated_profile_updates( from_id=1, to_id=2, limit=1, ) ) self.assertEqual( updates[0], ( 2, "@1234abcd:test", ProfileUpdateAction.UPDATE.value, {field_name}, ), ) fields_updates = self.get_success( self.store.get_profile_updates_for_fields( from_id=1, to_id=2, field_names={field_name}, ) ) self.assertEqual( fields_updates[0], ProfileUpdate( stream_id=2, user_id="@1234abcd:test", action=ProfileUpdateAction.UPDATE.value, affected_fields=frozenset({field_name}), ), ) self.get_success( self.handler.set_field( target_user=self.frank, requester=synapse.types.create_requester(self.frank), field_name=field_name, new_value="", ) ) delete_updates = self.get_success( self.store.get_updated_profile_updates( from_id=2, to_id=3, limit=1, ) ) self.assertEqual( delete_updates[0], (3, "@1234abcd:test", ProfileUpdateAction.UPDATE.value, {field_name}), ) @override_config({"include_profile_updates_in_sync": True}) def test_update_profile_set_field_writes_to_per_user_profile_tracking_table( self, ) -> None: """Test that profiles updates get recorded in the 'per user' profile updates stream tracking table, if MSC4429 is enabled.""" self.register_user("roger", "password") roger_token = self.login("roger", "password") self.register_user("millie", "password") millie_token = self.login("millie", "password") room_id = self.helper.create_room_as( room_creator=self.frank.to_string(), tok=self.frank_token, ) self.helper.join(room_id, "@roger:test", tok=roger_token) self.helper.join(room_id, "@millie:test", tok=millie_token) self.get_success( self.handler.set_field( target_user=self.frank, requester=synapse.types.create_requester(self.frank), field_name="m.status", new_value='{"text": "Holiday"}', ) ) per_user_updates = self.get_success( self.store.get_profile_updates_for_user_and_fields( from_id=0, to_id=10, user_id="@roger:test", field_names={"m.status"}, ) ) self.assertEqual( per_user_updates, [ ProfileUpdate( stream_id=3, user_id="@millie:test", action="joined_room", affected_fields=None, ), ProfileUpdate( stream_id=4, user_id=self.frank.to_string(), action="update", affected_fields=frozenset({"m.status"}), ), ], ) per_user_updates = self.get_success( self.store.get_profile_updates_for_user_and_fields( from_id=0, to_id=10, user_id="@millie:test", field_names={"m.status"}, ) ) self.assertEqual( per_user_updates, [ ProfileUpdate( stream_id=4, user_id=self.frank.to_string(), action="update", affected_fields=frozenset({"m.status"}), ), ], ) per_user_updates = self.get_success( self.store.get_profile_updates_for_user_and_fields( from_id=0, to_id=10, user_id=self.frank.to_string(), field_names={"m.status"}, ) ) self.assertEqual( per_user_updates, [ ProfileUpdate( stream_id=2, user_id="@roger:test", action="joined_room", affected_fields=None, ), ProfileUpdate( stream_id=3, user_id="@millie:test", action="joined_room", affected_fields=None, ), ProfileUpdate( stream_id=4, user_id=self.frank.to_string(), action="update", affected_fields=frozenset({"m.status"}), ), ], ) @override_config({"include_profile_updates_in_sync": True}) def test_membership_addition_to_room_adds_the_right_join_action_to_profile_streams( self, ) -> None: """Test that a membership event, which adds a user as joined to a room, adds the relevant joined action to the profile update stream tables. Here we consider join, knock and invite to all be additions to the room list of members for answering the question "which profiles should we send information about to clients based on memberships appearing". """ self.register_user("roger", "password") roger_token = self.login("roger", "password") room_id = self.helper.create_room_as( room_creator=self.frank.to_string(), tok=self.frank_token, ) self.helper.join(room_id, "@roger:test", tok=roger_token) per_user_updates = self.get_success( self.store.get_profile_updates_for_user_and_fields( from_id=0, to_id=10, user_id=self.frank.to_string(), field_names={"m.status"}, ) ) self.assertEqual( per_user_updates, [ ProfileUpdate( stream_id=2, user_id="@roger:test", action="joined_room", affected_fields=None, ), ], ) @override_config({"include_profile_updates_in_sync": True}) def test_previous_profile_updates_stream_rows_cleared_if_no_longer_sharing_a_room( self, ) -> None: """Test that previous profile update stream rows are removed for a user if the user no longer shares rooms with another user, if MSC4429 is enabled. This test ensures that when a user leaves a room, we clear all old profile update rows of users who the user no longer shares rooms with, to avoid leaking any further profile field updates from those users. """ self.register_user("roger", "password") roger_token = self.login("roger", "password") self.register_user("millie", "password") millie_token = self.login("millie", "password") self.register_user("gracie", "password") gracie_token = self.login("gracie", "password") room_id = self.helper.create_room_as( room_creator=self.frank.to_string(), tok=self.frank_token, ) room_with_millie_id = self.helper.create_room_as( room_creator=self.frank.to_string(), tok=self.frank_token, ) self.helper.join(room_id, "@roger:test", tok=roger_token) self.helper.join(room_with_millie_id, "@millie:test", tok=millie_token) self.helper.join(room_id, "@gracie:test", tok=gracie_token) self.get_success( self.handler.set_field( target_user=self.frank, requester=synapse.types.create_requester(self.frank), field_name="m.status", new_value='{"text": "Holiday"}', ) ) per_user_updates = self.get_success( self.store.get_profile_updates_for_user_and_fields( from_id=0, to_id=10, user_id="@roger:test", field_names={"m.status"}, ) ) self.assertEqual( per_user_updates, [ ProfileUpdate( stream_id=4, user_id="@gracie:test", action="joined_room", affected_fields=None, ), ProfileUpdate( stream_id=5, user_id=self.frank.to_string(), action="update", affected_fields=frozenset({"m.status"}), ), ], ) per_user_updates = self.get_success( self.store.get_profile_updates_for_user_and_fields( from_id=0, to_id=10, user_id="@millie:test", field_names={"m.status"}, ) ) self.assertEqual( per_user_updates, [ ProfileUpdate( stream_id=5, user_id=self.frank.to_string(), action="update", affected_fields=frozenset({"m.status"}), ), ], ) # Make frank leave room and verify only the "left room" + gracies join exists # for roger self.helper.leave(room_id, self.frank.to_string(), tok=self.frank_token) per_user_updates = self.get_success( self.store.get_profile_updates_for_user_and_fields( from_id=0, to_id=10, user_id="@roger:test", field_names={"m.status"}, ) ) self.assertEqual( per_user_updates, [ ProfileUpdate( stream_id=4, user_id="@gracie:test", action="joined_room", affected_fields=None, ), ProfileUpdate( stream_id=6, user_id=self.frank.to_string(), action="left_room", affected_fields=None, ), ], ) # Make gracie leave room and verify only the "left room"'s self.helper.leave(room_id, "@gracie:test", tok=gracie_token) per_user_updates = self.get_success( self.store.get_profile_updates_for_user_and_fields( from_id=0, to_id=10, user_id="@roger:test", field_names={"m.status"}, ) ) self.assertEqual( per_user_updates, [ ProfileUpdate( stream_id=6, user_id=self.frank.to_string(), action="left_room", affected_fields=None, ), ProfileUpdate( stream_id=7, user_id="@gracie:test", action="left_room", affected_fields=None, ), ], ) # Sanity check we didn't clear any rows for millie per_user_updates = self.get_success( self.store.get_profile_updates_for_user_and_fields( from_id=0, to_id=10, user_id="@millie:test", field_names={"m.status"}, ) ) self.assertEqual( per_user_updates, [ ProfileUpdate( stream_id=5, user_id=self.frank.to_string(), action="update", affected_fields=frozenset({"m.status"}), ), ], ) @parameterized.expand( [ ["displayname", "Frank"], ["avatar_url", "mxc://foobar"], ["m.status", '{"text": "Holiday", "emoji": "🏖"}'], ] ) @override_config({"include_profile_updates_in_sync": True}) def test_update_profile_notifies_notifier_on_set_field( self, field_name: str, new_value: str, ) -> None: """Test that profile updates wake up the profile updates stream on profile field updates, if MSC4429 is enabled.""" self.helper.create_room_as( room_creator=self.frank.to_string(), tok=self.frank_token, ) self.get_success( self.handler.set_field( target_user=self.frank, requester=synapse.types.create_requester(self.frank), field_name=field_name, new_value=new_value, ) ) calls_found = [ call for call in self.on_new_event.mock_calls if call.args[0] == StreamKeyType.PROFILE_UPDATES ] self.assertEqual(len(calls_found), 1) def test_background_update_room_membership_on_set_displayname(self) -> None: """Test that `set_displayname` returns immediately and that room membership updates are still done in background.""" 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", ) ) room_id = self.helper.create_room_as( self.frank.to_string(), tok=self.frank_token ) original_update_membership = self.hs.get_room_member_handler().update_membership async def slow_update_membership(*args: Any, **kwargs: Any) -> tuple[str, int]: await self.clock.sleep(Duration(milliseconds=10)) return await original_update_membership(*args, **kwargs) with patch.object( self.hs.get_room_member_handler(), "update_membership", side_effect=slow_update_membership, ): state_tuple = (EventTypes.Member, self.frank.to_string()) 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 Jr.", ) ) membership = self.get_success( self.storage_controllers.state.get_current_state( room_id, StateFilter.from_types([state_tuple]) ) ) self.assertEqual(membership[state_tuple].content["displayname"], "Frank") # Let's be sure we are over the delay introduced by slow_update_membership self.reactor.advance(Duration(milliseconds=20).as_secs()) membership = self.get_success( self.storage_controllers.state.get_current_state( room_id, StateFilter.from_types([state_tuple]) ) ) self.assertEqual( membership[state_tuple].content["displayname"], "Frank Jr." ) def test_background_update_room_membership_resume_after_restart(self) -> None: """Test that room membership updates triggered by changing the avatar or the display name are resumed after a restart.""" initial_displayname = "Frank" updated_displayname = "Frank Jr." self.get_success( self.handler.set_field( target_user=self.frank, requester=synapse.types.create_requester(self.frank), field_name=ProfileFields.DISPLAYNAME, new_value=initial_displayname, ) ) room_id_1 = self.helper.create_room_as( self.frank.to_string(), tok=self.frank_token ) room_id_2 = self.helper.create_room_as( self.frank.to_string(), tok=self.frank_token ) room_id_3 = self.helper.create_room_as( self.frank.to_string(), tok=self.frank_token ) # Set read receipts with different timestamps (simulate different read times) # Room 1 should be most recent, then room 2, then room 3 event_3 = self.helper.send(room_id_3, "Hello 3", tok=self.frank_token) event_2 = self.helper.send(room_id_2, "Hello 2", tok=self.frank_token) event_1 = self.helper.send(room_id_1, "Hello 1", tok=self.frank_token) self.get_success( self.store.insert_receipt( room_id_3, ReceiptTypes.READ, user_id=self.frank.to_string(), event_ids=[event_3["event_id"]], thread_id=None, data={"ts": 100}, ) ) self.get_success( self.store.insert_receipt( room_id_2, ReceiptTypes.READ, user_id=self.frank.to_string(), event_ids=[event_2["event_id"]], thread_id=None, data={"ts": 200}, ) ) self.get_success( self.store.insert_receipt( room_id_1, ReceiptTypes.READ, user_id=self.frank.to_string(), event_ids=[event_1["event_id"]], thread_id=None, data={"ts": 300}, ) ) original_update_membership = self.hs.get_room_member_handler().update_membership room_1_updated = False async def potentially_slow_update_membership( *args: Any, **kwargs: Any ) -> tuple[str, int]: if args[2] == room_id_2 or args[2] == room_id_3: await self.clock.sleep(Duration(milliseconds=10)) if args[2] == room_id_1: nonlocal room_1_updated room_1_updated = True return await original_update_membership(*args, **kwargs) with patch.object( self.hs.get_room_member_handler(), "update_membership", side_effect=potentially_slow_update_membership, ): state_tuple = (EventTypes.Member, self.frank.to_string()) self.get_success( self.handler.set_field( target_user=self.frank, requester=synapse.types.create_requester(self.frank), field_name=ProfileFields.DISPLAYNAME, new_value=updated_displayname, ) ) # Check that the displayname is updated immediately for the first room membership = self.get_success( self.storage_controllers.state.get_current_state( room_id_1, StateFilter.from_types([state_tuple]) ) ) self.assertEqual( membership[state_tuple].content["displayname"], updated_displayname ) # Simulate a synapse restart by emptying the list of running tasks # and canceling the deferred _, deferred = self.task_scheduler._running_tasks.popitem() deferred.cancel() # Let's reset the flag to track whether room 1 was updated after the restart room_1_updated = False # Let's be sure we are over the delay introduced by slow_update_membership # and that the task was not executed as expected self.reactor.advance(Duration(milliseconds=20).as_secs()) membership = self.get_success( self.storage_controllers.state.get_current_state( room_id_2, StateFilter.from_types([state_tuple]) ) ) self.assertEqual( membership[state_tuple].content["displayname"], initial_displayname ) cancelled_task = self.get_success( self.task_scheduler.get_tasks( actions=["update_join_states"], statuses=[TaskStatus.CANCELLED] ) )[0] self.get_success( self.task_scheduler.update_task( cancelled_task.id, status=TaskStatus.ACTIVE ) ) # Wait for the `TaskScheduler.SCHEDULE_INTERVAL` self.reactor.advance(Duration(minutes=1).as_secs()) # Let's be sure we are over the delay introduced by slow_update_membership self.reactor.advance(Duration(milliseconds=20).as_secs()) # Updates should have been resumed from room 2 after the restart # so room 1 should not have been updated this time self.assertFalse(room_1_updated) membership = self.get_success( self.storage_controllers.state.get_current_state( room_id_2, StateFilter.from_types([state_tuple]) ) ) self.assertEqual( membership[state_tuple].content["displayname"], updated_displayname ) membership = self.get_success( self.storage_controllers.state.get_current_state( room_id_3, StateFilter.from_types([state_tuple]) ) ) self.assertEqual( membership[state_tuple].content["displayname"], updated_displayname ) def test_room_update_ordering_by_read_receipt(self) -> None: """Test that rooms are updated in order of most recent read receipt.""" self.get_success( self.handler.set_displayname( self.frank, synapse.types.create_requester(self.frank), "Frank" ) ) # Create three rooms room_id_1 = self.helper.create_room_as( self.frank.to_string(), tok=self.frank_token ) room_id_2 = self.helper.create_room_as( self.frank.to_string(), tok=self.frank_token ) room_id_3 = self.helper.create_room_as( self.frank.to_string(), tok=self.frank_token ) # Send an event in each room to create something to mark as read event_1 = self.helper.send(room_id_1, "Hello 1", tok=self.frank_token) event_2 = self.helper.send(room_id_2, "Hello 2", tok=self.frank_token) event_3 = self.helper.send(room_id_3, "Hello 3", tok=self.frank_token) # Set read receipts with different timestamps (simulate different read times) # Room 3 is the most recent, then room 2, then room 1 self.get_success( self.store.insert_receipt( room_id_1, ReceiptTypes.READ, user_id=self.frank.to_string(), event_ids=[event_1["event_id"]], thread_id=None, data={"ts": 100}, ) ) self.get_success( self.store.insert_receipt( room_id_2, ReceiptTypes.READ, user_id=self.frank.to_string(), event_ids=[event_2["event_id"]], thread_id=None, data={"ts": 200}, ) ) self.get_success( self.store.insert_receipt( room_id_3, ReceiptTypes.READ, user_id=self.frank.to_string(), event_ids=[event_3["event_id"]], thread_id=None, data={"ts": 300}, ) ) # Track the order in which rooms are updated room_update_order = [] original_update_membership = self.hs.get_room_member_handler().update_membership async def track_update_membership(*args: Any, **kwargs: Any) -> tuple[str, int]: room_id = args[2] room_update_order.append(room_id) return await original_update_membership(*args, **kwargs) with patch.object( self.hs.get_room_member_handler(), "update_membership", side_effect=track_update_membership, ): self.get_success( self.handler.set_displayname( self.frank, synapse.types.create_requester(self.frank), "Frank Updated", ) ) # Wait for background task to complete self.get_success(self.clock.sleep(Duration(milliseconds=50)), by=1) # Get receipts to understand the actual stream ordering user_receipts = self.get_success( self.store.get_receipts_for_user_with_orderings( self.frank.to_string(), [ReceiptTypes.READ, ReceiptTypes.READ_PRIVATE], ) ) # Sort rooms by stream_ordering (descending) to get expected order rooms_by_stream_ordering = sorted( user_receipts.keys(), key=lambda room_id: -user_receipts[room_id]["stream_ordering"], ) # Verify rooms were updated in order of most recent read receipt (highest stream_ordering first) self.assertEqual(len(room_update_order), 3) self.assertEqual(room_update_order, rooms_by_stream_ordering) def test_room_update_ordering_with_no_receipts_fallback(self) -> None: """Test that rooms without read receipts fall back to alphabetical ordering.""" self.get_success( self.handler.set_displayname( self.frank, synapse.types.create_requester(self.frank), "Frank" ) ) # Create two rooms - ensure we know their alphabetical order room_id_a = self.helper.create_room_as( self.frank.to_string(), tok=self.frank_token ) room_id_b = self.helper.create_room_as( self.frank.to_string(), tok=self.frank_token ) # Ensure room_id_a comes before room_id_b alphabetically if room_id_a > room_id_b: room_id_a, room_id_b = room_id_b, room_id_a # Don't set any read receipts - should fall back to alphabetical # Track the order in which rooms are updated room_update_order = [] original_update_membership = self.hs.get_room_member_handler().update_membership async def track_update_membership(*args: Any, **kwargs: Any) -> tuple[str, int]: room_id = args[2] room_update_order.append(room_id) return await original_update_membership(*args, **kwargs) with patch.object( self.hs.get_room_member_handler(), "update_membership", side_effect=track_update_membership, ): self.get_success( self.handler.set_displayname( self.frank, synapse.types.create_requester(self.frank), "Frank Updated", ) ) # Wait for background task to complete self.get_success(self.clock.sleep(Duration(milliseconds=50)), by=1) # Verify rooms were updated in alphabetical order self.assertEqual(len(room_update_order), 2) self.assertEqual(room_update_order[0], room_id_a) # Alphabetically first self.assertEqual(room_update_order[1], room_id_b) # Alphabetically second def test_room_update_ordering_mixed_receipts_and_no_receipts(self) -> None: """Test ordering when some rooms have receipts and others don't.""" self.get_success( self.handler.set_displayname( self.frank, synapse.types.create_requester(self.frank), "Frank" ) ) # Create three rooms room_with_receipt = self.helper.create_room_as( self.frank.to_string(), tok=self.frank_token ) room_without_receipt_1 = self.helper.create_room_as( self.frank.to_string(), tok=self.frank_token ) room_without_receipt_2 = self.helper.create_room_as( self.frank.to_string(), tok=self.frank_token ) # Ensure we know the alphabetical order of rooms without receipts if room_without_receipt_1 > room_without_receipt_2: room_without_receipt_1, room_without_receipt_2 = ( room_without_receipt_2, room_without_receipt_1, ) # Send an event and set a read receipt for one room only event = self.helper.send(room_with_receipt, "Hello", tok=self.frank_token) self.get_success( self.store.insert_receipt( room_with_receipt, ReceiptTypes.READ, user_id=self.frank.to_string(), event_ids=[event["event_id"]], thread_id=None, data={"ts": 100}, ) ) # Track the order in which rooms are updated room_update_order = [] original_update_membership = self.hs.get_room_member_handler().update_membership async def track_update_membership(*args: Any, **kwargs: Any) -> tuple[str, int]: room_id = args[2] room_update_order.append(room_id) return await original_update_membership(*args, **kwargs) with patch.object( self.hs.get_room_member_handler(), "update_membership", side_effect=track_update_membership, ): self.get_success( self.handler.set_displayname( self.frank, synapse.types.create_requester(self.frank), "Frank Updated", ) ) # Wait for background task to complete self.get_success(self.clock.sleep(Duration(milliseconds=50)), by=1) # Verify ordering: room with receipt first, then others alphabetically self.assertEqual(len(room_update_order), 3) self.assertEqual(room_update_order[0], room_with_receipt) # Has receipt - first self.assertEqual( room_update_order[1], room_without_receipt_1 ) # No receipt - alphabetically first self.assertEqual( room_update_order[2], room_without_receipt_2 ) # No receipt - alphabetically second @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_field( user_id=self.frank, field_name=ProfileFields.DISPLAYNAME, new_value="Frank", ) ) self.assertEqual( (self.get_success(self.store.get_profile_displayname(self.frank))), "Frank", ) # Setting displayname a second time is forbidden self.get_failure( 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_field( target_user=self.frank, requester=synapse.types.create_requester(self.bob), field_name=ProfileFields.DISPLAYNAME, new_value="Frank Jr.", ), AuthError, ) def test_get_other_name(self) -> None: self.mock_federation.make_query.return_value = {"displayname": "Alice"} displayname = self.get_success(self.handler.get_displayname(self.alice)) self.assertEqual(displayname, "Alice") self.mock_federation.make_query.assert_called_with( destination="remote", query_type="profile", args={"user_id": "@alice:remote", "field": "displayname"}, ignore_backoff=True, ) def test_incoming_fed_query(self) -> None: self.get_success( self.store.create_profile(UserID.from_string("@caroline:test")) ) self.get_success( 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", ) ) response = self.get_success( self.query_handlers["profile"]( { "user_id": "@caroline:test", "field": "displayname", "origin": "servername.tld", } ) ) self.assertEqual({"displayname": "Caroline"}, response) def test_get_my_avatar(self) -> None: self.get_success( 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_field( user_id=self.frank, field_name=ProfileFields.DISPLAYNAME, new_value=None, ) ) self.get_success( self.store.set_profile_field( user_id=self.frank, field_name=ProfileFields.AVATAR_URL, new_value="http://my.server/me.png", ) ) profile = self.get_success(self.handler.get_profile(self.frank.to_string())) self.assertEqual("http://my.server/me.png", profile["avatar_url"]) def test_set_my_avatar(self) -> None: self.get_success( 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", ) ) self.assertEqual( (self.get_success(self.store.get_profile_avatar_url(self.frank))), "http://my.server/pic.gif", ) # Set avatar again self.get_success( 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( (self.get_success(self.store.get_profile_avatar_url(self.frank))), "http://my.server/me.png", ) # Set avatar to an empty string self.get_success( self.handler.set_field( target_user=self.frank, requester=synapse.types.create_requester(self.frank), field_name=ProfileFields.AVATAR_URL, new_value="", ) ) self.assertIsNone( (self.get_success(self.store.get_profile_avatar_url(self.frank))), ) @override_config({"enable_set_avatar_url": False}) def test_set_my_avatar_if_disabled(self) -> None: # Setting displayname for the first time is allowed self.get_success( 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( (self.get_success(self.store.get_profile_avatar_url(self.frank))), "http://my.server/me.png", ) # Set avatar a second time is forbidden self.get_failure( 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, ) def test_avatar_constraints_no_config(self) -> None: """Tests that the method to check an avatar against configured constraints skips all of its check if no constraint is configured. """ # The first check that's done by this method is whether the file exists; if we # don't get an error on a non-existing file then it means all of the checks were # successfully skipped. res = self.get_success( self.handler.check_avatar_size_and_mime_type("mxc://test/unknown_file") ) self.assertTrue(res) @unittest.override_config({"max_avatar_size": 50}) def test_avatar_constraints_allow_empty_avatar_url(self) -> None: """An empty avatar is always permitted.""" res = self.get_success(self.handler.check_avatar_size_and_mime_type("")) self.assertTrue(res) @unittest.override_config({"max_avatar_size": 50}) def test_avatar_constraints_missing(self) -> None: """Tests that an avatar isn't allowed if the file at the given MXC URI couldn't be found. """ res = self.get_success( self.handler.check_avatar_size_and_mime_type("mxc://test/unknown_file") ) self.assertFalse(res) @unittest.override_config({"max_avatar_size": 50}) def test_avatar_constraints_file_size(self) -> None: """Tests that a file that's above the allowed file size is forbidden but one that's below it is allowed. """ self._setup_local_files( { "small": {"size": 40}, "big": {"size": 60}, } ) res = self.get_success( self.handler.check_avatar_size_and_mime_type("mxc://test/small") ) self.assertTrue(res) res = self.get_success( self.handler.check_avatar_size_and_mime_type("mxc://test/big") ) self.assertFalse(res) @unittest.override_config({"allowed_avatar_mimetypes": ["image/png"]}) def test_avatar_constraint_mime_type(self) -> None: """Tests that a file with an unauthorised MIME type is forbidden but one with an authorised content type is allowed. """ self._setup_local_files( { "good": {"mimetype": "image/png"}, "bad": {"mimetype": "application/octet-stream"}, } ) res = self.get_success( self.handler.check_avatar_size_and_mime_type("mxc://test/good") ) self.assertTrue(res) res = self.get_success( self.handler.check_avatar_size_and_mime_type("mxc://test/bad") ) self.assertFalse(res) @unittest.override_config( {"server_name": "test:8888", "allowed_avatar_mimetypes": ["image/png"]} ) def test_avatar_constraint_on_local_server_with_port(self) -> None: """Test that avatar metadata is correctly fetched when the media is on a local server and the server has an explicit port. (This was previously a bug) """ local_server_name = self.hs.config.server.server_name media_id = "local" local_mxc = f"mxc://{local_server_name}/{media_id}" # mock up the existence of the avatar file self._setup_local_files({media_id: {"mimetype": "image/png"}}) # and now check that check_avatar_size_and_mime_type is happy self.assertTrue( self.get_success(self.handler.check_avatar_size_and_mime_type(local_mxc)) ) @parameterized.expand([("remote",), ("remote:1234",)]) @unittest.override_config({"allowed_avatar_mimetypes": ["image/png"]}) def test_check_avatar_on_remote_server(self, remote_server_name: str) -> None: """Test that avatar metadata is correctly fetched from a remote server""" media_id = "remote" remote_mxc = f"mxc://{remote_server_name}/{media_id}" # if the media is remote, check_avatar_size_and_mime_type just checks the # media cache, so we don't need to instantiate a real remote server. It is # sufficient to poke an entry into the db. self.get_success( self.hs.get_datastores().main.store_cached_remote_media( media_id=media_id, media_type="image/png", media_length=50, origin=remote_server_name, time_now_ms=self.clock.time_msec(), upload_name=None, filesystem_id="xyz", sha256="abcdefg12345", ) ) self.assertTrue( self.get_success(self.handler.check_avatar_size_and_mime_type(remote_mxc)) ) def _setup_local_files(self, names_and_props: dict[str, dict[str, Any]]) -> None: """Stores metadata about files in the database. Args: names_and_props: A dictionary with one entry per file, with the key being the file's name, and the value being a dictionary of properties. Supported properties are "mimetype" (for the file's type) and "size" (for the file's size). """ store = self.hs.get_datastores().main for name, props in names_and_props.items(): self.get_success( store.store_local_media( media_id=name, media_type=props.get("mimetype", "image/png"), time_now_ms=self.clock.time_msec(), upload_name=None, media_length=props.get("size", 50), user_id=UserID.from_string("@rin:test"), ) )