From bbb4f644019eef05611ff17dc9f1f4ef3ba545af Mon Sep 17 00:00:00 2001 From: Max Ghenis Date: Sun, 27 Sep 2026 15:50:27 -0400 Subject: [PATCH] Read only periods the current branch can see when uprating or carrying over Holder.get_known_periods() lists the periods of every stored key, with the branch name stripped, but Holder.get_array() reads only the requested branch, its parent_branch ancestors and "default". Simulation._calculate took the latest known period from the unscoped list, so a period stored only under an unrelated branch read back as None: uprating raised TypeError and auto-carry-over cached NaN. - Holder._readable_branch_names() is the one definition of what a branch can read; get_array() and the new get_known_periods(branch_name) both use it. - _calculate uses get_known_periods(self.branch_name). - OnDiskStorage splits "_" keys on the last "_": branch names like "no_salt" raised ValueError, "y_2019" listed the wrong period, and delete(None, "pre_tcja") also wiped "pre_tcja_ctc". - dump_simulation saves the values the dumped branch reads instead of reading every period under "default" and saving None. Co-Authored-By: Claude Opus 5.5 --- .../fix-uprating-branch-visibility.fixed.md | 1 + .../data_storage/on_disk_storage.py | 22 +- policyengine_core/holders/holder.py | 82 ++-- policyengine_core/simulations/simulation.py | 7 +- policyengine_core/tools/simulation_dumper.py | 12 +- .../test_known_periods_branch_visibility.py | 422 ++++++++++++++++++ 6 files changed, 504 insertions(+), 42 deletions(-) create mode 100644 changelog.d/fix-uprating-branch-visibility.fixed.md create mode 100644 tests/core/test_known_periods_branch_visibility.py diff --git a/changelog.d/fix-uprating-branch-visibility.fixed.md b/changelog.d/fix-uprating-branch-visibility.fixed.md new file mode 100644 index 00000000..36eb6f76 --- /dev/null +++ b/changelog.d/fix-uprating-branch-visibility.fixed.md @@ -0,0 +1 @@ +Uprating and auto-carry-over now use only periods the current branch can read (its own, its parent branches' and `default`), so a period stored only under an unrelated branch no longer raises `TypeError` or caches `NaN`; on-disk storage parses branch names that contain `_`, and `dump_simulation` saves the values the dumped branch reads. diff --git a/policyengine_core/data_storage/on_disk_storage.py b/policyengine_core/data_storage/on_disk_storage.py index 3563c805..3de90e95 100644 --- a/policyengine_core/data_storage/on_disk_storage.py +++ b/policyengine_core/data_storage/on_disk_storage.py @@ -9,6 +9,17 @@ from policyengine_core.periods import Period +def _split_key(key: str) -> tuple: + """Split a ``f"{branch_name}_{period}"`` file key into its two parts. + + Branch names often contain ``_`` (policyengine-us uses ``no_salt`` and + ``mtr_for_adult_1``) but a period's string form never does, so the key + splits on its last ``_``. + """ + branch_name, period = key.rsplit("_", 1) + return branch_name, period + + class OnDiskStorage: """ Low-level class responsible for storing and retrieving calculated vectors on disk @@ -82,12 +93,13 @@ def delete(self, period: Period = None, branch_name: str = "default") -> None: if period is None: # Only wipe files belonging to the requested branch (previously # this wiped every branch regardless of ``branch_name`` — same - # class of bug as C2 in InMemoryStorage). - branch_prefix = f"{branch_name}_" + # class of bug as C2 in InMemoryStorage). Compare the parsed + # branch name, not a prefix: deleting ``pre_tcja`` must not + # also wipe ``pre_tcja_ctc``. self._files = { period_item: value for period_item, value in self._files.items() - if not period_item.startswith(branch_prefix) + if _split_key(period_item)[0] != branch_name } return @@ -103,12 +115,12 @@ def delete(self, period: Period = None, branch_name: str = "default") -> None: } def get_known_periods(self) -> list: - return list([periods.period(x.split("_")[1]) for x in self._files.keys()]) + return [period for _, period in self.get_known_branch_periods()] def get_known_branch_periods(self) -> list: return [ (branch_name, periods.period(period)) - for branch_name, period in map(lambda x: x.split("_"), self._files.keys()) + for branch_name, period in map(_split_key, self._files.keys()) ] def restore(self) -> None: diff --git a/policyengine_core/holders/holder.py b/policyengine_core/holders/holder.py index 52ae8129..96b4675b 100644 --- a/policyengine_core/holders/holder.py +++ b/policyengine_core/holders/holder.py @@ -106,45 +106,53 @@ def _get_array_from_storage( value = self._disk_storage.get(period, branch_name) return value + def _readable_branch_names(self, branch_name: str = "default") -> List[str]: + """ + Branches whose stored values ``get_array(period, branch_name)`` can + return, in lookup order. + + That is ``branch_name`` itself, then (unless it is ``default``) each + ``simulation.parent_branch`` ancestor, then ``default``. Nested + branches inherit values from their parent (e.g. a ``no_salt`` branch + cloned from an ``itemizing`` branch still sees ``tax_unit_itemizes`` + set on the ``itemizing`` branch). Previously the fallback returned + the first branch in dict-insertion order (bug C1) — silently swapping + values between unrelated sibling branches (reform vs baseline) and + producing wrong reform deltas. The post-C1 behavior only fell back to + ``default``, which broke country-package nested-branch patterns that + relied on the ancestor's input being visible. + """ + if branch_name == "default": + return ["default"] + branch_names = [branch_name] + simulation = getattr(self, "simulation", None) + parent = getattr(simulation, "parent_branch", None) if simulation else None + while parent is not None: + branch_names.append(parent.branch_name) + parent = getattr(parent, "parent_branch", None) + branch_names.append("default") + return list(dict.fromkeys(branch_names)) + def get_array(self, period: Period, branch_name: str = "default") -> ArrayLike: """ Get the value of the variable for the given period. - If the value is not known, return ``None``. + Values stored under ``branch_name``, its ``parent_branch`` ancestors + and ``default`` are visible, in that order. If the value is not + known on any of them, return ``None``. """ if self.variable.is_neutralized: return self.default_array() + # The branch's own value is the common case: look it up before + # walking the ancestors. value = self._get_array_from_storage(period, branch_name) if value is not None: return value - if value is None and branch_name != "default": - # Walk up ``simulation.parent_branch`` so nested branches inherit - # values from their parent (e.g. a ``no_salt`` branch cloned - # from an ``itemizing`` branch still sees ``tax_unit_itemizes`` - # set on the ``itemizing`` branch). Fall back to ``default`` - # only if no ancestor branch has a value. Previously the - # fallback returned the first branch in dict-insertion order - # (bug C1) — silently swapping values between unrelated - # sibling branches (reform vs baseline) and producing wrong - # reform deltas. The post-C1 behavior only fell back to - # ``default``, which broke country-package nested-branch - # patterns that relied on the ancestor's input being visible. - parent = ( - getattr(self.simulation, "parent_branch", None) - if self.simulation - else None - ) - while parent is not None: - ancestor_value = self._get_array_from_storage( - period, - parent.branch_name, - ) - if ancestor_value is not None: - return ancestor_value - parent = getattr(parent, "parent_branch", None) - default_value = self._get_array_from_storage(period, "default") - if default_value is not None: - return default_value + for readable_branch_name in self._readable_branch_names(branch_name)[1:]: + value = self._get_array_from_storage(period, readable_branch_name) + if value is not None: + return value + return None def get_memory_usage(self) -> dict: """ @@ -189,11 +197,23 @@ def get_memory_usage(self) -> dict: return usage - def get_known_periods(self) -> List[Period]: + def get_known_periods(self, branch_name: str = None) -> List[Period]: """ Get the list of periods the variable value is known for. - """ + With ``branch_name``, list only the periods ``get_array(period, + branch_name)`` can read: those stored under that branch, its + ``parent_branch`` ancestors or ``default``. Without it, list the + periods stored under every branch in this holder, some of which that + branch may not be able to read. + """ + if branch_name is not None: + readable_branch_names = set(self._readable_branch_names(branch_name)) + return [ + period + for stored_branch_name, period in self.get_known_branch_periods() + if stored_branch_name in readable_branch_names + ] return list(self._memory_storage.get_known_periods()) + list( (self._disk_storage.get_known_periods() if self._disk_storage else []) ) diff --git a/policyengine_core/simulations/simulation.py b/policyengine_core/simulations/simulation.py index 566407a5..bc3a1348 100644 --- a/policyengine_core/simulations/simulation.py +++ b/policyengine_core/simulations/simulation.py @@ -864,8 +864,11 @@ def _calculate(self, variable_name: str, period: Period = None) -> ArrayLike: # If no result, use the default value and cache it if array is None: - # Check if the variable has a previously defined value - known_periods = holder.get_known_periods() + # Check if the variable has a previously defined value. + # Only periods this branch can read count: a period stored + # only under an unrelated branch would read back as ``None`` + # and reach the arithmetic below. + known_periods = holder.get_known_periods(self.branch_name) earlier_known_periods = [ known_period for known_period in known_periods diff --git a/policyengine_core/tools/simulation_dumper.py b/policyengine_core/tools/simulation_dumper.py index c3db0c4f..b22da1a6 100644 --- a/policyengine_core/tools/simulation_dumper.py +++ b/policyengine_core/tools/simulation_dumper.py @@ -32,7 +32,7 @@ def dump_simulation(simulation, directory): # Dump variable values for holder in entity._holders.values(): - _dump_holder(holder, directory) + _dump_holder(holder, directory, simulation.branch_name) def restore_simulation(directory, tax_benefit_system, **kwargs): @@ -64,10 +64,14 @@ def restore_simulation(directory, tax_benefit_system, **kwargs): return simulation -def _dump_holder(holder, directory): +def _dump_holder(holder, directory, branch_name="default"): + # Dump the values the simulation's branch reads. A holder can also store + # periods under branches this one cannot see; reading those back under + # this branch gives ``None``, which would be saved as an object array + # that ``restore_simulation`` cannot load. disk_storage = holder.create_disk_storage(directory, preserve=True) - for period in holder.get_known_periods(): - value = holder.get_array(period) + for period in dict.fromkeys(holder.get_known_periods(branch_name)): + value = holder.get_array(period, branch_name) disk_storage.put(value, period) diff --git a/tests/core/test_known_periods_branch_visibility.py b/tests/core/test_known_periods_branch_visibility.py new file mode 100644 index 00000000..43ae823d --- /dev/null +++ b/tests/core/test_known_periods_branch_visibility.py @@ -0,0 +1,422 @@ +"""Known periods a branch cannot read. + +Holder storage keys embed the branch name (``"no_salt:2018"`` in memory, +``"no_salt_2018"`` on disk), and a holder can hold keys for branches other +than the simulation's own. ``Holder.get_known_periods()`` lists the periods of +every key, while ``Holder.get_array(period, branch_name)`` reads only that +branch, its ``parent_branch`` ancestors and ``default``. Three defects came +from that gap: + +* ``Simulation._calculate`` took the latest known earlier period from the + unscoped list. If that period was stored only under a branch the simulation + cannot read, ``get_array`` returned ``None``: uprating raised ``TypeError`` + (``None * factor``) and auto-carry-over cached ``NaN``. +* ``OnDiskStorage`` parsed ``f"{branch}_{period}"`` keys with + ``split("_")[1]``. A branch name containing ``_`` raised ``ValueError`` + (``no_salt`` -> period ``"salt"``) or listed the wrong period + (``y_2019`` -> ``2019``). ``delete(None, "pre_tcja")`` also wiped + ``pre_tcja_ctc``, whose key shares the prefix. +* ``dump_simulation`` read every listed period under ``default``, saving + ``None`` for a period stored only on the dumped branch, which + ``restore_simulation`` could not load. +""" + +import itertools +import math +import os +import tempfile +import warnings + +import numpy as np +import pytest + +from policyengine_core import periods +from policyengine_core.country_template import CountryTaxBenefitSystem +from policyengine_core.country_template.entities import Person +from policyengine_core.data_storage import OnDiskStorage +from policyengine_core.experimental import MemoryConfig +from policyengine_core.model_api import YEAR, Variable +from policyengine_core.parameters import ParameterNode +from policyengine_core.simulations import SimulationBuilder +from policyengine_core.tools import simulation_dumper + +INDEX_START = 2015 +# The index grows 10% a year, so ratios between years are powers of 1.1. +INDEX = { + f"{year}-01-01": 100 * 1.1 ** (year - INDEX_START) for year in range(2015, 2023) +} + + +def growth(from_year: int, to_year: int) -> float: + return 1.1 ** (to_year - from_year) + + +def build_system(auto_carry_over: bool = False) -> CountryTaxBenefitSystem: + system = CountryTaxBenefitSystem() + system.auto_carry_over_input_variables = auto_carry_over + system.parameters.add_child( + "test_uprating", + ParameterNode("test_uprating", data={"index": {"values": INDEX}}), + ) + + class uprated_income(Variable): + value_type = float + entity = Person + definition_period = YEAR + label = "Uprated yearly income" + uprating = "test_uprating.index" + + class carried_income(Variable): + value_type = float + entity = Person + definition_period = YEAR + label = "Yearly income carried over without uprating" + + system.add_variable(uprated_income) + system.add_variable(carried_income) + return system + + +@pytest.fixture(scope="module") +def system(): + return build_system() + + +@pytest.fixture(scope="module") +def carry_over_system(): + return build_system(auto_carry_over=True) + + +def new_simulation(system): + return SimulationBuilder().build_default_simulation(system, count=1) + + +def store(simulation, variable, year, value, branch_name): + """Store ``value`` for ``year`` under ``branch_name`` in the holder of + ``simulation``, the way ``Holder.set_input`` does when a caller names a + branch other than the simulation's own.""" + simulation.get_holder(variable).set_input( + periods.period(year), np.array([value]), branch_name + ) + + +def only(array) -> float: + (value,) = array + return float(value) + + +# ----- Simulation._calculate -------------------------------------------------- + + +def test_uprating_skips_period_stored_only_on_unrelated_branch(system): + simulation = new_simulation(system) + store(simulation, "uprated_income", 2016, 1_000.0, "default") + store(simulation, "uprated_income", 2018, 5_000.0, "other") + + result = simulation.calculate("uprated_income", 2020) + + # Uprated from 2016, the latest year the default branch can read. + assert only(result) == pytest.approx(1_000 * growth(2016, 2020)) + + +def test_uprating_with_no_readable_earlier_period_uses_default(system): + simulation = new_simulation(system) + store(simulation, "uprated_income", 2018, 5_000.0, "other") + + assert only(simulation.calculate("uprated_income", 2020)) == 0 + + +def test_uprating_in_branch_skips_sibling_branch_period(system): + simulation = new_simulation(system) + simulation.set_input("uprated_income", 2016, [1_000.0]) + branch = simulation.get_branch("reform") + store(branch, "uprated_income", 2018, 5_000.0, "baseline") + + result = branch.calculate("uprated_income", 2020) + + assert only(result) == pytest.approx(1_000 * growth(2016, 2020)) + + +def test_uprating_in_nested_branch_reads_ancestor_periods(system): + simulation = new_simulation(system) + simulation.set_input("uprated_income", 2016, [1_000.0]) + itemizing = simulation.get_branch("itemizing") + itemizing.set_input("uprated_income", 2018, [5_000.0]) + no_salt = itemizing.get_branch("no_salt") + + result = no_salt.calculate("uprated_income", 2020) + + assert only(result) == pytest.approx(5_000 * growth(2018, 2020)) + + +def test_carry_over_skips_period_stored_only_on_unrelated_branch( + carry_over_system, +): + simulation = new_simulation(carry_over_system) + store(simulation, "carried_income", 2016, 1_000.0, "default") + store(simulation, "carried_income", 2018, 5_000.0, "other") + + result = simulation.calculate("carried_income", 2020) + + assert not math.isnan(only(result)) + assert only(result) == 1_000 + + +def test_carry_over_with_no_readable_period_uses_default(carry_over_system): + simulation = new_simulation(carry_over_system) + store(simulation, "carried_income", 2018, 5_000.0, "other") + + result = simulation.calculate("carried_income", 2020) + + assert not math.isnan(only(result)) + assert only(result) == 0 + + +def test_carry_over_ignores_later_period_on_unrelated_branch(carry_over_system): + simulation = new_simulation(carry_over_system) + store(simulation, "carried_income", 2016, 1_000.0, "default") + store(simulation, "carried_income", 2022, 9_000.0, "other") + + assert only(simulation.calculate("carried_income", 2020)) == 1_000 + + +# Invariant: values stored only under branches a simulation cannot read do not +# change what it calculates. Checked exhaustively against a simulation that +# never had those values, for every combination of readable years, unreadable +# years and requested year below. +READABLE_YEAR_SETS = [(), (2016,), (2018,), (2016, 2018)] +UNREADABLE_YEAR_SETS = [ + years + for size in range(1, 4) + for years in itertools.combinations((2015, 2017, 2019, 2021), size) +] +REQUESTED_YEARS = (2017, 2018, 2019, 2020) + + +def _calculate_with(system, variable, readable, unreadable, requested): + simulation = new_simulation(system) + branch = simulation.get_branch("reform") + for year in readable: + store(branch, variable, year, 1_000.0 * (year - 2000), "reform") + for year in unreadable: + store(branch, variable, year, 7_777.0, "baseline") + return only(branch.calculate(variable, requested)) + + +@pytest.mark.parametrize("requested", REQUESTED_YEARS) +@pytest.mark.parametrize("unreadable", UNREADABLE_YEAR_SETS) +@pytest.mark.parametrize("readable", READABLE_YEAR_SETS) +@pytest.mark.parametrize( + "variable, auto_carry_over", + [("uprated_income", False), ("carried_income", True)], +) +def test_unreadable_branch_periods_do_not_change_results( + system, + carry_over_system, + variable, + auto_carry_over, + readable, + unreadable, + requested, +): + tax_benefit_system = carry_over_system if auto_carry_over else system + + with_unreadable = _calculate_with( + tax_benefit_system, variable, readable, unreadable, requested + ) + without_unreadable = _calculate_with( + tax_benefit_system, variable, readable, (), requested + ) + + assert math.isfinite(with_unreadable) + assert with_unreadable == without_unreadable + + +# ----- Holder.get_known_periods(branch_name) agrees with get_array ------------ + +# default -> a -> a_b is one lineage; c is a sibling of a, and a +# branch named "default" nested under a reads only "default" in get_array. +STORED_BRANCHES = ("default", "a", "a_b", "c") + + +def _lineage(system): + simulation = new_simulation(system) + a = simulation.get_branch("a") + return { + "default": simulation, + "a": a, + "a_b": a.get_branch("a_b"), + "c": simulation.get_branch("c"), + "default under a": a.get_branch("default"), + } + + +@pytest.mark.parametrize("on_disk", [False, True]) +@pytest.mark.parametrize( + "stored", + [ + branches + for size in range(1, len(STORED_BRANCHES) + 1) + for branches in itertools.combinations(STORED_BRANCHES, size) + ], +) +def test_known_periods_for_branch_are_exactly_the_readable_ones( + system, tmp_path, stored, on_disk +): + """For every reader, ``get_known_periods(branch_name)`` lists a period if + and only if ``get_array(period, branch_name)`` returns a value.""" + for reader_name, reader in _lineage(system).items(): + holder = reader.get_holder("uprated_income") + if on_disk: + directory = tmp_path / reader_name.replace(" ", "-") + directory.mkdir() + holder._disk_storage = holder.create_disk_storage( + str(directory), preserve=True + ) + storage = holder._disk_storage if on_disk else holder._memory_storage + # One distinct year per stored branch, so each listed period traces + # back to exactly one branch. + for offset, branch_name in enumerate(stored): + storage.put(np.array([1.0]), periods.period(2016 + offset), branch_name) + + stored_periods = {period for _, period in holder.get_known_branch_periods()} + readable = { + period + for period in stored_periods + if holder.get_array(period, reader.branch_name) is not None + } + + listed = holder.get_known_periods(reader.branch_name) + + assert set(listed) == readable, reader_name + assert len(listed) == len(set(listed)), reader_name + + +def test_get_known_periods_without_branch_lists_every_branch(system): + simulation = new_simulation(system) + store(simulation, "uprated_income", 2016, 1.0, "default") + store(simulation, "uprated_income", 2018, 1.0, "other") + + holder = simulation.get_holder("uprated_income") + + assert sorted(holder.get_known_periods()) == [ + periods.period(2016), + periods.period(2018), + ] + assert holder.get_known_periods("default") == [periods.period(2016)] + + +# ----- OnDiskStorage key parsing ---------------------------------------------- + +BRANCH_NAMES = ( + "default", + "no_salt", + "mtr_for_adult_1", + "pre_tcja_ctc", + "y_2019", + "trailing_", +) +PERIODS = ( + periods.period(2025), + periods.period("2025-03"), + periods.period("2025-03-05"), + periods.period("month:2025-01:3"), + periods.period("year:2025-03"), + periods.period("year:2024:2"), +) + + +@pytest.fixture +def disk_storage(tmp_path): + return OnDiskStorage(str(tmp_path), preserve_storage_dir=True) + + +@pytest.mark.parametrize("period", PERIODS, ids=str) +@pytest.mark.parametrize("branch_name", BRANCH_NAMES) +def test_on_disk_keys_round_trip(disk_storage, branch_name, period): + disk_storage.put(np.array([3.0]), period, branch_name) + + assert disk_storage.get_known_branch_periods() == [(branch_name, period)] + assert disk_storage.get_known_periods() == [period] + np.testing.assert_array_equal(disk_storage.get(period, branch_name), [3.0]) + + # The same holds after rebuilding the index from the files on disk. + disk_storage.restore() + assert disk_storage.get_known_branch_periods() == [(branch_name, period)] + + +@pytest.mark.parametrize("branch_name", BRANCH_NAMES) +def test_on_disk_eternal_keys_round_trip(tmp_path, branch_name): + storage = OnDiskStorage(str(tmp_path), is_eternal=True, preserve_storage_dir=True) + storage.put(np.array([3.0]), periods.period(2025), branch_name) + + assert storage.get_known_branch_periods() == [ + (branch_name, periods.period(periods.ETERNITY)) + ] + + +def test_on_disk_delete_branch_leaves_branches_sharing_its_prefix(disk_storage): + period = periods.period(2025) + for branch_name in ("pre_tcja", "pre_tcja_ctc", "pre"): + disk_storage.put(np.array([1.0]), period, branch_name) + + disk_storage.delete(None, "pre_tcja") + + assert sorted(disk_storage.get_known_branch_periods()) == [ + ("pre", period), + ("pre_tcja_ctc", period), + ] + + +@pytest.mark.parametrize("branch_name", ["no_salt", "y_2019"]) +def test_disk_backed_branch_with_underscore_uprates(branch_name): + """Only public API: a variable added after ``memory_config`` is set gets + disk storage, and the branch's input is uprated from its disk key.""" + tax_benefit_system = build_system() + simulation = new_simulation(tax_benefit_system) + with warnings.catch_warnings(): + warnings.simplefilter("ignore") + simulation.memory_config = MemoryConfig(max_memory_occupation=0) + + class disk_income(Variable): + value_type = float + entity = Person + definition_period = YEAR + label = "Uprated yearly income stored on disk" + uprating = "test_uprating.index" + + tax_benefit_system.add_variable(disk_income) + branch = simulation.get_branch(branch_name) + branch.set_input("disk_income", 2018, [5_000.0]) + disk_keys = list(branch.get_holder("disk_income")._disk_storage._files) + assert disk_keys == [f"{branch_name}_2018"] + + result = branch.calculate("disk_income", 2020) + + assert only(result) == pytest.approx(5_000 * growth(2018, 2020)) + + +# ----- dump_simulation -------------------------------------------------------- + + +def test_dump_branch_saves_the_values_the_branch_reads(system): + simulation = new_simulation(system) + simulation.set_input("uprated_income", 2016, [1_000.0]) + simulation.set_input("uprated_income", 2017, [2_000.0]) + branch = simulation.get_branch("reform") + branch.set_input("uprated_income", 2017, [3_000.0]) + branch.set_input("uprated_income", 2018, [5_000.0]) + + directory = os.path.join(tempfile.mkdtemp(), "dump") + simulation_dumper.dump_simulation(branch, directory) + restored = simulation_dumper.restore_simulation(directory, system) + + holder = restored.get_holder("uprated_income") + assert sorted(holder.get_known_periods()) == [ + periods.period(2016), + periods.period(2017), + periods.period(2018), + ] + assert only(holder.get_array(2016)) == 1_000 + assert only(holder.get_array(2017)) == 3_000 + assert only(holder.get_array(2018)) == 5_000