"""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 and_, or_, 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, *, for_update: bool = False, ) -> Order | None: statement = select(Order).where(Order.order_uuid == order_uuid) if for_update: statement = statement.with_for_update() result = await session.execute(statement) 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_provider_order_registered( self, session: AsyncSession, order_uuid: str, provider_order_id: str, ) -> Order | None: order = await self.get_order_by_order_uuid(session, order_uuid) if order is None: return None order.provider_order_id = provider_order_id await session.flush() return order async def list_orders_pending_waybill( self, session: AsyncSession, *, limit: int, ) -> Sequence[Order]: cdek_pending = and_( Order.provider == "cdek", Order.provider_order_id.is_not(None), Order.provider_waybill_url.is_(None), ( Order.provider_order_status.is_(None) | Order.provider_order_status.not_in(TERMINAL_ORDER_STATUSES) ), ) cse_pending = and_( Order.provider == "cse", Order.provider_order_id.is_not(None), Order.provider_waybill_id.is_(None), ) statement = ( select(Order) .where(or_(cdek_pending, cse_pending)) .order_by(Order.provider_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.provider_order_status = order_status if waybill_uuid is not None and order.provider_waybill_id is None: order.provider_waybill_id = waybill_uuid order.provider_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.provider_waybill_url is None: order.provider_waybill_url = waybill_url order.provider_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]: cdek_ready = and_( Order.provider == "cdek", Order.provider_waybill_url.is_not(None), ) cse_ready = and_( Order.provider == "cse", Order.provider_waybill_id.is_not(None), ) statement = ( select(Order) .where( or_(cdek_ready, cse_ready), 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