import asyncio from decimal import Decimal import pytest from app.adapters.delivery_providers.base import ProviderRequestError from app.domain.price import normalize_delivery_request from app.schemas.request import ( DeliveryCalculationRequest, DeliveryEntity, ParcelType, ) from app.schemas.response import DeliveryPrice from app.services.aggregator import AggregatorService, InvalidDeliveryRequestError _MISSING = object() class StubProvider: def __init__( self, name: str, *, response: list[DeliveryPrice] | None = None, error: Exception | None = None, cache_ttl_seconds: int = 900, ) -> None: self.name = name self._response = response self._error = error self.cache_ttl_seconds = cache_ttl_seconds self.calls: list[DeliveryCalculationRequest] = [] async def get_prices( self, request: DeliveryCalculationRequest ) -> list[DeliveryPrice]: self.calls.append(request) if self._error is not None: raise self._error if self._response is None: raise RuntimeError("Stub provider has no response configured.") return self._response class StubCache: def __init__( self, *, forced_get_value: object = _MISSING, ) -> None: self._storage: dict[str, object] = {} self._forced_get_value = forced_get_value self.get_calls: list[str] = [] self.set_calls: list[tuple[str, object, int | None]] = [] async def get(self, key: str) -> object | None: self.get_calls.append(key) if self._forced_get_value is not _MISSING: return self._forced_get_value return self._storage.get(key) async def set(self, key: str, value: object, ttl: int | None = None) -> None: self.set_calls.append((key, value, ttl)) self._storage[key] = value def _make_request(**overrides: object) -> DeliveryCalculationRequest: payload: dict[str, object] = { "entity": DeliveryEntity.INDIVIDUAL, "from_city": 1, "to_city": 2, "weight_kg": 1.234, "length_cm": 30.0, "width_cm": 20.0, "height_cm": 10.0, "parcel_type": None, } payload.update(overrides) return DeliveryCalculationRequest(**payload) def _make_price(provider: str, price: str, *, service_name: str = "standard") -> DeliveryPrice: return DeliveryPrice( provider=provider, service_name=service_name, price=Decimal(price), currency="RUB", delivery_days_min=2, delivery_days_max=4, ) def test_get_all_prices_full_success_returns_sorted_and_updates_cache() -> None: provider_a = StubProvider( name="a", response=[ _make_price("a", "300.49", service_name="economy"), _make_price("a", "200.49", service_name="express"), ], cache_ttl_seconds=111, ) provider_b = StubProvider( name="b", response=[_make_price("b", "100.40")], cache_ttl_seconds=222, ) cache = StubCache() service = AggregatorService( [provider_a, provider_b], cache=cache, provider_price_multiplier=Decimal("1.1"), ) result = asyncio.run(service.get_all_prices(_make_request())) assert [price.provider for price in result] == ["b", "a", "a"] assert [price.service_name for price in result] == ["standard", "express", "economy"] assert [price.price for price in result] == [Decimal("110"), Decimal("221"), Decimal("331")] assert len(provider_a.calls) == 1 assert len(provider_b.calls) == 1 assert len(cache.get_calls) == 2 assert [len(payload) for _, payload, _ in cache.set_calls] == [2, 1] assert [ttl for _, _, ttl in cache.set_calls] == [111, 222] def test_get_all_prices_forwards_city_identifiers_and_uses_updated_cache_key() -> None: provider = StubProvider(name="cdek", response=[_make_price("cdek", "150.00")]) cache = StubCache() service = AggregatorService([provider], cache=cache) request = _make_request() result = asyncio.run(service.get_all_prices(request)) assert [price.provider for price in result] == ["cdek"] assert len(provider.calls) == 1 assert provider.calls[0].from_city == 1 assert provider.calls[0].to_city == 2 expected_key = AggregatorService._build_cache_key( provider_name=provider.name, request=normalize_delivery_request(request), ) assert cache.get_calls == [expected_key] def test_get_all_prices_partial_failure_excludes_failed_provider() -> None: provider_ok = StubProvider( name="ok", response=[ _make_price("ok", "150.00", service_name="economy"), _make_price("ok", "120.00", service_name="express"), ], ) provider_failed = StubProvider(name="failed", error=RuntimeError("provider unavailable")) service = AggregatorService([provider_ok, provider_failed], cache=StubCache()) result = asyncio.run(service.get_all_prices(_make_request())) assert [price.provider for price in result] == ["ok", "ok"] assert [price.service_name for price in result] == ["express", "economy"] assert len(provider_ok.calls) == 1 assert len(provider_failed.calls) == 1 def test_get_all_prices_all_failures_returns_empty_list() -> None: provider_a = StubProvider(name="a", error=RuntimeError("a down")) provider_b = StubProvider(name="b", error=RuntimeError("b down")) service = AggregatorService([provider_a, provider_b], cache=StubCache()) result = asyncio.run(service.get_all_prices(_make_request())) assert result == [] assert len(provider_a.calls) == 1 assert len(provider_b.calls) == 1 def test_get_all_prices_raises_invalid_request_for_provider_request_errors() -> None: provider = StubProvider( name="cdek", error=ProviderRequestError("city not found"), ) service = AggregatorService([provider], cache=StubCache()) with pytest.raises(InvalidDeliveryRequestError): asyncio.run(service.get_all_prices(_make_request())) assert len(provider.calls) == 1 def test_get_all_prices_cache_hit_skips_provider_call() -> None: cached_payload = [ _make_price("cdek", "100.40", service_name="express").model_dump(mode="json"), _make_price("cdek", "200.40", service_name="economy").model_dump(mode="json"), ] cache = StubCache(forced_get_value=cached_payload) provider = StubProvider(name="cdek", response=[_make_price("cdek", "150.00")]) service = AggregatorService( [provider], cache=cache, provider_price_multiplier=Decimal("1.1"), ) result = asyncio.run(service.get_all_prices(_make_request())) assert [price.provider for price in result] == ["cdek", "cdek"] assert [price.service_name for price in result] == ["express", "economy"] assert [price.price for price in result] == [Decimal("110"), Decimal("220")] assert provider.calls == [] assert len(cache.get_calls) == 1 assert cache.set_calls == [] @pytest.mark.parametrize( ("cache", "expected_provider_calls"), [ (StubCache(), 1), ( StubCache( forced_get_value=[ _make_price( "other", "100.40", service_name="DOCUMENT EXPRESS", ).model_dump(mode="json"), _make_price( "other", "200.40", service_name="Economy parcel", ).model_dump(mode="json"), ] ), 0, ), ], ) def test_get_all_prices_applies_same_parcel_type_filter_for_fresh_and_cached_results( cache: StubCache, expected_provider_calls: int, ) -> None: provider = StubProvider( name="other", response=[ _make_price("other", "100.40", service_name="DOCUMENT EXPRESS"), _make_price("other", "200.40", service_name="Economy parcel"), ], ) service = AggregatorService( [provider], cache=cache, provider_price_multiplier=Decimal("1.1"), ) result = asyncio.run( service.get_all_prices(_make_request(parcel_type=ParcelType.DOC)) ) assert [price.service_name for price in result] == ["DOCUMENT EXPRESS"] assert [price.price for price in result] == [Decimal("110")] assert len(provider.calls) == expected_provider_calls def test_get_all_prices_delegates_filtering_and_sorting_to_domain_logic() -> None: provider_a = StubProvider( name="a", response=[ _make_price("a", "300.00", service_name="economy"), _make_price("a", "200.00", service_name="express"), ], ) provider_b = StubProvider(name="b", response=[_make_price("b", "100.00")]) delegated_inputs: list[list[DeliveryPrice]] = [] delegated_multipliers: list[Decimal] = [] delegated_parcel_types: list[object | None] = [] def fake_filter_and_sort(prices, *, price_multiplier, parcel_type): price_list = list(prices) delegated_inputs.append(price_list) delegated_multipliers.append(price_multiplier) delegated_parcel_types.append(parcel_type) return [price_list[0]] service = AggregatorService( [provider_a, provider_b], cache=StubCache(), provider_price_multiplier=Decimal("1.23"), filter_and_sort_prices_fn=fake_filter_and_sort, ) result = asyncio.run(service.get_all_prices(_make_request(parcel_type=ParcelType.PARCEL))) assert [price.provider for price in delegated_inputs[0]] == ["a", "a", "b"] assert [price.service_name for price in delegated_inputs[0]] == [ "economy", "express", "standard", ] assert delegated_multipliers == [Decimal("1.23")] assert delegated_parcel_types == [ParcelType.PARCEL] assert [price.provider for price in result] == ["a"]