from __future__ import annotations import secrets import threading from collections.abc import Callable from datetime import date, datetime, timedelta, timezone from typing import Any from backend.bootstrap.config import ( SESSION_MAX_AGE, USERNAME_PATTERN, add_months, normalize_date, parse_iso_datetime, ) from backend.features.accounts.security import ( SecretVault, hash_password, token_hash, verify_password, ) class AccountService: """Preserved account, session, membership and birth-profile behavior.""" def __init__( self, database: Any, vault: SecretVault, current_user_supplier: Callable[[], int], access_supplier: Callable[[], dict[str, Any]], bind_user: Callable[[int], None], personal_field_builder: Callable[..., dict[str, Any]], auth_lock: threading.Lock, ) -> None: self.database = database self.vault = vault self.current_user_supplier = current_user_supplier self.access_supplier = access_supplier self.bind_user = bind_user self.personal_field_builder = personal_field_builder self.auth_lock = auth_lock @property def current_user_id(self) -> int: return int(self.current_user_supplier()) @staticmethod def membership_for_access(access: dict[str, Any]) -> dict[str, Any]: now = datetime.now(timezone.utc) starts = parse_iso_datetime(access.get("membership_starts_at")) expires = parse_iso_datetime(access.get("membership_expires_at")) subscribed = ( access.get("membership_status") == "active" and (not starts or starts <= now) and (not expires or expires > now) ) is_admin = str(access.get("role")) == "admin" active = is_admin or subscribed remaining_seconds = None if expires: remaining_seconds = max(0, int((expires - now).total_seconds())) return { "active": active, "subscribed": subscribed, "status": "active" if subscribed else str(access.get("membership_status") or "inactive"), "plan": str(access.get("membership_plan") or ""), "starts_at": str(access.get("membership_starts_at") or ""), "expires_at": str(access.get("membership_expires_at") or ""), "is_admin": is_admin, "remaining_seconds": remaining_seconds, "remaining_days": None if remaining_seconds is None else (remaining_seconds + 86399) // 86400, } def membership(self) -> dict[str, Any]: access = self.access_supplier() or self.database.user_access(self.current_user_id) or {} return self.membership_for_access(access) def register(self, username: str, password: str) -> dict[str, Any]: username = username.strip() self.validate_input(username, password) with self.auth_lock: salt, password_digest = hash_password(password) user = self.database.create_user(username, salt, password_digest) return self.create_session(user) def login(self, username: str, password: str) -> dict[str, Any]: username = username.strip() if not username or not password: raise ValueError("账号名和密码不能为空。") user = self.database.user_by_username(username) if not user or not verify_password( password, str(user.get("password_salt") or ""), str(user.get("password_hash") or ""), ): raise ValueError("账号名或密码不正确。") return self.create_session(user) def change_password(self, current_password: str, new_password: str) -> None: current_password = str(current_password or "") access = self.database.user_access(self.current_user_id) self.validate_input(str(access["username"]), new_password) credentials = self.database.user_password(self.current_user_id) if not credentials or not verify_password( current_password, str(credentials.get("password_salt") or ""), str(credentials.get("password_hash") or ""), ): raise ValueError("当前密码不正确。") salt, digest = hash_password(new_password) if not self.database.update_user_password(self.current_user_id, salt, digest): raise ValueError("账号不存在。") def create_session(self, user: dict[str, Any]) -> dict[str, Any]: session_token = secrets.token_urlsafe(32) csrf_token = secrets.token_urlsafe(24) expires = datetime.now(timezone.utc) + timedelta(seconds=SESSION_MAX_AGE) self.database.create_session( token_hash(session_token), int(user["id"]), csrf_token, expires.isoformat(timespec="seconds"), ) self.bind_user(int(user["id"])) access = self.database.user_access(int(user["id"])) or {} return { "user": { "id": int(user["id"]), "username": str(user["username"]), "role": str(access.get("role") or "user"), "membership": self.membership(), }, "session_token": session_token, "csrf_token": csrf_token, } @staticmethod def validate_input(username: str, password: str) -> None: if not USERNAME_PATTERN.fullmatch(username): raise ValueError("账号名应为 3 至 30 位中文、字母、数字、下划线或连字符。") if len(password) < 8 or len(password) > 128: raise ValueError("密码长度应为 8 至 128 位。") if password.isalpha() or password.isdigit(): raise ValueError("密码应同时包含字母、数字或符号中的至少两类。") def save_birth_profile(self, payload: dict[str, Any]) -> dict[str, Any]: birth_datetime = str(payload.get("birth_datetime") or "").strip() gender = str(payload.get("gender") or "unspecified").strip() current_date = normalize_date(str(payload.get("trade_date") or date.today().isoformat())) personal = self.personal_field_builder(birth_datetime, gender, current_date) encrypted = self.vault.encrypt_json( {"birth_datetime": birth_datetime, "gender": gender} ) self.database.save_user_birth_profile(self.current_user_id, encrypted) return self.public_personal_profile(personal) def stored_birth_profile(self) -> dict[str, str] | None: encrypted = self.database.get_user_birth_profile(self.current_user_id) if not encrypted: return None payload = self.vault.decrypt_json(encrypted) birth_datetime = str(payload.get("birth_datetime") or "").strip() if not birth_datetime: return None return { "birth_datetime": birth_datetime, "gender": str(payload.get("gender") or "unspecified"), } def personal_field( self, current_date: str, current_field: dict[str, Any], public: bool = False, ) -> dict[str, Any] | None: stored = self.stored_birth_profile() if not stored: return None personal = self.personal_field_builder( stored["birth_datetime"], stored["gender"], current_date, current_field, ) if public: return self.public_personal_profile(personal) personal.pop("birth", None) return personal @staticmethod def public_personal_profile(personal: dict[str, Any]) -> dict[str, Any]: allowed = { "day_master", "ten_god_tendency", "element_balance", "balance_tendency", "current", "notice", } return {key: value for key, value in personal.items() if key in allowed} def update_membership(self, payload: dict[str, Any]) -> None: try: user_id = int(payload.get("user_id")) except (TypeError, ValueError) as exc: raise ValueError("会员账号不正确。") from exc status = str(payload.get("status") or "inactive") if status not in {"active", "inactive", "suspended"}: raise ValueError("会员状态不正确。") access = self.database.user_access(user_id) if not access: raise ValueError("用户不存在。") starts_at = None expires_at = None plan = "" if status == "active": duration = str(payload.get("duration") or "").strip() durations = { "1_month": (1, "1个月"), "3_months": (3, "3个月"), "12_months": (12, "12个月"), "3_years": (36, "3年"), "permanent": (0, "永久"), } if duration not in durations: raise ValueError("请选择会员开通时长。") now = datetime.now(timezone.utc) existing_start = parse_iso_datetime(access.get("membership_starts_at")) existing_expiry = parse_iso_datetime(access.get("membership_expires_at")) starts = existing_start if existing_start and existing_start <= now else now months, plan = durations[duration] starts_at = starts.isoformat(timespec="seconds") if months: renewal_base = existing_expiry if existing_expiry and existing_expiry > now else now expires_at = add_months(renewal_base, months).isoformat(timespec="seconds") if not self.database.update_membership( user_id, status, plan, starts_at, expires_at ): raise ValueError("用户不存在。") def admin_users( self, usage_supplier: Callable[[int], int] ) -> list[dict[str, Any]]: rows = [] for user in self.database.list_users(): membership = self.membership_for_access(user) used = usage_supplier(int(user["id"])) if membership["active"] else 0 rows.append({ **user, "membership_active": membership["active"], "membership_subscribed": membership["subscribed"], "used_today": used, }) return rows