from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.config import get_settings 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 TokenResponse, UserPublic settings = get_settings() class AuthError(Exception): def __init__(self, message: str) -> None: self.message = message super().__init__(message) async def register_user( db: AsyncSession, *, username: str, password: str, ) -> TokenResponse: existing = await db.scalar(select(User).where(User.username == username)) if existing is not None: raise AuthError("用户名已存在") user = User(username=username, password_hash=hash_password(password)) 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, ) -> 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("用户名或密码错误") token = create_access_token(str(user.id)) return TokenResponse( access_token=token, user=UserPublic.model_validate(user), ) 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)