from datetime import datetime, timezone from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.config import get_settings from app.core.roles import ROLE_AGENT, ROLE_OEM 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 parse_id_header(value: str | None, *, field_name: str) -> int | None: if value is None or not str(value).strip(): return None try: parsed = int(str(value).strip()) except ValueError as exc: raise AuthError(f"{field_name} 无效") from exc if parsed < 1: raise AuthError(f"{field_name} 无效") return parsed async def _validate_register_owner_ids( db: AsyncSession, *, oem_id: int | None, agent_id: int | None, ) -> tuple[int | None, int | None]: resolved_oem_id = oem_id if agent_id is not None: agent = await db.get(User, agent_id) if agent is None or agent.role_id != ROLE_AGENT: raise AuthError("代理无效") if resolved_oem_id is None: resolved_oem_id = agent.oem_id elif agent.oem_id != resolved_oem_id: raise AuthError("代理与 OEM 不匹配") if resolved_oem_id is not None: oem = await db.get(User, resolved_oem_id) if oem is None or oem.role_id != ROLE_OEM: raise AuthError("OEM 无效") return resolved_oem_id, agent_id 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, oem_id: int | None = None, agent_id: int | 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("该电脑已绑定其它用户") resolved_oem_id, resolved_agent_id = await _validate_register_owner_ids( db, oem_id=oem_id, agent_id=agent_id ) user = User( username=username, password_hash=hash_password(password), role_id=0, phone=None, device_serial=serial, oem_id=resolved_oem_id, agent_id=resolved_agent_id, ) 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)