208 lines
6.7 KiB
Python
208 lines
6.7 KiB
Python
"""Background service that polls providers for waybill updates."""
|
|
|
|
import asyncio
|
|
from collections.abc import Callable, Mapping, Sequence
|
|
from contextlib import AbstractAsyncContextManager
|
|
from dataclasses import dataclass
|
|
from datetime import datetime, timezone
|
|
from typing import Protocol
|
|
|
|
import structlog
|
|
|
|
logger = structlog.get_logger(__name__)
|
|
|
|
|
|
class OrderInfoAdapterProtocol(Protocol):
|
|
async def get_order(self, provider_order_id: str) -> object: ...
|
|
|
|
|
|
class WaybillInfoAdapterProtocol(Protocol):
|
|
async def get_waybill(self, provider_waybill_id: str) -> object: ...
|
|
|
|
|
|
class OrderRecord(Protocol):
|
|
order_uuid: str
|
|
provider: str
|
|
provider_order_id: str | None
|
|
provider_waybill_id: str | None
|
|
|
|
|
|
class WaybillPollerRepositoryProtocol(Protocol):
|
|
def session(self) -> AbstractAsyncContextManager[object]: ...
|
|
|
|
async def list_orders_pending_waybill(
|
|
self, session: object, *, limit: int
|
|
) -> Sequence[OrderRecord]: ...
|
|
|
|
async def record_order_poll(
|
|
self,
|
|
session: object,
|
|
*,
|
|
order_uuid: str,
|
|
order_status: str | None,
|
|
waybill_uuid: str | None,
|
|
polled_at: datetime,
|
|
) -> object | None: ...
|
|
|
|
async def record_waybill_poll(
|
|
self,
|
|
session: object,
|
|
*,
|
|
order_uuid: str,
|
|
waybill_url: str | None,
|
|
polled_at: datetime,
|
|
) -> object | None: ...
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class PollBatchSummary:
|
|
processed: int
|
|
succeeded: int
|
|
failed: int
|
|
|
|
|
|
class WaybillPollerService:
|
|
def __init__(
|
|
self,
|
|
*,
|
|
order_repository: WaybillPollerRepositoryProtocol,
|
|
order_info_adapters: Mapping[str, OrderInfoAdapterProtocol],
|
|
waybill_info_adapters: Mapping[str, WaybillInfoAdapterProtocol],
|
|
batch_size: int,
|
|
datetime_now: Callable[[], datetime] = lambda: datetime.now(timezone.utc),
|
|
) -> None:
|
|
self._repository = order_repository
|
|
self._order_info_adapters = dict(order_info_adapters)
|
|
self._waybill_info_adapters = dict(waybill_info_adapters)
|
|
self._batch_size = batch_size
|
|
self._datetime_now = datetime_now
|
|
|
|
async def poll_once(self) -> PollBatchSummary:
|
|
async with self._repository.session() as session:
|
|
orders = await self._repository.list_orders_pending_waybill(
|
|
session, limit=self._batch_size
|
|
)
|
|
succeeded = 0
|
|
failed = 0
|
|
for order in orders:
|
|
try:
|
|
await self._handle_order(session, order)
|
|
succeeded += 1
|
|
except Exception:
|
|
failed += 1
|
|
logger.exception(
|
|
"waybill_poll_order_failed",
|
|
order_uuid=order.order_uuid,
|
|
provider=order.provider,
|
|
provider_order_id=order.provider_order_id,
|
|
provider_waybill_id=order.provider_waybill_id,
|
|
)
|
|
return PollBatchSummary(
|
|
processed=len(orders),
|
|
succeeded=succeeded,
|
|
failed=failed,
|
|
)
|
|
|
|
async def run_forever(
|
|
self,
|
|
*,
|
|
interval_seconds: float,
|
|
stop_event: asyncio.Event,
|
|
) -> None:
|
|
while not stop_event.is_set():
|
|
try:
|
|
summary = await self.poll_once()
|
|
logger.debug(
|
|
"waybill_poll_tick",
|
|
processed=summary.processed,
|
|
succeeded=summary.succeeded,
|
|
failed=summary.failed,
|
|
)
|
|
except Exception:
|
|
logger.exception("waybill_poll_tick_failed")
|
|
|
|
try:
|
|
await asyncio.wait_for(
|
|
stop_event.wait(), timeout=interval_seconds
|
|
)
|
|
except asyncio.TimeoutError:
|
|
continue
|
|
|
|
async def _handle_order(self, session: object, order: OrderRecord) -> None:
|
|
polled_at = self._datetime_now()
|
|
if order.provider_waybill_id is None:
|
|
provider_order_id = order.provider_order_id
|
|
if provider_order_id is None:
|
|
return
|
|
adapter = self._resolve_order_info_adapter(order.provider)
|
|
info = await adapter.get_order(provider_order_id)
|
|
status_code = _extract_status_code(info)
|
|
waybill_id = _extract_waybill_id(info)
|
|
await self._repository.record_order_poll(
|
|
session,
|
|
order_uuid=order.order_uuid,
|
|
order_status=status_code,
|
|
waybill_uuid=waybill_id,
|
|
polled_at=polled_at,
|
|
)
|
|
logger.info(
|
|
"waybill_poll_order_result",
|
|
order_uuid=order.order_uuid,
|
|
provider=order.provider,
|
|
provider_order_id=provider_order_id,
|
|
provider_order_status=status_code,
|
|
provider_waybill_id=waybill_id,
|
|
)
|
|
return
|
|
|
|
adapter = self._resolve_waybill_info_adapter(order.provider)
|
|
waybill = await adapter.get_waybill(order.provider_waybill_id)
|
|
waybill_url = _extract_waybill_url(waybill)
|
|
await self._repository.record_waybill_poll(
|
|
session,
|
|
order_uuid=order.order_uuid,
|
|
waybill_url=waybill_url,
|
|
polled_at=polled_at,
|
|
)
|
|
logger.info(
|
|
"waybill_poll_waybill_result",
|
|
order_uuid=order.order_uuid,
|
|
provider=order.provider,
|
|
provider_waybill_id=order.provider_waybill_id,
|
|
provider_waybill_url=waybill_url,
|
|
)
|
|
|
|
def _resolve_order_info_adapter(self, provider: str) -> OrderInfoAdapterProtocol:
|
|
adapter = self._order_info_adapters.get(provider)
|
|
if adapter is None:
|
|
raise RuntimeError(f"Order info adapter is not configured for {provider}.")
|
|
return adapter
|
|
|
|
def _resolve_waybill_info_adapter(
|
|
self, provider: str
|
|
) -> WaybillInfoAdapterProtocol:
|
|
adapter = self._waybill_info_adapters.get(provider)
|
|
if adapter is None:
|
|
raise RuntimeError(
|
|
f"Waybill info adapter is not configured for {provider}."
|
|
)
|
|
return adapter
|
|
|
|
|
|
def _extract_status_code(info: object) -> str | None:
|
|
value = getattr(info, "status_code", None)
|
|
return value if isinstance(value, str) and value else None
|
|
|
|
|
|
def _extract_waybill_id(info: object) -> str | None:
|
|
for attr in ("waybill_id", "waybill_uuid", "waybill_number"):
|
|
value = getattr(info, attr, None)
|
|
if isinstance(value, str) and value:
|
|
return value
|
|
return None
|
|
|
|
|
|
def _extract_waybill_url(info: object) -> str | None:
|
|
value = getattr(info, "url", None)
|
|
return value if isinstance(value, str) and value else None
|