mirror of
https://github.com/element-hq/synapse.git
synced 2026-09-02 11:23:47 +00:00
1455 lines
52 KiB
Python
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"),
|
|
)
|
|
)
|