auth_repo.py 3.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104
  1. from __future__ import annotations
  2. from datetime import datetime
  3. from sqlalchemy import delete, select
  4. from sqlalchemy.orm import Session
  5. from supply_infra.db.models.auth_session import AuthSession
  6. from supply_infra.db.models.auth_user import AuthUser
  7. class AuthRepository:
  8. """Persistence operations for local users and server-side sessions."""
  9. def __init__(self, session: Session) -> None:
  10. self.session = session
  11. def get_user(self, user_id: int) -> AuthUser | None:
  12. return self.session.get(AuthUser, user_id)
  13. def get_user_by_username(self, username: str) -> AuthUser | None:
  14. stmt = select(AuthUser).where(AuthUser.username == username)
  15. return self.session.scalar(stmt)
  16. def list_users(self) -> list[AuthUser]:
  17. stmt = select(AuthUser).order_by(AuthUser.created_at.asc(), AuthUser.id.asc())
  18. return list(self.session.scalars(stmt).all())
  19. def create_user(
  20. self,
  21. *,
  22. username: str,
  23. password_hash: str,
  24. display_name: str,
  25. role: str,
  26. status: str = "active",
  27. ) -> AuthUser:
  28. user = AuthUser(
  29. username=username,
  30. password_hash=password_hash,
  31. display_name=display_name,
  32. role=role,
  33. status=status,
  34. )
  35. self.session.add(user)
  36. self.session.flush()
  37. return user
  38. def delete_user(self, user: AuthUser) -> None:
  39. self.session.delete(user)
  40. self.session.flush()
  41. def create_session(
  42. self,
  43. *,
  44. token_hash: str,
  45. user_id: int,
  46. expires_at: datetime,
  47. ip_address: str | None,
  48. user_agent: str | None,
  49. ) -> AuthSession:
  50. auth_session = AuthSession(
  51. token_hash=token_hash,
  52. user_id=user_id,
  53. expires_at=expires_at,
  54. ip_address=ip_address,
  55. user_agent=user_agent,
  56. )
  57. self.session.add(auth_session)
  58. self.session.flush()
  59. return auth_session
  60. def get_active_session(
  61. self,
  62. *,
  63. token_hash: str,
  64. now: datetime,
  65. ) -> AuthUser | None:
  66. stmt = (
  67. select(AuthUser)
  68. .select_from(AuthSession)
  69. .join(AuthUser, AuthUser.id == AuthSession.user_id)
  70. .where(
  71. AuthSession.token_hash == token_hash,
  72. AuthSession.expires_at > now,
  73. AuthUser.status == "active",
  74. )
  75. )
  76. return self.session.scalar(stmt)
  77. def delete_session_by_hash(self, token_hash: str) -> None:
  78. self.session.execute(
  79. delete(AuthSession).where(AuthSession.token_hash == token_hash)
  80. )
  81. def delete_user_sessions(self, user_id: int) -> None:
  82. self.session.execute(
  83. delete(AuthSession).where(AuthSession.user_id == user_id)
  84. )
  85. def delete_expired_sessions(self, now: datetime) -> None:
  86. self.session.execute(
  87. delete(AuthSession).where(AuthSession.expires_at <= now)
  88. )