From f7aa4a8d062ff1ee1bff2259588d9f4d89105147 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E8=85=BE=E8=AE=AFWorkBuddy?= Date: Sun, 16 Aug 2026 01:48:58 +0800 Subject: [PATCH] =?UTF-8?q?B-40:=20=E6=AD=A3=E5=BC=8F=E7=99=BB=E5=BD=95?= =?UTF-8?q?=E4=B8=8E=E7=A7=9F=E6=88=B7=E9=9A=94=E7=A6=BB=E2=80=94=E2=80=94?= =?UTF-8?q?=E8=AE=A4=E8=AF=81=E3=80=81=E4=BC=9A=E8=AF=9D=E3=80=81=E6=9D=83?= =?UTF-8?q?=E9=99=90=E4=B8=8E=E5=AE=A1=E8=AE=A1?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 管理员/公司两类角色:PBKDF2 密码散列、随机初始密码、首次改密、 停用、重置密码与吊销会话。 - 会话只存令牌摘要,8 小时绝对过期;登录失败同 (账号,IP) 限流。 - 公司账号服务端绑定唯一公司;读取/上传/导出/主数据/审核接口逐项 服务端授权,跨公司访问返回 404 而非 403。 - 审计日志记录登录、改密、上传、导出、建公司、建/停/启用户。 - 决策记录见 docs/decisions/003-auth.md。 --- docs/decisions/003-auth.md | 69 +++++ src/bank_importer/auth.py | 309 ++++++++++++++++++++++ tests/test_auth.py | 247 ++++++++++++++++++ tests/test_server_auth.py | 521 +++++++++++++++++++++++++++++++++++++ 4 files changed, 1146 insertions(+) create mode 100644 docs/decisions/003-auth.md create mode 100644 src/bank_importer/auth.py create mode 100644 tests/test_auth.py create mode 100644 tests/test_server_auth.py diff --git a/docs/decisions/003-auth.md b/docs/decisions/003-auth.md new file mode 100644 index 0000000..812e332 --- /dev/null +++ b/docs/decisions/003-auth.md @@ -0,0 +1,69 @@ +# 003 身份认证与租户隔离技术决策 + +对应 Issue:B-40(`docs/issues/003-p0-auth-and-tenant-isolation.md`)。 + +## 密码散列:标准库 PBKDF2-HMAC-SHA256 + +- 只使用 Python 标准库(`hashlib.pbkdf2_hmac` + `hmac.compare_digest` + + `secrets`),不引入 bcrypt/argon2 等第三方依赖,与离线内网部署约束一致。 +- 260 000 次迭代、16 字节随机盐,存储格式 + `pbkdf2_sha256$<迭代数>$<盐hex>$<摘要hex>`,自描述、可平滑升级参数。 +- 明文密码永不入库、永不进日志、永不写入 `audit_log`; + `password_hash` 字段从不出现在任何 API 响应中(有测试断言)。 + +## 会话:数据库保存令牌摘要,8 小时绝对过期 + +- 令牌为 `secrets.token_urlsafe(32)`,数据库只存其 SHA-256 摘要; + 数据库泄露不直接暴露可用令牌。 +- 绝对过期 8 小时(一个工作班次),不做滑动续期,语义简单可测。 +- 会话可吊销:退出登录、停用账号、重置密码都会立即吊销该用户全部会话。 +- Cookie 名 `cw_session`,`HttpOnly; SameSite=Lax; Path=/`。 + 当前是纯 HTTP 的局域网部署,**刻意不加 `Secure`**(加了浏览器会直接拒发); + 若未来上 HTTPS,应补上 `Secure` 并配置反向代理。 + +## 登录限流:同一 (账号, IP) 10 分钟内失败 5 次即锁定 + +- 计数来自 `login_attempts` 表,窗口为滚动 10 分钟;触发后返回 429, + 且在窗口内不再记录新尝试,行为确定、可测试。 +- 失败提示统一为「账号或密码不正确」,不泄露是哪一部分错误。 + +## 角色与公司绑定 + +- `users.role` 仅 `admin` / `company`;数据库 CHECK 约束强制 + 公司账号必须绑定公司、管理员不得绑定公司。 +- 登录时前端选择工作端口(portal),服务端校验 portal 与角色一致, + 不匹配返回 403「账号与该工作端口不匹配」。 +- 管理员创建公司账号时,初始密码等于登录账号本身并置 + `must_change_password=1`(2026-08-08 产品决定,替代随机初始密码); + 首次登录必须改密,改密前所有业务 API 返回 403。初始密码未改前 + 等同账号名,因此创建用户的审计 detail 只记录公司,不重复账号名。 +- 「重置密码」仍生成随机一次性密码:重置发生在用户已改密之后, + 可预测的口令会让知道账号名的人直接接管账户。 +- 已知风险:初始密码可预测,账号创建后应尽快完成首登改密; + 创建到改密之间,知道账号名的人即可登录该账号。 + +## 租户隔离在服务端强制,404 优于 403 + +- 公司用户的批次列表、行明细、CSV 导出、上传归属全部由服务端按 + 会话中的 `company_id` 过滤;请求体/参数里的 `company_id` 对公司用户 + 一律忽略(跨公司上传防护)。 +- 访问他公司批次返回 404 而非 403:不暴露「该批次存在但属于别人」这一事实, + 避免 IDOR 探测。导出时显式指定他公司 `company_id` 仍返回 403, + 因为用户已声明知道该公司存在,此时给出明确拒绝更有操作性。 +- `GET /admin.html`、`/company.html` 的 302 跳转只是 UX 层引导, + 不是权限边界;真正的边界全部在 API 上。 + +## 引导管理员(bootstrap) + +- 服务启动迁移后若无任何 admin 账号,创建 `APP_ADMIN_USERNAME` + (默认 `group-admin`);密码取 `APP_BOOTSTRAP_ADMIN_PASSWORD`, + 未设置则生成随机初始密码并**只打印一次到 stdout**(不写日志文件), + `must_change_password=1`。 +- 不在迁移或代码中预置任何公司或用户;公司与公司账号全部由管理员 + 通过 `/api/admin/companies`、`/api/admin/users` 动态创建。 + +## 审计 + +- `audit_log` 记录登录成功/失败、退出、改密、上传、导出(含公司范围与行数)、 + 建公司、建/停/启用户、重置密码等动作,含操作者、目标、IP、时间。 +- 初始密码、新旧密码均不进入审计内容。 diff --git a/src/bank_importer/auth.py b/src/bank_importer/auth.py new file mode 100644 index 0000000..da5941d --- /dev/null +++ b/src/bank_importer/auth.py @@ -0,0 +1,309 @@ +"""Authentication, sessions, rate limiting and audit logging. + +Passwords are hashed with PBKDF2-HMAC-SHA256 (stdlib ``hashlib.pbkdf2_hmac``) +and per-user random salts; plaintext passwords are never stored or logged. +Session tokens are random URL-safe strings; only their SHA-256 digest is +persisted, so a database leak does not expose usable tokens. Every login +attempt and every privileged action lands in ``audit_log``. The reasoning +behind these choices is recorded in ``docs/decisions/003-auth.md``. +""" + +from __future__ import annotations + +from datetime import datetime, timedelta, timezone +import hashlib +import hmac +import secrets +import sqlite3 +import string + +from .db import utc_now + + +MIN_PASSWORD_LENGTH = 8 +PBKDF2_ITERATIONS = 260_000 +SESSION_TTL_HOURS = 8 +RATE_LIMIT_MAX_FAILURES = 5 +RATE_LIMIT_WINDOW_MINUTES = 10 +INITIAL_PASSWORD_LENGTH = 12 + + +def hash_password(password: str) -> str: + """Hash ``password`` as ``pbkdf2_sha256$$$``.""" + salt = secrets.token_bytes(16) + digest = hashlib.pbkdf2_hmac( + "sha256", password.encode("utf-8"), salt, PBKDF2_ITERATIONS + ) + return f"pbkdf2_sha256${PBKDF2_ITERATIONS}${salt.hex()}${digest.hex()}" + + +def verify_password(password: str, stored: str) -> bool: + """Constant-time check of ``password`` against a stored hash string.""" + try: + scheme, iterations, salt_hex, hash_hex = stored.split("$") + if scheme != "pbkdf2_sha256": + return False + salt = bytes.fromhex(salt_hex) + expected = bytes.fromhex(hash_hex) + digest = hashlib.pbkdf2_hmac( + "sha256", password.encode("utf-8"), salt, int(iterations) + ) + except (ValueError, TypeError): + return False + return hmac.compare_digest(digest, expected) + + +def generate_initial_password() -> str: + """Generate a 12-char initial password with upper, lower and digit chars.""" + alphabet = string.ascii_letters + string.digits + while True: + password = "".join( + secrets.choice(alphabet) for _ in range(INITIAL_PASSWORD_LENGTH) + ) + if ( + any(char.isupper() for char in password) + and any(char.islower() for char in password) + and any(char.isdigit() for char in password) + ): + return password + + +def validate_password_policy(password: str, username: str) -> str | None: + """Return an error message when ``password`` violates policy, else None.""" + if len(password) < MIN_PASSWORD_LENGTH: + return f"密码长度至少为 {MIN_PASSWORD_LENGTH} 位。" + if password.lower() == username.lower(): + return "密码不能与账号相同。" + if not any(char.isalpha() for char in password) or not any( + char.isdigit() for char in password + ): + return "密码必须同时包含字母和数字。" + return None + + +def create_user( + connection: sqlite3.Connection, + username: str, + password: str, + role: str, + company_id: int | None = None, + must_change_password: bool = True, +) -> int: + """Create a user, enforcing the role/company binding rules. Returns the id.""" + username = username.strip() + if not username: + raise ValueError("用户名不能为空。") + if role not in ("admin", "company"): + raise ValueError("角色必须是 admin 或 company。") + if role == "company": + if company_id is None: + raise ValueError("公司账号必须绑定公司。") + company = connection.execute( + "SELECT id FROM companies WHERE id = ?", (company_id,) + ).fetchone() + if company is None: + raise ValueError("绑定的公司不存在。") + elif company_id is not None: + raise ValueError("管理员账号不能绑定公司。") + + now = utc_now() + try: + with connection: + cursor = connection.execute( + """ + INSERT INTO users ( + username, password_hash, role, company_id, + must_change_password, created_at, updated_at + ) VALUES (?, ?, ?, ?, ?, ?, ?) + """, + ( + username, + hash_password(password), + role, + company_id, + 1 if must_change_password else 0, + now, + now, + ), + ) + except sqlite3.IntegrityError as exc: + raise ValueError("用户名已存在。") from exc + return int(cursor.lastrowid) + + +def authenticate( + connection: sqlite3.Connection, username: str, password: str, ip: str +) -> tuple[sqlite3.Row | None, str | None]: + """Verify credentials; returns ``(user_row, None)`` or ``(None, reason)``. + + ``reason`` is one of ``rate_limited``, ``disabled``, ``bad_credentials``. + Every non-rate-limited attempt is recorded in ``login_attempts`` and + ``audit_log``; the password itself is never stored anywhere. + """ + window_start = ( + datetime.now(timezone.utc) - timedelta(minutes=RATE_LIMIT_WINDOW_MINUTES) + ).isoformat() + failures = connection.execute( + """ + SELECT COUNT(*) AS n FROM login_attempts + WHERE username = ? AND ip = ? AND success = 0 AND created_at >= ? + """, + (username, ip, window_start), + ).fetchone() + if failures["n"] >= RATE_LIMIT_MAX_FAILURES: + return None, "rate_limited" + + user = connection.execute( + "SELECT * FROM users WHERE username = ?", (username,) + ).fetchone() + + if user is not None and user["status"] == "disabled": + _record_attempt(connection, username, ip, success=False) + audit( + connection, + "login_failed", + actor=user, + detail="账号已停用", + ip=ip, + ) + return None, "disabled" + + if user is None or not verify_password(password, user["password_hash"]): + _record_attempt(connection, username, ip, success=False) + audit(connection, "login_failed", actor=user, detail="账号或密码不正确", ip=ip) + return None, "bad_credentials" + + _record_attempt(connection, username, ip, success=True) + audit(connection, "login_success", actor=user, ip=ip) + return user, None + + +def create_session( + connection: sqlite3.Connection, user_id: int, ttl_hours: int = SESSION_TTL_HOURS +) -> str: + """Create a session with absolute expiry; returns the raw token.""" + token = secrets.token_urlsafe(32) + token_hash = hashlib.sha256(token.encode("utf-8")).hexdigest() + now = datetime.now(timezone.utc) + with connection: + connection.execute( + """ + INSERT INTO sessions (token_hash, user_id, created_at, expires_at) + VALUES (?, ?, ?, ?) + """, + ( + token_hash, + user_id, + now.isoformat(), + (now + timedelta(hours=ttl_hours)).isoformat(), + ), + ) + return token + + +def resolve_session(connection: sqlite3.Connection, token: str) -> sqlite3.Row | None: + """Return the user row for a live session token, else None. + + Expired or revoked sessions and disabled users are all rejected. + """ + token_hash = hashlib.sha256(token.encode("utf-8")).hexdigest() + return connection.execute( + """ + SELECT u.*, s.id AS session_id + FROM sessions s + JOIN users u ON u.id = s.user_id + WHERE s.token_hash = ? + AND s.revoked_at IS NULL + AND s.expires_at > ? + AND u.status = 'active' + """, + (token_hash, utc_now()), + ).fetchone() + + +def revoke_session(connection: sqlite3.Connection, token: str) -> None: + token_hash = hashlib.sha256(token.encode("utf-8")).hexdigest() + with connection: + connection.execute( + "UPDATE sessions SET revoked_at = ? WHERE token_hash = ? AND revoked_at IS NULL", + (utc_now(), token_hash), + ) + + +def revoke_user_sessions(connection: sqlite3.Connection, user_id: int) -> None: + with connection: + connection.execute( + "UPDATE sessions SET revoked_at = ? WHERE user_id = ? AND revoked_at IS NULL", + (utc_now(), user_id), + ) + + +def change_password( + connection: sqlite3.Connection, + user_id: int, + old_password: str, + new_password: str, +) -> str | None: + """Change a user's password; returns an error message or None on success.""" + user = connection.execute( + "SELECT * FROM users WHERE id = ?", (user_id,) + ).fetchone() + if user is None: + return "用户不存在。" + if not verify_password(old_password, user["password_hash"]): + return "原密码不正确。" + error = validate_password_policy(new_password, user["username"]) + if error is not None: + return error + with connection: + connection.execute( + """ + UPDATE users + SET password_hash = ?, must_change_password = 0, updated_at = ? + WHERE id = ? + """, + (hash_password(new_password), utc_now(), user_id), + ) + audit(connection, "password_change", actor=user, target=f"user:{user_id}") + return None + + +def audit( + connection: sqlite3.Connection, + action: str, + actor: sqlite3.Row | None = None, + target: str | None = None, + detail: str | None = None, + ip: str | None = None, +) -> None: + """Append an audit log entry. Never pass passwords in ``detail``.""" + with connection: + connection.execute( + """ + INSERT INTO audit_log ( + actor_user_id, actor_username, action, target, detail, ip, created_at + ) VALUES (?, ?, ?, ?, ?, ?, ?) + """, + ( + actor["id"] if actor is not None else None, + actor["username"] if actor is not None else None, + action, + target, + detail, + ip, + utc_now(), + ), + ) + + +def _record_attempt( + connection: sqlite3.Connection, username: str, ip: str, success: bool +) -> None: + with connection: + connection.execute( + """ + INSERT INTO login_attempts (username, ip, success, created_at) + VALUES (?, ?, ?, ?) + """, + (username, ip, 1 if success else 0, utc_now()), + ) diff --git a/tests/test_auth.py b/tests/test_auth.py new file mode 100644 index 0000000..c70823d --- /dev/null +++ b/tests/test_auth.py @@ -0,0 +1,247 @@ +from __future__ import annotations + +from datetime import datetime, timedelta, timezone +import hashlib +import sqlite3 +import unittest + +from bank_importer import auth +from bank_importer.db import connect, migrate, utc_now + + +class AuthTestCase(unittest.TestCase): + def setUp(self) -> None: + self.connection = connect(":memory:") + self.addCleanup(self.connection.close) + migrate(self.connection) + now = utc_now() + with self.connection: + self.connection.execute( + "INSERT INTO companies (name, created_at, updated_at) VALUES ('甲公司', ?, ?)", + (now, now), + ) + self.company_id = int(self.connection.execute("SELECT id FROM companies").fetchone()["id"]) + + def create_company_user(self, username: str = "cashier-a", password: str = "Init1234") -> int: + return auth.create_user( + self.connection, username, password, "company", company_id=self.company_id + ) + + def create_admin(self, username: str = "group-admin", password: str = "Admin1234") -> int: + return auth.create_user(self.connection, username, password, "admin") + + +class PasswordHashTests(unittest.TestCase): + def test_hash_format_and_verify_roundtrip(self) -> None: + stored = auth.hash_password("Secret123") + scheme, iterations, salt_hex, hash_hex = stored.split("$") + self.assertEqual("pbkdf2_sha256", scheme) + self.assertEqual(str(auth.PBKDF2_ITERATIONS), iterations) + self.assertEqual(32, len(salt_hex)) + self.assertEqual(64, len(hash_hex)) + self.assertTrue(auth.verify_password("Secret123", stored)) + + def test_wrong_password_is_rejected(self) -> None: + stored = auth.hash_password("Secret123") + self.assertFalse(auth.verify_password("Secret124", stored)) + + def test_same_password_gets_different_salts(self) -> None: + self.assertNotEqual(auth.hash_password("Secret123"), auth.hash_password("Secret123")) + + def test_malformed_stored_hash_is_rejected(self) -> None: + for stored in ("", "plain", "pbkdf2_sha256$x$y$z", "bcrypt$1$aa$bb"): + with self.subTest(stored=stored): + self.assertFalse(auth.verify_password("Secret123", stored)) + + +class PasswordPolicyTests(unittest.TestCase): + def test_policy_rejects_short_username_like_and_single_charset(self) -> None: + self.assertIsNotNone(auth.validate_password_policy("Ab1", "cashier")) + self.assertIsNotNone(auth.validate_password_policy("Cashier99", "cashier99")) + self.assertIsNotNone(auth.validate_password_policy("abcdefgh", "cashier")) + self.assertIsNotNone(auth.validate_password_policy("12345678", "cashier")) + self.assertIsNone(auth.validate_password_policy("GoodPass1", "cashier")) + + def test_initial_password_generator_guarantees_charset(self) -> None: + for _ in range(50): + password = auth.generate_initial_password() + self.assertEqual(12, len(password)) + self.assertTrue(any(char.isupper() for char in password)) + self.assertTrue(any(char.islower() for char in password)) + self.assertTrue(any(char.isdigit() for char in password)) + self.assertIsNone(auth.validate_password_policy(auth.generate_initial_password(), "x")) + + +class CreateUserTests(AuthTestCase): + def test_company_role_requires_company(self) -> None: + with self.assertRaises(ValueError): + auth.create_user(self.connection, "cashier-x", "Init1234", "company") + + def test_company_must_exist(self) -> None: + with self.assertRaises(ValueError): + auth.create_user(self.connection, "cashier-x", "Init1234", "company", company_id=999) + + def test_admin_cannot_bind_company(self) -> None: + with self.assertRaises(ValueError): + auth.create_user( + self.connection, "admin-x", "Init1234", "admin", company_id=self.company_id + ) + + def test_invalid_role_rejected(self) -> None: + with self.assertRaises(ValueError): + auth.create_user(self.connection, "nobody", "Init1234", "superuser") + + def test_duplicate_username_rejected(self) -> None: + self.create_company_user() + with self.assertRaises(ValueError): + self.create_company_user() + + def test_new_company_user_must_change_password_by_default(self) -> None: + user_id = self.create_company_user() + user = self.connection.execute( + "SELECT must_change_password, status FROM users WHERE id = ?", (user_id,) + ).fetchone() + self.assertEqual(1, user["must_change_password"]) + self.assertEqual("active", user["status"]) + + +class AuthenticateTests(AuthTestCase): + def test_success_returns_user_and_records_attempt(self) -> None: + self.create_company_user() + user, reason = auth.authenticate(self.connection, "cashier-a", "Init1234", "127.0.0.1") + self.assertIsNotNone(user) + self.assertIsNone(reason) + attempt = self.connection.execute("SELECT success FROM login_attempts").fetchone() + self.assertEqual(1, attempt["success"]) + entry = self.connection.execute( + "SELECT action FROM audit_log WHERE action = 'login_success'" + ).fetchone() + self.assertIsNotNone(entry) + + def test_bad_credentials_do_not_leak_which_part_failed(self) -> None: + self.create_company_user() + for username, password in (("cashier-a", "Wrong999"), ("no-such-user", "Init1234")): + with self.subTest(username=username): + user, reason = auth.authenticate(self.connection, username, password, "127.0.0.1") + self.assertIsNone(user) + self.assertEqual("bad_credentials", reason) + + def test_disabled_user_gets_disabled_reason(self) -> None: + user_id = self.create_company_user() + with self.connection: + self.connection.execute( + "UPDATE users SET status = 'disabled', updated_at = ? WHERE id = ?", + (utc_now(), user_id), + ) + user, reason = auth.authenticate(self.connection, "cashier-a", "Init1234", "127.0.0.1") + self.assertIsNone(user) + self.assertEqual("disabled", reason) + + def test_five_failures_within_window_trigger_rate_limit(self) -> None: + self.create_company_user() + for _ in range(auth.RATE_LIMIT_MAX_FAILURES): + user, reason = auth.authenticate(self.connection, "cashier-a", "Wrong999", "10.0.0.1") + self.assertEqual("bad_credentials", reason) + user, reason = auth.authenticate(self.connection, "cashier-a", "Init1234", "10.0.0.1") + self.assertIsNone(user) + self.assertEqual("rate_limited", reason) + # A different IP is not limited. + user, reason = auth.authenticate(self.connection, "cashier-a", "Init1234", "10.0.0.2") + self.assertIsNotNone(user) + + def test_failures_outside_window_do_not_count(self) -> None: + self.create_company_user() + old = (datetime.now(timezone.utc) - timedelta(minutes=30)).isoformat() + with self.connection: + for _ in range(auth.RATE_LIMIT_MAX_FAILURES): + self.connection.execute( + "INSERT INTO login_attempts (username, ip, success, created_at) VALUES (?, ?, 0, ?)", + ("cashier-a", "10.0.0.9", old), + ) + user, reason = auth.authenticate(self.connection, "cashier-a", "Init1234", "10.0.0.9") + self.assertIsNotNone(user) + self.assertIsNone(reason) + + +class SessionTests(AuthTestCase): + def test_create_and_resolve_roundtrip(self) -> None: + user_id = self.create_company_user() + token = auth.create_session(self.connection, user_id) + user = auth.resolve_session(self.connection, token) + self.assertIsNotNone(user) + self.assertEqual(user_id, user["id"]) + # Only the digest is stored, never the raw token. + row = self.connection.execute("SELECT token_hash FROM sessions").fetchone() + self.assertEqual(hashlib.sha256(token.encode()).hexdigest(), row["token_hash"]) + self.assertNotIn(token, row["token_hash"]) + + def test_expired_session_is_rejected(self) -> None: + user_id = self.create_company_user() + token = auth.create_session(self.connection, user_id) + past = (datetime.now(timezone.utc) - timedelta(hours=1)).isoformat() + with self.connection: + self.connection.execute( + "UPDATE sessions SET expires_at = ? WHERE token_hash = ?", + (past, hashlib.sha256(token.encode()).hexdigest()), + ) + self.assertIsNone(auth.resolve_session(self.connection, token)) + + def test_revoked_session_is_rejected(self) -> None: + user_id = self.create_company_user() + token = auth.create_session(self.connection, user_id) + auth.revoke_session(self.connection, token) + self.assertIsNone(auth.resolve_session(self.connection, token)) + + def test_revoke_user_sessions_kills_all_sessions(self) -> None: + user_id = self.create_company_user() + first = auth.create_session(self.connection, user_id) + second = auth.create_session(self.connection, user_id) + auth.revoke_user_sessions(self.connection, user_id) + self.assertIsNone(auth.resolve_session(self.connection, first)) + self.assertIsNone(auth.resolve_session(self.connection, second)) + + def test_disabled_user_session_is_rejected(self) -> None: + user_id = self.create_company_user() + token = auth.create_session(self.connection, user_id) + with self.connection: + self.connection.execute( + "UPDATE users SET status = 'disabled', updated_at = ? WHERE id = ?", + (utc_now(), user_id), + ) + self.assertIsNone(auth.resolve_session(self.connection, token)) + + +class ChangePasswordTests(AuthTestCase): + def test_change_password_clears_flag_and_updates_hash(self) -> None: + user_id = self.create_company_user() + error = auth.change_password(self.connection, user_id, "Init1234", "NewPass99") + self.assertIsNone(error) + user = self.connection.execute( + "SELECT password_hash, must_change_password FROM users WHERE id = ?", (user_id,) + ).fetchone() + self.assertEqual(0, user["must_change_password"]) + self.assertTrue(auth.verify_password("NewPass99", user["password_hash"])) + self.assertFalse(auth.verify_password("Init1234", user["password_hash"])) + entry = self.connection.execute( + "SELECT detail FROM audit_log WHERE action = 'password_change'" + ).fetchone() + self.assertIsNotNone(entry) + self.assertNotIn("NewPass99", entry["detail"] or "") + + def test_wrong_old_password_rejected(self) -> None: + user_id = self.create_company_user() + error = auth.change_password(self.connection, user_id, "Wrong999", "NewPass99") + self.assertIsNotNone(error) + user = self.connection.execute( + "SELECT must_change_password FROM users WHERE id = ?", (user_id,) + ).fetchone() + self.assertEqual(1, user["must_change_password"]) + + def test_policy_violation_rejected(self) -> None: + user_id = self.create_company_user() + error = auth.change_password(self.connection, user_id, "Init1234", "short") + self.assertIsNotNone(error) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_server_auth.py b/tests/test_server_auth.py new file mode 100644 index 0000000..d06cca1 --- /dev/null +++ b/tests/test_server_auth.py @@ -0,0 +1,521 @@ +"""HTTP integration tests for authentication, RBAC and tenant isolation. + +Spins up a real ``ThreadingHTTPServer`` with a temp database/storage and +drives it with stdlib ``http.client`` (cookies handled by hand). The server +module reads ``APP_DB_PATH`` / ``APP_STORAGE_DIR`` from module globals at +request time, so tests patch them per class. +""" + +from __future__ import annotations + +from http.client import HTTPConnection +from http.cookies import SimpleCookie +import json +import os +from pathlib import Path +import tempfile +import threading +import unittest + +from bank_importer import auth +from bank_importer.db import connect, migrate, utc_now + +import server + + +ROOT = Path(__file__).resolve().parents[1] +SAMPLES = ROOT / "流水模板" +CCB_SAMPLE = SAMPLES / "中国建设银行账户流水.xls" +CITIC_SAMPLE = SAMPLES / "中信银行账户流水.xlsx" + +BOOTSTRAP_PASSWORD = "BootAdmin123" +ADMIN_PASSWORD = "AdminPass123" +CASHIER_A_PASSWORD = "CashierA123" + + +class Client: + """Minimal HTTP client with a cookie jar.""" + + def __init__(self, host: str, port: int) -> None: + self.host = host + self.port = port + self.cookies: dict[str, str] = {} + + def request( + self, + method: str, + path: str, + body: bytes | None = None, + headers: dict[str, str] | None = None, + ) -> tuple[int, dict[str, str], bytes]: + connection = HTTPConnection(self.host, self.port) + request_headers = dict(headers or {}) + if self.cookies: + request_headers["Cookie"] = "; ".join( + f"{key}={value}" for key, value in self.cookies.items() + ) + connection.request(method, path, body=body, headers=request_headers) + response = connection.getresponse() + data = response.read() + response_headers = {key.lower(): value for key, value in response.getheaders()} + set_cookie = response_headers.get("set-cookie") + if set_cookie: + cookie = SimpleCookie() + cookie.load(set_cookie) + for key, morsel in cookie.items(): + if morsel.value: + self.cookies[key] = morsel.value + else: + self.cookies.pop(key, None) + status = response.status + connection.close() + return status, response_headers, data + + def get(self, path: str) -> tuple[int, dict[str, str], bytes]: + return self.request("GET", path) + + def post_json(self, path: str, payload: dict) -> tuple[int, dict, bytes]: + return self.request( + "POST", + path, + body=json.dumps(payload).encode("utf-8"), + headers={"Content-Type": "application/json"}, + ) + + def post_multipart( + self, path: str, fields: dict[str, str], filename: str, content: bytes + ) -> tuple[int, dict, bytes]: + boundary = "----cwtestboundary7f3a9c1e" + parts: list[bytes] = [] + for name, value in fields.items(): + parts.append( + f'--{boundary}\r\nContent-Disposition: form-data; name="{name}"\r\n\r\n{value}\r\n'.encode() + ) + parts.append( + f'--{boundary}\r\nContent-Disposition: form-data; name="file"; filename="{filename}"\r\n' + "Content-Type: application/octet-stream\r\n\r\n".encode() + + content + + b"\r\n" + ) + parts.append(f"--{boundary}--\r\n".encode()) + return self.request( + "POST", + path, + body=b"".join(parts), + headers={"Content-Type": f"multipart/form-data; boundary={boundary}"}, + ) + + +def as_json(data: bytes) -> dict: + return json.loads(data.decode("utf-8")) + + +class ServerAuthMatrixTests(unittest.TestCase): + """One live server; setUpClass builds the shared fixture via the API.""" + + @classmethod + def setUpClass(cls) -> None: + cls.temp_dir = tempfile.TemporaryDirectory() + root = Path(cls.temp_dir.name) + cls.db_path = root / "app.db" + cls.storage = root / "files" + + cls._old_db_path = server.DB_PATH + cls._old_storage = server.STORAGE_DIR + server.DB_PATH = cls.db_path + server.STORAGE_DIR = cls.storage + + os.environ["APP_BOOTSTRAP_ADMIN_PASSWORD"] = BOOTSTRAP_PASSWORD + connection = connect(cls.db_path) + migrate(connection) + generated = server.ensure_bootstrap_admin(connection) + assert generated is None, "env password set, nothing should be generated" + connection.close() + + class QuietHandler(server.AppHandler): + def log_message(self, *args) -> None: # silence per-request logs + pass + + cls.httpd = server.ThreadingHTTPServer(("127.0.0.1", 0), QuietHandler) + cls.port = cls.httpd.server_address[1] + cls.thread = threading.Thread(target=cls.httpd.serve_forever, daemon=True) + cls.thread.start() + + cls.known_passwords = {BOOTSTRAP_PASSWORD, ADMIN_PASSWORD, CASHIER_A_PASSWORD} + + # --- Admin bootstrap: must_change_password gate, then change. --- + cls.admin = Client("127.0.0.1", cls.port) + status, _, data = cls.admin.post_json( + "/api/login", + {"username": "group-admin", "password": BOOTSTRAP_PASSWORD, "portal": "admin"}, + ) + assert status == 200, data + assert as_json(data)["must_change_password"] is True + status, _, data = cls.admin.get("/api/batches") + assert status == 403, "must_change_password must block API access" + status, _, data = cls.admin.post_json( + "/api/password/change", + {"old_password": BOOTSTRAP_PASSWORD, "new_password": ADMIN_PASSWORD}, + ) + assert status == 200, data + status, _, _ = cls.admin.get("/api/batches") + assert status == 200 + + # --- Companies A and B. --- + status, _, data = cls.admin.post_json("/api/admin/companies", {"name": "甲公司"}) + assert status == 200, data + cls.company_a = as_json(data)["company_id"] + status, _, data = cls.admin.post_json("/api/admin/companies", {"name": "乙公司"}) + assert status == 200, data + cls.company_b = as_json(data)["company_id"] + + # --- Company user A: initial password shown once, forced change. --- + status, _, data = cls.admin.post_json( + "/api/admin/users", {"username": "cashier-a", "company_id": cls.company_a} + ) + assert status == 200, data + payload = as_json(data) + # Product decision: registration initial password equals the username. + assert payload["initial_password"] == "cashier-a" + cls.initial_password_a = payload["initial_password"] + cls.known_passwords.add(cls.initial_password_a) + + cls.cashier_a = Client("127.0.0.1", cls.port) + status, _, data = cls.cashier_a.post_json( + "/api/login", + { + "username": "cashier-a", + "password": cls.initial_password_a, + "portal": "company", + }, + ) + assert status == 200, data + assert as_json(data)["must_change_password"] is True + status, _, _ = cls.cashier_a.get("/api/batches") + assert status == 403, "must_change_password must block company API access" + status, _, data = cls.cashier_a.post_json( + "/api/password/change", + {"old_password": cls.initial_password_a, "new_password": CASHIER_A_PASSWORD}, + ) + assert status == 200, data + + # --- B gets a batch (admin upload), A gets its own batch. --- + status, _, data = cls.admin.post_multipart( + "/api/parse", + {"company_id": str(cls.company_b)}, + CCB_SAMPLE.name, + CCB_SAMPLE.read_bytes(), + ) + assert status == 200, data + cls.b_batch_id = as_json(data)["batch_id"] + + status, _, data = cls.cashier_a.post_multipart( + "/api/parse", {}, CITIC_SAMPLE.name, CITIC_SAMPLE.read_bytes() + ) + assert status == 200, data + cls.a_batch_id = as_json(data)["batch_id"] + + @classmethod + def tearDownClass(cls) -> None: + cls.httpd.shutdown() + cls.httpd.server_close() + server.DB_PATH = cls._old_db_path + server.STORAGE_DIR = cls._old_storage + os.environ.pop("APP_BOOTSTRAP_ADMIN_PASSWORD", None) + cls.temp_dir.cleanup() + + # ------------------------------------------------------------------ + # Helpers + # ------------------------------------------------------------------ + + def fresh_client(self) -> Client: + return Client("127.0.0.1", self.port) + + def create_company_user(self, username: str) -> tuple[Client, str, int]: + status, _, data = self.admin.post_json( + "/api/admin/users", {"username": username, "company_id": self.company_a} + ) + self.assertEqual(200, status, data) + payload = as_json(data) + initial = payload["initial_password"] + self.known_passwords.add(initial) + client = self.fresh_client() + status, _, data = client.post_json( + "/api/login", + {"username": username, "password": initial, "portal": "company"}, + ) + self.assertEqual(200, status, data) + new_password = "Changed123" + self.known_passwords.add(new_password) + status, _, data = client.post_json( + "/api/password/change", + {"old_password": initial, "new_password": new_password}, + ) + self.assertEqual(200, status, data) + return client, new_password, payload["user_id"] + + # ------------------------------------------------------------------ + # Unauthenticated matrix + # ------------------------------------------------------------------ + + def test_unauthenticated_api_calls_return_401(self) -> None: + anon = self.fresh_client() + for method_check in ( + lambda: anon.post_multipart("/api/parse", {}, "x.xls", b"data"), + lambda: anon.get("/api/batches"), + lambda: anon.get(f"/api/batches/{self.a_batch_id}/rows"), + lambda: anon.get("/api/export.csv"), + lambda: anon.get("/api/admin/users"), + lambda: anon.get("/api/admin/companies"), + lambda: anon.get("/api/admin/audit-log"), + lambda: anon.get("/api/me"), + ): + status, _, data = method_check() + self.assertEqual(401, status, data) + self.assertEqual("error", as_json(data)["status"]) + + def test_unauthenticated_portal_pages_redirect(self) -> None: + anon = self.fresh_client() + for page in ("/admin.html", "/company.html"): + status, headers, _ = anon.get(page) + self.assertEqual(302, status, page) + self.assertEqual("/", headers.get("location")) + + def test_wrong_password_returns_generic_401(self) -> None: + anon = self.fresh_client() + status, _, data = anon.post_json( + "/api/login", + {"username": "group-admin", "password": "WrongPass1", "portal": "admin"}, + ) + self.assertEqual(401, status) + message = as_json(data)["message"] + self.assertNotIn("密码不正确", message.replace("账号或密码不正确", "")) + + def test_rate_limit_after_five_failures(self) -> None: + anon = self.fresh_client() + for _ in range(5): + status, _, _ = anon.post_json( + "/api/login", + {"username": "ghost-user", "password": "WrongPass1", "portal": "admin"}, + ) + self.assertEqual(401, status) + status, _, _ = anon.post_json( + "/api/login", + {"username": "ghost-user", "password": "WrongPass1", "portal": "admin"}, + ) + self.assertEqual(429, status) + + def test_portal_mismatch_returns_403(self) -> None: + anon = self.fresh_client() + status, _, data = anon.post_json( + "/api/login", + {"username": "group-admin", "password": ADMIN_PASSWORD, "portal": "company"}, + ) + self.assertEqual(403, status) + self.assertIn("端口", as_json(data)["message"]) + + def test_logout_revokes_session(self) -> None: + client, _, _ = self.create_company_user("cashier-logout") + status, _, _ = client.get("/api/me") + self.assertEqual(200, status) + status, _, _ = client.request("POST", "/api/logout") + self.assertEqual(200, status) + self.assertNotIn("cw_session", client.cookies) + status, _, _ = client.get("/api/me") + self.assertEqual(401, status) + + def test_expired_session_returns_401(self) -> None: + client, _, user_id = self.create_company_user("cashier-expired") + connection = connect(self.db_path) + try: + with connection: + connection.execute( + "UPDATE sessions SET expires_at = ? WHERE user_id = ?", + ("2000-01-01T00:00:00+00:00", user_id), + ) + finally: + connection.close() + status, _, _ = client.get("/api/me") + self.assertEqual(401, status) + + # ------------------------------------------------------------------ + # Admin and company lifecycle + # ------------------------------------------------------------------ + + def test_duplicate_company_name_returns_409(self) -> None: + status, _, _ = self.admin.post_json("/api/admin/companies", {"name": "甲公司"}) + self.assertEqual(409, status) + + def test_me_returns_profile_without_password_material(self) -> None: + status, _, data = self.cashier_a.get("/api/me") + self.assertEqual(200, status) + payload = as_json(data) + self.assertEqual("cashier-a", payload["username"]) + self.assertEqual("company", payload["role"]) + self.assertEqual(self.company_a, payload["company_id"]) + self.assertEqual("甲公司", payload["company_name"]) + self.assertFalse(payload["must_change_password"]) + + def test_company_user_forbidden_on_all_admin_endpoints(self) -> None: + calls = ( + lambda: self.cashier_a.get("/api/admin/companies"), + lambda: self.cashier_a.post_json("/api/admin/companies", {"name": "丙公司"}), + lambda: self.cashier_a.get("/api/admin/users"), + lambda: self.cashier_a.post_json( + "/api/admin/users", {"username": "x", "company_id": self.company_a} + ), + lambda: self.cashier_a.request("POST", "/api/admin/users/1/disable"), + lambda: self.cashier_a.request("POST", "/api/admin/users/1/enable"), + lambda: self.cashier_a.request("POST", "/api/admin/users/1/reset-password"), + lambda: self.cashier_a.get("/api/admin/audit-log"), + ) + for call in calls: + status, _, data = call() + self.assertEqual(403, status, data) + + def test_admin_upload_requires_company_id(self) -> None: + status, _, data = self.admin.post_multipart( + "/api/parse", {}, CCB_SAMPLE.name, CCB_SAMPLE.read_bytes() + ) + self.assertEqual(400, status, data) + + def test_admin_batches_filter_by_company(self) -> None: + status, _, data = self.admin.get(f"/api/batches?company_id={self.company_b}") + self.assertEqual(200, status) + batches = as_json(data)["batches"] + self.assertTrue(batches) + for batch in batches: + self.assertEqual(self.company_b, batch["company_id"]) + self.assertEqual("乙公司", batch["company_name"]) + + # ------------------------------------------------------------------ + # Tenant isolation + # ------------------------------------------------------------------ + + def test_batches_scoped_to_own_company(self) -> None: + status, _, data = self.cashier_a.get("/api/batches") + self.assertEqual(200, status) + batches = as_json(data)["batches"] + self.assertTrue(batches) + for batch in batches: + self.assertEqual(self.company_a, batch["company_id"]) + self.assertNotIn(self.b_batch_id, [batch["id"] for batch in batches]) + + def test_idor_batch_rows_of_other_company_return_404(self) -> None: + status, _, _ = self.cashier_a.get(f"/api/batches/{self.b_batch_id}/rows") + self.assertEqual(404, status) + status, _, data = self.cashier_a.get(f"/api/batches/{self.a_batch_id}/rows") + self.assertEqual(200, status) + self.assertTrue(as_json(data)["rows"]) + + def test_export_csv_forced_to_own_company(self) -> None: + status, _, _ = self.cashier_a.get(f"/api/export.csv?company_id={self.company_b}") + self.assertEqual(403, status) + + status, headers, data = self.cashier_a.get("/api/export.csv") + self.assertEqual(200, status) + self.assertEqual("text/csv; charset=utf-8", headers.get("content-type")) + text = data.decode("utf-8-sig") + lines = [line for line in text.splitlines() if line] + self.assertGreater(len(lines), 1) + for line in lines[1:]: + self.assertEqual(str(self.a_batch_id), line.split(",", 1)[0]) + + def test_cross_company_upload_is_recorded_under_own_company(self) -> None: + # A uploads B's file bytes while claiming company B in the form; the + # server must bind the new (duplicate) batch to A from the session. + status, _, data = self.cashier_a.post_multipart( + "/api/parse", + {"company_id": str(self.company_b)}, + CCB_SAMPLE.name, + CCB_SAMPLE.read_bytes(), + ) + self.assertEqual(200, status, data) + self.assertEqual("duplicate", as_json(data)["status"]) + connection = connect(self.db_path) + try: + duplicate = connection.execute( + "SELECT company_id FROM import_batches WHERE status = 'duplicate'" + ).fetchone() + finally: + connection.close() + self.assertIsNotNone(duplicate) + self.assertEqual(self.company_a, duplicate["company_id"]) + + # ------------------------------------------------------------------ + # Disable / reset flows + # ------------------------------------------------------------------ + + def test_disabled_user_session_and_login_rejected(self) -> None: + client, password, user_id = self.create_company_user("cashier-disable") + status, _, _ = client.get("/api/me") + self.assertEqual(200, status) + + status, _, data = self.admin.request("POST", f"/api/admin/users/{user_id}/disable") + self.assertEqual(200, status, data) + + status, _, _ = client.get("/api/me") + self.assertEqual(401, status) + + fresh = self.fresh_client() + status, _, data = fresh.post_json( + "/api/login", + {"username": "cashier-disable", "password": password, "portal": "company"}, + ) + self.assertEqual(403, status, data) + + def test_reset_password_returns_once_and_revokes_sessions(self) -> None: + client, old_password, user_id = self.create_company_user("cashier-reset") + status, _, data = self.admin.request( + "POST", f"/api/admin/users/{user_id}/reset-password" + ) + self.assertEqual(200, status, data) + new_password = as_json(data)["initial_password"] + self.known_passwords.add(new_password) + + status, _, _ = client.get("/api/me") + self.assertEqual(401, status) + + fresh = self.fresh_client() + status, _, data = fresh.post_json( + "/api/login", + {"username": "cashier-reset", "password": new_password, "portal": "company"}, + ) + self.assertEqual(200, status, data) + self.assertTrue(as_json(data)["must_change_password"]) + + # ------------------------------------------------------------------ + # Secrets hygiene + # ------------------------------------------------------------------ + + def test_no_response_contains_password_hash(self) -> None: + bodies = [] + status, _, data = self.admin.get("/api/admin/users") + self.assertEqual(200, status) + bodies.append(data) + status, _, data = self.admin.get("/api/admin/audit-log?limit=100") + self.assertEqual(200, status) + bodies.append(data) + status, _, data = self.cashier_a.get("/api/me") + bodies.append(data) + status, _, data = self.cashier_a.get("/api/batches") + bodies.append(data) + for body in bodies: + self.assertNotIn("password_hash", body.decode("utf-8")) + + def test_audit_log_contains_no_plaintext_passwords(self) -> None: + connection = connect(self.db_path) + try: + rows = connection.execute( + "SELECT detail, target FROM audit_log" + ).fetchall() + finally: + connection.close() + for password in self.known_passwords: + for row in rows: + self.assertNotIn(password, row["detail"] or "") + self.assertNotIn(password, row["target"] or "") + + +if __name__ == "__main__": + unittest.main()