300 lines
9.9 KiB
Python
300 lines
9.9 KiB
Python
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(
|
|
"cdek",
|
|
"100.40",
|
|
service_name="DOCUMENT EXPRESS",
|
|
).model_dump(mode="json"),
|
|
_make_price(
|
|
"cdek",
|
|
"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="cdek",
|
|
response=[
|
|
_make_price("cdek", "100.40", service_name="DOCUMENT EXPRESS"),
|
|
_make_price("cdek", "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"]
|