from datetime import datetime, timezone from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.config import get_settings from app.core.heartbeat_verify import compute_server_verify_code from app.core.security import ( create_access_token, hash_password, session_redis_key, verify_password, ) from app.models.user import User from app.schemas.auth import HeartbeatResponse, TokenResponse, UserPublic settings = get_settings() class AuthError(Exception): def __init__(self, message: str) -> None: self.message = message super().__init__(message) def _normalize_device_serial(value: str | None) -> str: if not value: return "" return value.strip()[:255] def _vip_is_active(vip_end_time: datetime | None) -> bool: if vip_end_time is None: return False now = datetime.now(timezone.utc) end = vip_end_time if end.tzinfo is None: end = end.replace(tzinfo=timezone.utc) return end > now async def _sync_device_serial( db: AsyncSession, user: User, incoming: str, *, mismatch_message: str, ) -> None: stored = user.device_serial or "" if not stored: if incoming: user.device_serial = incoming await db.flush() elif incoming != stored: raise AuthError(mismatch_message) async def register_user( db: AsyncSession, *, username: str, password: str, device_serial: str | None = None, ) -> TokenResponse: existing = await db.scalar(select(User).where(User.username == username)) if existing is not None: raise AuthError("用户名已存在") serial = _normalize_device_serial(device_serial) if serial: bound = await db.scalar(select(User).where(User.device_serial == serial)) if bound is not None: raise AuthError("该电脑已绑定其它用户") user = User( username=username, password_hash=hash_password(password), role_id=0, phone=None, device_serial=serial, ) db.add(user) await db.flush() await db.refresh(user) token = create_access_token(str(user.id)) return TokenResponse( access_token=token, user=UserPublic.model_validate(user), ) async def login_user( db: AsyncSession, *, username: str, password: str, device_serial: str | None = None, ) -> TokenResponse: user = await db.scalar(select(User).where(User.username == username)) if user is None or not verify_password(password, user.password_hash): raise AuthError("用户名或密码错误") incoming = _normalize_device_serial(device_serial) await _sync_device_serial( db, user, incoming, mismatch_message="设备不匹配,无法登录" ) token = create_access_token(str(user.id)) return TokenResponse( access_token=token, user=UserPublic.model_validate(user), ) async def process_heartbeat( db: AsyncSession, user: User, device_serial: str | None, verify_code: str, ) -> HeartbeatResponse: incoming = _normalize_device_serial(device_serial) await _sync_device_serial(db, user, incoming, mismatch_message="设备不匹配") if not _vip_is_active(user.vip_end_time): raise AuthError("会员已过期或未开通") code = verify_code.strip() if not code: raise AuthError("verify_code 无效") return HeartbeatResponse( vip_end_time=user.vip_end_time, vip_active=True, server_verify_code=compute_server_verify_code(code), ) async def store_session(redis_client, token: str, user_id: int) -> None: ttl = settings.access_token_expire_minutes * 60 await redis_client.setex(session_redis_key(token), ttl, str(user_id)) async def revoke_session(redis_client, token: str) -> None: await redis_client.delete(session_redis_key(token)) async def session_user_id(redis_client, token: str) -> int | None: raw = await redis_client.get(session_redis_key(token)) if raw is None: return None return int(raw) async def get_user_by_id(db: AsyncSession, user_id: int) -> User | None: return await db.get(User, user_id)