Merge commit '01af446be7cbd63b38f2fe35c2c9a25fac4fdef8' as 'modules/restricted-guests/synapse'

This commit is contained in:
Andrew Ferrazzutti
2025-02-04 09:21:06 -05:00
parent c37b2daf5f
commit 386073c0e8
20 changed files with 1330 additions and 0 deletions
@@ -0,0 +1,134 @@
# Copyright 2023 Nordeck IT + Consulting GmbH
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import sqlite3
from asyncio import Future
from typing import Any, Awaitable, Callable, Tuple, TypeVar
from unittest.mock import Mock
from synapse.http.client import SimpleHttpClient
from synapse.module_api import ModuleApi
from synapse_guest_module import GuestModule
RV = TypeVar("RV")
TV = TypeVar("TV")
class SQLiteStore:
"""In-memory SQLite store. We can't just use a run_db_interaction function that opens
its own connection, since we need to use the same connection for all queries in a
test.
"""
def __init__(self) -> None:
self.conn = sqlite3.connect(":memory:")
async def run_db_interaction(
self, desc: str, f: Callable[..., RV], *args: Any, **kwargs: Any
) -> RV:
cur = CursorWrapper(self.conn.cursor())
try:
res = f(cur, *args, **kwargs)
self.conn.commit()
return res
except Exception:
self.conn.rollback()
raise
class CursorWrapper:
"""Wrapper around a SQLite cursor."""
def __init__(self, cursor: sqlite3.Cursor) -> None:
self.cur = cursor
def execute(self, sql: str, args: Any) -> None:
self.cur.execute(sql, args)
@property
def rowcount(self) -> Any:
return self.cur.rowcount
def fetchone(self) -> Any:
return self.cur.fetchone()
def fetchall(self) -> Any:
return self.cur.fetchall()
def __iter__(self) -> Any:
return self.cur.__iter__()
def __next__(self) -> Any:
return self.cur.__next__()
def make_awaitable(result: TV) -> Awaitable[TV]:
"""
Makes an awaitable, suitable for mocking an `async` function.
This uses Futures as they can be awaited multiple times so can be returned
to multiple callers.
This function has been copied directly from Synapse's tests code.
"""
future = Future() # type: ignore
future.set_result(result)
return future
def get_qualified_user_id(username: str) -> str:
return f"@{username}:matrix.local"
async def register_user(localpart: str, admin: bool = False) -> str:
return f"@{localpart}:matrix.local"
def create_module() -> Tuple[GuestModule, Mock, SQLiteStore]:
store = SQLiteStore()
_setup_db(store.conn)
client = Mock(spec=SimpleHttpClient)
client.post_json_get_json.return_value = make_awaitable(None)
# Create a mock based on the ModuleApi spec, but override some mocked functions
# because some capabilities are needed for running the tests.
module_api = Mock(spec=ModuleApi)
module_api.http_client = client
module_api.server_name = "matrix.local"
module_api.public_baseurl = "https://matrix.local:1234/"
module_api.run_db_interaction.side_effect = store.run_db_interaction
module_api.get_qualified_user_id.side_effect = get_qualified_user_id
module_api.check_user_exists.return_value = make_awaitable(False)
module_api.register_user.side_effect = register_user
module_api.register_device.return_value = make_awaitable(
("DEVICEID", "syn_registered_token", None, None)
)
# If necessary, give parse_config some configuration to parse.
config = GuestModule.parse_config(
{
"enable_user_reaper": False,
}
)
module = GuestModule(config, module_api)
return module, module_api, store
def _setup_db(conn: sqlite3.Connection) -> None:
conn.execute("CREATE TABLE access_tokens(user_id text, token text)")
conn.execute(
"CREATE TABLE users(name text, deactivated smallint, creation_ts bigint)"
)
@@ -0,0 +1,202 @@
# Copyright 2023 Nordeck IT + Consulting GmbH
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import aiounittest
from synapse.module_api import ProfileInfo, UserProfile
from synapse.module_api.errors import ConfigError
from synapse.types import UserID
from synapse_guest_module.config import GuestModuleConfig
from synapse_guest_module.guest_module import GuestModule
from tests import create_module
class GuestModuleTest(aiounittest.AsyncTestCase):
async def test_parse_config_empty(self) -> None:
config = GuestModule.parse_config({})
self.assertEqual(
config,
GuestModuleConfig(
user_id_prefix="guest-",
display_name_suffix=" (Guest)",
enable_user_reaper=True,
user_expiration_seconds=24 * 60 * 60,
),
)
async def test_parse_config_custom(self) -> None:
config = GuestModule.parse_config(
{
"user_id_prefix": "tmp-",
"display_name_suffix": " (Temporary)",
"enable_user_reaper": False,
"user_expiration_seconds": 100,
}
)
self.assertEqual(
config,
GuestModuleConfig(
user_id_prefix="tmp-",
display_name_suffix=" (Temporary)",
enable_user_reaper=False,
user_expiration_seconds=100,
),
)
async def test_parse_config_fail_user_id_prefix(self) -> None:
with self.assertRaisesRegex(
ConfigError, "Config option 'user_id_prefix' must be a string"
):
GuestModule.parse_config(
{
"user_id_prefix": 1234,
}
)
async def test_parse_config_fail_display_name_suffix(self) -> None:
with self.assertRaisesRegex(
ConfigError, "Config option 'display_name_suffix' must be a string"
):
GuestModule.parse_config(
{
"display_name_suffix": 1234,
}
)
async def test_parse_config_fail_enable_user_reaper(self) -> None:
with self.assertRaisesRegex(
ConfigError, "Config option 'enable_user_reaper' must be a bool"
):
GuestModule.parse_config(
{
"enable_user_reaper": "False",
}
)
async def test_parse_config_fail_user_expiration_seconds(self) -> None:
with self.assertRaisesRegex(
ConfigError, "Config option 'user_expiration_seconds' must be a number"
):
GuestModule.parse_config(
{
"user_expiration_seconds": "1",
}
)
async def test_profile_update_no_guest(self) -> None:
module, module_api, _ = create_module()
await module.profile_update(
"@my-user:matrix.local",
ProfileInfo(display_name="My User", avatar_url=None),
True,
False,
)
module_api.set_displayname.assert_not_called()
async def test_profile_update_guest_keep(self) -> None:
module, module_api, _ = create_module()
await module.profile_update(
"@guest-asdf:matrix.local",
ProfileInfo(display_name="My User (Guest)", avatar_url=None),
True,
False,
)
module_api.set_displayname.assert_not_called()
async def test_profile_update_guest_add_and_trim(self) -> None:
module, module_api, _ = create_module()
await module.profile_update(
"@guest-asdf:matrix.local",
ProfileInfo(display_name="My User ", avatar_url=None),
True,
False,
)
module_api.set_displayname.assert_awaited_once_with(
UserID.from_string("@guest-asdf:matrix.local"),
"My User (Guest)",
)
async def test_callback_user_may_create_room_no_guest(self) -> None:
module, _, _ = create_module()
allow = await module.callback_user_may_create_room(
"@my-user:matrix.local",
)
self.assertTrue(allow)
async def test_callback_user_may_create_room_guest(self) -> None:
module, _, _ = create_module()
allow = await module.callback_user_may_create_room(
"@guest-asdf:matrix.local",
)
self.assertFalse(allow)
async def test_callback_user_may_invite_no_guest(self) -> None:
module, _, _ = create_module()
allow = await module.callback_user_may_invite(
"@my-user:matrix.local",
"@inviter:matrix.local",
"!room:matrix.local",
)
self.assertTrue(allow)
async def test_callback_user_may_invite_guest(self) -> None:
module, _, _ = create_module()
allow = await module.callback_user_may_invite(
"@guest-asdf:matrix.local",
"@inviter:matrix.local",
"!room:matrix.local",
)
self.assertFalse(allow)
async def test_callback_check_username_for_spam_no_guest(self) -> None:
module, _, _ = create_module()
allow = await module.callback_check_username_for_spam(
UserProfile(
user_id="@my-user:matrix.local",
display_name=None,
avatar_url=None,
),
)
self.assertFalse(allow)
async def test_callback_check_username_for_spam_guest(self) -> None:
module, _, _ = create_module()
allow = await module.callback_check_username_for_spam(
UserProfile(
user_id="@guest-asdf:matrix.local",
display_name=None,
avatar_url=None,
),
)
self.assertTrue(allow)
@@ -0,0 +1,91 @@
# Copyright 2023 Nordeck IT + Consulting GmbH
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import io
from typing import cast
from unittest.mock import ANY
import aiounittest
from twisted.web.server import Request
from twisted.web.test.requesthelper import DummyRequest
from tests import create_module, make_awaitable
class GuestUserReaperTest(aiounittest.AsyncTestCase):
async def test_async_render_POST_missing_displayname(self) -> None:
module, _, _ = create_module()
request = cast(Request, DummyRequest([]))
request.content = io.BytesIO(b"{}")
status, response = await module.registration_servlet._async_render_POST(request)
self.assertEqual(status, 400)
self.assertEqual(
response, {"msg": "You must provide a 'displayname' as a string"}
)
async def test_async_render_POST_empty_displayname(self) -> None:
module, _, _ = create_module()
request = cast(Request, DummyRequest([]))
request.content = io.BytesIO(b'{"displayname":" "}')
status, response = await module.registration_servlet._async_render_POST(request)
self.assertEqual(status, 400)
self.assertEqual(
response, {"msg": "You must provide a 'displayname' as a string"}
)
async def test_async_render_POST_no_free_username(self) -> None:
module, module_api, _ = create_module()
request = cast(Request, DummyRequest([]))
request.content = io.BytesIO(b'{"displayname":"My Name"}')
module_api.check_user_exists.return_value = make_awaitable(True)
status, response = await module.registration_servlet._async_render_POST(request)
self.assertEqual(status, 500)
self.assertEqual(
response, {"msg": "Internal error: Could not find a free username"}
)
self.assertEqual(module_api.check_user_exists.call_count, 10)
async def test_async_render_POST_success(self) -> None:
module, module_api, _ = create_module()
request = cast(Request, DummyRequest([]))
request.content = io.BytesIO(b'{"displayname":"My Name "}')
status, response = await module.registration_servlet._async_render_POST(request)
self.assertEqual(status, 201)
self.assertRegex(response.pop("userId"), r"^@guest-[A-Za-z0-9]+:matrix.local$")
self.assertDictEqual(
response,
{
"accessToken": "syn_registered_token",
"deviceId": "DEVICEID",
"homeserverUrl": "https://matrix.local:1234/",
# "userId" was already checked by self.assertRegex and was removed from the object
},
)
module_api.register_user.assert_called_with(ANY, "My Name (Guest)")
@@ -0,0 +1,128 @@
# Copyright 2023 Nordeck IT + Consulting GmbH
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import time
from unittest.mock import call
import aiounittest
from tests import create_module, make_awaitable
class GuestUserReaperTest(aiounittest.AsyncTestCase):
async def test_get_admin_token_register(self) -> None:
module, module_api, _ = create_module()
token = await module.reaper.get_admin_token()
module_api.check_user_exists.assert_called_with("guest-reaper")
module_api.register_user.assert_called_with("guest-reaper", admin=True)
module_api.register_device.assert_called_with("@guest-reaper:matrix.local")
self.assertEqual(token, "syn_registered_token")
async def test_get_admin_token_create_device(self) -> None:
module, module_api, _ = create_module()
module_api.check_user_exists.return_value = make_awaitable(True)
token = await module.reaper.get_admin_token()
module_api.check_user_exists.assert_called_with("guest-reaper")
module_api.register_user.assert_not_called()
module_api.register_device.assert_called_with("@guest-reaper:matrix.local")
self.assertEqual(token, "syn_registered_token")
async def test_get_admin_token_read_from_db(self) -> None:
module, module_api, store = create_module()
store.conn.execute(
"INSERT INTO access_tokens VALUES ('@guest-reaper:matrix.local', 'syn_db_token')"
)
module_api.check_user_exists.return_value = make_awaitable(True)
token = await module.reaper.get_admin_token()
module_api.check_user_exists.assert_not_called()
module_api.register_user.assert_not_called()
module_api.register_device.assert_not_called()
self.assertEqual(token, "syn_db_token")
async def test_deactivate_expired_guest_users_success(self) -> None:
module, module_api, store = create_module()
now = int(time.time())
store.conn.executemany(
"INSERT INTO users VALUES (?, ?, ?)",
[
["@user-1:matrix.local", 0, 0],
["@guest-reaper:matrix.local", 0, 0],
["@guest-active:matrix.local", 0, now],
["@guest-deactivated:matrix.local", 1, 0],
["@guest-old-1:matrix.local", 0, 0],
["@guest-old-2:matrix.local", 0, 0],
],
)
await module.reaper.deactivate_expired_guest_users()
self.assertEqual(module_api.http_client.post_json_get_json.await_count, 2)
module_api.http_client.post_json_get_json.assert_has_awaits(
[
call(
uri="http://localhost:8008/_synapse/admin/v1/deactivate/@guest-old-1:matrix.local",
post_json={},
headers={"Authorization": ["Bearer syn_registered_token"]},
),
call(
uri="http://localhost:8008/_synapse/admin/v1/deactivate/@guest-old-2:matrix.local",
post_json={},
headers={"Authorization": ["Bearer syn_registered_token"]},
),
]
)
async def test_deactivate_expired_guest_users_with_failure(self) -> None:
module, module_api, store = create_module()
store.conn.executemany(
"INSERT INTO users VALUES (?, ?, ?)",
[
["@guest-old-1:matrix.local", 0, 0],
["@guest-old-2:matrix.local", 0, 0],
],
)
module_api.http_client.post_json_get_json.side_effect = Exception("")
await module.reaper.deactivate_expired_guest_users()
module_api.http_client.post_json_get_json.assert_has_awaits(
[
call(
uri="http://localhost:8008/_synapse/admin/v1/deactivate/@guest-old-1:matrix.local",
post_json={},
headers={"Authorization": ["Bearer syn_registered_token"]},
),
call(
uri="http://localhost:8008/_synapse/admin/v1/deactivate/@guest-old-2:matrix.local",
post_json={},
headers={"Authorization": ["Bearer syn_registered_token"]},
),
]
)