205 lines
6.0 KiB
Python
205 lines
6.0 KiB
Python
"""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,
|
|
*,
|
|
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]:
|
|
statement = (
|
|
select(Order)
|
|
.where(
|
|
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)
|
|
),
|
|
)
|
|
.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]:
|
|
statement = (
|
|
select(Order)
|
|
.where(
|
|
Order.provider == "cdek",
|
|
Order.provider_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
|