""" 验证码仓储 SQLAlchemy 实现 """ from __future__ import annotations from datetime import datetime, timezone from typing import Optional from sqlalchemy.orm import Session from packages.adapters.sqlalchemy_impl.models import VerificationCodeModel from packages.domain.verification_code import VerificationCode from packages.ports.verification_code_repository import VerificationCodeRepository class SQLAlchemyVerificationCodeRepository(VerificationCodeRepository): def __init__(self, session: Session): self.session = session def save(self, code: VerificationCode) -> None: model = self.session.get(VerificationCodeModel, code.id) if model is None: model = VerificationCodeModel(id=code.id) self.session.add(model) model.recipient = code.recipient model.code = code.code model.code_type = code.code_type model.expires_at = code.expires_at model.used_at = code.used_at model.attempts = code.attempts model.created_at = code.created_at self.session.commit() self.session.refresh(model) def find_latest(self, recipient: str, code_type: str) -> Optional[VerificationCode]: model = ( self.session.query(VerificationCodeModel) .filter( VerificationCodeModel.recipient == recipient.strip(), VerificationCodeModel.code_type == code_type, ) .order_by(VerificationCodeModel.created_at.desc()) .first() ) return self._to_entity(model) def find_by_id(self, code_id: str) -> Optional[VerificationCode]: return self._to_entity(self.session.get(VerificationCodeModel, code_id)) def count_today(self, recipient: str, code_type: str) -> int: now = datetime.now(timezone.utc) start_of_day = now.replace(hour=0, minute=0, second=0, microsecond=0) return ( self.session.query(VerificationCodeModel) .filter( VerificationCodeModel.recipient == recipient.strip(), VerificationCodeModel.code_type == code_type, VerificationCodeModel.created_at >= start_of_day, ) .count() ) @staticmethod def _to_entity(model: VerificationCodeModel | None) -> VerificationCode | None: if model is None: return None # SQLAlchemy 从数据库读出的 DateTime 是 naive(不带时区), # 领域模型期望 aware datetime(带 timezone.utc),直接用会报 # "can't compare offset-naive and offset-aware datetimes" def _ensure_aware(dt: datetime | None) -> datetime | None: if dt is None: return None if dt.tzinfo is None: return dt.replace(tzinfo=timezone.utc) return dt return VerificationCode( id=model.id, recipient=model.recipient, code=model.code, code_type=model.code_type, expires_at=_ensure_aware(model.expires_at), used_at=_ensure_aware(model.used_at), attempts=model.attempts, created_at=_ensure_aware(model.created_at), )