Files
synapse/tests/handlers/test_profile.py
T

1455 lines
52 KiB
Python

#
# 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:
# <https://www.gnu.org/licenses/agpl-3.0.html>.
#
# Originally licensed under the Apache License, Version 2.0:
# <http://www.apache.org/licenses/LICENSE-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"),
)
)