diff --git a/changelog.d/19974.feature b/changelog.d/19974.feature new file mode 100644 index 00000000000..4a9290ea763 --- /dev/null +++ b/changelog.d/19974.feature @@ -0,0 +1 @@ +Add experimental support for MSC4502: Targeted and unrestricted room member queries. diff --git a/rust/src/config/mod.rs b/rust/src/config/mod.rs index d79d12a83ae..b0ccbbb049b 100644 --- a/rust/src/config/mod.rs +++ b/rust/src/config/mod.rs @@ -73,4 +73,5 @@ pub struct ExperimentalConfig { pub msc4491_enabled: bool, pub msc4143_enabled: bool, pub msc4446_enabled: bool, + pub msc4502_enabled: bool, } diff --git a/rust/src/handlers/versions.rs b/rust/src/handlers/versions.rs index 25da9d23fc2..588818a12ea 100644 --- a/rust/src/handlers/versions.rs +++ b/rust/src/handlers/versions.rs @@ -269,6 +269,9 @@ pub struct UnstableFeatureMap { /// MSC4446: Allow moving the fully read marker backwards. #[serde(rename = "com.beeper.msc4446")] msc4446_enabled: bool, + /// MSC4502: Targeted and unrestricted room member queries + #[serde(rename = "io.element.msc4502")] + msc4502: bool, // Whether new rooms will be set to encrypted or not (based on presets). #[serde(rename = "io.element.e2ee_forced.public")] @@ -320,6 +323,7 @@ pub fn synapse_config_to_global_unstable_feature_map( msc4491_enabled: config.experimental.msc4491_enabled, msc4143_enabled: config.experimental.msc4143_enabled, msc4446_enabled: config.experimental.msc4446_enabled, + msc4502: config.experimental.msc4502_enabled, e2ee_forced_public: config .room .encryption_enabled_by_default_for_room_presets diff --git a/synapse/appservice/__init__.py b/synapse/appservice/__init__.py index c55a83a8799..7e95076eeee 100644 --- a/synapse/appservice/__init__.py +++ b/synapse/appservice/__init__.py @@ -62,6 +62,10 @@ # user ID -> {device ID -> [algorithm]} TransactionUnusedFallbackKeys = dict[str, dict[str, list[str]]] +# Scopes assignable to application services for extended privileges. +SCOPE_QUERY_ROOM_MEMBERSHIP = "urn:matrix:client:io.element.msc4502:rooms:is_joined" +KNOWN_SCOPES = frozenset({SCOPE_QUERY_ROOM_MEMBERSHIP}) + class ApplicationServiceState(Enum): DOWN = "down" @@ -104,6 +108,7 @@ def __init__( supports_unstable_ephemeral: bool = False, msc3202_transaction_extensions: bool = False, msc4190_device_management: bool = False, + scopes: Iterable[str] | None = None, ): self.token = token self.url = ( @@ -140,6 +145,11 @@ def __init__( else: self.protocols = set() + self.scopes = set(scopes) if scopes else set() + unknown_scopes = self.scopes - KNOWN_SCOPES + if unknown_scopes: + raise ValueError(f"Unknown application service scope(s): {unknown_scopes}") + self.rate_limited = rate_limited def _check_namespaces( @@ -379,6 +389,9 @@ def is_exclusive_user(self, user_id: str) -> bool: def is_interested_in_protocol(self, protocol: str) -> bool: return protocol in self.protocols + def has_scope(self, scope: str) -> bool: + return scope in self.scopes + def is_exclusive_alias(self, alias: str) -> bool: return self._is_exclusive(ApplicationService.NS_ALIASES, alias) diff --git a/synapse/config/appservice.py b/synapse/config/appservice.py index 7a629d10bf6..4e61ef694f0 100644 --- a/synapse/config/appservice.py +++ b/synapse/config/appservice.py @@ -199,6 +199,14 @@ def _load_appservice( "The `io.element.msc4190` option should be true or false if specified." ) + # Opt-in list of scopes granted to this appservice for restricted C-S API + # functionality. + scopes = as_info.get("io.element.msc4502.scopes", []) + if not isinstance(scopes, list) or not all(isinstance(s, str) for s in scopes): + raise ValueError( + "The `io.element.msc4502.scopes` option should be a list of strings if specified." + ) + return ApplicationService( token=as_info["as_token"], url=as_info["url"], @@ -213,4 +221,5 @@ def _load_appservice( supports_ephemeral=supports_ephemeral, msc3202_transaction_extensions=msc3202_transaction_extensions, msc4190_device_management=msc4190_enabled, + scopes=scopes, ) diff --git a/synapse/config/experimental.py b/synapse/config/experimental.py index 6ad9f535179..97dc803ab7c 100644 --- a/synapse/config/experimental.py +++ b/synapse/config/experimental.py @@ -203,6 +203,9 @@ def read_config( # See https://github.com/element-hq/synapse/issues/19524 self.msc4370_enabled = experimental.get("msc4370_enabled", False) + # MSC4502: Targeted and unrestricted room member queries + self.msc4502_enabled: bool = experimental.get("msc4502_enabled", False) + auth_delegated = (config.get("matrix_authentication_service") or {}).get( "enabled", False ) diff --git a/synapse/rest/__init__.py b/synapse/rest/__init__.py index a56a81a8e91..c8ede662aa9 100644 --- a/synapse/rest/__init__.py +++ b/synapse/rest/__init__.py @@ -60,6 +60,7 @@ retention, room, room_keys, + room_membership, room_upgrade_rest_servlet, sendtodevice, sync, @@ -128,6 +129,7 @@ rendezvous.register_servlets, auth_metadata.register_servlets, thread_subscriptions.register_servlets, + room_membership.register_servlets, ) SERVLET_GROUPS: dict[str, Iterable[RegisterServletsFunc]] = { diff --git a/synapse/rest/client/room_membership.py b/synapse/rest/client/room_membership.py new file mode 100644 index 00000000000..3193294900c --- /dev/null +++ b/synapse/rest/client/room_membership.py @@ -0,0 +1,133 @@ +# +# This file is licensed under the Affero General Public License (AGPL) version 3. +# +# Copyright (C) 2026 Element Creations 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: +# . +# + +import logging +import re +from http import HTTPStatus +from typing import TYPE_CHECKING + +from synapse.api.constants import EventTypes, Membership +from synapse.api.errors import Codes, SynapseError +from synapse.appservice import SCOPE_QUERY_ROOM_MEMBERSHIP +from synapse.http.server import HttpServer +from synapse.http.servlet import RestServlet, parse_string +from synapse.http.site import SynapseRequest +from synapse.types import JsonDict, RoomID, UserID +from synapse.util.stringutils import parse_and_validate_server_name + +if TYPE_CHECKING: + from synapse.server import HomeServer + +logger = logging.getLogger(__name__) + + +class AppserviceRoomMembershipRestServlet(RestServlet): + PATTERNS = [ + re.compile( + r"^/_matrix/client/unstable/io\.element\.msc4502/rooms/(?P[^/]*)/is_joined$" + ) + ] + CATEGORY = "Client API requests" + + def __init__(self, hs: "HomeServer"): + super().__init__() + self.auth = hs.get_auth() + self.store = hs.get_datastores().main + self.storage_controllers = hs.get_storage_controllers() + self.is_mine_id = hs.is_mine_id + self.is_mine_server_name = hs.is_mine_server_name + + async def on_GET( + self, request: SynapseRequest, room_id: str + ) -> tuple[int, JsonDict]: + requester = await self.auth.get_user_by_req(request, allow_guest=False) + + app_service = ( + self.store.get_app_service_by_id(requester.app_service_id) + if requester.app_service_id + else None + ) + + # Users and appservices can call this endpoint if they have the scope. + has_scope = ( + app_service is not None + and app_service.has_scope(SCOPE_QUERY_ROOM_MEMBERSHIP) + ) or SCOPE_QUERY_ROOM_MEMBERSHIP in requester.scope + + if not has_scope: + raise SynapseError( + HTTPStatus.FORBIDDEN, + f"Missing {SCOPE_QUERY_ROOM_MEMBERSHIP} scope", + Codes.FORBIDDEN, + ) + + if not RoomID.is_valid(room_id): + raise SynapseError( + HTTPStatus.BAD_REQUEST, "Invalid room ID", Codes.INVALID_PARAM + ) + + mxid = parse_string(request, "mxid") + server_name = parse_string(request, "server_name") + + if (mxid is None) == (server_name is None): + raise SynapseError( + HTTPStatus.BAD_REQUEST, + "Exactly one of 'mxid' or 'server_name' query parameters must be given", + Codes.MISSING_PARAM, + ) + + if mxid is not None: + if not UserID.is_valid(mxid): + raise SynapseError( + HTTPStatus.BAD_REQUEST, + "Invalid MXID: %s" % (mxid,), + Codes.INVALID_PARAM, + ) + if self.is_mine_id(mxid): + ( + membership, + _, + ) = await self.store.get_local_current_membership_for_user_in_room( + mxid, room_id + ) + joined = membership == Membership.JOIN + else: + event = await self.storage_controllers.state.get_current_state_event( + room_id, EventTypes.Member, mxid + ) + joined = ( + event is not None + and event.content.get("membership") == Membership.JOIN + ) + else: + assert server_name is not None + try: + parse_and_validate_server_name(server_name) + except ValueError: + raise SynapseError( + HTTPStatus.BAD_REQUEST, + "Invalid server name: %s" % (server_name,), + Codes.INVALID_PARAM, + ) + if self.is_mine_server_name(server_name): + joined = await self.store.is_locally_joined(room_id) + else: + joined = await self.store.is_host_joined(room_id, server_name) + + return HTTPStatus.OK, {"joined": joined} + + +def register_servlets(hs: "HomeServer", http_server: HttpServer) -> None: + if hs.config.experimental.msc4502_enabled: + AppserviceRoomMembershipRestServlet(hs).register(http_server) diff --git a/synapse/storage/databases/main/roommember.py b/synapse/storage/databases/main/roommember.py index 667ad1ace86..00e35399ced 100644 --- a/synapse/storage/databases/main/roommember.py +++ b/synapse/storage/databases/main/roommember.py @@ -595,6 +595,21 @@ async def get_local_users_in_room(self, room_id: str) -> Sequence[str]: desc="get_local_users_in_room", ) + @cached(max_entries=10000) + async def is_locally_joined(self, room_id: str) -> bool: + """ + Checks if any local user is currently joined to the given room. + """ + sql = """ + SELECT 1 FROM local_current_membership + WHERE room_id = ? AND membership = ? + LIMIT 1 + """ + rows = await self.db_pool.execute( + "is_locally_joined", sql, room_id, Membership.JOIN + ) + return bool(rows) + async def get_local_users_related_to_room( self, room_id: str ) -> list[tuple[str, str]]: diff --git a/tests/appservice/test_appservice.py b/tests/appservice/test_appservice.py index 620c2b907b2..8c0e9636f12 100644 --- a/tests/appservice/test_appservice.py +++ b/tests/appservice/test_appservice.py @@ -24,7 +24,11 @@ from twisted.internet import defer -from synapse.appservice import ApplicationService, Namespace +from synapse.appservice import ( + SCOPE_QUERY_ROOM_MEMBERSHIP, + ApplicationService, + Namespace, +) from synapse.types import UserID from tests import unittest @@ -257,3 +261,33 @@ def test_member_list_match(self) -> Generator["defer.Deferred[Any]", object, Non ) ) ) + + +class ApplicationServiceScopesTestCase(unittest.TestCase): + def test_has_no_scopes_by_default(self) -> None: + service = ApplicationService( + id="unique_identifier", + sender=UserID.from_string("@as:test"), + token="some_token", + ) + self.assertEqual(len(service.scopes), 0) + self.assertFalse(service.has_scope(SCOPE_QUERY_ROOM_MEMBERSHIP)) + + def test_has_valid_scope_if_specified(self) -> None: + service = ApplicationService( + id="unique_identifier", + sender=UserID.from_string("@as:test"), + token="some_token", + scopes=[SCOPE_QUERY_ROOM_MEMBERSHIP], + ) + self.assertEqual(len(service.scopes), 1) + self.assertTrue(service.has_scope(SCOPE_QUERY_ROOM_MEMBERSHIP)) + + def test_unknown_scope_raises(self) -> None: + with self.assertRaises(ValueError): + ApplicationService( + id="unique_identifier", + sender=UserID.from_string("@as:test"), + token="some_token", + scopes=["does:not:exist"], + ) diff --git a/tests/rest/client/test_room_membership.py b/tests/rest/client/test_room_membership.py new file mode 100644 index 00000000000..f71f01c282a --- /dev/null +++ b/tests/rest/client/test_room_membership.py @@ -0,0 +1,233 @@ +# +# This file is licensed under the Affero General Public License (AGPL) version 3. +# +# Copyright (C) 2026 Element Creations 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: +# . +# + +from http import HTTPStatus +from unittest.mock import AsyncMock, patch + +from twisted.internet.testing import MemoryReactor + +from synapse.appservice import SCOPE_QUERY_ROOM_MEMBERSHIP, ApplicationService +from synapse.rest import admin +from synapse.rest.client import login, room, room_membership +from synapse.server import HomeServer +from synapse.types import JsonDict, UserID, create_requester +from synapse.util.clock import Clock + +from tests import unittest +from tests.test_utils import event_injection +from tests.unittest import override_config + +AS_TOKEN = "i_am_an_app_service" +AS_TOKEN_NO_SCOPE = "i_am_an_app_service_without_scope" + + +class AppserviceRoomMembershipRestServletTestCase(unittest.HomeserverTestCase): + servlets = [ + admin.register_servlets_for_client_rest_resource, + login.register_servlets, + room.register_servlets, + room_membership.register_servlets, + ] + + def default_config(self) -> JsonDict: + config = super().default_config() + config["experimental_features"] = { + "msc4502_enabled": True, + **config.get("experimental_features", {}), + } + return config + + def prepare( + self, reactor: MemoryReactor, clock: Clock, homeserver: HomeServer + ) -> None: + self.creator = self.register_user("owner", "pass") + self.creator_tok = self.login("owner", "pass") + self.room_id = self.helper.create_room_as(self.creator, tok=self.creator_tok) + + self.joined_user = self.register_user("joined_user", "pass") + self.joined_user_tok = self.login("joined_user", "pass") + self.helper.join(self.room_id, self.joined_user, tok=self.joined_user_tok) + + self.not_joined_user = self.register_user("not_joined_user", "pass") + self.not_joined_user_tok = self.login("not_joined_user", "pass") + + self.remote_server = "elsewhere.com" + self.remote_user = UserID.from_string(f"@joined_user:{self.remote_server}") + self.get_success( + event_injection.inject_member_event( + self.hs, self.room_id, self.remote_user.to_string(), "join" + ) + ) + self.not_joined_remote_user = UserID.from_string( + f"@not_joined_user:{self.remote_server}" + ) + + self.unknown_server = "unknown.org" + self.unknown_room_id = "!unknown:unknown.org" + + main_store = self.hs.get_datastores().main + main_store.services_cache.append( + ApplicationService( + AS_TOKEN, + id="as_with_scope", + sender=UserID.from_string("@as:test"), + scopes=[SCOPE_QUERY_ROOM_MEMBERSHIP], + ) + ) + main_store.services_cache.append( + ApplicationService( + AS_TOKEN_NO_SCOPE, + id="as_without_scope", + sender=UserID.from_string("@as2:test"), + ) + ) + + def _get_joined( + self, room_id: str, params: str, access_token: str | None + ) -> tuple[int, JsonDict]: + channel = self.make_request( + "GET", + f"/_matrix/client/unstable/io.element.msc4502/rooms/{room_id}/is_joined?{params}", + access_token=access_token, + ) + return channel.code, channel.json_body + + def test_invalid_room_id_format(self) -> None: + code, body = self._get_joined( + "not-a-room-id", f"mxid={self.joined_user}", AS_TOKEN + ) + self.assertEqual(code, HTTPStatus.BAD_REQUEST, body) + self.assertEqual(body["errcode"], "M_INVALID_PARAM") + + def test_both_mxid_and_server_name_given(self) -> None: + code, body = self._get_joined( + self.room_id, + f"mxid={self.joined_user}&server_name={self.hs.hostname}", + AS_TOKEN, + ) + self.assertEqual(code, HTTPStatus.BAD_REQUEST, body) + self.assertEqual(body["errcode"], "M_MISSING_PARAM") + + def test_neither_mxid_nor_server_name_given(self) -> None: + code, body = self._get_joined(self.room_id, "", AS_TOKEN) + self.assertEqual(code, HTTPStatus.BAD_REQUEST, body) + self.assertEqual(body["errcode"], "M_MISSING_PARAM") + + def test_invalid_mxid_format(self) -> None: + code, body = self._get_joined(self.room_id, "mxid=not-a-userid", AS_TOKEN) + self.assertEqual(code, HTTPStatus.BAD_REQUEST, body) + self.assertEqual(body["errcode"], "M_INVALID_PARAM") + + def test_invalid_server_name_format(self) -> None: + code, body = self._get_joined(self.room_id, "server_name=foo_bar", AS_TOKEN) + self.assertEqual(code, HTTPStatus.BAD_REQUEST, body) + self.assertEqual(body["errcode"], "M_INVALID_PARAM") + + def test_local_user_joined(self) -> None: + code, body = self._get_joined( + self.room_id, f"mxid={self.joined_user}", AS_TOKEN + ) + self.assertEqual(code, HTTPStatus.OK, body) + self.assertEqual(body, {"joined": True}) + + def test_local_user_not_joined(self) -> None: + code, body = self._get_joined( + self.room_id, f"mxid={self.not_joined_user}", AS_TOKEN + ) + self.assertEqual(code, HTTPStatus.OK, body) + self.assertEqual(body, {"joined": False}) + + def test_remote_user_joined(self) -> None: + code, body = self._get_joined( + self.room_id, f"mxid={self.remote_user.to_string()}", AS_TOKEN + ) + self.assertEqual(code, HTTPStatus.OK, body) + self.assertEqual(body, {"joined": True}) + + def test_remote_user_not_joined(self) -> None: + code, body = self._get_joined( + self.room_id, f"mxid={self.not_joined_remote_user.to_string()}", AS_TOKEN + ) + self.assertEqual(code, HTTPStatus.OK, body) + self.assertEqual(body, {"joined": False}) + + def test_local_server_name_joined(self) -> None: + code, body = self._get_joined( + self.room_id, f"server_name={self.hs.hostname}", AS_TOKEN + ) + self.assertEqual(code, HTTPStatus.OK, body) + self.assertEqual(body, {"joined": True}) + + def test_remote_server_name_joined(self) -> None: + code, body = self._get_joined( + self.room_id, f"server_name={self.remote_server}", AS_TOKEN + ) + self.assertEqual(code, HTTPStatus.OK, body) + self.assertEqual(body, {"joined": True}) + + def test_remote_server_name_not_joined(self) -> None: + code, body = self._get_joined( + self.room_id, f"server_name={self.unknown_server}", AS_TOKEN + ) + self.assertEqual(code, HTTPStatus.OK, body) + self.assertEqual(body, {"joined": False}) + + def test_nonexistent_room_returns_false(self) -> None: + code, body = self._get_joined( + self.unknown_room_id, f"server_name={self.unknown_server}", AS_TOKEN + ) + self.assertEqual(code, HTTPStatus.OK, body) + self.assertEqual(body, {"joined": False}) + + def test_no_token_unauthorized(self) -> None: + code, body = self._get_joined(self.room_id, f"mxid={self.joined_user}", None) + self.assertEqual(code, HTTPStatus.UNAUTHORIZED, body) + self.assertEqual(body["errcode"], "M_MISSING_TOKEN") + + def test_normal_user_token_forbidden(self) -> None: + code, body = self._get_joined( + self.room_id, f"mxid={self.joined_user}", self.creator_tok + ) + self.assertEqual(code, HTTPStatus.FORBIDDEN, body) + self.assertEqual(body["errcode"], "M_FORBIDDEN") + + def test_same_user_token_forbidden(self) -> None: + code, body = self._get_joined( + self.room_id, f"mxid={self.joined_user}", self.joined_user_tok + ) + self.assertEqual(code, HTTPStatus.FORBIDDEN, body) + self.assertEqual(body["errcode"], "M_FORBIDDEN") + + def test_user_with_oauth_scope_allowed(self) -> None: + requester = create_requester(self.creator, scope={SCOPE_QUERY_ROOM_MEMBERSHIP}) + with patch.object( + self.hs.get_auth(), "get_user_by_req", AsyncMock(return_value=requester) + ): + code, body = self._get_joined( + self.room_id, f"mxid={self.joined_user}", "doesnt-matter" + ) + self.assertEqual(code, HTTPStatus.OK, body) + self.assertEqual(body, {"joined": True}) + + def test_appservice_without_scope_forbidden(self) -> None: + code, body = self._get_joined( + self.room_id, f"mxid={self.joined_user}", AS_TOKEN_NO_SCOPE + ) + self.assertEqual(code, HTTPStatus.FORBIDDEN, body) + self.assertEqual(body["errcode"], "M_FORBIDDEN") + + @override_config({"experimental_features": {"msc4502_enabled": False}}) + def test_unreachable_when_experimental_flag_disabled(self) -> None: + code, _ = self._get_joined(self.room_id, f"mxid={self.joined_user}", AS_TOKEN) + self.assertEqual(code, HTTPStatus.NOT_FOUND) diff --git a/tests/rest/client/test_versions.py b/tests/rest/client/test_versions.py index bbdbe38e072..1ed6bb145b4 100644 --- a/tests/rest/client/test_versions.py +++ b/tests/rest/client/test_versions.py @@ -153,6 +153,17 @@ def test_msc4446_true_if_enabled(self) -> None: self.assertEqual(channel.code, 200, channel.result) self.assertTrue(channel.json_body["unstable_features"]["com.beeper.msc4446"]) + def test_msc4502_false_by_default(self) -> None: + channel = self.make_request("GET", "/_matrix/client/versions") + self.assertEqual(channel.code, 200, channel.result) + self.assertFalse(channel.json_body["unstable_features"]["io.element.msc4502"]) + + @unittest.override_config({"experimental_features": {"msc4502_enabled": True}}) + def test_msc4502_true_if_enabled(self) -> None: + channel = self.make_request("GET", "/_matrix/client/versions") + self.assertEqual(channel.code, 200, channel.result) + self.assertTrue(channel.json_body["unstable_features"]["io.element.msc4502"]) + def _sanity_check_versions_response(self, versions_response: JsonDict) -> None: """ Make sure this looks like a `/_matrix/client/versions` response diff --git a/tests/storage/test_appservice.py b/tests/storage/test_appservice.py index 4b9d069d6a3..a199e06e85e 100644 --- a/tests/storage/test_appservice.py +++ b/tests/storage/test_appservice.py @@ -21,7 +21,7 @@ import json import os import tempfile -from typing import cast +from typing import Any, cast from unittest.mock import AsyncMock, Mock import yaml @@ -29,7 +29,11 @@ from twisted.internet import defer from twisted.internet.testing import MemoryReactor -from synapse.appservice import ApplicationService, ApplicationServiceState +from synapse.appservice import ( + SCOPE_QUERY_ROOM_MEMBERSHIP, + ApplicationService, + ApplicationServiceState, +) from synapse.config._base import ConfigError from synapse.events import EventBase from synapse.server import HomeServer @@ -38,7 +42,7 @@ ApplicationServiceStore, ApplicationServiceTransactionStore, ) -from synapse.types import DeviceListUpdates +from synapse.types import DeviceListUpdates, JsonDict from synapse.util.clock import Clock from tests import unittest @@ -479,8 +483,8 @@ def __init__( class ApplicationServiceStoreConfigTestCase(unittest.HomeserverTestCase): - def _write_config(self, suffix: str, **kwargs: str) -> str: - vals = { + def _write_config(self, suffix: str, **kwargs: Any) -> str: + vals: JsonDict = { "id": "id" + suffix, "url": "url" + suffix, "as_token": "as_token" + suffix, @@ -566,3 +570,68 @@ def test_duplicate_as_tokens(self) -> None: self.assertIn(f1, str(e)) self.assertIn(f2, str(e)) self.assertIn("as_token", str(e)) + + def test_invalid_scopes_raises(self) -> None: + f = self._write_config( + suffix="1", **{"io.element.msc4502.scopes": "not-a-list"} + ) + + self.hs.config.appservice.app_service_config_files = [f] + self.hs.config.caches.event_cache_size = 1 + + server_name = self.hs.hostname + database = self.hs.get_datastores().databases[0] + with self.assertRaises(ValueError): + ApplicationServiceStore( + database, + make_conn( + db_config=database._database_config, + engine=database.engine, + default_txn_name="test", + server_name=server_name, + ), + self.hs, + ) + + def test_known_scope_works(self) -> None: + f = self._write_config( + suffix="1", **{"io.element.msc4502.scopes": [SCOPE_QUERY_ROOM_MEMBERSHIP]} + ) + + self.hs.config.appservice.app_service_config_files = [f] + self.hs.config.caches.event_cache_size = 1 + + server_name = self.hs.hostname + database = self.hs.get_datastores().databases[0] + ApplicationServiceStore( + database, + make_conn( + db_config=database._database_config, + engine=database.engine, + default_txn_name="test", + server_name=server_name, + ), + self.hs, + ) + + def test_unknown_scope_raises(self) -> None: + f = self._write_config( + suffix="1", **{"io.element.msc4502.scopes": ["does:not:exist"]} + ) + + self.hs.config.appservice.app_service_config_files = [f] + self.hs.config.caches.event_cache_size = 1 + + server_name = self.hs.hostname + database = self.hs.get_datastores().databases[0] + with self.assertRaises(ValueError): + ApplicationServiceStore( + database, + make_conn( + db_config=database._database_config, + engine=database.engine, + default_txn_name="test", + server_name=server_name, + ), + self.hs, + ) diff --git a/tests/storage/test_roommember.py b/tests/storage/test_roommember.py index 711c9ba4ac6..5064050fdde 100644 --- a/tests/storage/test_roommember.py +++ b/tests/storage/test_roommember.py @@ -79,6 +79,14 @@ def test_one_member(self) -> None: self.assertEqual([self.room], [m.room_id for m in rooms_for_user]) + def test_is_locally_joined(self) -> None: + self.room = self.helper.create_room_as(self.u_alice, tok=self.t_alice) + + self.assertTrue(self.get_success(self.store.is_locally_joined(self.room))) + self.assertFalse( + self.get_success(self.store.is_locally_joined("!doesnotexist:test")) + ) + def test_count_known_servers(self) -> None: """ _count_known_servers will calculate how many servers are in a room.