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 00000000..716d56bb --- /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.changed.md b/changelog.d/non-numeric-uprating-defined-for.changed.md new file mode 100644 index 00000000..5a1f87ec --- /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/replace-variable-rollback.fixed.md b/changelog.d/replace-variable-rollback.fixed.md new file mode 100644 index 00000000..9bc94160 --- /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 new file mode 100644 index 00000000..9eace881 --- /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, with a warning. diff --git a/policyengine_core/holders/holder.py b/policyengine_core/holders/holder.py index 11a8ae50..b4293873 100644 --- a/policyengine_core/holders/holder.py +++ b/policyengine_core/holders/holder.py @@ -105,6 +105,7 @@ 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 _get_array_from_storage( self, period: Period, branch_name: str = "default" @@ -370,6 +371,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() @@ -396,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] diff --git a/policyengine_core/simulations/simulation.py b/policyengine_core/simulations/simulation.py index 2652d972..e780fd57 100644 --- a/policyengine_core/simulations/simulation.py +++ b/policyengine_core/simulations/simulation.py @@ -874,6 +874,17 @@ def _calculate(self, variable_name: str, period: Period = None) -> ArrayLike: self._check_period_consistency(period, variable) if variable.defined_for is not None: + # 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 @@ -902,6 +913,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 c381b11c..dd03950c 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,15 @@ 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) + # 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) + 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 +226,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. @@ -243,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 c3db0c4f..1320bd07 100644 --- a/policyengine_core/tools/simulation_dumper.py +++ b/policyengine_core/tools/simulation_dumper.py @@ -2,13 +2,22 @@ import os +import warnings 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 +35,27 @@ 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, with a + warning: ``apply_reform`` then keeps its calculated values as dumped. """ simulation = Simulation( tax_benefit_system, tax_benefit_system.instantiate_entities() @@ -58,17 +76,58 @@ 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 -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): @@ -123,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 @@ -135,6 +195,34 @@ 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) + return input_periods is not None + + +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 629a1f4f..00912db7 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,71 @@ 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. + + 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 " + f"bool, int and float variables can be uprated. {advice}" + ) + + 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 00000000..2e78e6a1 --- /dev/null +++ b/tests/core/test_holder_write_fast_cache.py @@ -0,0 +1,218 @@ +"""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 policyengine_core.experimental import MemoryConfig +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") + + # 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] + + 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): + # ``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 + + +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/test_holder_write_fast_cache_property.py b/tests/core/test_holder_write_fast_cache_property.py new file mode 100644 index 00000000..d5e22c95 --- /dev/null +++ b/tests/core/test_holder_write_fast_cache_property.py @@ -0,0 +1,285 @@ +"""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. 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. +""" + +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] + ) + + 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": + 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 00000000..050915ab --- /dev/null +++ b/tests/core/test_restore_input_registry.py @@ -0,0 +1,225 @@ +"""``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 warnings + +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_an_input_set_on_a_branch_is_not_recorded_for_the_default_value( + system, tmp_path +): + # 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]) + + 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 (directory / variable).iterdir(): + if file.suffix != ".npy": + file.unlink() + + +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) + 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_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. + 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 00000000..722f4a17 --- /dev/null +++ b/tests/core/test_restore_input_registry_property.py @@ -0,0 +1,218 @@ +"""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, 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 + 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, 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), +) + + +@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)) + 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)) + + 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 00000000..27c86d57 --- /dev/null +++ b/tests/core/variables/test_non_numeric_uprating_and_defined_for.py @@ -0,0 +1,412 @@ +"""``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) 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 + # ``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" + + +@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("condition", **NON_NUMERIC[type_name]), + _variable("amount", value_type=float), + ) + system.variables["amount"].defined_for = "condition" + simulation = _simulation(system, amount={"2012": 5.0}) + + 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 + + +@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. + 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") diff --git a/tests/fixtures/uprated_inputs.py b/tests/fixtures/uprated_inputs.py new file mode 100644 index 00000000..ce100731 --- /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]