| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104 |
- from __future__ import annotations
- from datetime import datetime
- from sqlalchemy import delete, select
- from sqlalchemy.orm import Session
- from supply_infra.db.models.auth_session import AuthSession
- from supply_infra.db.models.auth_user import AuthUser
- class AuthRepository:
- """Persistence operations for local users and server-side sessions."""
- def __init__(self, session: Session) -> None:
- self.session = session
- def get_user(self, user_id: int) -> AuthUser | None:
- return self.session.get(AuthUser, user_id)
- def get_user_by_username(self, username: str) -> AuthUser | None:
- stmt = select(AuthUser).where(AuthUser.username == username)
- return self.session.scalar(stmt)
- def list_users(self) -> list[AuthUser]:
- stmt = select(AuthUser).order_by(AuthUser.created_at.asc(), AuthUser.id.asc())
- return list(self.session.scalars(stmt).all())
- def create_user(
- self,
- *,
- username: str,
- password_hash: str,
- display_name: str,
- role: str,
- status: str = "active",
- ) -> AuthUser:
- user = AuthUser(
- username=username,
- password_hash=password_hash,
- display_name=display_name,
- role=role,
- status=status,
- )
- self.session.add(user)
- self.session.flush()
- return user
- def delete_user(self, user: AuthUser) -> None:
- self.session.delete(user)
- self.session.flush()
- def create_session(
- self,
- *,
- token_hash: str,
- user_id: int,
- expires_at: datetime,
- ip_address: str | None,
- user_agent: str | None,
- ) -> AuthSession:
- auth_session = AuthSession(
- token_hash=token_hash,
- user_id=user_id,
- expires_at=expires_at,
- ip_address=ip_address,
- user_agent=user_agent,
- )
- self.session.add(auth_session)
- self.session.flush()
- return auth_session
- def get_active_session(
- self,
- *,
- token_hash: str,
- now: datetime,
- ) -> AuthUser | None:
- stmt = (
- select(AuthUser)
- .select_from(AuthSession)
- .join(AuthUser, AuthUser.id == AuthSession.user_id)
- .where(
- AuthSession.token_hash == token_hash,
- AuthSession.expires_at > now,
- AuthUser.status == "active",
- )
- )
- return self.session.scalar(stmt)
- def delete_session_by_hash(self, token_hash: str) -> None:
- self.session.execute(
- delete(AuthSession).where(AuthSession.token_hash == token_hash)
- )
- def delete_user_sessions(self, user_id: int) -> None:
- self.session.execute(
- delete(AuthSession).where(AuthSession.user_id == user_id)
- )
- def delete_expired_sessions(self, now: datetime) -> None:
- self.session.execute(
- delete(AuthSession).where(AuthSession.expires_at <= now)
- )
|