"""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 os 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 _env_flag(value: str | None) -> bool: """Parse a yes/no style environment flag; blank/absent means False.""" return (value or "").strip().lower() in {"1", "true", "yes", "on"} # Ops switch for internal test environments: APP_LOGIN_RATE_LIMIT_DISABLED=1 # turns off the login-failure lockout entirely. The default (unset) keeps the # production policy — 5 failures within 10 minutes lock the (账号, IP) pair. RATE_LIMIT_DISABLED = _env_flag(os.environ.get("APP_LOGIN_RATE_LIMIT_DISABLED")) 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(exclude: str | None = None) -> str: """Generate a 12-char initial password with upper, lower and digit chars. When ``exclude`` is given, the result is guaranteed to differ from it (case-insensitive) so a fresh account never starts with a password equal to its own username. """ 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) and (exclude is None or password.lower() != exclude.lower()) ): 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. """ if not RATE_LIMIT_DISABLED: 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()), )