from __future__ import annotations import threading import unittest from concurrent.futures import ThreadPoolExecutor from pathlib import Path from tempfile import TemporaryDirectory from backend.features.accounts.security import SecretVault from backend.features.accounts.service import AccountService from database import ReviewDatabase class InviteRegistrationTests(unittest.TestCase): def setUp(self) -> None: self.temp = TemporaryDirectory() self.addCleanup(self.temp.cleanup) self.database = ReviewDatabase(Path(self.temp.name) / "review.db") self.bound_user_id = 0 self.service = AccountService( database=self.database, vault=SecretVault(SecretVault.generate_key()), current_user_supplier=lambda: self.bound_user_id, access_supplier=lambda: self.database.user_access(self.bound_user_id) or {}, bind_user=self._bind, personal_field_builder=lambda *args, **kwargs: {}, auth_lock=threading.Lock(), ) def _bind(self, user_id: int) -> None: self.bound_user_id = int(user_id) def _bootstrap_admin(self) -> None: self.service.register("root_admin", "Password123") def _one_code(self) -> str: return self.service.generate_invite_codes(1)[0]["code"] def test_first_account_is_created_without_an_invite_code(self) -> None: result = self.service.register("root_admin", "Password123") self.assertEqual(result["user"]["role"], "admin") def test_registration_requires_an_invite_code_once_an_account_exists(self) -> None: self._bootstrap_admin() with self.assertRaises(ValueError) as error: self.service.register("second_user", "Password123") self.assertIn("邀请码", str(error.exception)) self.assertEqual(self.database.count_users(), 1) def test_unknown_used_and_revoked_codes_are_all_rejected(self) -> None: self._bootstrap_admin() with self.assertRaises(ValueError): self.service.register("second_user", "Password123", "", "XB-AAAA-AAAA-AAAA") code = self._one_code() self.service.register("second_user", "Password123", "", code) with self.assertRaises(ValueError) as used: self.service.register("third_user", "Password123", "", code) self.assertIn("已被使用", str(used.exception)) revoked = self._one_code() self.service.revoke_invite_code(revoked) with self.assertRaises(ValueError) as gone: self.service.register("fourth_user", "Password123", "", revoked) self.assertIn("作废", str(gone.exception)) self.assertEqual(self.database.count_users(), 2) def test_invite_code_is_accepted_with_or_without_separators(self) -> None: self._bootstrap_admin() code = self._one_code() self.service.register("second_user", "Password123", "", code.replace("-", "").lower()) record = self.database.invite_code(code) self.assertEqual(record["status"], "used") self.assertTrue(record["used_at"]) def test_concurrent_registrations_consume_one_code_once(self) -> None: self._bootstrap_admin() code = self._one_code() def attempt(index: int) -> str: try: self.service.register(f"racer_{index}", "Password123", "", code) return "ok" except ValueError as exc: return str(exc) with ThreadPoolExecutor(max_workers=6) as pool: outcomes = list(pool.map(attempt, range(6))) self.assertEqual(outcomes.count("ok"), 1) self.assertEqual(self.database.count_users(), 2) self.assertEqual(self.database.count_invite_codes()["used"], 1) def test_failed_account_creation_keeps_the_code_available(self) -> None: self._bootstrap_admin() code = self._one_code() with self.assertRaises(ValueError): self.service.register("root_admin", "Password123", "", code) self.assertEqual(self.database.invite_code(code)["status"], "unused") self.service.register("second_user", "Password123", "", code) self.assertEqual(self.database.invite_code(code)["status"], "used") def test_used_code_cannot_be_revoked_and_stays_reported(self) -> None: self._bootstrap_admin() code = self._one_code() self.service.register("second_user", "Password123", "", code) with self.assertRaises(ValueError): self.service.revoke_invite_code(code) overview = self.service.invite_overview() # total 让页面能直接显示"共 N 个",不必自己加总 self.assertEqual(overview["summary"], {"unused": 0, "used": 1, "revoked": 0}) row = overview["codes"][0] self.assertEqual(row["used_by_username"], "second_user") self.assertNotIn(code, row["code_masked"]) self.assertTrue(row["code_masked"].endswith("••••")) self.assertEqual(row["code_id"], AccountService.invite_handle(code)) def test_codes_can_be_revoked_through_their_public_handle(self) -> None: self._bootstrap_admin() code = self._one_code() self.service.revoke_invite_code(AccountService.invite_handle(code)) self.assertEqual(self.database.invite_code(code)["status"], "revoked") def test_batch_generation_is_bounded(self) -> None: self._bootstrap_admin() with self.assertRaises(ValueError): self.service.generate_invite_codes(AccountService.INVITE_MAX_BATCH + 1) created = self.service.generate_invite_codes(3, "内部测试") self.assertEqual(len({item["code"] for item in created}), 3) self.assertEqual(self.database.count_invite_codes()["unused"], 3) self.assertEqual(self.service.invite_overview()["codes"][0]["note"], "内部测试") if __name__ == "__main__": unittest.main()