diff --git a/modules/restricted-guests/synapse/synapse_guest_module/guest_module.py b/modules/restricted-guests/synapse/synapse_guest_module/guest_module.py index 281cf9dd6b..7f53f5052e 100644 --- a/modules/restricted-guests/synapse/synapse_guest_module/guest_module.py +++ b/modules/restricted-guests/synapse/synapse_guest_module/guest_module.py @@ -218,6 +218,7 @@ class GuestModule: """ CREATE TABLE IF NOT EXISTS guest_module_mas_users ( mas_user_id TEXT PRIMARY KEY, + user_id TEXT, created_at BIGINT NOT NULL ) """, diff --git a/modules/restricted-guests/synapse/synapse_guest_module/guest_registration_servlet.py b/modules/restricted-guests/synapse/synapse_guest_module/guest_registration_servlet.py index fcf1469161..41257e5726 100644 --- a/modules/restricted-guests/synapse/synapse_guest_module/guest_registration_servlet.py +++ b/modules/restricted-guests/synapse/synapse_guest_module/guest_registration_servlet.py @@ -104,7 +104,7 @@ class GuestRegistrationServlet(DirectServeJsonResource): displayname + self._config.display_name_suffix, ) - await self._store_mas_user(mas_user_id, int(time.time())) + await self._store_mas_user(mas_user_id, user_id, int(time.time())) # Determine how long to keep the access token valid for. # @@ -135,7 +135,7 @@ class GuestRegistrationServlet(DirectServeJsonResource): return 500, {"msg": "Internal error: Could not find a free username"} - async def _store_mas_user(self, mas_user_id: str, created_at: int) -> None: + async def _store_mas_user(self, mas_user_id: str, user_id: str, created_at: int) -> None: if self._mas_tables_ready is not None: await self._mas_tables_ready.wait() @@ -145,6 +145,7 @@ class GuestRegistrationServlet(DirectServeJsonResource): table="guest_module_mas_users", values={ "mas_user_id": mas_user_id, + "user_id": user_id, "created_at": created_at, }, ) diff --git a/modules/restricted-guests/synapse/tests/__init__.py b/modules/restricted-guests/synapse/tests/__init__.py index d79a60a891..b97dd1c45b 100644 --- a/modules/restricted-guests/synapse/tests/__init__.py +++ b/modules/restricted-guests/synapse/tests/__init__.py @@ -160,5 +160,5 @@ def _setup_db(conn: sqlite3.Connection) -> None: "CREATE TABLE users(name text, deactivated smallint, creation_ts bigint)" ) conn.execute( - "CREATE TABLE guest_module_mas_users(mas_user_id text, created_at bigint)" + "CREATE TABLE guest_module_mas_users(mas_user_id text, user_id text, created_at bigint)" ) diff --git a/modules/restricted-guests/synapse/tests/test_guest_user_reaper.py b/modules/restricted-guests/synapse/tests/test_guest_user_reaper.py index 06815e1427..960e3f999f 100644 --- a/modules/restricted-guests/synapse/tests/test_guest_user_reaper.py +++ b/modules/restricted-guests/synapse/tests/test_guest_user_reaper.py @@ -136,10 +136,10 @@ class GuestUserReaperTest(aiounittest.AsyncTestCase): now = int(time.time()) store.conn.executemany( - "INSERT INTO guest_module_mas_users VALUES (?, ?)", + "INSERT INTO guest_module_mas_users VALUES (?, ?, ?)", [ - ["mas-old-1", 0], - ["mas-active", now], + ["mas-old-1", "@old-1:localhost", 0], + ["mas-active", "@active:localhost", now], ], ) @@ -161,6 +161,6 @@ class GuestUserReaperTest(aiounittest.AsyncTestCase): ) remaining_users = store.conn.execute( - "SELECT mas_user_id FROM guest_module_mas_users" + "SELECT mas_user_id, user_id FROM guest_module_mas_users" ).fetchall() - self.assertEqual(remaining_users, [("mas-active",)]) + self.assertEqual(remaining_users, [("mas-active", "@active:localhost")])