"""PostgreSQL order repository.""" from collections.abc import Sequence from contextlib import AbstractAsyncContextManager from dataclasses import dataclass from datetime import datetime from typing import Any from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker from app.domain.cdek_polling import TERMINAL_ORDER_STATUSES from app.repositories.order.models import Order @dataclass(frozen=True) class OrderData: order_uuid: str payment_url: str price: int tariff_code: str provider: str account_email: str payload: dict[str, Any] class OrderRepository: def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None: self._session_factory = session_factory def session(self) -> AbstractAsyncContextManager[AsyncSession]: return self._session_factory.begin() async def create_order(self, session: AsyncSession, order_data: OrderData) -> Order: order = Order( order_uuid=order_data.order_uuid, payment_url=order_data.payment_url, price=order_data.price, tariff_code=order_data.tariff_code, provider=order_data.provider, account_email=order_data.account_email, payload=order_data.payload, ) session.add(order) await session.flush() await session.refresh(order) return order async def get_order_by_order_uuid( self, session: AsyncSession, order_uuid: str, ) -> Order | None: result = await session.execute( select(Order).where(Order.order_uuid == order_uuid) ) return result.scalar_one_or_none() async def mark_payment_status( self, session: AsyncSession, order_uuid: str, status: str, payment_id: int, ) -> Order | None: order = await self.get_order_by_order_uuid(session, order_uuid) if order is None: return None order.payment_status = status order.tbank_payment_id = payment_id await session.flush() return order async def mark_cdek_order_registered( self, session: AsyncSession, order_uuid: str, cdek_order_uuid: str, ) -> Order | None: order = await self.get_order_by_order_uuid(session, order_uuid) if order is None: return None order.cdek_order_uuid = cdek_order_uuid await session.flush() return order async def mark_cse_order_registered( self, session: AsyncSession, order_uuid: str, cse_order_number: str, ) -> Order | None: order = await self.get_order_by_order_uuid(session, order_uuid) if order is None: return None order.cse_order_number = cse_order_number await session.flush() return order async def list_orders_pending_waybill( self, session: AsyncSession, *, limit: int, ) -> Sequence[Order]: statement = ( select(Order) .where( Order.cdek_order_uuid.is_not(None), Order.cdek_waybill_url.is_(None), ( Order.cdek_order_status.is_(None) | Order.cdek_order_status.not_in(TERMINAL_ORDER_STATUSES) ), ) .order_by(Order.cdek_polled_at.asc().nulls_first()) .limit(limit) .with_for_update(skip_locked=True) ) result = await session.execute(statement) return result.scalars().all() async def record_order_poll( self, session: AsyncSession, *, order_uuid: str, order_status: str | None, waybill_uuid: str | None, polled_at: datetime, ) -> Order | None: order = await self.get_order_by_order_uuid(session, order_uuid) if order is None: return None order.cdek_order_status = order_status if waybill_uuid is not None and order.cdek_waybill_uuid is None: order.cdek_waybill_uuid = waybill_uuid order.cdek_polled_at = polled_at await session.flush() return order async def record_waybill_poll( self, session: AsyncSession, *, order_uuid: str, waybill_url: str | None, polled_at: datetime, ) -> Order | None: order = await self.get_order_by_order_uuid(session, order_uuid) if order is None: return None if waybill_url is not None and order.cdek_waybill_url is None: order.cdek_waybill_url = waybill_url order.cdek_polled_at = polled_at await session.flush() return order async def record_payment_email_sent( self, session: AsyncSession, *, order_uuid: str, sent_at: datetime, ) -> Order | None: order = await self.get_order_by_order_uuid(session, order_uuid) if order is None: return None if order.payment_email_sent_at is None: order.payment_email_sent_at = sent_at await session.flush() return order async def list_orders_pending_waybill_email( self, session: AsyncSession, *, limit: int, ) -> Sequence[Order]: statement = ( select(Order) .where( Order.cdek_waybill_url.is_not(None), Order.waybill_email_sent_at.is_(None), ) .order_by(Order.created_at.asc()) .limit(limit) .with_for_update(skip_locked=True) ) result = await session.execute(statement) return result.scalars().all() async def record_waybill_email_sent( self, session: AsyncSession, *, order_uuid: str, sent_at: datetime, ) -> Order | None: order = await self.get_order_by_order_uuid(session, order_uuid) if order is None: return None if order.waybill_email_sent_at is None: order.waybill_email_sent_at = sent_at await session.flush() return order