From e2db11ee9853fc1959ca1772d13f0018b7b3d337 Mon Sep 17 00:00:00 2001 From: Max Ghenis Date: Fri, 2 Oct 2026 16:36:27 -0400 Subject: [PATCH 1/7] Restore inputs as inputs, evict the fast cache on holder writes, reject non-numeric uprating and defined_for Three bugs found by the review of #563, all present on master b78b0ba9: - restore_simulation put every value back with put_in_cache, which records no input, so apply_reform on a restored simulation dropped its inputs. dump_simulation now lists each variable's input periods in inputs.txt and restore_simulation records exactly those in _user_input_keys. - Holder.set_input, put_in_cache and delete_arrays left the simulation's fast cache untouched, so calculate kept returning the replaced value. Every holder write and delete now drops the entries it changes, in the holder's own simulation and only for branches that simulation reads. - An Enum, str or date variable with uprating, and a variable defined_for one, raised TypeError in the middle of a calculation. Both are now rejected when the variables are registered. Co-Authored-By: Claude Opus 5.5 --- changelog.d/holder-write-fast-cache.fixed.md | 1 + .../non-numeric-uprating-defined-for.fixed.md | 1 + .../restore-simulation-input-record.fixed.md | 1 + policyengine_core/holders/holder.py | 51 +++ policyengine_core/simulations/simulation.py | 16 +- .../taxbenefitsystems/tax_benefit_system.py | 46 ++- policyengine_core/tools/simulation_dumper.py | 76 ++++- policyengine_core/variables/variable.py | 69 +++- tests/core/test_holder_write_fast_cache.py | 169 ++++++++++ .../test_holder_write_fast_cache_property.py | 278 ++++++++++++++++ tests/core/test_restore_input_registry.py | 186 +++++++++++ .../test_restore_input_registry_property.py | 212 ++++++++++++ ...st_non_numeric_uprating_and_defined_for.py | 307 ++++++++++++++++++ tests/fixtures/uprated_inputs.py | 165 ++++++++++ 14 files changed, 1569 insertions(+), 9 deletions(-) create mode 100644 changelog.d/holder-write-fast-cache.fixed.md create mode 100644 changelog.d/non-numeric-uprating-defined-for.fixed.md create mode 100644 changelog.d/restore-simulation-input-record.fixed.md create mode 100644 tests/core/test_holder_write_fast_cache.py create mode 100644 tests/core/test_holder_write_fast_cache_property.py create mode 100644 tests/core/test_restore_input_registry.py create mode 100644 tests/core/test_restore_input_registry_property.py create mode 100644 tests/core/variables/test_non_numeric_uprating_and_defined_for.py create mode 100644 tests/fixtures/uprated_inputs.py diff --git a/changelog.d/holder-write-fast-cache.fixed.md b/changelog.d/holder-write-fast-cache.fixed.md new file mode 100644 index 000000000..716d56bb8 --- /dev/null +++ b/changelog.d/holder-write-fast-cache.fixed.md @@ -0,0 +1 @@ +Writing to or deleting from a holder directly (`Holder.set_input`, `Holder.put_in_cache`, `Holder.delete_arrays`) now drops the `Simulation.calculate` fast-cache entries it replaces, so `calculate` no longer returns the value the holder held before. diff --git a/changelog.d/non-numeric-uprating-defined-for.fixed.md b/changelog.d/non-numeric-uprating-defined-for.fixed.md new file mode 100644 index 000000000..cbfed0853 --- /dev/null +++ b/changelog.d/non-numeric-uprating-defined-for.fixed.md @@ -0,0 +1 @@ +An Enum, `str` or date variable with `uprating`, and a variable whose `defined_for` names one, are now rejected with a `ValueError` when registered (or when `uprating` is assigned) instead of raising `TypeError` in the middle of a calculation; `defined_for` must name a `bool`, `int` or `float` variable. diff --git a/changelog.d/restore-simulation-input-record.fixed.md b/changelog.d/restore-simulation-input-record.fixed.md new file mode 100644 index 000000000..ac96e8344 --- /dev/null +++ b/changelog.d/restore-simulation-input-record.fixed.md @@ -0,0 +1 @@ +`restore_simulation` now restores as inputs the values the dumped simulation stored as inputs (`dump_simulation` lists their periods in an `inputs.txt` beside each variable's arrays), so `apply_reform` on a restored simulation no longer drops them; a dump written before this has no such list, so every value in it is restored as an input. diff --git a/policyengine_core/holders/holder.py b/policyengine_core/holders/holder.py index 11a8ae50b..db45fb9ab 100644 --- a/policyengine_core/holders/holder.py +++ b/policyengine_core/holders/holder.py @@ -105,6 +105,56 @@ def delete_arrays( self._memory_storage.delete(period, branch_name) if self._disk_storage: self._disk_storage.delete(period, branch_name) + self._evict_fast_cache(period, branch_name, contained=True) + + def _evict_fast_cache( + self, period: Period, branch_name: str, contained: bool = False + ) -> None: + """Drop the simulation's ``_fast_cache`` entries a storage write or delete makes stale. + + ``Simulation.calculate`` answers a repeated request from + ``_fast_cache``, keyed by ``(variable name, requested period)``, + before it reads this holder. So every write into, or delete from, + this holder's storage drops the entries for the periods it changes; + otherwise ``calculate`` kept returning the value the storage no + longer held (for example after ``holder.set_input``). + + The fast cache belongs to this holder's simulation: a branch has its + own holders and its own fast cache, and keeps the values it started + with, so nothing outside this simulation is touched. Nor is anything + here when ``branch_name`` is a branch this simulation does not read. + + A write changes one storage key: ``period``, or, for an ETERNITY + variable, the one value every period reads. A delete (``contained``) + removes every period ``period`` contains, or every period when + ``period`` is ``None``. + """ + simulation = self.simulation + fast_cache = getattr(simulation, "_fast_cache", None) + if not fast_cache: + return + name = self.variable.name + drop_all = period is None or self.variable.definition_period == periods.ETERNITY + if not drop_all: + period = periods.period(period) + if not contained and (name, period) not in fast_cache: + return + visible_branch_names = getattr(simulation, "_get_visible_branch_names", None) + if visible_branch_names is not None and branch_name not in ( + visible_branch_names() + ): + return + if not drop_all and not contained: + del fast_cache[(name, period)] + return + stale_keys = [ + key + for key in fast_cache + if key[0] == name + and (drop_all or not isinstance(key[1], Period) or period.contains(key[1])) + ] + for key in stale_keys: + del fast_cache[key] def _get_array_from_storage( self, period: Period, branch_name: str = "default" @@ -370,6 +420,7 @@ def _set( self._disk_storage.put(value, period, branch_name) else: self._memory_storage.put(value, period, branch_name) + self._evict_fast_cache(period, branch_name) if user_input_contexts: if not hasattr(simulation, "_user_input_keys"): simulation._user_input_keys = set() diff --git a/policyengine_core/simulations/simulation.py b/policyengine_core/simulations/simulation.py index 2652d9723..ed32635e0 100644 --- a/policyengine_core/simulations/simulation.py +++ b/policyengine_core/simulations/simulation.py @@ -874,10 +874,20 @@ def _calculate(self, variable_name: str, period: Period = None) -> ArrayLike: self._check_period_consistency(period, variable) if variable.defined_for is not None: - mask = ( - self.calculate(variable.defined_for, period, map_to=variable.entity.key) - > 0 + defined_for_values = self.calculate( + variable.defined_for, period, map_to=variable.entity.key ) + if isinstance(defined_for_values, EnumArray) or ( + getattr(defined_for_values, "dtype", np.dtype(float)).kind not in "biuf" + ): + # Registration rejects these (see + # ``TaxBenefitSystem._check_defined_for``); this catches a + # defined_for set or a variable replaced afterwards, with + # the same message instead of a TypeError from ``> 0``. + variable.check_defined_for_variable( + self.tax_benefit_system.get_variable(variable.defined_for) + ) + mask = defined_for_values > 0 if np.all(~mask): array = holder.default_array() array = self._cast_formula_result(array, variable) diff --git a/policyengine_core/taxbenefitsystems/tax_benefit_system.py b/policyengine_core/taxbenefitsystems/tax_benefit_system.py index c381b11c4..e730a2e7c 100644 --- a/policyengine_core/taxbenefitsystems/tax_benefit_system.py +++ b/policyengine_core/taxbenefitsystems/tax_benefit_system.py @@ -51,6 +51,7 @@ from policyengine_core.periods import Instant, Period from policyengine_core.populations import GroupPopulation, Population from policyengine_core.variables import Variable +from policyengine_core.variables.variable import NUMERIC_VALUE_TYPES log = logging.getLogger(__name__) @@ -92,6 +93,8 @@ class TaxBenefitSystem: """Short list of basic inputs to get medium accuracy.""" modelled_policies: str = None """A YAML filepath containing metadata describing the modelled policies.""" + _defined_for_checks_deferred: bool = False + """Whether ``load_variable`` leaves ``defined_for`` checks to a later pass over every variable.""" def __init__(self, entities: Sequence[Entity] = None, reform=None) -> None: if entities is None: @@ -119,7 +122,14 @@ def __init__(self, entities: Sequence[Entity] = None, reform=None) -> None: self.variable_module_metadata = {} if self.variables_dir is not None: - self.add_variables_from_directory(self.variables_dir) + # A variable's defined_for variable may be in a file loaded after + # it, so check every variable once the whole directory is loaded. + self._defined_for_checks_deferred = True + try: + self.add_variables_from_directory(self.variables_dir) + finally: + self._defined_for_checks_deferred = False + self._check_defined_for_variables() self.data_modified = False if self.parameters_dir is not None: @@ -215,10 +225,44 @@ def load_variable( ) variable = variable_class(baseline_variable=baseline_variable) + if not self._defined_for_checks_deferred: + self._check_defined_for(variable) self.variables[variable.name] = variable return variable + def _check_defined_for(self, variable: Variable) -> None: + """Check the ``defined_for`` links that ``variable`` takes part in. + + Both directions: the variable ``variable`` is defined for, if it is + registered, and, when ``variable`` cannot be compared with zero, every + registered variable defined for it. See + :meth:`Variable.check_defined_for_variable`. + """ + defined_for = variable.defined_for + if defined_for is not None: + defined_for_variable = ( + variable + if defined_for == variable.name + else self.variables.get(defined_for) + ) + if defined_for_variable is not None: + variable.check_defined_for_variable(defined_for_variable) + if variable.value_type in NUMERIC_VALUE_TYPES: + return + for other in self.variables.values(): + if other.defined_for == variable.name and other.name != variable.name: + other.check_defined_for_variable(variable) + + def _check_defined_for_variables(self) -> None: + """Check every registered variable's ``defined_for`` variable.""" + for variable in self.variables.values(): + if variable.defined_for is None: + continue + defined_for_variable = self.variables.get(variable.defined_for) + if defined_for_variable is not None: + variable.check_defined_for_variable(defined_for_variable) + def add_variable(self, variable: Type[Variable]) -> Variable: """Adds an OpenFisca variable to the tax and benefit system. diff --git a/policyengine_core/tools/simulation_dumper.py b/policyengine_core/tools/simulation_dumper.py index c3db0c4fa..08edc1b3b 100644 --- a/policyengine_core/tools/simulation_dumper.py +++ b/policyengine_core/tools/simulation_dumper.py @@ -6,9 +6,17 @@ import numpy as np from policyengine_core.data_storage import OnDiskStorage +from policyengine_core import periods from policyengine_core.periods import ETERNITY from policyengine_core.simulations import Simulation +# Next to each variable's arrays: the periods, one per line, whose dumped +# value was an input (stored through ``set_input``). ``restore_simulation`` +# registers exactly these as inputs, so ``apply_reform``, which keeps inputs +# and drops calculated values, keeps the same values in the restored +# simulation as in the dumped one. +INPUT_PERIODS_FILE = "inputs.txt" + def dump_simulation(simulation, directory): """ @@ -26,18 +34,26 @@ def dump_simulation(simulation, directory): entities_dump_dir = os.path.join(directory, "__entities__") os.mkdir(entities_dump_dir) + input_keys = _input_storage_keys(simulation) for entity in simulation.populations.values(): # Dump entity structure _dump_entity(entity, entities_dump_dir) # Dump variable values for holder in entity._holders.values(): - _dump_holder(holder, directory) + _dump_holder(holder, directory, input_keys) def restore_simulation(directory, tax_benefit_system, **kwargs): """ Restore simulation from directory + + Values the dumped simulation stored as inputs are restored as inputs + (recorded in ``_user_input_keys``, as ``set_input`` records them), and + every other value as a calculated one, so ``apply_reform`` keeps and drops + the same values it would have in the dumped simulation. A dump written + before inputs were recorded (no ``inputs.txt``) does not say which values + were inputs, so every value in it is restored as an input. """ simulation = Simulation( tax_benefit_system, tax_benefit_system.instantiate_entities() @@ -64,11 +80,38 @@ def restore_simulation(directory, tax_benefit_system, **kwargs): return simulation -def _dump_holder(holder, directory): +def _dump_holder(holder, directory, input_keys=frozenset()): disk_storage = holder.create_disk_storage(directory, preserve=True) + input_periods = [] for period in holder.get_known_periods(): value = holder.get_array(period) disk_storage.put(value, period) + # The input record of exactly the value dumped: ``get_array`` above + # reads the default branch. + if (holder.variable.name, "default", str(period)) in input_keys: + input_periods.append(str(period)) + path = os.path.join(disk_storage.storage_dir, INPUT_PERIODS_FILE) + with open(path, "w") as file: + file.write("".join(f"{period}\n" for period in dict.fromkeys(input_periods))) + + +def _input_storage_keys(simulation): + """The storage keys ``_user_input_keys`` records as inputs. + + Each record entry becomes ``(variable, branch, period)`` with the period + as storage writes it, as ``Simulation._invalidate_all_caches`` reads the + record back through the storage: an ETERNITY variable's one value is an + input whatever period its entry names. + """ + input_keys = set() + for name, branch_name, period in getattr(simulation, "_user_input_keys", ()): + variable = simulation.tax_benefit_system.get_variable(name) + if variable is not None and variable.definition_period == ETERNITY: + period = ETERNITY + elif period is None: + continue + input_keys.add((name, branch_name, str(periods.period(period)))) + return input_keys def _dump_entity(population, directory): @@ -135,6 +178,33 @@ def _restore_holder(simulation, variable, directory): holder = simulation.get_holder(variable) + input_periods_path = os.path.join(storage_dir, INPUT_PERIODS_FILE) + if os.path.exists(input_periods_path): + with open(input_periods_path) as file: + input_periods = set(file.read().split()) + else: + # Dumped before inputs were recorded: nothing says which values were + # calculated, so keep every value as an input. + input_periods = None + for period in disk_storage.get_known_periods(): value = disk_storage.get(period) - holder.put_in_cache(value, period) + if input_periods is None or str(period) in input_periods: + _restore_input(simulation, holder, period, value) + else: + holder.put_in_cache(value, period) + + +def _restore_input(simulation, holder, period, value): + """Store ``value`` as an input, recorded as ``set_input`` records one. + + ``Holder.set_input`` would also run the variable's ``set_input`` helper, + but a dump holds values already split into the variable's own periods. + """ + if not hasattr(simulation, "_user_input_contexts"): + simulation._user_input_contexts = [] + simulation._user_input_contexts.append(simulation.branch_name) + try: + holder._set(period, value, simulation.branch_name) + finally: + simulation._user_input_contexts.pop() diff --git a/policyengine_core/variables/variable.py b/policyengine_core/variables/variable.py index 629a1f4f8..ab0d2bb32 100644 --- a/policyengine_core/variables/variable.py +++ b/policyengine_core/variables/variable.py @@ -42,6 +42,21 @@ class VariableCategory: DEMOGRAPHIC = "demographic" +NUMERIC_VALUE_TYPES = (bool, int, float) +"""Value types whose arrays support arithmetic and comparison with numbers. + +Only these can be uprated (multiplied by an index ratio) or used as a +``defined_for`` mask (compared with zero). Enum arrays allow only ``==`` and +``!=``, and ``str`` and date arrays cannot be multiplied by a float or +compared with a number, so either use raised a ``TypeError`` in the middle +of a calculation. +""" + + +def _value_type_name(value_type: type) -> str: + return getattr(value_type, "__name__", str(value_type)) + + class Variable: """ A `variable `_ of the legislation. @@ -105,7 +120,7 @@ class Variable: """Categorical attribute describing whether the variable is a stock or a flow.""" defined_for: str = None - """The name of another variable, nonzero values of which are used to define the set of entities for which this variable is defined.""" + """The name of another variable, nonzero values of which are used to define the set of entities for which this variable is defined. That variable must be a ``bool``, ``int`` or ``float`` variable; the tax-benefit system rejects an Enum, ``str`` or date one when both variables are registered. An Enum member is also accepted, and names the variable called after its value (``StateCode.CA`` names ``CA``).""" metadata: dict = None """Free dictionary field used to store any metadata.""" @@ -123,7 +138,7 @@ class Variable: """List of variables that are subtracted from the variable. Alternatively, can be a parameter name.""" uprating: str = None - """Name of a parameter used to uprate the variable. When the variable has no value for a requested period, its value in the latest known earlier period (an input, or a value already calculated or defaulted and cached) is multiplied by the ratio of this parameter at the two period starts, or carried over unchanged if the parameter is zero at the earlier start. Where the parameter has no value it is held flat: before its first value it takes that first value, and after an explicit null it keeps the last value before the null. The variable therefore carries over unchanged across any span the parameter does not cover.""" + """Name of a parameter used to uprate the variable. Only ``bool``, ``int`` and ``float`` variables can be uprated; ``uprating`` on an Enum, ``str`` or date variable raises a ``ValueError`` when the variable is defined or the attribute is assigned. When the variable has no value for a requested period, its value in the latest known earlier period (an input, or a value already calculated or defaulted and cached) is multiplied by the ratio of this parameter at the two period starts, or carried over unchanged if the parameter is zero at the earlier start. Where the parameter has no value it is held flat: before its first value it takes that first value, and after an explicit null it keeps the last value before the null. The variable therefore carries over unchanged across any span the parameter does not cover.""" hidden_input: bool = False """Whether the variable is hidden from the input screen entirely on PolicyEngine.""" @@ -327,6 +342,7 @@ def __init__(self, baseline_variable=None): self.formulas = self.set_formulas(formulas_attr) self.check_computation_modes() + self.check_uprating_value_type() check_formula_determinism(self) if unexpected_attrs: @@ -357,6 +373,7 @@ def uprating(self, value): self._explicit_attribute_names = old_explicit | {"uprating"} try: self.check_computation_modes() + self.check_uprating_value_type() except ValueError: self._uprating = old_value self._explicit_attribute_names = old_explicit @@ -403,6 +420,54 @@ def check_computation_modes(self): "input or constant variables should use none." ) + def check_uprating_value_type(self): + """Reject ``uprating`` on a variable whose values cannot be multiplied. + + Uprating multiplies the latest earlier value by a ratio of index + values. Enum arrays allow only ``==`` and ``!=``, and ``str`` and date + arrays cannot be multiplied by a float, so ``uprating`` on such a + variable raised a ``TypeError`` the first time a later period was + uprated. The check covers an ``uprating`` inherited from a baseline + variable too, since the inherited one is what ``calculate`` uses. + """ + if self.uprating is None: + return + value_type = getattr(self, "value_type", None) + if value_type is None or value_type in NUMERIC_VALUE_TYPES: + return + raise ValueError( + f'Variable "{self.name}" has uprating "{self.uprating}", but its ' + f"value_type is {_value_type_name(value_type)}. Uprating " + "multiplies the latest earlier value by an index ratio, so only " + "bool, int and float variables can be uprated. Remove uprating: " + "with auto_carry_over_input_variables, an input without it " + "carries over to later periods unchanged." + ) + + def check_defined_for_variable(self, defined_for_variable: "Variable") -> None: + """Reject a ``defined_for`` variable that cannot be compared with zero. + + ``calculate`` keeps this variable's formula result where the + ``defined_for`` variable is greater than zero, and uses the default + value elsewhere. That needs a ``bool``, ``int`` or ``float`` variable: + an Enum array allows only ``==`` and ``!=``, and ``str`` and date + arrays cannot be compared with a number, so the comparison raised a + ``TypeError`` in the middle of a calculation. For an Enum condition, + define a ``bool`` variable that compares it with a member and name + that variable instead. + """ + value_type = getattr(defined_for_variable, "value_type", None) + if value_type is None or value_type in NUMERIC_VALUE_TYPES: + return + raise ValueError( + f'Variable "{self.name}" is defined_for "{defined_for_variable.name}", ' + f"whose value_type is {_value_type_name(value_type)}. defined_for " + "keeps values where that variable is greater than zero, so it must " + "name a bool, int or float variable. For an Enum condition, define " + "a bool variable that compares it with a member (for example " + "`state == State.present`) and name that variable instead." + ) + def set( self, attributes, diff --git a/tests/core/test_holder_write_fast_cache.py b/tests/core/test_holder_write_fast_cache.py new file mode 100644 index 000000000..f22aac07c --- /dev/null +++ b/tests/core/test_holder_write_fast_cache.py @@ -0,0 +1,169 @@ +"""Writing to or deleting from a holder drops the fast-cache entries it changes. + +``Simulation.calculate`` answers a repeated request from ``_fast_cache`` +before it reads the holder. ``Simulation.set_input`` and +``Simulation.delete_arrays`` dropped the matching entry, but the holder's own +methods did not: after ``holder.set_input(period, [300, 400])`` the holder +held ``[300, 400]`` while ``calculate`` still returned the value it had +calculated before. + +Every holder write (``set_input``, ``put_in_cache``, a ``set_input`` helper's +stores) and delete now drops the entries it makes stale, in the holder's own +simulation only. The property test is +``test_holder_write_fast_cache_property.py``. +""" + +from __future__ import annotations + +import numpy as np +import pytest + +from policyengine_core import periods +from tests.fixtures.uprated_inputs import build_simulation, build_system + + +@pytest.fixture(scope="module") +def system(): + return build_system() + + +def _fast_cached(simulation, variable): + return sorted( + str(period) for name, period in simulation._fast_cache if name == variable + ) + + +def test_holder_set_input_replaces_a_calculated_value(system): + # The reported case: master kept returning the uprated [1038, 79]. + simulation = build_simulation(system, [("uprated_count", "2012", [1001, 77])]) + assert simulation.calculate("uprated_count", "2013").tolist() == [1038, 79] + holder = simulation.get_holder("uprated_count") + + holder.set_input(periods.period("2013"), [300, 400]) + + assert holder.get_array(periods.period("2013")).tolist() == [300, 400] + assert simulation.calculate("uprated_count", "2013").tolist() == [300, 400] + + +def test_holder_set_input_split_into_months_replaces_each_month(system): + simulation = build_simulation(system) + simulation.calculate("monthly_doubled", "2013-02") + assert _fast_cached(simulation, "monthly_doubled") == ["2013-02"] + holder = simulation.get_holder("monthly_doubled") + + # The helper stores twelve months; February is one of them. + holder.delete_arrays(periods.period("2013")) + holder.set_input(periods.period("2013"), [1200, 2400]) + + assert simulation.calculate("monthly_doubled", "2013-02").tolist() == [100, 200] + + +def test_holder_set_input_of_an_eternity_variable_replaces_every_period(system): + simulation = build_simulation(system) + # One stored value answers every period; each request is cached apart. + assert simulation.calculate("eternal_code_plus_one", "2012").tolist() == [1, 1] + assert _fast_cached(simulation, "eternal_code_plus_one") == ["2012"] + + simulation.get_holder("eternal_code_plus_one").set_input( + periods.period("2015"), [8, 9] + ) + + assert simulation.calculate("eternal_code_plus_one", "2012").tolist() == [8, 9] + assert simulation.calculate("eternal_code_plus_one", "2015").tolist() == [8, 9] + + +def test_holder_put_in_cache_replaces_a_calculated_value(system): + simulation = build_simulation(system, [("uprated_amount", "2012", [100.0, 50.0])]) + simulation.calculate("doubled_amount", "2012") + + simulation.get_holder("doubled_amount").put_in_cache( + np.array([7.0, 8.0], dtype=np.float32), periods.period("2012") + ) + + assert simulation.calculate("doubled_amount", "2012").tolist() == [7.0, 8.0] + + +def test_holder_delete_arrays_drops_the_periods_it_deletes(system): + simulation = build_simulation(system, [("monthly_amount", "2013", [1200, 2400])]) + simulation.calculate("monthly_doubled", "2013-02") + simulation.calculate("monthly_doubled", "2014-06") + simulation.calculate("doubled_amount", "2013") + + # Deleting a year deletes the months in it, and nothing else. + simulation.get_holder("monthly_doubled").delete_arrays(periods.period("2013")) + + assert _fast_cached(simulation, "monthly_doubled") == ["2014-06"] + assert _fast_cached(simulation, "doubled_amount") == ["2013"] + simulation.get_holder("monthly_amount").delete_arrays(periods.period("2013")) + simulation.set_input("monthly_amount", "2013", [2400, 4800]) + assert simulation.calculate("monthly_doubled", "2013-02").tolist() == [400, 800] + + +def test_holder_delete_arrays_without_a_period_drops_every_period(system): + simulation = build_simulation(system, [("monthly_amount", "2013", [1200, 2400])]) + simulation.calculate("monthly_doubled", "2013-02") + simulation.calculate("monthly_doubled", "2014-06") + + simulation.get_holder("monthly_doubled").delete_arrays() + + assert _fast_cached(simulation, "monthly_doubled") == [] + + +def test_write_under_a_branch_the_simulation_does_not_read_keeps_the_entry(system): + simulation = build_simulation(system, [("uprated_count", "2012", [1001, 77])]) + calculated = simulation.calculate("uprated_count", "2013") + + simulation.get_holder("uprated_count").set_input( + periods.period("2013"), [300, 400], "sibling" + ) + + # Still answered from the fast cache: the same array object. + assert simulation.calculate("uprated_count", "2013") is calculated + assert calculated.tolist() == [1038, 79] + + +def test_write_under_default_replaces_what_a_branch_read_from_it(system): + simulation = build_simulation(system) + branch = simulation.get_branch("reform") + # No value under the branch's own name, so it reads (and caches) the + # default; a later store under ``default`` in its holder replaces that. + assert branch.calculate("doubled_amount", "2013").tolist() == [0, 0] + holder = branch.get_holder("doubled_amount") + holder.delete_arrays(periods.period("2013"), "reform") + + holder.put_in_cache( + np.array([7.0, 8.0], dtype=np.float32), periods.period("2013"), "default" + ) + + assert branch.calculate("doubled_amount", "2013").tolist() == [7.0, 8.0] + + +def test_a_branch_write_leaves_the_parent_and_the_parent_write_leaves_the_branch( + system, +): + simulation = build_simulation(system, [("uprated_count", "2012", [1001, 77])]) + parent_value = simulation.calculate("uprated_count", "2013") + branch = simulation.get_branch("reform") + branch_value = branch.calculate("uprated_count", "2014") + + # A branch has its own holders and its own fast cache. + branch.get_holder("uprated_count").set_input( + periods.period("2013"), [300, 400], "reform" + ) + assert simulation.calculate("uprated_count", "2013") is parent_value + assert branch.calculate("uprated_count", "2013").tolist() == [300, 400] + + simulation.get_holder("uprated_count").set_input(periods.period("2014"), [5, 6]) + assert simulation.calculate("uprated_count", "2014").tolist() == [5, 6] + assert branch.calculate("uprated_count", "2014") is branch_value + + +def test_storing_a_calculated_value_keeps_it_in_the_fast_cache(system): + # ``calculate`` stores its result in the holder and then caches it: the + # store must not drop the entry the same request is about to rely on. + simulation = build_simulation(system, [("uprated_amount", "2012", [100.0, 50.0])]) + + first = simulation.calculate("doubled_amount", "2012") + + assert _fast_cached(simulation, "doubled_amount") == ["2012"] + assert simulation.calculate("doubled_amount", "2012") is first diff --git a/tests/core/test_holder_write_fast_cache_property.py b/tests/core/test_holder_write_fast_cache_property.py new file mode 100644 index 000000000..1dc8f8041 --- /dev/null +++ b/tests/core/test_holder_write_fast_cache_property.py @@ -0,0 +1,278 @@ +"""Properties of the fast cache under holder writes and deletes. + +Random sequences of operations run on a simulation and the branches made +from it: ``calculate``; ``set_input``, ``put_in_cache`` and ``delete_arrays`` +on a holder, under the simulation's own branch name, ``default``, or a +branch it does not read; ``Simulation.set_input`` and +``Simulation.delete_arrays``; and creating a branch. + +1. **The fast cache never changes what ``calculate`` returns.** A second + tree runs the same operations but empties every fast cache before each + ``calculate``. Every result agrees (bytes, or error), and so does every + stored value at the end. +2. **A holder write or delete drops only what it changes.** The entries it + removes all belong to its own simulation and variable, were stored under + a branch that simulation reads, and are for the period written (any + period, for an ETERNITY variable) or for a period the deleted one + contains. Every other simulation's fast cache is untouched. + +``test_holder_write_fast_cache.py`` pins the same behaviour with examples. +""" + +from __future__ import annotations + +import numpy as np +import pytest + +# The country-package smoke job installs no dev dependencies. +hypothesis = pytest.importorskip("hypothesis") +st = hypothesis.strategies + +from policyengine_core import periods +from policyengine_core.enums import EnumArray +from tests.fixtures.uprated_inputs import ( + INPUT_VALUE_TYPES, + PERSON_COUNT, + build_simulation, + build_system, +) + +YEARS = ["2012", "2013", "2014"] +MONTHS = ["2013-01", "2013-02", "2014-06"] +VARIABLES = { + # name: (periods to calculate, periods to write or delete) + "uprated_count": (YEARS, YEARS), + "uprated_amount": (YEARS, YEARS), + "doubled_amount": (YEARS, YEARS), + "masked_amount": (YEARS, YEARS), + "eligible": (YEARS, YEARS), + "monthly_amount": (MONTHS, MONTHS + ["2013"]), + "monthly_doubled": (MONTHS, MONTHS + ["2013"]), + "eternal_code": (YEARS, YEARS + ["ETERNITY"]), + "eternal_code_plus_one": (YEARS, YEARS + ["ETERNITY"]), +} +NAMES = sorted(VARIABLES) +BRANCH_NAMES = ["a", "b"] +# Where a holder operation stores or deletes: the simulation's own branch +# name, the default one, or one the simulation does not read. +TARGETS = ["own", "default", "unread"] + +_node = st.integers(0, 5) +_name = st.sampled_from(NAMES) +_index = st.integers(0, 3) +_seed = st.integers(0, 9) +_operation = st.one_of( + st.tuples(st.just("calculate"), _node, _name, _index), + st.tuples(st.just("calculate"), _node, _name, _index), + st.tuples( + st.just("holder_set_input"), + _node, + _name, + _index, + st.sampled_from(TARGETS), + _seed, + ), + st.tuples( + st.just("holder_put_in_cache"), + _node, + _name, + _index, + st.sampled_from(TARGETS), + _seed, + ), + st.tuples( + st.just("holder_delete"), + _node, + _name, + st.one_of(st.none(), _index), + st.sampled_from(TARGETS), + ), + st.tuples(st.just("set_input"), _node, _name, _index, _seed), + st.tuples(st.just("delete_arrays"), _node, _name, st.one_of(st.none(), _index)), + st.tuples(st.just("branch"), _node, st.sampled_from(BRANCH_NAMES)), +) + + +@pytest.fixture(scope="module") +def system(): + return build_system() + + +def _values(name, seed): + value_type = INPUT_VALUE_TYPES.get(name, float) + if value_type is bool: + return np.array([seed % 2 == 1, seed % 3 != 0]) + base = np.arange(PERSON_COUNT) + 1 + if value_type is int: + return (base * (seed + 3)).astype(np.int32) + return (base * (seed + 0.5) * 12).astype(np.float32) + + +def _result(function): + try: + value = function() + except Exception as error: # Compare failures as well as values. + return ("error", type(error).__name__, str(error)) + if value is None: + return ("none",) + if isinstance(value, EnumArray): + value = value.view(np.ndarray) + value = np.asarray(value) + return ("array", value.dtype.str, value.shape, value.tobytes()) + + +class _Tree: + """A root simulation and the branches made from it, in creation order.""" + + def __init__(self, system, empty_fast_caches): + self.nodes = [build_simulation(system)] + self.empty_fast_caches = empty_fast_caches + + def node(self, index): + return self.nodes[index % len(self.nodes)] + + def fast_cache_keys(self): + return [set(simulation._fast_cache) for simulation in self.nodes] + + +def _target_branch(simulation, target): + if target == "own": + return simulation.branch_name + if target == "default": + return "default" + return "unread" + + +def _apply(tree, operation): + """Run ``operation``; return its result and what it may drop. + + The second value is ``None``, or ``(node index, predicate)`` for a holder + write or delete: the predicate says which of that simulation's + fast-cache keys the operation is allowed to remove. + """ + kind = operation[0] + simulation = tree.node(operation[1]) + node_index = tree.nodes.index(simulation) + name = operation[2] if len(operation) > 2 else None + if kind == "branch": + branch = simulation.get_branch(name) + if all(branch is not node for node in tree.nodes): + tree.nodes.append(branch) + return ("branch", tree.nodes.index(branch)), None + + calculate_periods, write_periods = VARIABLES[name] + variable = simulation.tax_benefit_system.get_variable(name) + eternal = variable.definition_period == periods.ETERNITY + if kind == "calculate": + period = calculate_periods[operation[3] % len(calculate_periods)] + if tree.empty_fast_caches: + for node in tree.nodes: + node._fast_cache.clear() + return _result(lambda: simulation.calculate(name, period)), None + + period = ( + None + if operation[3] is None + else periods.period(write_periods[operation[3] % len(write_periods)]) + ) + if kind == "set_input": + values = _values(name, operation[4]) + return _result(lambda: simulation.set_input(name, period, values)), None + if kind == "delete_arrays": + return _result(lambda: simulation.delete_arrays(name, period)), None + + holder = simulation.get_holder(name) + branch_name = _target_branch(simulation, operation[4]) + read = branch_name in simulation._get_visible_branch_names() + if kind == "holder_delete": + + def may_drop(key): + return ( + read + and key[0] == name + and (period is None or eternal or period.contains(key[1])) + ) + + return ( + _result(lambda: holder.delete_arrays(period, branch_name)), + (node_index, may_drop), + ) + + values = _values(name, operation[5]) + # A ``set_input`` helper stores every month of an annual input. + written = ( + period.get_subperiods(variable.definition_period) + if kind == "holder_set_input" + and not eternal + and period.unit != variable.definition_period + else [period] + ) + + def may_drop(key): + return read and key[0] == name and (eternal or key[1] in written) + + if kind == "holder_set_input": + result = _result(lambda: holder.set_input(period, values, branch_name)) + else: + result = _result(lambda: holder.put_in_cache(values, period, branch_name)) + return result, (node_index, may_drop) + + +def _stored(tree): + stored = [] + for simulation in tree.nodes: + values = {} + for population in simulation.populations.values(): + for name, holder in population._holders.items(): + for branch_name, period in holder.get_known_branch_periods(): + value = holder._memory_storage.get(period, branch_name) + values[(name, branch_name, str(period))] = ( + value.dtype.str, + value.tobytes(), + ) + stored.append(values) + return stored + + +@hypothesis.settings( + max_examples=500, + deadline=None, + suppress_health_check=[ + hypothesis.HealthCheck.too_slow, + hypothesis.HealthCheck.data_too_large, + hypothesis.HealthCheck.function_scoped_fixture, + ], +) +@hypothesis.given(operations=st.lists(_operation, max_size=40)) +def test_fast_cache_is_transparent_and_drops_only_what_changes(system, operations): + cached = _Tree(system, empty_fast_caches=False) + uncached = _Tree(system, empty_fast_caches=True) + + for step, operation in enumerate(operations): + before = cached.fast_cache_keys() + result, drop = _apply(cached, operation) + after = cached.fast_cache_keys() + + # 1. The fast cache never changes what ``calculate`` returns. + assert result == _apply(uncached, operation)[0], (step, operation) + + # 2. A holder write or delete drops only what it changes. + if drop is not None: + node_index, may_drop = drop + for index, (keys_before, keys_after) in enumerate(zip(before, after)): + removed = keys_before - keys_after + if index != node_index: + assert not removed, (step, operation, index, removed) + for key in removed: + assert may_drop(key), (step, operation, key) + + assert _stored(cached) == _stored(uncached) + for index, simulation in enumerate(cached.nodes): + twin = uncached.nodes[index] + twin._fast_cache.clear() + for name, (calculate_periods, _) in VARIABLES.items(): + for period in calculate_periods: + assert _result(lambda: simulation.calculate(name, period)) == _result( + lambda: twin.calculate(name, period) + ), (index, name, period) + twin._fast_cache.clear() diff --git a/tests/core/test_restore_input_registry.py b/tests/core/test_restore_input_registry.py new file mode 100644 index 000000000..b21dbe283 --- /dev/null +++ b/tests/core/test_restore_input_registry.py @@ -0,0 +1,186 @@ +"""``restore_simulation`` restores inputs as inputs. + +``apply_reform`` drops every cached value except the inputs recorded in +``_user_input_keys``. ``restore_simulation`` used to put every value back +with ``put_in_cache``, which records nothing, so ``apply_reform`` on a +restored simulation dropped its inputs too: an uprated input calculated +three years later came out as ``[0, 0]`` instead of ``[1116, 85]``. + +``dump_simulation`` now writes the periods whose values are inputs next to +each variable's arrays (``inputs.txt``), and ``restore_simulation`` records +exactly those as inputs. The property test is +``test_restore_input_registry_property.py``. +""" + +from __future__ import annotations + +import os + +import numpy as np +import pytest + +from policyengine_core import periods +from policyengine_core.tools.simulation_dumper import ( + dump_simulation, + restore_simulation, +) +from tests.fixtures.uprated_inputs import NoOp, build_simulation, build_system + +# Part of the dump format: policyengine-core#560 writes the same file. +INPUT_PERIODS_FILE = "inputs.txt" + + +@pytest.fixture(scope="module") +def system(): + return build_system() + + +def _dump_and_restore(simulation, directory): + dump_simulation(simulation, str(directory)) + return restore_simulation(str(directory), simulation.tax_benefit_system) + + +def _inputs(simulation, variable): + return sorted( + str(period) + for name, branch_name, period in simulation._user_input_keys + if name == variable and branch_name == "default" + ) + + +def test_apply_reform_after_restore_keeps_restored_inputs(system, tmp_path): + # The reported case: master gave [0, 0]. + simulation = build_simulation(system, [("uprated_count", "2012", [1001, 77])]) + simulation.calculate("uprated_count", "2013") + restored = _dump_and_restore(simulation, tmp_path) + + restored.apply_reform(NoOp) + + fresh = build_simulation(system, [("uprated_count", "2012", [1001, 77])]) + fresh.apply_reform(NoOp) + assert fresh.calculate("uprated_count", "2015").tolist() == [1116, 85] + assert restored.calculate("uprated_count", "2015").tolist() == [1116, 85] + assert restored.calculate("uprated_count", "2012").tolist() == [1001, 77] + + +def test_restore_records_inputs_but_not_calculated_values(system, tmp_path): + simulation = build_simulation(system, [("uprated_count", "2012", [1001, 77])]) + simulation.calculate("uprated_count", "2013") + simulation.calculate("uprated_count", "2014") + restored = _dump_and_restore(simulation, tmp_path) + + assert _inputs(restored, "uprated_count") == ["2012"] + # Every dumped value is back before any reform. + for year in ("2012", "2013", "2014"): + np.testing.assert_array_equal( + restored.get_holder("uprated_count").get_array(periods.period(year)), + simulation.get_holder("uprated_count").get_array(periods.period(year)), + ) + + restored.apply_reform(NoOp) + + holder = restored.get_holder("uprated_count") + assert holder.get_array(periods.period("2013")) is None + assert holder.get_array(periods.period("2014")) is None + assert holder.get_array(periods.period("2012")).tolist() == [1001, 77] + + +def test_dump_writes_the_input_periods(system, tmp_path): + simulation = build_simulation( + system, + [ + ("uprated_count", "2012", [1001, 77]), + ("uprated_count", "2014", [2001, 177]), + ], + ) + simulation.calculate("uprated_count", "2013") + simulation.calculate("doubled_amount", "2013") + dump_simulation(simulation, str(tmp_path)) + + def read(variable): + with open(tmp_path / variable / INPUT_PERIODS_FILE) as file: + return sorted(file.read().split()) + + assert read("uprated_count") == ["2012", "2014"] + # A calculated variable and a defaulted input record no inputs. + assert read("doubled_amount") == [] + assert read("uprated_amount") == [] + + +def test_restore_records_months_split_from_an_annual_input(system, tmp_path): + # ``set_input`` helpers store the twelve months; each is an input. + simulation = build_simulation(system, [("monthly_amount", "2013", [1200, 2400])]) + simulation.calculate("monthly_doubled", "2013-02") + restored = _dump_and_restore(simulation, tmp_path) + + assert _inputs(restored, "monthly_amount") == sorted( + str(month) for month in periods.period("2013").get_subperiods(periods.MONTH) + ) + assert _inputs(restored, "monthly_doubled") == [] + + restored.apply_reform(NoOp) + + assert restored.calculate("monthly_amount", "2013-07").tolist() == [100, 200] + assert restored.get_holder("monthly_doubled").get_known_periods() == [] + + +def test_restore_records_an_eternity_input_set_for_a_year(system, tmp_path): + # The record names the period given to ``set_input``; storage keeps one + # ETERNITY value, which is the input. + simulation = build_simulation(system, [("eternal_code", "2013", [5, 6])]) + simulation.calculate("eternal_code_plus_one", "2015") + restored = _dump_and_restore(simulation, tmp_path) + + assert len(_inputs(restored, "eternal_code")) == 1 + assert _inputs(restored, "eternal_code_plus_one") == [] + + restored.apply_reform(NoOp) + + assert restored.calculate("eternal_code", "2020").tolist() == [5, 6] + assert restored.get_holder("eternal_code_plus_one").get_known_periods() == [] + assert restored.calculate("eternal_code_plus_one", "2020").tolist() == [6, 7] + + +def test_restore_of_a_dump_without_an_input_record_keeps_every_value(system, tmp_path): + # Dumps written before inputs were recorded say nothing about which + # values were calculated, so every value is restored as an input. + simulation = build_simulation(system, [("uprated_count", "2012", [1001, 77])]) + simulation.calculate("uprated_count", "2013") + dump_simulation(simulation, str(tmp_path)) + for variable in os.listdir(tmp_path): + record = tmp_path / variable / INPUT_PERIODS_FILE + if record.exists(): + record.unlink() + + restored = restore_simulation(str(tmp_path), system) + + assert _inputs(restored, "uprated_count") == ["2012", "2013"] + restored.apply_reform(NoOp) + assert restored.calculate("uprated_count", "2012").tolist() == [1001, 77] + np.testing.assert_array_equal( + restored.calculate("uprated_count", "2013"), + simulation.calculate("uprated_count", "2013"), + ) + + +def test_restored_simulation_exports_its_inputs(system, tmp_path): + # ``to_input_dataframe`` exports what the record names: before, a + # restored simulation exported no input variable at all. + simulation = build_simulation( + system, + [ + ("uprated_count", "2012", [1001, 77]), + ("eligible", "2012", [True, False]), + ], + ) + simulation.calculate("uprated_count", "2013") + restored = _dump_and_restore(simulation, tmp_path) + + exported = restored.to_input_dataframe() + expected = simulation.to_input_dataframe() + + assert sorted(exported.columns) == sorted(expected.columns) + assert "uprated_count__2012" in exported.columns + assert "uprated_count__2013" not in exported.columns + for column in expected.columns: + np.testing.assert_array_equal(exported[column], expected[column]) diff --git a/tests/core/test_restore_input_registry_property.py b/tests/core/test_restore_input_registry_property.py new file mode 100644 index 000000000..058babe5e --- /dev/null +++ b/tests/core/test_restore_input_registry_property.py @@ -0,0 +1,212 @@ +"""Property: a restored simulation keeps and drops what the dumped one would. + +For random inputs (including inputs replaced after calculating, inputs split +by a ``set_input`` helper, and ETERNITY inputs set for a year) and random +calculations before the dump: + +1. **Round trip.** The restored simulation stores exactly the dumped + simulation's values (default branch), byte for byte, and records the same + storage keys as inputs. +2. **Same answers.** Calculating any sequence of requests on the restored + simulation gives what the dumped simulation gives. +3. **Reforms keep the same values.** After ``apply_reform`` with a reform + that changes nothing, the restored simulation and the dumped simulation + store the same values and give the same answers to the same requests, + byte for byte. So does a new simulation given only the inputs, except in + the one case below. + +The exception is an existing difference between the dumped simulation and a +new one, not a restore one: a ``set_input`` helper splitting an annual input +counts a month the simulation already calculated as already set, so the +split depends on what was calculated before the input (chip task_fd52e7d5). +The comparison with a new simulation leaves those examples out. + +``test_restore_input_registry.py`` pins the same behaviour with examples. +""" + +from __future__ import annotations + +import tempfile + +import numpy as np +import pytest + +# The country-package smoke job installs no dev dependencies. +hypothesis = pytest.importorskip("hypothesis") +st = hypothesis.strategies + +from policyengine_core import periods +from policyengine_core.enums import EnumArray +from policyengine_core.tools.simulation_dumper import ( + dump_simulation, + restore_simulation, +) +from tests.fixtures.uprated_inputs import ( + NoOp, + build_simulation, + build_system, + input_values, +) + +YEARS = ["2012", "2013", "2014", "2015"] +MONTHS = ["2013-01", "2013-02", "2014-06"] + +INPUT_PERIODS = { + "uprated_count": YEARS, + "uprated_amount": YEARS, + "masked_amount": YEARS, + "eligible": YEARS, + # Annual inputs go through the helper that splits them into months. + "monthly_amount": MONTHS + ["2013"], + "eternal_code": ["2013", "ETERNITY"], +} +REQUESTS = [ + ("uprated_count", YEARS), + ("uprated_amount", YEARS), + ("doubled_amount", YEARS), + ("masked_amount", YEARS), + ("eligible", YEARS), + ("monthly_amount", MONTHS), + ("monthly_doubled", MONTHS), + ("eternal_code", ["2013"]), + ("eternal_code_plus_one", ["2015"]), +] + +_input = st.tuples( + st.sampled_from(sorted(INPUT_PERIODS)), + st.integers(0, 3), + st.integers(0, 20), +).map( + lambda draw: ( + draw[0], + INPUT_PERIODS[draw[0]][draw[1] % len(INPUT_PERIODS[draw[0]])], + input_values(draw[0], draw[2]), + ) +) +_request = st.tuples(st.integers(0, len(REQUESTS) - 1), st.integers(0, 3)).map( + lambda draw: ( + REQUESTS[draw[0]][0], + REQUESTS[draw[0]][1][draw[1] % len(REQUESTS[draw[0]][1])], + ) +) +# Before the dump: set an input or calculate, in any order. +_step = st.one_of( + st.tuples(st.just("input"), _input), + st.tuples(st.just("calculate"), _request), +) + + +@pytest.fixture(scope="module") +def system(): + return build_system() + + +def _result(function): + try: + value = function() + except Exception as error: # Compare failures as well as values. + return ("error", type(error).__name__, str(error)) + if isinstance(value, EnumArray): + value = value.view(np.ndarray) + value = np.asarray(value) + return ("array", value.dtype.str, value.shape, value.tobytes()) + + +def _stored(simulation): + """Every default-branch value, as bytes, keyed by variable and period.""" + stored = {} + for population in simulation.populations.values(): + for name, holder in population._holders.items(): + for branch_name, period in holder.get_known_branch_periods(): + if branch_name != "default": + continue + value = holder.get_array(period) + stored[(name, str(period))] = (value.dtype.str, value.tobytes()) + return stored + + +def _split_after_calculating(steps): + """Whether an annual input to the monthly variable follows a calculation.""" + calculated = False + for kind, step in steps: + if kind == "calculate": + calculated = True + elif calculated and step[0] == "monthly_amount" and step[1] == "2013": + return True + return False + + +def _input_keys(simulation): + """Storage keys the input record names, as storage writes them.""" + keys = set() + for name, branch_name, period in simulation._user_input_keys: + if simulation.get_holder(name).variable.definition_period == periods.ETERNITY: + period = periods.ETERNITY + keys.add((name, branch_name, str(periods.period(period)))) + return keys + + +@hypothesis.settings( + max_examples=200, + deadline=None, + suppress_health_check=[ + hypothesis.HealthCheck.too_slow, + hypothesis.HealthCheck.data_too_large, + hypothesis.HealthCheck.function_scoped_fixture, + ], +) +@hypothesis.given( + steps=st.lists(_step, max_size=12), + requests=st.lists(_request, max_size=10), +) +def test_restored_simulation_keeps_and_drops_what_the_dumped_one_would( + system, steps, requests +): + dumped = build_simulation(system) + inputs = [] + for kind, step in steps: + if kind == "input": + # Some inputs are refused (an annual input twice for a variable + # split into months): the new simulation gets the same refusal. + inputs.append(step) + _result(lambda: dumped.set_input(*step)) + else: + _result(lambda: dumped.calculate(*step)) + + with tempfile.TemporaryDirectory(prefix="core-restore-") as directory: + dump_simulation(dumped, directory) + restored = restore_simulation(directory, system) + + # 1. Round trip. + assert _stored(restored) == _stored(dumped) + stored_keys = {(name, "default", period) for name, period in _stored(dumped).keys()} + assert _input_keys(restored) == _input_keys(dumped) & stored_keys + + # 2. Same answers. + for request in requests: + assert _result(lambda: restored.calculate(*request)) == _result( + lambda: dumped.calculate(*request) + ), request + + # 3. Reforms keep the same values. + fresh = build_simulation(system) + for step in inputs: + _result(lambda: fresh.set_input(*step)) + compare_with_fresh = not _split_after_calculating(steps) + for simulation in (dumped, restored, fresh): + simulation.apply_reform(NoOp) + assert _stored(restored) == _stored(dumped) + if compare_with_fresh: + assert _stored(restored) == _stored(fresh) + for request in requests: + answers = { + label: _result(lambda: simulation.calculate(*request)) + for label, simulation in ( + ("dumped", dumped), + ("restored", restored), + ("fresh", fresh), + ) + } + assert answers["restored"] == answers["dumped"], (request, answers) + if compare_with_fresh: + assert answers["restored"] == answers["fresh"], (request, answers) diff --git a/tests/core/variables/test_non_numeric_uprating_and_defined_for.py b/tests/core/variables/test_non_numeric_uprating_and_defined_for.py new file mode 100644 index 000000000..88e791aa3 --- /dev/null +++ b/tests/core/variables/test_non_numeric_uprating_and_defined_for.py @@ -0,0 +1,307 @@ +"""``uprating`` and ``defined_for`` need values that support arithmetic. + +Uprating multiplies a variable's earlier value by an index ratio, and +``defined_for`` keeps values where another variable is greater than zero. +Enum arrays allow only ``==`` and ``!=``, and ``str`` and date arrays cannot +be multiplied by a float or compared with a number. So an Enum, ``str`` or +date variable with ``uprating``, or a variable ``defined_for`` one, used to +register without complaint and then raise ``TypeError: Forbidden operation`` +(or a numpy error) in the middle of a calculation. + +Both are now rejected when the variables are registered, with a message that +names them. +""" + +from __future__ import annotations + +import datetime +import textwrap + +import pytest + +from policyengine_core.country_template import CountryTaxBenefitSystem +from policyengine_core.country_template import entities as template_entities +from policyengine_core.country_template.entities import Person +from policyengine_core.enums import Enum +from policyengine_core.model_api import YEAR, Reform, Variable +from policyengine_core.parameters import ParameterNode +from policyengine_core.simulations import SimulationBuilder +from policyengine_core.taxbenefitsystems import TaxBenefitSystem + + +class State(Enum): + absent = "Absent" + present = "Present" + + +class Flag(Enum): + # The value names a variable: ``defined_for = Flag.in_scope``. + in_scope = "in_scope" + + +NON_NUMERIC = { + "Enum": dict(value_type=Enum, possible_values=State, default_value=State.absent), + "str": dict(value_type=str), + "date": dict(value_type=datetime.date), +} +NUMERIC = { + "bool": dict(value_type=bool), + "int": dict(value_type=int), + "float": dict(value_type=float), +} + + +def _variable(name, **attributes): + return type( + name, + (Variable,), + dict(entity=Person, definition_period=YEAR, label=name, **attributes), + ) + + +def _system() -> CountryTaxBenefitSystem: + system = CountryTaxBenefitSystem() + system.auto_carry_over_input_variables = True + system.parameters.add_child( + "probe", + ParameterNode( + "probe", + data={"index": {"values": {"2010-01-01": 100, "2015-01-01": 200}}}, + ), + ) + return system + + +def _simulation(system, **inputs): + return SimulationBuilder().build_from_entities(system, {"persons": {"p": inputs}}) + + +# --- uprating --------------------------------------------------------------- + + +@pytest.mark.parametrize("type_name", NON_NUMERIC) +def test_uprating_on_a_non_numeric_variable_is_rejected(type_name): + system = _system() + variable = _variable("uprated", uprating="probe.index", **NON_NUMERIC[type_name]) + + with pytest.raises(ValueError) as error: + system.add_variable(variable) + + message = str(error.value) + assert 'Variable "uprated" has uprating "probe.index"' in message + assert f"value_type is {type_name}" in message + assert "uprated" not in system.variables + + +@pytest.mark.parametrize("type_name", NUMERIC) +def test_uprating_on_a_numeric_variable_is_accepted(type_name): + system = _system() + system.add_variable( + _variable("uprated", uprating="probe.index", **NUMERIC[type_name]) + ) + simulation = _simulation(system, uprated={"2012": 1}) + + # The index doubles between 2012 and 2015. + expected = {"bool": True, "int": 2, "float": 2.0}[type_name] + assert simulation.calculate("uprated", "2015").tolist() == [expected] + + +@pytest.mark.parametrize("type_name", NON_NUMERIC) +def test_assigning_uprating_to_a_non_numeric_variable_is_rejected(type_name): + # Country packages assign ``variable.uprating`` after loading (default + # uprating for dollar inputs), so the setter checks as well. + system = _system() + variable = system.add_variable(_variable("plain", **NON_NUMERIC[type_name])) + + with pytest.raises(ValueError, match='Variable "plain" has uprating'): + variable.uprating = "probe.index" + + assert variable.uprating is None + + +def test_a_reform_cannot_turn_an_uprated_variable_into_an_enum(): + # The reform's class declares no uprating, but inherits the baseline's. + system = _system() + system.add_variable(_variable("uprated", value_type=float, uprating="probe.index")) + + with pytest.raises(ValueError, match="value_type is Enum"): + system.update_variable(_variable("uprated", **NON_NUMERIC["Enum"])) + + +def test_an_enum_input_without_uprating_carries_over(): + # What the error message advises: with auto-carry-over, an input without + # ``uprating`` keeps its value in later periods. + system = _system() + system.add_variable(_variable("state", **NON_NUMERIC["Enum"])) + simulation = _simulation(system, state={"2012": "present"}) + + assert simulation.calculate("state", "2015").decode_to_str().tolist() == ["present"] + + +# --- defined_for ------------------------------------------------------------ + + +@pytest.mark.parametrize("target_first", [True, False]) +@pytest.mark.parametrize("type_name", NON_NUMERIC) +def test_defined_for_a_non_numeric_variable_is_rejected(type_name, target_first): + system = _system() + target = _variable("condition", **NON_NUMERIC[type_name]) + dependent = _variable("amount", value_type=float, defined_for="condition") + + with pytest.raises(ValueError) as error: + # Whichever is added second completes the pair. + system.add_variables( + *((target, dependent) if target_first else (dependent, target)) + ) + + message = str(error.value) + assert 'Variable "amount" is defined_for "condition"' in message + assert f"value_type is {type_name}" in message + + +@pytest.mark.parametrize( + "type_name, value, expected", + [("bool", True, 5.0), ("int", 3, 5.0), ("float", 0.5, 5.0), ("int", 0, 0.0)], +) +def test_defined_for_a_numeric_variable_keeps_nonzero_entities( + type_name, value, expected +): + system = _system() + system.add_variables( + _variable("condition", **NUMERIC[type_name]), + _variable("amount", value_type=float, defined_for="condition"), + ) + simulation = _simulation(system, condition={"2015": value}, amount={"2012": 5.0}) + + # No 2015 input: the 2012 one carries over where the condition is nonzero. + assert simulation.calculate("amount", "2015").tolist() == [expected] + + +def test_defined_for_an_enum_member_names_the_variable_called_after_its_value(): + system = _system() + system.add_variables( + _variable("in_scope", value_type=bool), + _variable("amount", value_type=float, defined_for=Flag.in_scope), + ) + assert system.variables["amount"].defined_for == "in_scope" + simulation = _simulation(system, in_scope={"2015": True}, amount={"2012": 5.0}) + + assert simulation.calculate("amount", "2015").tolist() == [5.0] + + +def test_defined_for_a_variable_that_is_not_registered_is_left_alone(): + # As before: it only fails if the variable is still missing when + # calculated. + system = _system() + system.add_variable( + _variable("amount", value_type=float, defined_for="not_defined") + ) + + assert system.variables["amount"].defined_for == "not_defined" + + +def test_a_reform_cannot_turn_a_defined_for_variable_into_an_enum(): + system = _system() + system.add_variables( + _variable("condition", value_type=bool), + _variable("amount", value_type=float, defined_for="condition"), + ) + + class make_condition_an_enum(Reform): + def apply(self): + self.update_variable(_variable("condition", **NON_NUMERIC["Enum"])) + + with pytest.raises(ValueError, match='"amount" is defined_for "condition"'): + make_condition_an_enum(system) + + +VARIABLE_FILE = textwrap.dedent( + """ + from policyengine_core.country_template.entities import Person + from policyengine_core.enums import Enum + from policyengine_core.model_api import YEAR, Variable + + + class State(Enum): + absent = "Absent" + present = "Present" + """ +) +CONDITION = textwrap.dedent( + """ + + class condition(Variable): + value_type = {value_type} + {extra} + entity = Person + definition_period = YEAR + label = "Condition" + """ +) +AMOUNT = textwrap.dedent( + """ + + class amount(Variable): + value_type = float + entity = Person + definition_period = YEAR + defined_for = "condition" + label = "Amount" + """ +) + + +def _directory_system(directory): + class DirectorySystem(TaxBenefitSystem): + entities = template_entities.entities + variables_dir = str(directory) + + return DirectorySystem + + +@pytest.mark.parametrize("condition_loaded_first", [True, False]) +def test_a_variables_directory_is_checked_once_it_is_loaded( + tmp_path, condition_loaded_first +): + # Files in a directory load before its subdirectories, so each layout + # fixes which of the two variables is registered first. + enum_condition = CONDITION.format( + value_type="Enum", + extra="possible_values = State\n default_value = State.absent", + ) + first, second = ( + (enum_condition, AMOUNT) if condition_loaded_first else (AMOUNT, enum_condition) + ) + (tmp_path / "first.py").write_text(VARIABLE_FILE + first) + (tmp_path / "later").mkdir() + (tmp_path / "later" / "second.py").write_text(VARIABLE_FILE + second) + + with pytest.raises(ValueError, match='"amount" is defined_for "condition"'): + _directory_system(tmp_path)() + + +def test_a_variables_directory_with_a_bool_condition_loads(tmp_path): + (tmp_path / "first.py").write_text(VARIABLE_FILE + AMOUNT) + (tmp_path / "later").mkdir() + (tmp_path / "later" / "second.py").write_text( + VARIABLE_FILE + CONDITION.format(value_type="bool", extra="") + ) + + system = _directory_system(tmp_path)() + + assert system.variables["amount"].defined_for == "condition" + + +def test_defined_for_changed_after_registration_fails_with_the_same_message(): + # Registration cannot see an attribute assigned afterwards; ``calculate`` + # then raises the same error instead of a TypeError from ``> 0``. + system = _system() + system.add_variables( + _variable("state", **NON_NUMERIC["Enum"]), + _variable("amount", value_type=float), + ) + system.variables["amount"].defined_for = "state" + simulation = _simulation(system, state={"2015": "present"}, amount={"2012": 5.0}) + + with pytest.raises(ValueError, match='"amount" is defined_for "state"'): + simulation.calculate("amount", "2015") diff --git a/tests/fixtures/uprated_inputs.py b/tests/fixtures/uprated_inputs.py new file mode 100644 index 000000000..ce100731c --- /dev/null +++ b/tests/fixtures/uprated_inputs.py @@ -0,0 +1,165 @@ +"""A small tax-benefit system for tests of inputs, uprating and caches. + +Two people, each in their own group entities. The variables cover the +storage paths the dump/restore and fast-cache tests need: uprated int and +float inputs, an input with a ``set_input`` helper (annual values split into +months), ETERNITY inputs and formulas, formulas reading other variables, and +an uprated input masked by ``defined_for``. +""" + +from policyengine_core import periods +from policyengine_core.country_template import CountryTaxBenefitSystem, entities +from policyengine_core.parameters import Parameter +from policyengine_core.reforms import Reform +from policyengine_core.simulations import SimulationBuilder +from policyengine_core.variables import Variable + +PERSON_COUNT = 2 + + +class uprated_count(Variable): + value_type = int + entity = entities.Person + definition_period = periods.YEAR + uprating = "index" + label = "Uprated integer input" + + +class uprated_amount(Variable): + value_type = float + entity = entities.Person + definition_period = periods.YEAR + uprating = "index" + label = "Uprated float input" + + +class doubled_amount(Variable): + value_type = float + entity = entities.Person + definition_period = periods.YEAR + label = "Twice the uprated float input" + + def formula(person, period): + return 2 * person("uprated_amount", period) + + +class monthly_amount(Variable): + value_type = float + entity = entities.Person + definition_period = periods.MONTH + label = "Monthly input; an annual input is split into twelve months" + + +class monthly_doubled(Variable): + value_type = float + entity = entities.Person + definition_period = periods.MONTH + label = "Twice the monthly input" + + def formula(person, period): + return 2 * person("monthly_amount", period) + + +class eternal_code(Variable): + value_type = int + entity = entities.Person + definition_period = periods.ETERNITY + label = "ETERNITY input" + + +class eternal_code_plus_one(Variable): + value_type = int + entity = entities.Person + definition_period = periods.ETERNITY + label = "ETERNITY formula" + + def formula(person, period): + return person("eternal_code", period) + 1 + + +class eligible(Variable): + value_type = bool + entity = entities.Person + definition_period = periods.YEAR + label = "Eligibility" + + +class masked_amount(Variable): + value_type = float + entity = entities.Person + definition_period = periods.YEAR + uprating = "index" + defined_for = "eligible" + label = "Uprated float input, kept only where eligible" + + +VARIABLES = ( + uprated_count, + uprated_amount, + doubled_amount, + monthly_amount, + monthly_doubled, + eternal_code, + eternal_code_plus_one, + eligible, + masked_amount, +) + +# Value type of each input variable, for building input arrays. +INPUT_VALUE_TYPES = { + "uprated_count": int, + "uprated_amount": float, + "monthly_amount": float, + "eternal_code": int, + "eligible": bool, + "masked_amount": float, +} + + +class NoOp(Reform): + """A reform that changes nothing: ``apply_reform`` then only drops caches.""" + + def apply(self): + pass + + +def build_system() -> CountryTaxBenefitSystem: + system = CountryTaxBenefitSystem() + system.auto_carry_over_input_variables = True + system.parameters.add_child( + "index", + Parameter( + "index", + data={ + "values": { + f"{year}-01-01": 100 * 1.037 ** (year - 2010) + for year in range(2010, 2021) + } + }, + ), + ) + system.add_variables(*VARIABLES) + return system + + +def build_simulation(system, inputs=None): + """A two-person simulation with ``inputs`` set in order. + + ``inputs`` is a list of ``(variable, period, values)``. + """ + simulation = SimulationBuilder().build_default_simulation( + system, count=PERSON_COUNT + ) + for variable, period, values in inputs or (): + simulation.set_input(variable, period, values) + return simulation + + +def input_values(variable: str, seed: int): + """Two input values for ``variable``, varied by ``seed``.""" + value_type = INPUT_VALUE_TYPES[variable] + if value_type is bool: + return [seed % 2 == 1, seed % 3 != 0] + if value_type is int: + return [1001 + 37 * seed, 77 + seed] + return [1000.5 + 37.25 * seed, 12.75 + seed] From e03a5da4e314eb7928795cc1cc8234f22fc18e41 Mon Sep 17 00:00:00 2001 From: Max Ghenis Date: Fri, 2 Oct 2026 16:43:01 -0400 Subject: [PATCH 2/7] Make the legacy-dump test strip every sidecar file, not only inputs.txt A dump written before inputs were recorded holds the arrays and nothing else, so the test stays right when another change adds its own sidecar. Co-Authored-By: Claude Opus 5.5 --- tests/core/test_restore_input_registry.py | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/tests/core/test_restore_input_registry.py b/tests/core/test_restore_input_registry.py index b21dbe283..74aa28d70 100644 --- a/tests/core/test_restore_input_registry.py +++ b/tests/core/test_restore_input_registry.py @@ -147,10 +147,13 @@ def test_restore_of_a_dump_without_an_input_record_keeps_every_value(system, tmp simulation = build_simulation(system, [("uprated_count", "2012", [1001, 77])]) simulation.calculate("uprated_count", "2013") dump_simulation(simulation, str(tmp_path)) + # Such a dump holds the arrays and nothing else. for variable in os.listdir(tmp_path): - record = tmp_path / variable / INPUT_PERIODS_FILE - if record.exists(): - record.unlink() + if variable == "__entities__": + continue + for file in (tmp_path / variable).iterdir(): + if file.suffix != ".npy": + file.unlink() restored = restore_simulation(str(tmp_path), system) From 8f56e0f8f7ac6d3c08f3213c536d6167357ead32 Mon Sep 17 00:00:00 2001 From: Max Ghenis Date: Fri, 2 Oct 2026 16:43:36 -0400 Subject: [PATCH 3/7] Define the fast-cache eviction helper at the end of Holder Next to delete_arrays, git merged it with other changes to that method without a conflict but left their lines inside the new helper. Co-Authored-By: Claude Opus 5.5 --- policyengine_core/holders/holder.py | 98 ++++++++++++++--------------- 1 file changed, 49 insertions(+), 49 deletions(-) diff --git a/policyengine_core/holders/holder.py b/policyengine_core/holders/holder.py index db45fb9ab..b42938735 100644 --- a/policyengine_core/holders/holder.py +++ b/policyengine_core/holders/holder.py @@ -107,55 +107,6 @@ def delete_arrays( self._disk_storage.delete(period, branch_name) self._evict_fast_cache(period, branch_name, contained=True) - def _evict_fast_cache( - self, period: Period, branch_name: str, contained: bool = False - ) -> None: - """Drop the simulation's ``_fast_cache`` entries a storage write or delete makes stale. - - ``Simulation.calculate`` answers a repeated request from - ``_fast_cache``, keyed by ``(variable name, requested period)``, - before it reads this holder. So every write into, or delete from, - this holder's storage drops the entries for the periods it changes; - otherwise ``calculate`` kept returning the value the storage no - longer held (for example after ``holder.set_input``). - - The fast cache belongs to this holder's simulation: a branch has its - own holders and its own fast cache, and keeps the values it started - with, so nothing outside this simulation is touched. Nor is anything - here when ``branch_name`` is a branch this simulation does not read. - - A write changes one storage key: ``period``, or, for an ETERNITY - variable, the one value every period reads. A delete (``contained``) - removes every period ``period`` contains, or every period when - ``period`` is ``None``. - """ - simulation = self.simulation - fast_cache = getattr(simulation, "_fast_cache", None) - if not fast_cache: - return - name = self.variable.name - drop_all = period is None or self.variable.definition_period == periods.ETERNITY - if not drop_all: - period = periods.period(period) - if not contained and (name, period) not in fast_cache: - return - visible_branch_names = getattr(simulation, "_get_visible_branch_names", None) - if visible_branch_names is not None and branch_name not in ( - visible_branch_names() - ): - return - if not drop_all and not contained: - del fast_cache[(name, period)] - return - stale_keys = [ - key - for key in fast_cache - if key[0] == name - and (drop_all or not isinstance(key[1], Period) or period.contains(key[1])) - ] - for key in stale_keys: - del fast_cache[key] - def _get_array_from_storage( self, period: Period, branch_name: str = "default" ) -> ArrayLike: @@ -447,3 +398,52 @@ def default_array(self) -> ArrayLike: """ return self.variable.default_array(self.population.count) + + def _evict_fast_cache( + self, period: Period, branch_name: str, contained: bool = False + ) -> None: + """Drop the simulation's ``_fast_cache`` entries a storage write or delete makes stale. + + ``Simulation.calculate`` answers a repeated request from + ``_fast_cache``, keyed by ``(variable name, requested period)``, + before it reads this holder. So every write into, or delete from, + this holder's storage drops the entries for the periods it changes; + otherwise ``calculate`` kept returning the value the storage no + longer held (for example after ``holder.set_input``). + + The fast cache belongs to this holder's simulation: a branch has its + own holders and its own fast cache, and keeps the values it started + with, so nothing outside this simulation is touched. Nor is anything + here when ``branch_name`` is a branch this simulation does not read. + + A write changes one storage key: ``period``, or, for an ETERNITY + variable, the one value every period reads. A delete (``contained``) + removes every period ``period`` contains, or every period when + ``period`` is ``None``. + """ + simulation = self.simulation + fast_cache = getattr(simulation, "_fast_cache", None) + if not fast_cache: + return + name = self.variable.name + drop_all = period is None or self.variable.definition_period == periods.ETERNITY + if not drop_all: + period = periods.period(period) + if not contained and (name, period) not in fast_cache: + return + visible_branch_names = getattr(simulation, "_get_visible_branch_names", None) + if visible_branch_names is not None and branch_name not in ( + visible_branch_names() + ): + return + if not drop_all and not contained: + del fast_cache[(name, period)] + return + stale_keys = [ + key + for key in fast_cache + if key[0] == name + and (drop_all or not isinstance(key[1], Period) or period.contains(key[1])) + ] + for key in stale_keys: + del fast_cache[key] From 3c7df2980c4f74e41fba24e10592c39dc0ebfe94 Mon Sep 17 00:00:00 2001 From: Max Ghenis Date: Fri, 2 Oct 2026 16:49:10 -0400 Subject: [PATCH 4/7] Keep the branch fast-cache tests valid if a branch input also drops what was calculated from it Co-Authored-By: Claude Opus 5.5 --- tests/core/test_holder_write_fast_cache.py | 3 ++- tests/core/test_holder_write_fast_cache_property.py | 9 ++++++++- 2 files changed, 10 insertions(+), 2 deletions(-) diff --git a/tests/core/test_holder_write_fast_cache.py b/tests/core/test_holder_write_fast_cache.py index f22aac07c..ca91d580f 100644 --- a/tests/core/test_holder_write_fast_cache.py +++ b/tests/core/test_holder_write_fast_cache.py @@ -144,7 +144,6 @@ def test_a_branch_write_leaves_the_parent_and_the_parent_write_leaves_the_branch simulation = build_simulation(system, [("uprated_count", "2012", [1001, 77])]) parent_value = simulation.calculate("uprated_count", "2013") branch = simulation.get_branch("reform") - branch_value = branch.calculate("uprated_count", "2014") # A branch has its own holders and its own fast cache. branch.get_holder("uprated_count").set_input( @@ -153,9 +152,11 @@ def test_a_branch_write_leaves_the_parent_and_the_parent_write_leaves_the_branch assert simulation.calculate("uprated_count", "2013") is parent_value assert branch.calculate("uprated_count", "2013").tolist() == [300, 400] + branch_value = branch.calculate("uprated_count", "2014") simulation.get_holder("uprated_count").set_input(periods.period("2014"), [5, 6]) assert simulation.calculate("uprated_count", "2014").tolist() == [5, 6] assert branch.calculate("uprated_count", "2014") is branch_value + assert branch_value.tolist() == [311, 414] def test_storing_a_calculated_value_keeps_it_in_the_fast_cache(system): diff --git a/tests/core/test_holder_write_fast_cache_property.py b/tests/core/test_holder_write_fast_cache_property.py index 1dc8f8041..d5e22c95e 100644 --- a/tests/core/test_holder_write_fast_cache_property.py +++ b/tests/core/test_holder_write_fast_cache_property.py @@ -14,7 +14,10 @@ removes all belong to its own simulation and variable, were stored under a branch that simulation reads, and are for the period written (any period, for an ETERNITY variable) or for a period the deleted one - contains. Every other simulation's fast cache is untouched. + contains. Every other simulation's fast cache is untouched. One case is + left open: an input set on a branch may drop more of that branch's + entries, because whether it also drops the values calculated from the + input it replaces is a separate question (policyengine-core#560). ``test_holder_write_fast_cache.py`` pins the same behaviour with examples. """ @@ -208,7 +211,11 @@ def may_drop(key): else [period] ) + on_branch = simulation.parent_branch is not None + def may_drop(key): + if kind == "holder_set_input" and on_branch: + return True return read and key[0] == name and (eternal or key[1] in written) if kind == "holder_set_input": From 1c24abd6952d2b9e0de55486fa7d8584edab2a5e Mon Sep 17 00:00:00 2001 From: Max Ghenis Date: Fri, 2 Oct 2026 17:40:24 -0400 Subject: [PATCH 5/7] Address review: branch inputs in dumps, replace rollback, calculate-time checks, legacy-dump warning - Test that an input set on a branch (which shares the input record) does not mark the default branch's calculated value as an input, with a branch step in the restore property. - replace_variable keeps the existing variable when the replacement is rejected. - calculate repeats the uprating check before it uprates, for an uprating assigned on a class that declares uprating itself; the defined_for message is tested for str and date as well as Enum. - The message for an inherited uprating points to replace_variable. - restore_simulation warns when a dump does not record its inputs. - Changelog: say that such systems no longer load, including a group variable defined_for a person Enum, which master masked on summed indices. Co-Authored-By: Claude Opus 5.5 --- ...on-numeric-uprating-defined-for.changed.md | 1 + .../non-numeric-uprating-defined-for.fixed.md | 1 - .../replace-variable-rollback.fixed.md | 1 + .../restore-simulation-input-record.fixed.md | 2 +- policyengine_core/simulations/simulation.py | 3 + .../taxbenefitsystems/tax_benefit_system.py | 16 ++-- policyengine_core/tools/simulation_dumper.py | 24 +++++- policyengine_core/variables/variable.py | 23 +++++- tests/core/test_restore_input_registry.py | 54 +++++++++++-- .../test_restore_input_registry_property.py | 12 ++- ...st_non_numeric_uprating_and_defined_for.py | 78 +++++++++++++++++-- 11 files changed, 185 insertions(+), 30 deletions(-) create mode 100644 changelog.d/non-numeric-uprating-defined-for.changed.md delete mode 100644 changelog.d/non-numeric-uprating-defined-for.fixed.md create mode 100644 changelog.d/replace-variable-rollback.fixed.md diff --git a/changelog.d/non-numeric-uprating-defined-for.changed.md b/changelog.d/non-numeric-uprating-defined-for.changed.md new file mode 100644 index 000000000..5a1f87ec1 --- /dev/null +++ b/changelog.d/non-numeric-uprating-defined-for.changed.md @@ -0,0 +1 @@ +`uprating` on an Enum, `str` or date variable, and a `defined_for` that names one, are now rejected with a `ValueError` when the variables are registered (or when `uprating` is assigned), so a tax-benefit system containing either no longer loads; before, it loaded and raised `TypeError` when the variable was uprated or calculated, or, for a group variable `defined_for` a person Enum, masked on the sum of its members' Enum indices. `defined_for` must name a `bool`, `int` or `float` variable. diff --git a/changelog.d/non-numeric-uprating-defined-for.fixed.md b/changelog.d/non-numeric-uprating-defined-for.fixed.md deleted file mode 100644 index cbfed0853..000000000 --- a/changelog.d/non-numeric-uprating-defined-for.fixed.md +++ /dev/null @@ -1 +0,0 @@ -An Enum, `str` or date variable with `uprating`, and a variable whose `defined_for` names one, are now rejected with a `ValueError` when registered (or when `uprating` is assigned) instead of raising `TypeError` in the middle of a calculation; `defined_for` must name a `bool`, `int` or `float` variable. diff --git a/changelog.d/replace-variable-rollback.fixed.md b/changelog.d/replace-variable-rollback.fixed.md new file mode 100644 index 000000000..9bc941603 --- /dev/null +++ b/changelog.d/replace-variable-rollback.fixed.md @@ -0,0 +1 @@ +`TaxBenefitSystem.replace_variable` now keeps the existing variable when the replacement is rejected, instead of leaving the system without it. diff --git a/changelog.d/restore-simulation-input-record.fixed.md b/changelog.d/restore-simulation-input-record.fixed.md index ac96e8344..9eace8811 100644 --- a/changelog.d/restore-simulation-input-record.fixed.md +++ b/changelog.d/restore-simulation-input-record.fixed.md @@ -1 +1 @@ -`restore_simulation` now restores as inputs the values the dumped simulation stored as inputs (`dump_simulation` lists their periods in an `inputs.txt` beside each variable's arrays), so `apply_reform` on a restored simulation no longer drops them; a dump written before this has no such list, so every value in it is restored as an input. +`restore_simulation` now restores as inputs the values the dumped simulation stored as inputs (`dump_simulation` lists their periods in an `inputs.txt` beside each variable's arrays), so `apply_reform` on a restored simulation no longer drops them; a dump written before this has no such list, so every value in it is restored as an input, with a warning. diff --git a/policyengine_core/simulations/simulation.py b/policyengine_core/simulations/simulation.py index ed32635e0..25cd938e6 100644 --- a/policyengine_core/simulations/simulation.py +++ b/policyengine_core/simulations/simulation.py @@ -912,6 +912,9 @@ def _calculate(self, variable_name: str, period: Period = None) -> ArrayLike: and known_period.start < period.start ] if variable.uprating is not None and len(earlier_known_periods) > 0: + # Registration rejects these; an ``uprating`` assigned + # past the setter gets the same message here. + variable.check_uprating_value_type() # Take the latest period from the filtered list itself. # Indexing ``known_periods`` with a position in the # filtered list picked the wrong period whenever a later diff --git a/policyengine_core/taxbenefitsystems/tax_benefit_system.py b/policyengine_core/taxbenefitsystems/tax_benefit_system.py index e730a2e7c..dd03950ca 100644 --- a/policyengine_core/taxbenefitsystems/tax_benefit_system.py +++ b/policyengine_core/taxbenefitsystems/tax_benefit_system.py @@ -122,8 +122,9 @@ def __init__(self, entities: Sequence[Entity] = None, reform=None) -> None: self.variable_module_metadata = {} if self.variables_dir is not None: - # A variable's defined_for variable may be in a file loaded after - # it, so check every variable once the whole directory is loaded. + # Check every variable once, after the whole directory is loaded, + # instead of scanning the variables loaded so far each time a + # non-numeric one is added. self._defined_for_checks_deferred = True try: self.add_variables_from_directory(self.variables_dir) @@ -287,9 +288,14 @@ def replace_variable(self, variable: Type[Variable]) -> None: :param Variable variable: New variable class to add. Must be a subclass of Variable. """ name = variable.__name__ - if self.variables.get(name) is not None: - del self.variables[name] - self.load_variable(variable, update=False) + replaced = self.variables.pop(name, None) + try: + self.load_variable(variable, update=False) + except Exception: + # A rejected replacement leaves the system as it was. + if replaced is not None: + self.variables[name] = replaced + raise self.data_modified = True def update_variable(self, variable: Type[Variable]) -> Variable: diff --git a/policyengine_core/tools/simulation_dumper.py b/policyengine_core/tools/simulation_dumper.py index 08edc1b3b..1320bd073 100644 --- a/policyengine_core/tools/simulation_dumper.py +++ b/policyengine_core/tools/simulation_dumper.py @@ -2,6 +2,7 @@ import os +import warnings import numpy as np @@ -53,7 +54,8 @@ def restore_simulation(directory, tax_benefit_system, **kwargs): every other value as a calculated one, so ``apply_reform`` keeps and drops the same values it would have in the dumped simulation. A dump written before inputs were recorded (no ``inputs.txt``) does not say which values - were inputs, so every value in it is restored as an input. + were inputs, so every value in it is restored as an input, with a + warning: ``apply_reform`` then keeps its calculated values as dumped. """ simulation = Simulation( tax_benefit_system, tax_benefit_system.instantiate_entities() @@ -74,8 +76,22 @@ def restore_simulation(directory, tax_benefit_system, **kwargs): variables_to_restore = ( variable for variable in os.listdir(directory) if variable != "__entities__" ) - for variable in variables_to_restore: - _restore_holder(simulation, variable, directory) + without_input_record = [ + variable + for variable in variables_to_restore + if not _restore_holder(simulation, variable, directory) + ] + if without_input_record: + warnings.warn( + f"The simulation dump in {directory} does not record which values " + f"were inputs ({len(without_input_record)} variables have no " + f"{INPUT_PERIODS_FILE}; it was written by an earlier version of " + "policyengine-core). Every value in it is restored as an input, " + "so apply_reform keeps the calculated values as dumped instead of " + "recalculating them. Dump the simulation again to record its " + "inputs.", + stacklevel=2, + ) return simulation @@ -166,6 +182,7 @@ def _restore_entity(population, directory): def _restore_holder(simulation, variable, directory): + """Restore one variable's values; return whether its inputs were recorded.""" storage_dir = os.path.join(directory, variable) is_variable_eternal = ( simulation.tax_benefit_system.get_variable(variable).definition_period @@ -193,6 +210,7 @@ def _restore_holder(simulation, variable, directory): _restore_input(simulation, holder, period, value) else: holder.put_in_cache(value, period) + return input_periods is not None def _restore_input(simulation, holder, period, value): diff --git a/policyengine_core/variables/variable.py b/policyengine_core/variables/variable.py index ab0d2bb32..00912db76 100644 --- a/policyengine_core/variables/variable.py +++ b/policyengine_core/variables/variable.py @@ -429,19 +429,36 @@ def check_uprating_value_type(self): variable raised a ``TypeError`` the first time a later period was uprated. The check covers an ``uprating`` inherited from a baseline variable too, since the inherited one is what ``calculate`` uses. + + It runs when the variable is defined and when ``uprating`` is + assigned. A class that declares ``uprating`` itself (even as + ``None``) replaces the property that checks assignments, so + ``calculate`` runs it again before it uprates. """ if self.uprating is None: return value_type = getattr(self, "value_type", None) if value_type is None or value_type in NUMERIC_VALUE_TYPES: return + explicit = getattr(self, "_explicit_attribute_names", frozenset()) + if self.baseline_variable is not None and "uprating" not in explicit: + # ``update_variable`` keeps every attribute the new class leaves + # unset, so the new class cannot drop the uprating itself. + advice = ( + "The uprating is inherited from the variable this one " + "updates: use replace_variable, which inherits nothing, to " + "change its value_type." + ) + else: + advice = ( + "Remove uprating: with auto_carry_over_input_variables, an " + "input without it carries over to later periods unchanged." + ) raise ValueError( f'Variable "{self.name}" has uprating "{self.uprating}", but its ' f"value_type is {_value_type_name(value_type)}. Uprating " "multiplies the latest earlier value by an index ratio, so only " - "bool, int and float variables can be uprated. Remove uprating: " - "with auto_carry_over_input_variables, an input without it " - "carries over to later periods unchanged." + f"bool, int and float variables can be uprated. {advice}" ) def check_defined_for_variable(self, defined_for_variable: "Variable") -> None: diff --git a/tests/core/test_restore_input_registry.py b/tests/core/test_restore_input_registry.py index 74aa28d70..f949addd9 100644 --- a/tests/core/test_restore_input_registry.py +++ b/tests/core/test_restore_input_registry.py @@ -15,6 +15,7 @@ from __future__ import annotations import os +import warnings import numpy as np import pytest @@ -141,21 +142,48 @@ def test_restore_records_an_eternity_input_set_for_a_year(system, tmp_path): assert restored.calculate("eternal_code_plus_one", "2020").tolist() == [6, 7] -def test_restore_of_a_dump_without_an_input_record_keeps_every_value(system, tmp_path): - # Dumps written before inputs were recorded say nothing about which - # values were calculated, so every value is restored as an input. +def test_an_input_set_on_a_branch_is_not_recorded_for_the_default_value( + system, tmp_path +): + # A branch shares its parent's input record. Its input for 2013 must not + # make the parent's calculated 2013 value, which is what gets dumped, an + # input. simulation = build_simulation(system, [("uprated_count", "2012", [1001, 77])]) simulation.calculate("uprated_count", "2013") - dump_simulation(simulation, str(tmp_path)) - # Such a dump holds the arrays and nothing else. - for variable in os.listdir(tmp_path): + simulation.get_branch("reform").set_input("uprated_count", "2013", [5, 6]) + assert ("uprated_count", "reform", periods.period("2013")) in ( + simulation._user_input_keys + ) + + restored = _dump_and_restore(simulation, tmp_path) + + assert _inputs(restored, "uprated_count") == ["2012"] + restored.apply_reform(NoOp) + assert ( + restored.get_holder("uprated_count").get_array(periods.period("2013")) is None + ) + + +def _dump_without_input_record(simulation, directory): + """A dump as earlier versions wrote it: the arrays and nothing else.""" + dump_simulation(simulation, str(directory)) + for variable in os.listdir(directory): if variable == "__entities__": continue - for file in (tmp_path / variable).iterdir(): + for file in (directory / variable).iterdir(): if file.suffix != ".npy": file.unlink() - restored = restore_simulation(str(tmp_path), system) + +def test_restore_of_a_dump_without_an_input_record_keeps_every_value(system, tmp_path): + # Dumps written before inputs were recorded say nothing about which + # values were calculated, so every value is restored as an input. + simulation = build_simulation(system, [("uprated_count", "2012", [1001, 77])]) + simulation.calculate("uprated_count", "2013") + _dump_without_input_record(simulation, tmp_path) + + with pytest.warns(UserWarning, match="does not record which values were inputs"): + restored = restore_simulation(str(tmp_path), system) assert _inputs(restored, "uprated_count") == ["2012", "2013"] restored.apply_reform(NoOp) @@ -166,6 +194,16 @@ def test_restore_of_a_dump_without_an_input_record_keeps_every_value(system, tmp ) +def test_restore_of_a_dump_with_an_input_record_does_not_warn(system, tmp_path): + simulation = build_simulation(system, [("uprated_count", "2012", [1001, 77])]) + simulation.calculate("uprated_count", "2013") + dump_simulation(simulation, str(tmp_path)) + + with warnings.catch_warnings(): + warnings.simplefilter("error") + restore_simulation(str(tmp_path), system) + + def test_restored_simulation_exports_its_inputs(system, tmp_path): # ``to_input_dataframe`` exports what the record names: before, a # restored simulation exported no input variable at all. diff --git a/tests/core/test_restore_input_registry_property.py b/tests/core/test_restore_input_registry_property.py index 058babe5e..722f4a17e 100644 --- a/tests/core/test_restore_input_registry_property.py +++ b/tests/core/test_restore_input_registry_property.py @@ -1,8 +1,8 @@ """Property: a restored simulation keeps and drops what the dumped one would. For random inputs (including inputs replaced after calculating, inputs split -by a ``set_input`` helper, and ETERNITY inputs set for a year) and random -calculations before the dump: +by a ``set_input`` helper, ETERNITY inputs set for a year, and inputs set on +a branch of the dumped simulation) and random calculations before the dump: 1. **Round trip.** The restored simulation stores exactly the dumped simulation's values (default branch), byte for byte, and records the same @@ -89,10 +89,12 @@ REQUESTS[draw[0]][1][draw[1] % len(REQUESTS[draw[0]][1])], ) ) -# Before the dump: set an input or calculate, in any order. +# Before the dump: set an input, calculate, or set an input on a branch (a +# branch shares the simulation's input record), in any order. _step = st.one_of( st.tuples(st.just("input"), _input), st.tuples(st.just("calculate"), _request), + st.tuples(st.just("branch_input"), _input), ) @@ -170,6 +172,10 @@ def test_restored_simulation_keeps_and_drops_what_the_dumped_one_would( # split into months): the new simulation gets the same refusal. inputs.append(step) _result(lambda: dumped.set_input(*step)) + elif kind == "branch_input": + # Stays on the branch: the dumped simulation's own values and + # inputs are what a new simulation is given. + _result(lambda: dumped.get_branch("reform").set_input(*step)) else: _result(lambda: dumped.calculate(*step)) diff --git a/tests/core/variables/test_non_numeric_uprating_and_defined_for.py b/tests/core/variables/test_non_numeric_uprating_and_defined_for.py index 88e791aa3..0de05694c 100644 --- a/tests/core/variables/test_non_numeric_uprating_and_defined_for.py +++ b/tests/core/variables/test_non_numeric_uprating_and_defined_for.py @@ -124,9 +124,17 @@ def test_a_reform_cannot_turn_an_uprated_variable_into_an_enum(): system = _system() system.add_variable(_variable("uprated", value_type=float, uprating="probe.index")) - with pytest.raises(ValueError, match="value_type is Enum"): + with pytest.raises(ValueError) as error: system.update_variable(_variable("uprated", **NON_NUMERIC["Enum"])) + # ``update_variable`` cannot drop an inherited attribute, so the message + # points to the way that can. + assert "value_type is Enum" in str(error.value) + assert "use replace_variable" in str(error.value) + + system.replace_variable(_variable("uprated", **NON_NUMERIC["Enum"])) + assert system.variables["uprated"].uprating is None + def test_an_enum_input_without_uprating_carries_over(): # What the error message advises: with auto-carry-over, an input without @@ -292,16 +300,74 @@ def test_a_variables_directory_with_a_bool_condition_loads(tmp_path): assert system.variables["amount"].defined_for == "condition" -def test_defined_for_changed_after_registration_fails_with_the_same_message(): +@pytest.mark.parametrize("type_name", NON_NUMERIC) +def test_defined_for_changed_after_registration_fails_with_the_same_message( + type_name, +): # Registration cannot see an attribute assigned afterwards; ``calculate`` # then raises the same error instead of a TypeError from ``> 0``. system = _system() system.add_variables( - _variable("state", **NON_NUMERIC["Enum"]), + _variable("condition", **NON_NUMERIC[type_name]), _variable("amount", value_type=float), ) - system.variables["amount"].defined_for = "state" - simulation = _simulation(system, state={"2015": "present"}, amount={"2012": 5.0}) + system.variables["amount"].defined_for = "condition" + simulation = _simulation(system, amount={"2012": 5.0}) - with pytest.raises(ValueError, match='"amount" is defined_for "state"'): + with pytest.raises(ValueError) as error: simulation.calculate("amount", "2015") + + message = str(error.value) + assert 'Variable "amount" is defined_for "condition"' in message + assert f"value_type is {type_name}" in message + + +def test_a_group_variable_defined_for_a_person_enum_is_rejected(): + # Mapped to the group, the Enum's indices were summed over its members, + # so this masked on a number with no meaning instead of raising. + system = _system() + system.add_variable(_variable("state", **NON_NUMERIC["Enum"])) + household_amount = type( + "household_amount", + (Variable,), + dict( + value_type=float, + entity=template_entities.Household, + definition_period=YEAR, + label="household_amount", + defined_for="state", + ), + ) + + with pytest.raises(ValueError, match='"household_amount" is defined_for "state"'): + system.add_variable(household_amount) + + +def test_a_rejected_replacement_leaves_the_variable_as_it_was(): + system = _system() + system.add_variables( + _variable("condition", value_type=bool), + _variable("amount", value_type=float, defined_for="condition"), + ) + original = system.variables["condition"] + + with pytest.raises(ValueError, match='"amount" is defined_for "condition"'): + system.replace_variable(_variable("condition", **NON_NUMERIC["Enum"])) + + assert system.variables["condition"] is original + simulation = _simulation(system, condition={"2015": True}, amount={"2012": 5.0}) + assert simulation.calculate("amount", "2015").tolist() == [5.0] + + +def test_uprating_assigned_past_the_setter_fails_when_it_would_uprate(): + # A class that declares ``uprating`` itself, even as ``None``, replaces + # the property that checks assignments. + system = _system() + variable = system.add_variable( + _variable("plain", uprating=None, **NON_NUMERIC["Enum"]) + ) + variable.uprating = "probe.index" + simulation = _simulation(system, plain={"2012": "present"}) + + with pytest.raises(ValueError, match='Variable "plain" has uprating'): + simulation.calculate("plain", "2015") From a6949c7c2ffb167a8fa9f418a6e5ead8d15f8d62 Mon Sep 17 00:00:00 2001 From: Max Ghenis Date: Fri, 2 Oct 2026 18:14:10 -0400 Subject: [PATCH 6/7] Test the branch-input case by its outcome, not by whether the record is shared policyengine-core#561 makes a branch copy the input record instead of sharing it; the test now holds either way. Co-Authored-By: Claude Opus 5.5 --- tests/core/test_restore_input_registry.py | 10 ++++------ 1 file changed, 4 insertions(+), 6 deletions(-) diff --git a/tests/core/test_restore_input_registry.py b/tests/core/test_restore_input_registry.py index f949addd9..050915abc 100644 --- a/tests/core/test_restore_input_registry.py +++ b/tests/core/test_restore_input_registry.py @@ -145,15 +145,13 @@ def test_restore_records_an_eternity_input_set_for_a_year(system, tmp_path): def test_an_input_set_on_a_branch_is_not_recorded_for_the_default_value( system, tmp_path ): - # A branch shares its parent's input record. Its input for 2013 must not - # make the parent's calculated 2013 value, which is what gets dumped, an - # input. + # A branch shares its parent's input record (unless it copies it, as + # policyengine-core#561 makes it), so its input for 2013 may be in the + # record the dump reads. It must not make the parent's calculated 2013 + # value, which is what gets dumped, an input. simulation = build_simulation(system, [("uprated_count", "2012", [1001, 77])]) simulation.calculate("uprated_count", "2013") simulation.get_branch("reform").set_input("uprated_count", "2013", [5, 6]) - assert ("uprated_count", "reform", periods.period("2013")) in ( - simulation._user_input_keys - ) restored = _dump_and_restore(simulation, tmp_path) From 83597c22b1eeb31257e3e9fd70acbc86afcb3b6d Mon Sep 17 00:00:00 2001 From: Max Ghenis Date: Sat, 3 Oct 2026 07:30:10 -0400 Subject: [PATCH 7/7] Address review round 2: check defined_for by type before mapping; pin three cache paths - calculate checks the defined_for variable's type before calculating and mapping it. Mapped to a group entity, an Enum's indices were summed into numbers and str or date values failed inside the mapping, so the values-based check missed a cross-entity defined_for set after registration. - Regression tests for three cache paths no test pinned (each one a surviving mutant in GPT-6.1 Sol's review): a write under an ancestor branch name in a nested branch, a write to disk storage, and an ETERNITY value requested without a period. Co-Authored-By: Claude Opus 5.5 --- policyengine_core/simulations/simulation.py | 27 ++++++----- tests/core/test_holder_write_fast_cache.py | 48 +++++++++++++++++++ ...st_non_numeric_uprating_and_defined_for.py | 39 +++++++++++++++ 3 files changed, 101 insertions(+), 13 deletions(-) diff --git a/policyengine_core/simulations/simulation.py b/policyengine_core/simulations/simulation.py index 25cd938e6..e780fd57a 100644 --- a/policyengine_core/simulations/simulation.py +++ b/policyengine_core/simulations/simulation.py @@ -874,20 +874,21 @@ def _calculate(self, variable_name: str, period: Period = None) -> ArrayLike: self._check_period_consistency(period, variable) if variable.defined_for is not None: - defined_for_values = self.calculate( - variable.defined_for, period, map_to=variable.entity.key + # Registration rejects a non-numeric defined_for variable (see + # ``TaxBenefitSystem._check_defined_for``). This catches one set, + # or a variable replaced, afterwards. It reads the variable's + # type, not the values: mapped to a group entity, an Enum's + # indices are summed into numbers, and str or date values fail + # inside the mapping. + defined_for_variable = self.tax_benefit_system.get_variable( + variable.defined_for + ) + if defined_for_variable is not None: + variable.check_defined_for_variable(defined_for_variable) + mask = ( + self.calculate(variable.defined_for, period, map_to=variable.entity.key) + > 0 ) - if isinstance(defined_for_values, EnumArray) or ( - getattr(defined_for_values, "dtype", np.dtype(float)).kind not in "biuf" - ): - # Registration rejects these (see - # ``TaxBenefitSystem._check_defined_for``); this catches a - # defined_for set or a variable replaced afterwards, with - # the same message instead of a TypeError from ``> 0``. - variable.check_defined_for_variable( - self.tax_benefit_system.get_variable(variable.defined_for) - ) - mask = defined_for_values > 0 if np.all(~mask): array = holder.default_array() array = self._cast_formula_result(array, variable) diff --git a/tests/core/test_holder_write_fast_cache.py b/tests/core/test_holder_write_fast_cache.py index ca91d580f..2e78e6a16 100644 --- a/tests/core/test_holder_write_fast_cache.py +++ b/tests/core/test_holder_write_fast_cache.py @@ -19,6 +19,7 @@ import pytest from policyengine_core import periods +from policyengine_core.experimental import MemoryConfig from tests.fixtures.uprated_inputs import build_simulation, build_system @@ -168,3 +169,50 @@ def test_storing_a_calculated_value_keeps_it_in_the_fast_cache(system): assert _fast_cached(simulation, "doubled_amount") == ["2012"] assert simulation.calculate("doubled_amount", "2012") is first + + +def test_write_under_an_ancestor_branch_name_replaces_what_a_nested_branch_read(): + # A nested branch reads its own key, then each ancestor's, then the + # default. With the variable kept out of storage (cache blacklist), the + # nested branch's fast cache is all that holds its calculated value, so + # a write under the parent branch's name must drop it. + system = build_system() + system.cache_blacklist = ["doubled_amount"] + simulation = build_simulation(system) + simulation.opt_out_cache = True + nested = simulation.get_branch("a").get_branch("b") + assert nested.calculate("doubled_amount", "2013").tolist() == [0, 0] + assert _fast_cached(nested, "doubled_amount") == ["2013"] + + nested.get_holder("doubled_amount").set_input( + periods.period("2013"), [7.0, 8.0], "a" + ) + + assert nested.calculate("doubled_amount", "2013").tolist() == [7.0, 8.0] + + +def test_a_write_to_disk_storage_replaces_a_calculated_value(system): + simulation = build_simulation(system) + with pytest.warns(Warning): + simulation.memory_config = MemoryConfig(max_memory_occupation=0) + holder = simulation.get_holder("doubled_amount") + holder._disk_storage = holder.create_disk_storage() + holder._on_disk_storable = True + assert simulation.calculate("doubled_amount", "2013").tolist() == [0, 0] + assert holder._memory_storage.get(periods.period("2013")) is None + + holder.put_in_cache(np.array([7.0, 8.0], dtype=np.float32), periods.period("2013")) + + assert holder._disk_storage.get(periods.period("2013")).tolist() == [7.0, 8.0] + assert simulation.calculate("doubled_amount", "2013").tolist() == [7.0, 8.0] + + +def test_a_write_replaces_an_eternity_value_requested_without_a_period(system): + # ``calculate`` with no period caches the result under ``None``. + simulation = build_simulation(system) + assert simulation.calculate("eternal_code").tolist() == [0, 0] + assert ("eternal_code", None) in simulation._fast_cache + + simulation.get_holder("eternal_code").set_input(periods.period("2015"), [8, 9]) + + assert simulation.calculate("eternal_code").tolist() == [8, 9] diff --git a/tests/core/variables/test_non_numeric_uprating_and_defined_for.py b/tests/core/variables/test_non_numeric_uprating_and_defined_for.py index 0de05694c..27c86d578 100644 --- a/tests/core/variables/test_non_numeric_uprating_and_defined_for.py +++ b/tests/core/variables/test_non_numeric_uprating_and_defined_for.py @@ -322,6 +322,45 @@ def test_defined_for_changed_after_registration_fails_with_the_same_message( assert f"value_type is {type_name}" in message +@pytest.mark.parametrize("type_name", NON_NUMERIC) +def test_a_group_variable_defined_for_changed_after_registration_fails_too( + type_name, +): + # Mapped to the household, an Enum's indices are summed into numbers and + # str or date values fail inside the mapping, so the check reads the + # variable's type, not the values. + system = _system() + system.add_variables( + _variable("condition", **NON_NUMERIC[type_name]), + type( + "household_amount", + (Variable,), + dict( + value_type=float, + entity=template_entities.Household, + definition_period=YEAR, + label="household_amount", + ), + ), + ) + system.variables["household_amount"].defined_for = "condition" + values = {"Enum": "present", "str": "yes", "date": "2000-01-01"}[type_name] + simulation = SimulationBuilder().build_from_entities( + system, + { + "persons": {"a": {"condition": {"2015": values}}}, + "households": {"h": {"parents": ["a"], "household_amount": {"2012": 20.0}}}, + }, + ) + + with pytest.raises(ValueError) as error: + simulation.calculate("household_amount", "2015") + + message = str(error.value) + assert 'Variable "household_amount" is defined_for "condition"' in message + assert f"value_type is {type_name}" in message + + def test_a_group_variable_defined_for_a_person_enum_is_rejected(): # Mapped to the group, the Enum's indices were summed over its members, # so this masked on a number with no meaning instead of raising.