"""Aggregator service orchestrating providers, cache and domain rules.""" import asyncio import hashlib import json from collections.abc import Iterable, Sequence from decimal import Decimal from typing import Protocol from app.adapters.delivery_providers.base import DeliveryProvider, ProviderRequestError from app.domain.price import ( DEFAULT_PROVIDER_PRICE_MULTIPLIER, DEFAULT_WEIGHT_ROUND_SCALE, NormalizedDeliveryRequest, filter_and_sort_prices, normalize_delivery_request, ) from app.schemas.order import OrderCreateRequest, OrderCreateResponse from app.schemas.request import DeliveryRequest from app.schemas.response import DeliveryPrice class AggregatorServiceError(RuntimeError): """Base exception for AggregatorService failures.""" class InvalidDeliveryRequestError(AggregatorServiceError): """Raised when provider rejects delivery request as invalid.""" class InvalidOrderCreateRequestError(AggregatorServiceError): """Raised when provider rejects order creation payload as invalid.""" class OrderCreationUnavailableError(AggregatorServiceError): """Raised when order creation cannot be completed.""" class PriceCacheProtocol(Protocol): async def get(self, key: str) -> object | None: ... async def set(self, key: str, value: object, ttl: int | None = None) -> None: ... class FilterAndSortPricesFn(Protocol): def __call__( self, prices: Iterable[object], *, price_multiplier: Decimal = DEFAULT_PROVIDER_PRICE_MULTIPLIER, parcel_type: object | None = None, ) -> list[object]: ... class OrderRegistrationAdapterProtocol(Protocol): async def register_order( self, request: OrderCreateRequest ) -> OrderCreateResponse: ... class AggregatorService: def __init__( self, providers: Sequence[DeliveryProvider], cache: PriceCacheProtocol | None = None, order_adapter: OrderRegistrationAdapterProtocol | None = None, *, weight_round_scale: int = DEFAULT_WEIGHT_ROUND_SCALE, provider_price_multiplier: Decimal = DEFAULT_PROVIDER_PRICE_MULTIPLIER, filter_and_sort_prices_fn: FilterAndSortPricesFn = filter_and_sort_prices, ) -> None: self._providers = tuple(providers) self._cache = cache self._order_adapter = order_adapter self._weight_round_scale = weight_round_scale self._provider_price_multiplier = provider_price_multiplier self._filter_and_sort_prices = filter_and_sort_prices_fn async def get_all_prices(self, request: DeliveryRequest) -> list[DeliveryPrice]: normalized_request = normalize_delivery_request( request, weight_round_scale=self._weight_round_scale ) provider_request = self._to_provider_request(normalized_request) provider_results = await asyncio.gather( *( self._get_provider_prices( provider=provider, request=provider_request, cache_key=self._build_cache_key( provider_name=provider.name, request=normalized_request, ), ) for provider in self._providers ), return_exceptions=True, ) successful_results = [ price for provider_prices in provider_results if isinstance(provider_prices, list) for price in provider_prices ] provider_errors = [ error for error in provider_results if isinstance(error, Exception) ] if ( not successful_results and provider_errors and any(isinstance(error, ProviderRequestError) for error in provider_errors) ): raise InvalidDeliveryRequestError( "Delivery request is invalid for configured providers." ) filtered_and_sorted = self._filter_and_sort_prices( successful_results, price_multiplier=self._provider_price_multiplier, parcel_type=request.parcel_type, ) return [self._coerce_delivery_price(price) for price in filtered_and_sorted] async def create_order(self, request: OrderCreateRequest) -> OrderCreateResponse: if self._order_adapter is None: raise OrderCreationUnavailableError( "Order registration adapter is not configured." ) try: created_order = await self._order_adapter.register_order(request) except ProviderRequestError as exc: raise InvalidOrderCreateRequestError( "Order request is invalid for the configured provider." ) from exc except Exception as exc: raise OrderCreationUnavailableError( "Order creation is temporarily unavailable." ) from exc return OrderCreateResponse.model_validate( created_order, from_attributes=True, ) async def _get_provider_prices( self, *, provider: DeliveryProvider, request: DeliveryRequest, cache_key: str, ) -> list[DeliveryPrice]: cached_prices = await self._get_cached_prices(cache_key) if cached_prices is not None: return cached_prices fresh_prices = await provider.get_prices(request) await self._set_cached_prices( cache_key, fresh_prices, ttl=getattr(provider, "cache_ttl_seconds", None), ) return fresh_prices async def _get_cached_prices(self, cache_key: str) -> list[DeliveryPrice] | None: if self._cache is None: return None try: payload = await self._cache.get(cache_key) except Exception: return None if payload is None: return None try: return self._coerce_delivery_prices(payload) except Exception: return None async def _set_cached_prices( self, cache_key: str, payload: list[DeliveryPrice], *, ttl: int | None, ) -> None: if self._cache is None: return try: await self._cache.set(cache_key, payload, ttl=ttl) except Exception: return @staticmethod def _to_provider_request(request: NormalizedDeliveryRequest) -> DeliveryRequest: return DeliveryRequest( entity=request.entity, from_city=request.from_city, to_city=request.to_city, country_code=request.country_code, weight_kg=request.weight_kg, length_cm=request.length_cm, width_cm=request.width_cm, height_cm=request.height_cm, ) @staticmethod def _build_cache_key(provider_name: str, request: NormalizedDeliveryRequest) -> str: cache_payload = { "provider": provider_name, "entity": request.entity, "from_city": request.from_city, "to_city": request.to_city, "country_code": request.country_code, "weight_kg": str(request.weight_kg), "length_cm": str(request.length_cm), "width_cm": str(request.width_cm), "height_cm": str(request.height_cm), } serialized_payload = json.dumps( cache_payload, ensure_ascii=True, sort_keys=True, separators=(",", ":"), ) cache_hash = hashlib.sha256(serialized_payload.encode("utf-8")).hexdigest() return f"delivery-price:{provider_name}:{cache_hash}" @staticmethod def _coerce_delivery_price(value: object) -> DeliveryPrice: return DeliveryPrice.model_validate(value, from_attributes=True) @classmethod def _coerce_delivery_prices(cls, value: object) -> list[DeliveryPrice]: if not isinstance(value, list): raise TypeError("Cached delivery prices must be a list.") return [cls._coerce_delivery_price(item) for item in value]