diff --git a/tariff_fetch/urdb/arcadia/rateutils.py b/tariff_fetch/urdb/arcadia/rateutils.py index f2371d7..7afe5c2 100644 --- a/tariff_fetch/urdb/arcadia/rateutils.py +++ b/tariff_fetch/urdb/arcadia/rateutils.py @@ -26,8 +26,10 @@ def tariff_iter_rates_for_dt( scenario: Scenario, library: Library, dt: datetime, + _seen: set[int] | None = None, ) -> Iterator[TariffRateExtended]: """Yield all rate entries that apply for a scenario at a given datetime.""" + _seen = _seen or set() rates = tariff.get("rates", []) for rate in rates: @@ -39,6 +41,10 @@ def tariff_iter_rates_for_dt( if not rate_is_applied_to_datetime(rate, dt): continue if rate["rate_bands"]: + rate_id = rate["tariff_rate_id"] + if rate_id in _seen: + continue + _seen.add(rate_id) yield rate elif rider_id := rate.get("rider_id"): try: @@ -49,7 +55,7 @@ def tariff_iter_rates_for_dt( f"Skipping inaccessible rider {rider_id} attached to rate {rate['tariff_rate_id']} ({rate['rate_name']})", ) continue - yield from tariff_iter_rates_for_dt(rider_tariff, scenario, library, dt) + yield from tariff_iter_rates_for_dt(rider_tariff, scenario, library, dt, _seen) # ================================ diff --git a/tests/test_arcadia_urdb_demand.py b/tests/test_arcadia_urdb_demand.py index d5394fd..a185bdf 100644 --- a/tests/test_arcadia_urdb_demand.py +++ b/tests/test_arcadia_urdb_demand.py @@ -153,8 +153,8 @@ def test_build_schedule_must_be_demand_based(): tariff: TariffExtended = { **KW_TARIFF, "rates": [ - {**RATE, "rate_bands": [{**BAND, "rate_amount": 15.0}]}, - {**DEMAND_RATE, "rate_bands": [{**BAND, "rate_amount": 10.0}]}, + {**RATE, "tariff_rate_id": 10000, "rate_bands": [{**BAND, "rate_amount": 15.0}]}, + {**DEMAND_RATE, "tariff_rate_id": 20000, "rate_bands": [{**BAND, "rate_amount": 10.0}]}, ], } scenario = make_stub_scenario(tariff) @@ -209,12 +209,14 @@ def test_build_demand_schedule_averages_sampled_datetimes(): "rates": [ { **RATE, + "tariff_rate_id": 10000, "charge_type": "DEMAND_BASED", "quantity_key": "base_kw", "rate_bands": [{**BAND, "rate_amount": 5.0}], }, { **RATE, + "tariff_rate_id": 20000, "charge_type": "DEMAND_BASED", "quantity_key": "seasonal_kw", "rate_bands": [{**BAND, "rate_amount": 10.0}], diff --git a/tests/test_arcadia_urdb_rateutils.py b/tests/test_arcadia_urdb_rateutils.py index 315b51c..35da961 100644 --- a/tests/test_arcadia_urdb_rateutils.py +++ b/tests/test_arcadia_urdb_rateutils.py @@ -10,7 +10,16 @@ from tariff_fetch.urdb.arcadia.exception import RateConversionError, TariffAccessDenied from tariff_fetch.urdb.arcadia.library import LibraryDebugStore, TariffLibrary from tariff_fetch.urdb.arcadia.scenario import Scenario -from tests.arcadia_urdb_fixtures import StubLibrary, make_band, make_consumption_rate, make_percentage_rate, make_rate +from tests.arcadia_urdb_fixtures import ( + BAND, + RATE, + TARIFF, + StubLibrary, + make_band, + make_consumption_rate, + make_percentage_rate, + make_rate, +) def test_rate_filter_bands_excludes_non_matching_choice_band(): @@ -272,3 +281,43 @@ def iter_pages(self, **kwargs): library.get_tariff(77) assert calls == [77] + + +def test_tariff_iter_rates_for_dt_deduplicates(): + duplicated_rate = { + **RATE, + "tariff_rate_id": 18603035, + "rate_name": "Universal Service Charge", + "rate_bands": [{**BAND, "tariff_rate_id": 18603035, "rate_amount": 0.32}], + } + rider_pointer_rate = { + **RATE, + "tariff_rate_id": 18603036, + "rate_name": "Universal Service Charge Rider", + "rate_bands": [], + "rider_id": 669, + } + parent_tariff = {**TARIFF, "rates": [duplicated_rate, rider_pointer_rate]} + rider_tariff = {**TARIFF, "tariff_id": 669, "master_tariff_id": 669, "rates": [duplicated_rate]} + rider_calls: list[int] = [] + + def get_tariff(tariff_id: int): + rider_calls.append(tariff_id) + return rider_tariff + + library = SimpleNamespace( + tariffs=SimpleNamespace(get_tariff=get_tariff), + record_issue=lambda key, message: None, + ) + + result = list( + ru.tariff_iter_rates_for_dt( + parent_tariff, # type: ignore[arg-type] + Scenario(1, 2025, False, {"SUPPLY"}), + library, # type: ignore[arg-type] + datetime(2025, 1, 1, 0, 30), + ) + ) + + assert result == [duplicated_rate] + assert rider_calls == [669]