diff --git a/changelog.d/branch-set-input-invalidation.added.md b/changelog.d/branch-set-input-invalidation.added.md new file mode 100644 index 00000000..b173c9e9 --- /dev/null +++ b/changelog.d/branch-set-input-invalidation.added.md @@ -0,0 +1 @@ +Simulation.drop_computed_arrays() deletes every value a simulation holds except inputs, for branches whose policy changes after they are created. diff --git a/changelog.d/branch-set-input-invalidation.changed.md b/changelog.d/branch-set-input-invalidation.changed.md new file mode 100644 index 00000000..d1dc7ef0 --- /dev/null +++ b/changelog.d/branch-set-input-invalidation.changed.md @@ -0,0 +1 @@ +Disk storage writes a new file for every store, so recalculating a value no longer changes what branches sharing the directory read; apply_reform keeps inputs by a flag on each stored array instead of replaying _user_input_keys; values a custom set_input handler calculates are no longer treated as inputs; dump_simulation records which values were inputs (inputs.txt), and restore_simulation restores only those as inputs; a calculation running when an input set on its simulation drops values is not cached (nor is anything in another simulation calculated from a value read before the change), and the outermost one in that simulation is run again from the new inputs until a run changes none (at most ten times); a branch stops reading macro-cache files once an input is set on it. diff --git a/changelog.d/branch-set-input-invalidation.fixed.md b/changelog.d/branch-set-input-invalidation.fixed.md new file mode 100644 index 00000000..146bf237 --- /dev/null +++ b/changelog.d/branch-set-input-invalidation.fixed.md @@ -0,0 +1 @@ +Simulation.set_input on a branch now drops the values the branch holds that may have been calculated from the value it replaces (every stored array carries a sequence number, and each simulation records what its values may have been calculated from), so a value the parent calculated before branching no longer shadows the branch's input. Simulation.derivative keeps inputs set after the simulation was built. An input a formula sets for the period it is calculating is no longer overwritten by the formula's result, and dump_simulation on a branch dumps the branch's own values (it read the default branch's, so the dump could not be restored). diff --git a/docs/_toc.yml b/docs/_toc.yml index c909f933..7835e4b6 100644 --- a/docs/_toc.yml +++ b/docs/_toc.yml @@ -7,6 +7,7 @@ parts: - caption: Using PolicyEngine Core chapters: - file: usage/simulation + - file: usage/branches - file: usage/country - file: usage/cli - file: usage/parameters diff --git a/docs/usage/branches.md b/docs/usage/branches.md new file mode 100644 index 00000000..df3e45d4 --- /dev/null +++ b/docs/usage/branches.md @@ -0,0 +1,199 @@ +# Branches + +A branch is a copy of a simulation that calculates under different inputs or +policy without changing the simulation it came from. Formulas use branches to +compare alternatives (for example, tax liability if itemizing), and marginal +tax rates use them to recalculate with slightly higher earnings. + +```python +branch = simulation.get_branch("itemizing") +branch.set_input("tax_unit_itemizes", 2026, itemizes) +tax_if_itemizing = branch.calculate("income_tax", 2026) +``` + +## What a branch starts with + +The first `get_branch(name)` call returns a new branch holding every value its +parent holds at that moment: inputs and calculated values alike. Later calls +with the same name return that branch as it is. What the parent stores after +the branch is created does not reach the branch, and what the branch stores +does not reach the parent. A branch of a branch reads its own values first, +then its parent's, then those of the simulation at the root. + +## Inputs set on a branch + +`set_input` on a branch stores the new value and drops every value the branch +holds that may have been calculated from the value it replaces. What the branch +calculates next therefore uses the input, as a simulation given the input +before calculating anything would, whatever its parent had calculated before +the branch was created. + +To decide what to drop, every stored value carries a sequence number from one +counter shared by the whole process, so a larger number means a later store. +Each simulation also keeps a store history: for each variable and period, the +number of the first store its values may have been calculated from, and for +each variable, the first time one of its values was uprated or carried over +from another period. The history records: + +- the simulation's own stores, including values a holder calculates but does + not keep (`variables_to_drop`, the cache blacklist), values read from the + macro cache, and the default a spiral returns; +- the number from which values may have been calculated from anything: the + first value read from the macro cache (what it was calculated from was + never calculated here) and the values restored from a dump (see below), so + an input for any variable drops them; +- a copy of its parent's history, taken when the branch is created; +- the history of any simulation its formulas calculate in, taken in each time + `calculate` there returns or raises (a formula that branches, sets an input + and calculates in the branch hands back values calculated there); a branch + calculated from a thread with no formula context hands its history to every + ancestor with a calculation running. + +When `set_input(variable, period, value)` is called on a branch, the branch +drops each value it holds, other than an input, whose number is at least the +earliest recorded store of the variable for a period that shares a day with +`period`, or the earliest recorded uprated or carried-over value of the +variable. It then forgets the records numbered from there on (its remaining +values were all stored earlier) and records again the inputs it keeps, unless a +formula is still running in the branch: such a formula may hold, in its own +variables, a value it read before the drop, so the records stay. A +calculation that was running then (one whose formula sets an input, say) may +have read the replaced value, so its result is kept neither in storage nor in +the macro cache. Neither is any result calculated from it, in any simulation: +each calculation notes, for every other simulation it got a value from +(directly or through the calculations it called), that simulation's count of +such input changes when the value's calculation began, and keeps its own +result only if none has changed since. So a parent formula calculating in a +branch whose formula calls back into the parent and then changes the branch's +input does not keep what it got, nor does a formula that read a branch and +then calculated there something that changed the branch's input. A +calculation whose call into another simulation settled before returning keeps +its result, and unrelated simulations and other threads keep caching. The +outermost calculation +running in the simulation whose input changed (a `calculate`, or a direct +`calculate_add`, whose terms run within it) then runs again from the new +inputs, inner calculations included, until a run changes no input, and keeps +that result, so uprating and carry-over find its period as they would had the +inputs come first. After ten reruns it stops: a formula that keeps changing +inputs returns its last result without keeping it, and a later uprating or +carry-over may then not find that period. The budget is per simulation whose +input changes, so calculations nested across several such simulations can +rerun more. An input stored for the very period being +calculated after the calculation began (by its own formula, say; under the +branch's name or any it reads, such as `default` through `Holder.set_input`) +is the result, as it would be had it been set first. + +A custom `set_input` handler that calculates values between its own stores +calculates them from inputs it has not yet replaced, so if it calculated +anything (through `calculate`, `calculate_add` or `_calculate`), the branch +drops again, by the same rule, once the handler returns or raises (the inputs +it stored before raising stay). + +If there is no such record, the branch drops nothing. That is the case when a +formula creates the branch while it is still calculating the variable the +branch overrides, unless another branch the formula created for the same +comparison already handed back a value calculated from that variable: then the +branch drops what came back and what was calculated after it. + +A value written into a holder's storage without going through `set_input` or a +calculation has no number; an input for a variable holding such a value drops +every calculated value. + +### Why this is sound + +A formula's result is stored after everything the formula read was stored or +recorded, so a calculated value carries a larger number than each value it was +calculated from, directly or through other calculated values. Every value a +simulation holds was calculated there, inherited from its parent, handed back +by another simulation's `calculate`, or restored from a dump. In the first +three cases, for every variable it was calculated from, the simulation's +history holds a record of that variable for an overlapping period numbered no +later than the value read (for a value summed or divided from other periods, +the record may be of those periods), unless it was calculated from a +macro-cache read; restored values and values calculated from a macro-cache +read are covered by the record that values from a number on may depend on +anything. So any value that depends on the overridden variable at an +overlapping period was stored after the earliest record that applies. Uprating +and carry-over read which periods hold values at all, which is why the first +uprated or carried-over value also counts. A result whose calculation was +running when an input changed is not kept unless it was calculated again from +the new inputs. + +The rule can drop more than it needs to (a value stored later that does not +depend on the input is calculated again) but not less, within these limits: + +- **Policy changes are not tracked.** A branch whose tax-benefit system or + parameters differ from its parent's still holds the parent's values + calculated under the parent's policy. Call `branch.drop_computed_arrays()` + after changing them. +- **Formulas must not write into arrays they read.** A formula that changes a + cached array in place (`array += x`) changes a value after its sequence + number was assigned, so values calculated from it may be dropped too late or + not at all. +- **Formulas should calculate, not inspect storage.** A formula that reads + `holder.get_known_periods()`, `simulation.get_array()` or another + simulation's storage directly, and acts on what it finds, depends on values + that are not recorded. +- **Unrelated simulations in other threads.** A formula that calculates in a + simulation other than its own branches, from a thread it starts without + copying its context (`contextvars.copy_context`, which `asyncio.to_thread` + does), does not take in that simulation's history. Its own branches are + covered: their history goes to every ancestor with a calculation running. +- **A branch a formula keeps between calls is a snapshot.** It holds what its + parent held when it was created, so inputs set on the parent afterwards do + not reach what the formula reads from it. +- **Inputs on the root simulation drop nothing.** `set_input` on a simulation + that is not a branch keeps its earlier behaviour: values it already + calculated stay. +- **Existing child branches keep their values.** An input set on a branch does + not reach branches already created from it. +- **Values already read stay read.** A formula that read a value from another + simulation and then calculates there a formula that changes that + simulation's input keeps what it read before the change. Its result is + returned but not kept, and it is not run again: running it again would + create and read its branches the same way. Likewise a formula that catches + an error raised after such a change returns its own fallback, unkept. + +With disk storage (`MemoryConfig`), every store writes a new file, named with +the sequence number and a token for the process (a forked child gets its own), +so a value recalculated in one simulation does not change a file another +simulation, or another process, still maps; files stay until the storage +directory is removed. `OnDiskStorage.restore` takes each key's most recently +written file, and on a timestamp tie the current process's own; between two +other processes' files written within one clock tick it cannot tell which came +last. + +A simulation dump (`dump_simulation`) holds the values the simulation reads +(on a branch, its own and those it inherited) and records which were inputs. +`restore_simulation` restores those as inputs and every other value as +calculated under one later number. The dump does not say what each value was +calculated from, nor what was read without being kept, so the restored +simulation records that values from that number on may depend on anything: +an input set for any variable on a branch of it drops all of them. A dump +written before inputs were recorded is restored with every value as an +input, as before, so such values never drop. With disk storage, a branch whose +name contains `_` cannot be dumped yet: `OnDiskStorage.get_known_periods` +splits its keys on every `_` (as on master). + +Two related behaviours: a branch stops reading macro-cache files once an input +is set on it, since they are keyed by branch name and period but not by inputs +(a file for its name may have been written by this branch before the input, or +by another simulation's branch of the same name), and a branch that read one +drops, on any input, everything calculated from the first read on; and +`requires_computation_after` is satisfied by a prerequisite requested before a +drop removed its values. + +## Dropping calculated values + +`simulation.drop_computed_arrays()` deletes every value the simulation holds +except inputs (the dataset or situation it was built from, values set with +`set_input` on it, and, for a branch, values set on the simulations it was +created from before it was created), and returns how many arrays it deleted. +Values a custom `set_input` handler calculates are not inputs. Use it on a +branch after changing its policy: + +```python +branch = simulation.get_branch("pre_reform_rules", clone_system=True) +branch.tax_benefit_system.parameters.gov.some_rate.update(period="2026", value=0.2) +branch.drop_computed_arrays() +``` diff --git a/policyengine_core/data_storage/in_memory_storage.py b/policyengine_core/data_storage/in_memory_storage.py index a1bddf76..61643a98 100644 --- a/policyengine_core/data_storage/in_memory_storage.py +++ b/policyengine_core/data_storage/in_memory_storage.py @@ -1,9 +1,13 @@ -from typing import Dict, Union +from typing import Dict, List, Optional, Set, Tuple, Union import numpy from numpy.typing import ArrayLike from policyengine_core import periods +from policyengine_core.data_storage.store_history import ( + advance_sequence_past, + next_sequence_number, +) from policyengine_core.periods import Period @@ -47,8 +51,23 @@ def __init__(self, is_eternal: bool): # with a copy the first time it is read. A key left here after code # outside this class empties ``_arrays`` costs one extra copy at most. self._shared = set() + # When each array was stored (see ``store_history``), and which were + # stored as inputs rather than calculated. Both describe the stored + # value, so ``clone`` copies them with the arrays. + self._sequence_numbers: Dict[str, int] = {} + self._input_keys: Set[str] = set() self.is_eternal = is_eternal + def __setstate__(self, state: dict) -> None: + # Storages pickled before stores were numbered have neither record. + state.setdefault("_sequence_numbers", {}) + state.setdefault("_input_keys", set()) + self.__dict__.update(state) + # Numbers from the process that pickled this storage must stay below + # those of stores made after unpickling it. + if self._sequence_numbers: + advance_sequence_past(max(self._sequence_numbers.values())) + def clone(self, share_arrays: bool = False) -> "InMemoryStorage": """Copy this storage. @@ -83,6 +102,8 @@ def clone(self, share_arrays: bool = False) -> "InMemoryStorage": clone._arrays[key] = array.copy() else: clone._arrays = {key: array.copy() for key, array in self._arrays.items()} + clone._sequence_numbers = dict(self._sequence_numbers) + clone._input_keys = set(self._input_keys) return clone def get(self, period: Period, branch_name: str = "default") -> ArrayLike: @@ -102,8 +123,19 @@ def get(self, period: Period, branch_name: str = "default") -> ArrayLike: return values def put( - self, value: ArrayLike, period: Period, branch_name: str = "default" + self, + value: ArrayLike, + period: Period, + branch_name: str = "default", + sequence_number: Optional[int] = None, + is_input: bool = False, ) -> None: + """Store ``value`` for ``period`` on ``branch_name``. + + ``sequence_number`` records when the value was stored (a new number + by default), and ``is_input`` whether it is an input rather than a + calculated value; see :meth:`drop_computed`. + """ if self.is_eternal: period = periods.period(periods.ETERNITY) period = periods.period(period) @@ -126,6 +158,54 @@ def put( key = f"{branch_name}:{period}" self._arrays[key] = value self._shared.discard(key) + self._sequence_numbers[key] = ( + next_sequence_number() if sequence_number is None else sequence_number + ) + if is_input: + self._input_keys.add(key) + else: + self._input_keys.discard(key) + + def drop_computed(self, *, since: Optional[int] = None) -> int: + """Delete stored values that are not inputs, and return how many. + + With ``since``, only values stored with that sequence number or a + later one are deleted (a value with no recorded number counts as + later). Inputs, which ``put`` received with ``is_input``, are kept + whatever their number. + """ + dropped = [ + key + for key in self._arrays + if key not in self._input_keys + and (since is None or self._sequence_numbers.get(key, since) >= since) + ] + for key in dropped: + del self._arrays[key] + self._sequence_numbers.pop(key, None) + self._shared.discard(key) + return len(dropped) + + def inputs_since(self, since: Optional[int] = None) -> List[Tuple[Period, int]]: + """The period and number of each input stored at ``since`` or later (or ever).""" + return [ + (periods.period(key.split(":", 1)[1]), self._sequence_numbers[key]) + for key in self._input_keys + if key in self._sequence_numbers + and (since is None or self._sequence_numbers[key] >= since) + ] + + def has_unnumbered_values(self) -> bool: + """Whether a value was stored without ``put`` (so without a number).""" + return any(key not in self._sequence_numbers for key in self._arrays) + + def _forget_deleted_keys(self) -> None: + self._sequence_numbers = { + key: number + for key, number in self._sequence_numbers.items() + if key in self._arrays + } + self._input_keys.intersection_update(self._arrays) def delete(self, period: Period = None, branch_name: str = "default") -> None: if period is None: @@ -138,6 +218,7 @@ def delete(self, period: Period = None, branch_name: str = "default") -> None: if not period_item.startswith(branch_prefix) } self._shared.intersection_update(self._arrays) + self._forget_deleted_keys() return if self.is_eternal: @@ -156,6 +237,7 @@ def delete(self, period: Period = None, branch_name: str = "default") -> None: ) } self._shared.intersection_update(self._arrays) + self._forget_deleted_keys() def get_known_periods(self) -> list: # Split on the first colon only: an anchored period's string form diff --git a/policyengine_core/data_storage/on_disk_storage.py b/policyengine_core/data_storage/on_disk_storage.py index 3563c805..4330e6d9 100644 --- a/policyengine_core/data_storage/on_disk_storage.py +++ b/policyengine_core/data_storage/on_disk_storage.py @@ -1,13 +1,33 @@ import os import shutil +import uuid +from typing import Dict, List, Optional, Set, Tuple import numpy from numpy.typing import ArrayLike from policyengine_core import periods +from policyengine_core.data_storage.store_history import ( + advance_sequence_past, + next_sequence_number, +) from policyengine_core.enums import EnumArray from policyengine_core.periods import Period +# Distinguishes the files this process writes from other processes' files in +# the same directory, whose sequence numbers may repeat this process's. +_PROCESS_TOKEN = uuid.uuid4().hex[:12] + + +def _new_process_token() -> None: + global _PROCESS_TOKEN + _PROCESS_TOKEN = uuid.uuid4().hex[:12] + + +# A forked child inherits the parent's token and counter: give it its own token. +if hasattr(os, "register_at_fork"): + os.register_at_fork(after_in_child=_new_process_token) + class OnDiskStorage: """ @@ -22,19 +42,33 @@ def __init__( ): self._files = {} self._enums = {} + # As in ``InMemoryStorage``: when each file was stored, and which were + # stored as inputs. + self._sequence_numbers: Dict[str, int] = {} + self._input_keys: Set[str] = set() self.is_eternal = is_eternal self.preserve_storage_dir = preserve_storage_dir self.storage_dir = storage_dir + def __setstate__(self, state: dict) -> None: + # Storages pickled before stores were numbered have neither record. + state.setdefault("_sequence_numbers", {}) + state.setdefault("_input_keys", set()) + self.__dict__.update(state) + # Numbers from the process that pickled this storage must stay below + # those of stores made after unpickling it. + if self._sequence_numbers: + advance_sequence_past(max(self._sequence_numbers.values())) + def clone(self) -> "OnDiskStorage": """Create a private metadata view over this storage directory. The file and enum mappings are copied so deleting or rewiring entries - through the clone does not mutate the source storage. The underlying - ``.npy`` files remain shared: writing the same ``{branch}_{period}`` - key from two views targets the same path and can overwrite the file. - Clones retain the original cleanup owner so the shared directory stays - alive, but never own cleanup themselves. + through the clone does not mutate the source storage. The directory + is shared, but every store writes a new file, so writing a key + through one view leaves the file another view maps. Files stay until + the directory is removed. Clones retain the original cleanup owner so + the shared directory stays alive, but never own cleanup themselves. """ clone = OnDiskStorage( self.storage_dir, @@ -43,6 +77,8 @@ def clone(self) -> "OnDiskStorage": ) clone._files = self._files.copy() clone._enums = self._enums.copy() + clone._sequence_numbers = dict(self._sequence_numbers) + clone._input_keys = set(self._input_keys) clone._storage_dir_owner = getattr(self, "_storage_dir_owner", self) return clone @@ -64,19 +100,84 @@ def get(self, period: Period, branch_name: str = "default") -> ArrayLike: return self._decode_file(values) def put( - self, value: ArrayLike, period: Period, branch_name: str = "default" + self, + value: ArrayLike, + period: Period, + branch_name: str = "default", + sequence_number: Optional[int] = None, + is_input: bool = False, ) -> None: + """Store ``value`` for ``period`` on ``branch_name``. + + ``sequence_number`` and ``is_input`` are as in + :meth:`InMemoryStorage.put`. + """ if self.is_eternal: period = periods.period(periods.ETERNITY) period = periods.period(period) filename = f"{branch_name}_{period}" - path = os.path.join(self.storage_dir, filename) + ".npy" + if sequence_number is None: + sequence_number = next_sequence_number() + # A new file for every store: clones share this directory and may + # still map an earlier file for the same key. The process token keeps + # files from processes whose counters restarted apart. + path = ( + os.path.join( + self.storage_dir, f"{filename}.{_PROCESS_TOKEN}.{sequence_number}" + ) + + ".npy" + ) if isinstance(value, EnumArray): self._enums[path] = value.possible_values value = value.view(numpy.ndarray) numpy.save(path, value) self._files[filename] = path + self._sequence_numbers[filename] = sequence_number + if is_input: + self._input_keys.add(filename) + else: + self._input_keys.discard(filename) + + def drop_computed(self, *, since: Optional[int] = None) -> int: + """Forget stored values that are not inputs, and return how many. + + As :meth:`InMemoryStorage.drop_computed`. The files stay on disk + (other views of this directory may still use them) until the storage + directory is removed. + """ + dropped = [ + key + for key in self._files + if key not in self._input_keys + and (since is None or self._sequence_numbers.get(key, since) >= since) + ] + for key in dropped: + del self._files[key] + self._sequence_numbers.pop(key, None) + return len(dropped) + + def inputs_since(self, since: Optional[int] = None) -> List[Tuple[Period, int]]: + """As :meth:`InMemoryStorage.inputs_since`.""" + # Period strings contain no "_"; branch names may. + return [ + (periods.period(key.rsplit("_", 1)[1]), self._sequence_numbers[key]) + for key in self._input_keys + if key in self._sequence_numbers + and (since is None or self._sequence_numbers[key] >= since) + ] + + def has_unnumbered_values(self) -> bool: + """As :meth:`InMemoryStorage.has_unnumbered_values`.""" + return any(key not in self._sequence_numbers for key in self._files) + + def _forget_deleted_keys(self) -> None: + self._sequence_numbers = { + key: number + for key, number in self._sequence_numbers.items() + if key in self._files + } + self._input_keys.intersection_update(self._files) def delete(self, period: Period = None, branch_name: str = "default") -> None: if period is None: @@ -89,6 +190,7 @@ def delete(self, period: Period = None, branch_name: str = "default") -> None: for period_item, value in self._files.items() if not period_item.startswith(branch_prefix) } + self._forget_deleted_keys() return if self.is_eternal: @@ -101,6 +203,7 @@ def delete(self, period: Period = None, branch_name: str = "default") -> None: for period_item, value in self._files.items() if not period_item == f"{branch_name}_{period}" } + self._forget_deleted_keys() def get_known_periods(self) -> list: return list([periods.period(x.split("_")[1]) for x in self._files.keys()]) @@ -113,13 +216,36 @@ def get_known_branch_periods(self) -> list: def restore(self) -> None: self._files = files = {} + self._sequence_numbers = {} + self._input_keys = set() # Restore self._files from content of storage_dir. + latest = {} for filename in os.listdir(self.storage_dir): if not filename.endswith(".npy"): continue path = os.path.join(self.storage_dir, filename) filename_core = filename.rsplit(".", 1)[0] - files[filename_core] = path + # Files are named "...npy" + # (each store writes a new file); older dumps are ".npy". + # Keep each key's most recently written file: sequence numbers + # restart in every process, so they only order one process's files. + key, token, number = filename_core, None, 0 + stem, _, last = filename_core.rpartition(".") + if stem and last.isdigit(): + key, number = stem, int(last) + stem, _, middle = key.rpartition(".") + if ( + stem + and len(middle) == 12 + and all(c in "0123456789abcdef" for c in middle) + ): + key, token = stem, middle + # On a timestamp tie (a coarse filesystem clock), this process's + # own file is the later one; numbers only order one process's. + order = (os.stat(path).st_mtime_ns, token == _PROCESS_TOKEN, number) + if key not in latest or order > latest[key]: + latest[key] = order + files[key] = path def __del__(self) -> None: if self.preserve_storage_dir: diff --git a/policyengine_core/data_storage/store_history.py b/policyengine_core/data_storage/store_history.py new file mode 100644 index 00000000..2640cef2 --- /dev/null +++ b/policyengine_core/data_storage/store_history.py @@ -0,0 +1,237 @@ +"""The order in which a simulation's values were stored. + +Every array a holder stores gets a sequence number from one process-wide +counter, so a larger number means a later store. A value calculated by a +formula is stored after everything the formula read, so it carries a larger +number than each stored value it was calculated from, directly or through +other calculated values. + +``Simulation.set_input`` on a branch uses this to drop only the values that +can depend on the input it replaces: see +:meth:`policyengine_core.simulations.Simulation.set_input`. +""" + +import itertools +import weakref +from typing import Dict, List, Optional, Tuple + +from policyengine_core import periods +from policyengine_core.periods import Period + +_sequence = itertools.count(1) + + +def next_sequence_number() -> int: + """Return a sequence number larger than every one returned before.""" + return next(_sequence) + + +def advance_sequence_past(number: int) -> None: + """Make every later sequence number larger than ``number``. + + Numbers come from a counter in each process, so values unpickled from + another process may carry larger numbers than this process has handed + out; stores made after them must still be numbered later. + """ + global _sequence + _sequence = itertools.count(max(next(_sequence), number + 1)) + + +def periods_overlap(first: Period, second: Period) -> bool: + """Whether two periods share at least one day.""" + if periods.ETERNITY in (first.unit, second.unit): + return True + return first.start <= second.stop and second.start <= first.stop + + +class StoreHistory: + """What the values a simulation holds may have been calculated from. + + Each simulation has its own history. For each variable it records: + + - for each period, the sequence number of the first value stored or + calculated for that period that a value the simulation holds may have + been calculated from, including values a holder calculates but does + not keep (``variables_to_drop``, the cache blacklist), whose readers + carry later numbers all the same; + - the sequence number of the first time such a value of the variable was + derived from its values for other periods: uprated or carried over (or + given the default because no period it could be uprated or carried over + from holds a value). Such a value depends on which other periods hold + values, so an input set for any period can change it. + + It also records the first sequence number from which held values may + have been calculated from anything (:meth:`record_unknown_sources`): + values restored from a dump, which does not say what each was + calculated from, and values calculated from a macro-cache read, whose + sources were never calculated here. An input for any variable can + change those. + + The simulation's own stores are recorded as they happen. A branch starts + with a copy of its parent's history, as it starts with a copy of its + parent's values. When a formula running in one simulation calculates a + value in another (a branch it created, for instance), the caller merges + the other's history into its own, since what the caller stores next may + be calculated from what it got back. A drop removes the records at or + after its sequence number (see :meth:`prune`). + """ + + def __init__(self): + self._first_stored: Dict[str, Dict[Period, int]] = {} + self._first_derived: Dict[str, int] = {} + # Values numbered from here on may have been calculated from anything. + self._unknown_sources_since: Optional[int] = None + # Every change to the records since this history was created or last + # pruned, in order, so a merge takes in only what changed since it + # last read this history (``period`` is None for a derived record, + # and ``variable_name`` too for an unknown-sources record). + self._journal: List[Tuple[Optional[str], Optional[Period], int]] = [] + self._generation = 0 # changes when a prune empties the journal + # For each history merged into this one, the generation and journal + # length read last. Keyed weakly: an entry goes with the history it + # describes (a formula's temporary branch, say). + self._merged: "weakref.WeakKeyDictionary[StoreHistory, Tuple[int, int]]" = ( + weakref.WeakKeyDictionary() + ) + + def __getstate__(self) -> dict: + state = dict(self.__dict__) + # Weak references do not pickle; merging again only repeats work. + state["_merged"] = {} + return state + + def __setstate__(self, state: dict) -> None: + state.setdefault("_journal", []) + state.setdefault("_generation", 0) + state.setdefault("_unknown_sources_since", None) + self.__dict__.update(state) + self._merged = weakref.WeakKeyDictionary() + numbers = [ + number + for stored in self._first_stored.values() + for number in stored.values() + ] + list(self._first_derived.values()) + if self._unknown_sources_since is not None: + numbers.append(self._unknown_sources_since) + if numbers: + advance_sequence_past(max(numbers)) + + def copy(self) -> "StoreHistory": + new = StoreHistory() + new._first_stored = { + variable_name: dict(stored) + for variable_name, stored in self._first_stored.items() + } + new._first_derived = dict(self._first_derived) + new._unknown_sources_since = self._unknown_sources_since + new._merged = weakref.WeakKeyDictionary(self._merged) + return new + + def record_store(self, variable_name: str, period: Period, sequence_number: int): + stored = self._first_stored.setdefault(variable_name, {}) + recorded = stored.get(period) + if recorded is None or sequence_number < recorded: + stored[period] = sequence_number + self._journal.append((variable_name, period, sequence_number)) + + def record_derived(self, variable_name: str, sequence_number: int): + recorded = self._first_derived.get(variable_name) + if recorded is None or sequence_number < recorded: + self._first_derived[variable_name] = sequence_number + self._journal.append((variable_name, None, sequence_number)) + + def record_unknown_sources(self, sequence_number: int): + """Values numbered ``sequence_number`` or later may depend on any input.""" + recorded = self._unknown_sources_since + if recorded is None or sequence_number < recorded: + self._unknown_sources_since = sequence_number + self._journal.append((None, None, sequence_number)) + + def merge(self, other: "StoreHistory") -> None: + """Take in ``other``'s records, keeping the earlier number of each.""" + if other is self: + return + # Note how far ``other`` had got before reading it: anything it + # records meanwhile (another thread) is read next time. + generation, end = other._generation, len(other._journal) + read = self._merged.get(other) + if read is not None and read[0] == generation: + changes = other._journal[read[1] : end] + else: + changes = [ + (variable_name, period, sequence_number) + for variable_name, stored in list(other._first_stored.items()) + for period, sequence_number in list(stored.items()) + ] + [ + (variable_name, None, sequence_number) + for variable_name, sequence_number in list(other._first_derived.items()) + ] + unknown_sources_since = other._unknown_sources_since + if unknown_sources_since is not None: + changes.append((None, None, unknown_sources_since)) + for variable_name, period, sequence_number in changes: + if variable_name is None: + self.record_unknown_sources(sequence_number) + elif period is None: + self.record_derived(variable_name, sequence_number) + else: + self.record_store(variable_name, period, sequence_number) + self._merged[other] = (generation, end) + + def prune(self, since: Optional[int] = None) -> None: + """Forget the records numbered ``since`` or later (all, without ``since``). + + Called after the simulation drops every value it holds, other than + inputs, numbered ``since`` or later: nothing it still holds was + calculated from what those records describe. The caller records again + the inputs it keeps. + """ + self._first_stored = { + variable_name: kept + for variable_name, stored in self._first_stored.items() + if ( + kept := { + period: sequence_number + for period, sequence_number in stored.items() + if since is not None and sequence_number < since + } + ) + } + self._first_derived = { + variable_name: sequence_number + for variable_name, sequence_number in self._first_derived.items() + if since is not None and sequence_number < since + } + if since is None or ( + self._unknown_sources_since is not None + and self._unknown_sources_since >= since + ): + self._unknown_sources_since = None + self._journal = [] + self._generation += 1 + # Merging a history again must restore what was forgotten here. + self._merged = weakref.WeakKeyDictionary() + + def earliest_dependency(self, variable_name: str, period: Period) -> Optional[int]: + """The first sequence number from which a value may depend on ``variable_name`` at ``period``. + + This is the earliest recorded store of the variable for any period + that shares a day with ``period``, the earliest recorded derivation + of one of its values from other periods, or the first number from + which values may depend on anything, whichever came first. + ``None`` means no value the simulation holds can depend on the + variable's value at ``period``. + """ + candidates = [ + sequence_number + for stored_period, sequence_number in self._first_stored.get( + variable_name, {} + ).items() + if periods_overlap(stored_period, period) + ] + derived = self._first_derived.get(variable_name) + if derived is not None: + candidates.append(derived) + if self._unknown_sources_since is not None: + candidates.append(self._unknown_sources_since) + return min(candidates) if candidates else None diff --git a/policyengine_core/holders/holder.py b/policyengine_core/holders/holder.py index 11a8ae50..f5a13ee6 100644 --- a/policyengine_core/holders/holder.py +++ b/policyengine_core/holders/holder.py @@ -1,6 +1,6 @@ import os import warnings -from typing import TYPE_CHECKING, Any, List, Tuple +from typing import TYPE_CHECKING, Any, List, Optional, Tuple import numpy import psutil @@ -8,6 +8,7 @@ from policyengine_core import commons, periods, tools from policyengine_core.data_storage import InMemoryStorage, OnDiskStorage +from policyengine_core.data_storage.store_history import next_sequence_number from policyengine_core.enums import Enum from policyengine_core.errors import PeriodMismatchError from policyengine_core.periods import Period @@ -106,6 +107,53 @@ def delete_arrays( if self._disk_storage: self._disk_storage.delete(period, branch_name) + def _drop_computed(self, since: Optional[int] = None) -> int: + """Delete every stored value that is not an input, and return how many. + + With ``since``, only values stored with that sequence number or a + later one are deleted. Inputs are values stored through + :meth:`set_input`, on any branch. + """ + dropped = self._memory_storage.drop_computed(since=since) + if self._disk_storage is not None: + dropped += self._disk_storage.drop_computed(since=since) + return dropped + + def _record_inputs(self, since: Optional[int] = None) -> None: + """Record in the simulation's history each input stored at ``since`` or later.""" + for storage in (self._memory_storage, self._disk_storage): + if storage is not None: + for period, sequence_number in storage.inputs_since(since): + self._record_store(period, sequence_number) + + def _is_input(self, period: Period, branch_name: str = "default") -> bool: + """Whether the value stored for ``period`` on ``branch_name`` is an input.""" + if self.variable.definition_period == periods.ETERNITY: + period = periods.period(periods.ETERNITY) + period = periods.period(period) + return f"{branch_name}:{period}" in self._memory_storage._input_keys or ( + self._disk_storage is not None + and f"{branch_name}_{period}" in self._disk_storage._input_keys + ) + + def _stored_sequence_number( + self, period: Period, branch_name: str = "default" + ) -> Optional[int]: + """The sequence number of the value stored for ``period`` on ``branch_name``.""" + if self.variable.definition_period == periods.ETERNITY: + period = periods.period(periods.ETERNITY) + period = periods.period(period) + number = self._memory_storage._sequence_numbers.get(f"{branch_name}:{period}") + if number is None and self._disk_storage is not None: + number = self._disk_storage._sequence_numbers.get(f"{branch_name}_{period}") + return number + + def _has_unnumbered_values(self) -> bool: + return self._memory_storage.has_unnumbered_values() or ( + self._disk_storage is not None + and self._disk_storage.has_unnumbered_values() + ) + def _get_array_from_storage( self, period: Period, branch_name: str = "default" ) -> ArrayLike: @@ -154,6 +202,24 @@ def get_array(self, period: Period, branch_name: str = "default") -> ArrayLike: if default_value is not None: return default_value + def _branch_holding(self, period: Period, branch_name: str) -> Optional[str]: + """The branch whose value :meth:`get_array` returns for ``period``, if any. + + As there: ``branch_name`` itself, else its nearest ancestor, else + ``default``. + """ + names = [branch_name] + if branch_name != "default": + parent = getattr(self.simulation, "parent_branch", None) + while parent is not None: + names.append(parent.branch_name) + parent = getattr(parent, "parent_branch", None) + names.append("default") + for name in names: + if self._get_array_from_storage(period, name) is not None: + return name + return None + def get_memory_usage(self) -> dict: """ Get data about the virtual memory usage of the holder. @@ -269,6 +335,12 @@ def set_input( self._raise_if_input_contains_nan(numpy.asarray(array)) simulation = getattr(self, "simulation", None) if simulation is not None: + # On a branch, drop what may have been calculated from the value + # this input replaces (see ``Simulation.set_input``). + simulation._drop_values_that_may_depend_on(self.variable.name, period) + # A calculation running meanwhile looks for inputs set for its own + # period (``Simulation._cache_result``). + simulation._inputs_set = getattr(simulation, "_inputs_set", 0) + 1 if not hasattr(simulation, "_user_input_keys"): simulation._user_input_keys = set() if not hasattr(simulation, "_user_input_contexts"): @@ -279,7 +351,20 @@ def set_input( self.variable.set_input and period.unit != self.variable.definition_period ): - return self.variable.set_input(self, period, array) + started = getattr(simulation, "_calculations_started", 0) + try: + return self.variable.set_input(self, period, array) + finally: + if ( + simulation is not None + and getattr(simulation, "_calculations_started", 0) != started + ): + # The handler calculated, maybe between its stores, + # from inputs it had not yet replaced (also if it + # then failed: the inputs it stored stay). + simulation._drop_values_that_may_depend_on( + self.variable.name, period + ) return self._set(period, array, branch_name, validate_nan=True) finally: if simulation is not None: @@ -347,10 +432,17 @@ def _set( value: ArrayLike, branch_name: str = "default", validate_nan: bool = False, + is_input: Optional[bool] = None, + sequence_number: Optional[int] = None, ) -> None: simulation = getattr(self, "simulation", None) user_input_contexts = getattr(simulation, "_user_input_contexts", None) - if user_input_contexts and branch_name == "default": + # A value is an input when stored while ``set_input`` runs, unless the + # caller says otherwise: ``put_in_cache`` stores calculated values, + # including those a custom ``set_input`` handler calculates. + if is_input is None: + is_input = bool(user_input_contexts) + if is_input and user_input_contexts and branch_name == "default": branch_name = user_input_contexts[-1] value = self._to_array(value, validate_nan=validate_nan) if self.variable.definition_period != periods.ETERNITY: @@ -366,11 +458,20 @@ def _set( >= self.simulation.memory_config.max_memory_occupation_pc ) - if should_store_on_disk: - self._disk_storage.put(value, period, branch_name) - else: - self._memory_storage.put(value, period, branch_name) - if user_input_contexts: + if sequence_number is None or is_input: + # Inputs are always numbered when stored: a calculation running + # meanwhile looks for inputs stored after it began. + sequence_number = next_sequence_number() + storage = self._disk_storage if should_store_on_disk else self._memory_storage + storage.put( + value, + period, + branch_name, + sequence_number=sequence_number, + is_input=is_input, + ) + self._record_store(period, sequence_number) + if is_input and simulation is not None: if not hasattr(simulation, "_user_input_keys"): simulation._user_input_keys = set() simulation._user_input_keys.add((self.variable.name, branch_name, period)) @@ -378,17 +479,26 @@ def _set( def put_in_cache( self, value: ArrayLike, period: Period, branch_name: str = "default" ) -> None: - if self._do_not_store: - return - - if ( + if self._do_not_store or ( self.simulation.opt_out_cache and self.simulation.tax_benefit_system.cache_blacklist and self.variable.name in self.simulation.tax_benefit_system.cache_blacklist ): + # Not kept, but whatever reads it is stored after it all the same. + self._record_store(period, next_sequence_number()) return - self._set(period, value, branch_name) + self._set(period, value, branch_name, is_input=False) + + def _record_store(self, period: Period, sequence_number: int) -> None: + simulation = getattr(self, "simulation", None) + if simulation is None: + return + if self.variable.definition_period == periods.ETERNITY: + period = periods.period(periods.ETERNITY) + simulation._get_store_history().record_store( + self.variable.name, periods.period(period), sequence_number + ) def default_array(self) -> ArrayLike: """ diff --git a/policyengine_core/simulations/simulation.py b/policyengine_core/simulations/simulation.py index 2652d972..b9143d08 100644 --- a/policyengine_core/simulations/simulation.py +++ b/policyengine_core/simulations/simulation.py @@ -1,8 +1,8 @@ import hashlib import tempfile +from contextlib import contextmanager from contextvars import ContextVar -from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Type, Union +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Type, Union import numpy as np import pandas as pd @@ -12,6 +12,10 @@ from policyengine_core import commons, periods from policyengine_core.data.dataset import Dataset +from policyengine_core.data_storage.store_history import ( + StoreHistory, + next_sequence_number, +) from policyengine_core.entities.entity import Entity from policyengine_core.enums import Enum, EnumArray from policyengine_core.errors import CycleError, SpiralError @@ -117,16 +121,76 @@ def _uprating_index_value(parameter, instant) -> Optional[float]: ) -@dataclass(frozen=True) -class PreservedUserInput: - variable_name: str - branch_name: str - period: Period - value: object - storage: str - disk_key: Optional[str] = None - disk_file: Optional[str] = None - disk_enum: object = None +# The simulation whose formula is running. When that formula calculates a +# value in another simulation (a branch it created, say), what the formula +# returns, and so what its simulation stores, may be calculated from the other +# simulation's values: ``calculate`` then merges the other's store history into +# this one's (see ``StoreHistory``). +_formula_simulation: ContextVar[Optional["Simulation"]] = ContextVar( + "_formula_simulation", default=None +) + +# How many times ``calculate`` runs a calculation again when an input set on +# the simulation while it ran (by a formula, say) dropped values. A formula +# that changes an input on every run is kept from looping; its last result is +# returned but not kept. +_RERUNS_AFTER_INPUT_CHANGE = 10 + +# The calculations running in this context, innermost last, as (simulation, +# frame) pairs. A frame notes ``start``, its simulation's ``_input_epoch`` when +# it (last) began, and ``reads``: for each other simulation it got a value +# from, directly or through the calculations it called, that simulation's +# ``_input_epoch`` when the calculation that produced the value began. If any +# of those has changed since (a drop while calculations ran there), the value +# may come from a replaced one, and so may what the frame calculates from it, +# which is then not kept (``_frame_is_stale``). +_calculation_frames: ContextVar[tuple] = ContextVar("_calculation_frames", default=()) + + +def _new_frame(simulation: "Simulation") -> dict: + return {"start": getattr(simulation, "_input_epoch", 0), "reads": {}} + + +def _frame_is_stale(simulation: "Simulation", frame: dict) -> bool: + """Whether what the frame calculates may come from a value since replaced.""" + return getattr(simulation, "_input_epoch", 0) != frame["start"] or any( + getattr(other, "_input_epoch", 0) != epoch + for other, epoch in frame["reads"].items() + ) + + +def _hand_to_caller(frames: tuple, simulation: "Simulation", frame: dict) -> None: + """Pass what a calculation in ``simulation`` read on to the calling frame. + + The calculation's own simulation counts at its ``start``: if an input + there changed while it ran, its caller's result is stale too. + """ + if not frames: + return + caller_simulation, caller = frames[-1] + reads = caller["reads"] + for other, epoch in frame["reads"].items(): + if other is not caller_simulation and other not in reads: + reads[other] = epoch + if simulation is not caller_simulation and simulation not in reads: + reads[simulation] = frame["start"] + + +@contextmanager +def _calculation_frame(simulation: "Simulation"): + """Run a calculation in a new frame of ``_calculation_frames``. + + Whatever way it ends (a value, an early default, an error), its caller + learns what it read and whether that may since have been replaced. + """ + frames = _calculation_frames.get() + frame = _new_frame(simulation) + token = _calculation_frames.set(frames + ((simulation, frame),)) + try: + yield frame + finally: + _calculation_frames.reset(token) + _hand_to_caller(frames, simulation, frame) class Simulation: @@ -210,6 +274,9 @@ def __init__( # post-``apply_reform`` cache wipe would also wipe the dataset the # simulation was loaded from. self._user_input_keys: set[tuple[str, str, Period]] = set() + # What this simulation's values may have been calculated from; each + # branch starts with a copy (see ``set_input``). + self._store_history = StoreHistory() self.debug: bool = False self.trace: bool = trace self.tracer: SimpleTracer = SimpleTracer() if not trace else FullTracer() @@ -334,83 +401,15 @@ def _invalidate_all_caches(self) -> None: Called after ``apply_reform`` and any other operation that changes the tax-benefit system underneath an already-calculated simulation. - Every (variable, branch, period) that was populated via - ``set_input`` is preserved — those are source data, not stale - formula output — so a structural reform applied after dataset - load doesn't silently discard the dataset. Everything else - (formula outputs, cached short-path results, on-disk caches) is - wiped so the next ``calculate`` recomputes under the new - tax-benefit system. + Every value stored through ``set_input`` is preserved — those are + source data, not stale formula output — so a structural reform + applied after dataset load doesn't silently discard the dataset. + Everything else (formula outputs, cached short-path results, on-disk + caches) is wiped, here and in every branch, so the next ``calculate`` + recomputes under the new tax-benefit system. """ - self._fast_cache = {} self.invalidated_caches = set() - # Snapshot user-provided inputs before wiping so they can be - # replayed into the fresh storage. Use the storage API instead of - # hand-building keys, since ETERNITY variables canonicalize every - # period to the single ETERNITY storage key. - preserved: list[PreservedUserInput] = [] - user_input_keys = getattr(self, "_user_input_keys", None) or set() - for variable_name, branch_name, period in user_input_keys: - holder = self.get_holder(variable_name) - stored_value = holder._memory_storage.get(period, branch_name) - if stored_value is not None: - preserved.append( - PreservedUserInput( - variable_name=variable_name, - branch_name=branch_name, - period=period, - value=stored_value, - storage="memory", - ) - ) - continue - if holder._disk_storage is not None: - disk_period = ( - periods.period(periods.ETERNITY) - if holder._disk_storage.is_eternal - else periods.period(period) - ) - disk_key = f"{branch_name}_{disk_period}" - disk_file = holder._disk_storage._files.get(disk_key) - if disk_file is not None: - preserved.append( - PreservedUserInput( - variable_name=variable_name, - branch_name=branch_name, - period=period, - value=None, - storage="disk", - disk_key=disk_key, - disk_file=disk_file, - disk_enum=holder._disk_storage._enums.get(disk_file), - ) - ) - # Iterate only over holders that already exist on each population — - # lazy-creating a holder for every variable in the tax-benefit - # system (thousands in policyengine-us) inflated the cost of - # ``apply_reform`` from milliseconds to seconds and broke the - # YAML full-suite on downstream repos. Untouched variables have - # no holder and therefore nothing to wipe. - for population in self.populations.values(): - for holder in population._holders.values(): - holder._memory_storage._arrays = {} - if holder._disk_storage is not None: - holder._disk_storage._files = {} - # Replay preserved user inputs so ``calculate`` still sees them. - for user_input in preserved: - holder = self.get_holder(user_input.variable_name) - if user_input.storage == "disk" and holder._disk_storage is not None: - holder._disk_storage._files[user_input.disk_key] = user_input.disk_file - if user_input.disk_enum is not None: - holder._disk_storage._enums[user_input.disk_file] = ( - user_input.disk_enum - ) - else: - holder._memory_storage.put( - user_input.value, - user_input.period, - user_input.branch_name, - ) + self._drop_computed() for branch in self.branches.values(): branch._invalidate_all_caches() @@ -661,6 +660,12 @@ def calculate( if _fast_cache is not None: _cached = _fast_cache.get(_fast_key) if _cached is not None: + self._share_store_history_with_caller() + frames = _calculation_frames.get() + if frames and frames[-1][0] is not self: + frames[-1][1]["reads"].setdefault( + self, getattr(self, "_input_epoch", 0) + ) return _cached self.tracer.record_calculation_start(variable_name, period, self.branch_name) @@ -670,8 +675,33 @@ def calculate( # check_formula_determinism), so there is nothing to make reproducible # here. + # Formulas running in this simulation may hold values they read in + # local variables; a drop meanwhile must not forget what they came + # from (see ``_drop_computed``). + # Only the outermost calculation in this simulation runs again after + # an input change here: running it again runs the inner ones too. + outermost = not getattr(self, "_calculations_in_flight", 0) + self._calculations_in_flight = getattr(self, "_calculations_in_flight", 0) + 1 + frame_context = _calculation_frame(self) + frame = frame_context.__enter__() try: + input_epoch = getattr(self, "_input_epoch", 0) result = self._calculate(variable_name, period) + for _ in range(_RERUNS_AFTER_INPUT_CHANGE if outermost else 0): + if input_epoch == getattr(self, "_input_epoch", 0): + break + # An input set here while it ran (by a formula, say) dropped + # values, so the result was not kept (``_cache_result``). + # Calculate it again from the new inputs, until a run changes + # none, and keep that result, as a simulation given those + # inputs first would: later uprating and carry-over look for + # the periods held. + input_epoch = getattr(self, "_input_epoch", 0) + frame.update(_new_frame(self)) + result = self._calculate(variable_name, period) + # Satisfies ``requires_computation_after`` from now on, even if a + # branch input later drops the values. + self._get_requested_variables().add(variable_name) if isinstance(result, EnumArray) and decode_enums: result = result.decode_to_str() self.tracer.record_calculation_result(result) @@ -682,8 +712,15 @@ def calculate( result = self.map_result(result, source_entity, map_to) return result finally: - self.tracer.record_calculation_end() - self.purge_cache_of_invalid_values() + try: + # Also when the calculation fails: whether it fails can depend + # on what it read, and a calling formula may catch the error. + self._share_store_history_with_caller() + finally: + self._calculations_in_flight -= 1 + frame_context.__exit__(None, None, None) + self.tracer.record_calculation_end() + self.purge_cache_of_invalid_values() def map_result( self, @@ -795,6 +832,9 @@ def _calculate(self, variable_name: str, period: Period = None) -> ArrayLike: """ if variable_name not in self.tax_benefit_system.variables: raise ValueError(f"Variable {variable_name} does not exist.") + input_state = self._calculation_start() + # Lets a custom ``set_input`` handler tell whether it calculated. + self._calculations_started = getattr(self, "_calculations_started", 0) + 1 population = self.get_variable_population(variable_name) holder = population.get_holder(variable_name) variable = self.tax_benefit_system.get_variable( @@ -837,6 +877,13 @@ def _calculate(self, variable_name: str, period: Period = None) -> ArrayLike: value = smc.get_cache_value(cache_path) if value is not None: + # Served without being stored: record it, so values + # calculated from it count as later (see ``set_input``). + # What it was calculated from was never calculated here, + # so values from here on may depend on any input. + sequence_number = next_sequence_number() + holder._record_store(period, sequence_number) + self._get_store_history().record_unknown_sources(sequence_number) return value if variable.requires_computation_after is not None: @@ -847,7 +894,16 @@ def _calculate(self, variable_name: str, period: Period = None) -> ArrayLike: required_is_known_periods = self.get_holder( variable.requires_computation_after ).get_known_periods() - if (not variable_in_stack) and (not len(required_is_known_periods) > 0): + # A branch input may have dropped the prerequisite's values; it + # was still requested. + required_was_requested = ( + variable.requires_computation_after in self._get_requested_variables() + ) + if ( + (not variable_in_stack) + and (not len(required_is_known_periods) > 0) + and not required_was_requested + ): raise ValueError( f"Variable {variable_name} requires {variable.requires_computation_after} to be requested first. That variable is known in: {required_is_known_periods}. The full stack is: {variables_in_stack}. {variable_in_stack, len(required_is_known_periods) > 0}" ) @@ -867,7 +923,7 @@ def _calculate(self, variable_name: str, period: Period = None) -> ArrayLike: values = self.calculate_divide(variable_name, period) if alternate_period_handling: - if is_cache_available: + if is_cache_available and self._may_keep(input_state): smc.set_cache_value(cache_path, values) return values @@ -881,18 +937,31 @@ def _calculate(self, variable_name: str, period: Period = None) -> ArrayLike: if np.all(~mask): array = holder.default_array() array = self._cast_formula_result(array, variable) - holder.put_in_cache(array, period, self.branch_name) - return array + return self._cache_result(holder, array, period, input_state) array = None # First, try to run a formula try: self._check_for_cycle(variable.name, period) - array = self._run_formula(variable, population, period) + token = _formula_simulation.set(self) + try: + array = self._run_formula(variable, population, period) + finally: + _formula_simulation.reset(token) # If no result, use the default value and cache it if array is None: + if variable.uprating is not None or ( + self.tax_benefit_system.auto_carry_over_input_variables + and variable.calculate_output is None + ): + # The value is uprated or carried over from another + # period, or defaults for lack of one: an input set later + # for another period of the variable can change it. + self._get_store_history().record_derived( + variable_name, next_sequence_number() + ) # Check if the variable has a previously defined value known_periods = holder.get_known_periods() earlier_known_periods = [ @@ -949,6 +1018,13 @@ def _calculate(self, variable_name: str, period: Period = None) -> ArrayLike: # win over a later "2024-06" monthly value (bug H1). last_known_period = max(known_periods, key=lambda p: p.start) if last_known_period.start > period.start: + stored_input = ( + self._input_set_meanwhile(holder, period, input_state[2]) + if input_state[1] != getattr(self, "_inputs_set", 0) + else None + ) + if stored_input is not None: + return stored_input return holder.default_array() # Pass branch_name through so auto-carry-over respects the # active branch instead of reaching for the "default" @@ -969,10 +1045,12 @@ def _calculate(self, variable_name: str, period: Period = None) -> ArrayLike: array = EnumArray(array, variable.possible_values) array = self._cast_formula_result(array, variable) - holder.put_in_cache(array, period, self.branch_name) + array = self._cache_result(holder, array, period, input_state) except SpiralError: array = holder.default_array() + # Not stored, but what reads it is (see ``set_input``). + holder._record_store(period, next_sequence_number()) except RecursionError as e: if isinstance(self.tracer, FullTracer): self.tracer.print_computation_log() @@ -987,14 +1065,80 @@ def _calculate(self, variable_name: str, period: Period = None) -> ArrayLike: f"RecursionError while calculating {variable_name} for period {period}. The full computation stack is:\n{stack_formatted}" ) - if is_cache_available: + # Neither cache keeps a result an input change may have made obsolete + # (see ``_cache_result``). + unchanged = self._may_keep(input_state) + if is_cache_available and unchanged: smc.set_cache_value(cache_path, array) - if hasattr(self, "_fast_cache"): + if hasattr(self, "_fast_cache") and unchanged: self._fast_cache[(variable_name, period)] = array return array + def _input_state(self) -> Tuple[int, int]: + """How many drops ran during calculations, and inputs were set, here.""" + return getattr(self, "_input_epoch", 0), getattr(self, "_inputs_set", 0) + + def _may_keep(self, input_state: Tuple[int, int, int]) -> bool: + """Whether a result calculated since ``input_state`` may be kept (see ``_cache_result``).""" + frames = _calculation_frames.get() + return input_state[:2] == self._input_state() and not ( + frames and _frame_is_stale(*frames[-1]) + ) + + def _calculation_start(self) -> Tuple[int, int, int]: + """:meth:`_input_state` when a calculation begins, and a sequence number then.""" + return (*self._input_state(), next_sequence_number()) + + def _input_set_meanwhile( + self, holder: Holder, period: Period, started_at: int + ) -> Optional[ArrayLike]: + """The input this simulation reads for ``period``, if it was stored after ``started_at``. + + That is an input set while the calculation that began at + ``started_at`` ran (by its own formula, say), under any branch name + the simulation reads. + """ + stored_on = holder._branch_holding(period, self.branch_name) + if ( + stored_on is not None + and holder._is_input(period, stored_on) + and (holder._stored_sequence_number(period, stored_on) or 0) > started_at + ): + return holder._get_array_from_storage(period, stored_on) + return None + + def _cache_result( + self, + holder: Holder, + array: ArrayLike, + period: Period, + input_state: Tuple[int, int, int], + ) -> ArrayLike: + """Cache a calculated value, and return the value to use for it. + + ``input_state`` is :meth:`_calculation_start` when the calculation began. + If an input for the same period was set meanwhile (by the formula + itself, say), that input is the value, as it would be had it been set + first. A calculation that was running when an input set on this + simulation dropped values may have read the replaced value, so its + result is returned but not kept (see ``_drop_computed``). + """ + epoch, inputs_set, started_at = input_state + if inputs_set != getattr(self, "_inputs_set", 0): + stored_input = self._input_set_meanwhile(holder, period, started_at) + if stored_input is not None: + return stored_input + frames = _calculation_frames.get() + stale = bool(frames) and _frame_is_stale(*frames[-1]) + if epoch == getattr(self, "_input_epoch", 0) and not stale: + holder.put_in_cache(array, period, self.branch_name) + else: + # Not kept, but whatever reads it is stored after it all the same. + holder._record_store(period, next_sequence_number()) + return array + def purge_cache_of_invalid_values(self) -> None: # We wait for the end of calculate(), signalled by an empty stack, before purging the cache if self.tracer.stack: @@ -1044,13 +1188,35 @@ def calculate_add( ) ) - result = sum( - self.calculate(variable_name, sub_period) - for sub_period in period.get_subperiods(variable.definition_period) - ) - holder = self.get_holder(variable.name) - holder.put_in_cache(result, period, self.branch_name) - return result + def total(): + return sum( + self.calculate(variable_name, sub_period) + for sub_period in period.get_subperiods(variable.definition_period) + ) + + # As in ``calculate``: only an outermost sum runs again, and while it + # sums it is in flight, so its terms do not run again on their own + # (and a drop meanwhile keeps the records of what they read). + outermost = not getattr(self, "_calculations_in_flight", 0) + self._calculations_in_flight = getattr(self, "_calculations_in_flight", 0) + 1 + try: + with _calculation_frame(self) as frame: + input_epoch = getattr(self, "_input_epoch", 0) + input_state = self._calculation_start() + result = total() + for _ in range(_RERUNS_AFTER_INPUT_CHANGE if outermost else 0): + if input_epoch == getattr(self, "_input_epoch", 0): + break + # An input changed while summing: earlier terms may be + # obsolete (see ``calculate``). Sum again. + input_epoch = getattr(self, "_input_epoch", 0) + frame.update(_new_frame(self)) + input_state = self._calculation_start() + result = total() + holder = self.get_holder(variable.name) + return self._cache_result(holder, result, period, input_state) + finally: + self._calculations_in_flight -= 1 def calculate_divide( self, @@ -1079,11 +1245,12 @@ def calculate_divide( ) if period.unit == periods.MONTH: - computation_period = period.this_year - result = self.calculate(variable_name, period=computation_period) / 12.0 - holder = self.get_holder(variable.name) - holder.put_in_cache(result, period, self.branch_name) - return result + with _calculation_frame(self): + input_state = self._calculation_start() + computation_period = period.this_year + result = self.calculate(variable_name, period=computation_period) / 12.0 + holder = self.get_holder(variable.name) + return self._cache_result(holder, result, period, input_state) elif period.unit == periods.YEAR: return self.calculate(variable_name, period) @@ -1393,6 +1560,28 @@ def delete_arrays(self, variable: str, period: Period = None) -> None: if _fast_cache is not None: _fast_cache.pop((variable, period), None) + def drop_computed_arrays(self) -> int: + """Delete every value this simulation holds except inputs. + + Inputs are the values stored through ``set_input``: the dataset or + situation the simulation was built from, inputs set on it and, for a + branch, inputs set on the simulations it was created from before it + was created. Values a custom ``set_input`` handler calculates are not + inputs. Every other value is calculated again when next requested, + and the simulation stops reading macro-cache files. + + Use this on a branch whose tax-benefit system or parameters differ + from its parent's, whose values the branch would otherwise inherit; + ``set_input`` on a branch drops what depends on the input by itself. + Branches already created from this simulation keep their values. + + Returns: + int: The number of arrays deleted. + """ + # Recalculate rather than read macro-cache files written before. + self.macro_cache_read = False + return self._drop_computed() + def get_known_periods(self, variable: str) -> List[Period]: """ Get a list variable's known period, i.e. the periods where a value has been initialized and @@ -1426,6 +1615,49 @@ def set_input(self, variable_name: str, period: Period, value: ArrayLike) -> Non array([12, 14], dtype=int32) If a ``set_input`` property has been set for the variable, this method may accept inputs for periods not matching the ``definition_period`` of the variable. To read more about this, check the `documentation `_. + + On a branch (see :meth:`get_branch`), the input also drops what the + branch holds that may have been calculated from the value it + replaces, so what the branch calculates next uses the input, as a + simulation given the input before calculating anything would. Every + value is stored after everything it was calculated from, and each + simulation records the first stores its values may have been + calculated from (see :mod:`policyengine_core.data_storage.store_history`): + its own, its parent's when it was created, and those of simulations + its formulas calculated in. The branch drops each value it holds, + other than an input, stored at or after the earliest recorded store of + ``variable_name`` for a period that shares a day with ``period``, or + the earliest recorded value of ``variable_name`` uprated or carried + over from another period (or given the default for want of one), or + the first value that may depend on anything (restored from a dump, + or calculated from a macro-cache read). If there is none of these, + nothing it holds depends on the value and nothing is dropped. That is the case when a formula creates the branch while + still calculating the variable it overrides, unless another branch + it created for the same comparison already returned a value + calculated from the variable; then only what came back from there, + and what was calculated after, is dropped. + + Once an input is set on it, the branch stops reading macro-cache + files, which are keyed by branch name and period but not by inputs. + A calculation that was running when its input changed, in the branch + or in a simulation calling into it, is not kept; the outermost one in + the branch is run again from the new inputs until a run changes none + (at most ten times; after that its result is returned but not kept), + and an input stored for the very period it calculates after it began + is its result. + + What this does not track: a branch given a different tax-benefit + system or parameters (call :meth:`drop_computed_arrays` on it); + formulas that write into an array they read instead of returning a + new one; formulas that test whether a value is stored + (``get_known_periods``, ``get_array``) or read another simulation's + storage directly, rather than calculating; a simulation other than + the formula's own branches calculated from a thread the formula + starts without copying its context; and branches a formula keeps + between calls, which hold what their parent held when they were + created. Inputs set on a simulation that is not a branch drop + nothing, as before, and branches already created from the branch + keep their values. """ period = periods.period(period) if self.start_instant is None or self.start_instant > period.start: @@ -1440,6 +1672,89 @@ def set_input(self, variable_name: str, period: Period, value: ArrayLike) -> Non if _fast_cache is not None: _fast_cache.pop((variable_name, period), None) + def _get_store_history(self) -> StoreHistory: + history = getattr(self, "_store_history", None) + if history is None: + history = self._store_history = StoreHistory() + return history + + def _get_requested_variables(self) -> set: + requested = getattr(self, "_requested_variables", None) + if requested is None: + requested = self._requested_variables = set() + return requested + + def _share_store_history_with_caller(self) -> None: + """Merge this simulation's store history into that of a formula calling it.""" + caller = _formula_simulation.get() + if caller is not None: + if caller is not self: + caller._get_store_history().merge(self._get_store_history()) + return + # No formula is visible here: the call comes from code outside any + # formula, or from a thread a formula started without copying its + # context. In the second case a formula in one of this simulation's + # ancestors is waiting for the result, so give it to every ancestor + # with a calculation running (more records only make drops broader). + ancestor = getattr(self, "parent_branch", None) + while ancestor is not None: + if getattr(ancestor, "_calculations_in_flight", 0): + ancestor._get_store_history().merge(self._get_store_history()) + ancestor = getattr(ancestor, "parent_branch", None) + + def _drop_values_that_may_depend_on( + self, variable_name: str, period: Period + ) -> int: + """On a branch, drop what may depend on ``variable_name`` at ``period``. + + Called before an input for ``variable_name`` at ``period`` is stored; + see :meth:`set_input`. Returns the number of arrays dropped. + """ + if getattr(self, "parent_branch", None) is None: + return 0 + # Macro-cache files are keyed by branch and period, not by inputs, so + # a file for this branch's name (written by this branch, or by another + # simulation's branch of the same name) may hold values calculated + # without this input, even before the branch has read any. + self.macro_cache_read = False + since = self._get_store_history().earliest_dependency( + variable_name, periods.period(period) + ) + if self.get_holder(variable_name)._has_unnumbered_values(): + # Written into storage directly, so neither numbered nor recorded: + # anything may have been calculated from it. + since = 0 + if since is None: + return 0 + return self._drop_computed(since) + + def _drop_computed(self, since: Optional[int] = None) -> int: + """Drop every non-input value numbered ``since`` or later (all, without it).""" + dropped = 0 + for population in self.populations.values(): + for holder in population._holders.values(): + dropped += holder._drop_computed(since) + # The fast cache can also hold values a holder does not keep. + self._fast_cache = {} + if getattr(self, "_calculations_in_flight", 0): + # A formula running here may still hold values calculated from + # what the records describe: keep the records, and keep none of + # the results running calculations return, here or in the + # simulations calling into this one (``_cache_result``); the + # outermost one here runs again (``calculate``). + # Calculations elsewhere that already read from this simulation + # see the change through their frames' ``reads``. + self._input_epoch = getattr(self, "_input_epoch", 0) + 1 + return dropped + # Nothing the simulation still holds was calculated from what the + # records numbered ``since`` or later describe, except the inputs it + # keeps, which are recorded again. + self._get_store_history().prune(since) + for population in self.populations.values(): + for holder in population._holders.values(): + holder._record_inputs(since) + return dropped + def get_variable_population(self, variable_name: str) -> Population: variable = self.tax_benefit_system.get_variable( variable_name, check_existence=True @@ -1523,6 +1838,12 @@ def clone( new.tax_benefit_system = self.tax_benefit_system new.debug = debug new.trace = trace + # The copy holds what this simulation holds, so it starts from what + # those values may have been calculated from, and diverges from there. + new._store_history = self._get_store_history().copy() + new._requested_variables = set(self._get_requested_variables()) + # Calculations running in this simulation are not running in the copy. + new._calculations_in_flight = 0 return new @@ -1553,6 +1874,13 @@ def get_branch( several threads at once: two first reads of the same array can each make a copy. + A new branch starts with the values this simulation holds; asked for + a name it already has, this returns that branch as it is. An input + set on the branch drops those that may have been calculated from the + value it replaces (see :meth:`set_input`). A branch whose + tax-benefit system or parameters are changed should call + :meth:`drop_computed_arrays` before calculating. + Args: name (str, optional): Name of the branch. Defaults to "branch". clone_system (bool, optional): Whether to clone the tax-benefit system. Use this if you're changing policy parameters. Defaults to False. @@ -1601,9 +1929,7 @@ def derivative( period = periods.period(self.default_calculation_period) alt_sim = self.clone() - for computed_variable in alt_sim.tax_benefit_system.variables: - if computed_variable not in self.input_variables: - alt_sim.delete_arrays(computed_variable) + alt_sim.drop_computed_arrays() alt_sim.set_input(wrt, period, self.calculate(wrt, period) + delta) original_value = self.calculate(variable, period) new_value = alt_sim.calculate(variable, period) @@ -1975,8 +2301,11 @@ def subsample( df = subset_df - # Update the dataset and rebuild the simulation + # Update the dataset and rebuild the simulation. Nothing stored before + # survives the rebuild, so start a new store history before the + # rebuilt inputs are recorded in it. self.dataset = Dataset.from_dataframe(df, self.dataset.time_period) + self._store_history = StoreHistory() self.build_from_dataset() # Purge ``_fast_cache`` entries populated by ``to_input_dataframe`` diff --git a/policyengine_core/tools/simulation_dumper.py b/policyengine_core/tools/simulation_dumper.py index c3db0c4f..4c26b89a 100644 --- a/policyengine_core/tools/simulation_dumper.py +++ b/policyengine_core/tools/simulation_dumper.py @@ -6,9 +6,13 @@ import numpy as np from policyengine_core.data_storage import OnDiskStorage +from policyengine_core.data_storage.store_history import next_sequence_number from policyengine_core.periods import ETERNITY from policyengine_core.simulations import Simulation +# Next to each variable's arrays: the periods whose values were inputs. +_INPUTS_FILE = "inputs.txt" + def dump_simulation(simulation, directory): """ @@ -55,20 +59,43 @@ def restore_simulation(directory, tax_benefit_system, **kwargs): _restore_entity(population, entities_dump_dir) population.count = person_count - variables_to_restore = ( + variables_to_restore = [ variable for variable in os.listdir(directory) if variable != "__entities__" - ) + ] + # Inputs first, then every calculated value under one later number. The + # dump does not say what each value was calculated from (nor what was read + # without being kept), so an input set for any variable on a branch of the + # restored simulation drops them all. for variable in variables_to_restore: - _restore_holder(simulation, variable, directory) + _restore_holder(simulation, variable, directory, inputs=True) + calculated_number = next_sequence_number() + restored = sum( + _restore_holder( + simulation, variable, directory, calculated_number=calculated_number + ) + for variable in variables_to_restore + ) + if restored: + simulation._get_store_history().record_unknown_sources(calculated_number) return simulation def _dump_holder(holder, directory): disk_storage = holder.create_disk_storage(directory, preserve=True) - for period in holder.get_known_periods(): - value = holder.get_array(period) - disk_storage.put(value, period) + branch_name = holder.simulation.branch_name + inputs = [] + for period in dict.fromkeys(holder.get_known_periods()): + # What the simulation itself reads: on a branch, its own value, else + # its nearest ancestor's or the default one. + stored_on = holder._branch_holding(period, branch_name) + if stored_on is None: + continue + disk_storage.put(holder._get_array_from_storage(period, stored_on), period) + if holder._is_input(period, stored_on): + inputs.append(str(period)) + with open(os.path.join(disk_storage.storage_dir, _INPUTS_FILE), "w") as file: + file.write("\n".join(inputs)) def _dump_entity(population, directory): @@ -122,7 +149,9 @@ def _restore_entity(population, directory): return person_count -def _restore_holder(simulation, variable, directory): +def _restore_holder( + simulation, variable, directory, inputs=False, calculated_number=None +): storage_dir = os.path.join(directory, variable) is_variable_eternal = ( simulation.tax_benefit_system.get_variable(variable).definition_period @@ -134,7 +163,25 @@ def _restore_holder(simulation, variable, directory): disk_storage.restore() holder = simulation.get_holder(variable) + inputs_path = os.path.join(storage_dir, _INPUTS_FILE) + if os.path.exists(inputs_path): + with open(inputs_path) as file: + input_periods = set(file.read().split()) + else: + # Dumped before inputs were recorded: keep every value, as inputs. + input_periods = None + restored = 0 for period in disk_storage.get_known_periods(): + is_input = input_periods is None or str(period) in input_periods + if is_input != inputs: + continue + restored += 1 value = disk_storage.get(period) - holder.put_in_cache(value, period) + holder._set( + period, + value, + is_input=is_input, + sequence_number=None if is_input else calculated_number, + ) + return restored diff --git a/tests/core/test_branch_input_invalidation.py b/tests/core/test_branch_input_invalidation.py new file mode 100644 index 00000000..5650e02e --- /dev/null +++ b/tests/core/test_branch_input_invalidation.py @@ -0,0 +1,2276 @@ +"""``set_input`` on a branch drops what may depend on the value it replaces. + +A branch starts with every value its parent holds. ``set_input`` on a branch +used to store the new value and keep everything calculated from the old one, +so a value the parent calculated before branching answered for the branch, +whatever the branch's inputs said, and the branch's results depended on what +had been calculated before it was created. + +Every stored value now carries a sequence number from one process-wide +counter, and a value is stored after everything it was calculated from. Each +simulation keeps a history of the first stores its values may have been +calculated from: a branch copies its parent's, and a formula that calculates +in another simulation takes in that simulation's history. An input set on a +branch drops the branch's non-input values stored at or after the earliest +recorded store, or derivation, of the overridden variable for an overlapping +period (``StoreHistory.earliest_dependency``), and the records the drop made +obsolete. + +The invariant tested throughout (and, over random sequences of calculations, +branches and inputs, by ``test_branch_input_invalidation_property.py``): +whatever the branch's family calculated before, and in whatever order, a +calculation in a branch equals the same calculation in a new simulation given +the branch's inputs, its own and those it inherited when it was created, +before calculating anything. +""" + +from __future__ import annotations + +import numpy as np +import pytest + +from policyengine_core import periods +from policyengine_core.data_storage import InMemoryStorage +from policyengine_core.data_storage.store_history import ( + StoreHistory, + next_sequence_number, +) +from policyengine_core.experimental import MemoryConfig +from policyengine_core.model_api import Reform +from policyengine_core.simulations import SimulationBuilder +import policyengine_core.simulations.simulation as simulation_module +from tests.fixtures.branch_input_invalidation import ( + CARRY_OVER_SYSTEM, + FORMULA_RUNS, + ROOT_INPUTS, + branching_simulation, + synthetic_simulation, +) + +JANUARY = periods.period("2017-01") +JANUARY_2022 = periods.period("2022-01") + +SITUATION = { + "persons": { + "a": {"birth": {"ETERNITY": "1980-01-01"}, "salary": {"2017-01": 4000}}, + "b": {"birth": {"ETERNITY": "1985-01-01"}, "salary": {"2017-01": 1000}}, + }, + "households": { + "h": { + "parents": ["a", "b"], + "accommodation_size": {"2017-01": 80}, + "housing_occupancy_status": {"2017-01": "tenant"}, + "rent": {"2017-01": 900}, + } + }, +} + + +def _build(tax_benefit_system, situation=SITUATION): + return SimulationBuilder().build_from_entities(tax_benefit_system, situation) + + +def _with_salary(tax_benefit_system, salary, period=JANUARY): + """A new simulation given ``salary`` before calculating anything.""" + simulation = _build(tax_benefit_system) + simulation.set_input("salary", period, np.asarray(salary, dtype=float)) + return simulation + + +def _stored_keys(simulation, variable): + return set(simulation.get_holder(variable)._memory_storage._arrays) + + +# ----- Examples on the country template ----- # + + +def test_branch_input_reaches_a_value_the_parent_calculated_first( + tax_benefit_system, +): + simulation = _build(tax_benefit_system) + simulation.calculate("income_tax", JANUARY) + simulation.calculate("disposable_income", JANUARY) + + branch = simulation.get_branch("raise") + branch.set_input("salary", JANUARY, np.array([5000.0, 1000.0])) + + fresh = _with_salary(tax_benefit_system, [5000.0, 1000.0]) + for variable in ("income_tax", "disposable_income", "household_income"): + assert np.array_equal( + branch.calculate(variable, JANUARY), fresh.calculate(variable, JANUARY) + ), variable + # The parent keeps its own values. + assert np.array_equal( + simulation.calculate("income_tax", JANUARY), + _build(tax_benefit_system).calculate("income_tax", JANUARY), + ) + + +def test_branch_keeps_values_stored_before_the_input_first_existed( + tax_benefit_system, +): + """Values stored before the overridden variable was first stored stay.""" + situation = { + "persons": {"a": {"birth": {"ETERNITY": "1980-01-01"}}}, + "households": { + "h": { + "parents": ["a"], + "accommodation_size": {"2017-01": 80}, + "housing_occupancy_status": {"2017-01": "tenant"}, + } + }, + } + simulation = _build(tax_benefit_system, situation) + # housing_tax is stored before salary has any value. + simulation.calculate("housing_tax", "2017") + simulation.set_input("salary", JANUARY, np.array([3000.0])) + simulation.calculate("income_tax", JANUARY) + + branch = simulation.get_branch("raise") + branch.set_input("salary", JANUARY, np.array([4000.0])) + + assert branch.calculate("income_tax", JANUARY) == pytest.approx( + 4000 * tax_benefit_system.parameters(JANUARY).taxes.income_tax_rate + ) + assert _stored_keys(branch, "housing_tax") == {"default:2017"} + + +def test_branch_input_with_no_history_drops_nothing(tax_benefit_system): + """The usual case costs nothing: no value can depend on a value never stored.""" + simulation = _build(tax_benefit_system) + simulation.calculate("disposable_income", JANUARY) + branch = simulation.get_branch("later") + before = { + variable: _stored_keys(branch, variable) + for variable in ("disposable_income", "income_tax", "salary") + } + + # salary has never been stored, read or derived for 2018-03. + dropped = branch._drop_values_that_may_depend_on( + "salary", periods.period("2018-03") + ) + branch.set_input("salary", "2018-03", np.array([1.0, 2.0])) + + assert dropped == 0 + for variable, keys in before.items(): + assert _stored_keys(branch, variable) >= keys, variable + + +def test_branch_keeps_inputs_of_the_dataset_and_of_ancestor_branches( + tax_benefit_system, +): + simulation = _build(tax_benefit_system) + simulation.calculate("disposable_income", JANUARY) + parent_branch = simulation.get_branch("parent") + parent_branch.set_input("rent", JANUARY, np.array([500.0])) + parent_branch.calculate("housing_allowance", JANUARY) + + child = parent_branch.get_branch("child") + child.set_input("salary", JANUARY, np.array([0.0, 0.0])) + + assert "default:2017-01" in _stored_keys(child, "accommodation_size") + assert "parent:2017-01" in _stored_keys(child, "rent") + assert child.calculate("rent", JANUARY) == 500.0 + fresh = _with_salary(tax_benefit_system, [0.0, 0.0]) + fresh.set_input("rent", JANUARY, np.array([500.0])) + for variable in ("housing_allowance", "household_income", "income_tax"): + assert np.array_equal( + child.calculate(variable, JANUARY), fresh.calculate(variable, JANUARY) + ), variable + + +def test_input_for_a_year_drops_values_calculated_from_a_month( + tax_benefit_system, +): + simulation = _build(tax_benefit_system) + simulation.calculate("income_tax", JANUARY) + + branch = simulation.get_branch("annual") + # salary divides a yearly input over the months the branch has not set. + branch.set_input("salary", "2017", np.array([12000.0, 24000.0])) + + rate = tax_benefit_system.parameters(JANUARY).taxes.income_tax_rate + assert np.allclose( + branch.calculate("income_tax", JANUARY), [1000 * rate, 2000 * rate] + ) + + +def test_input_set_again_on_a_branch_drops_what_the_first_input_fed( + tax_benefit_system, +): + simulation = _build(tax_benefit_system) + branch = simulation.get_branch("raise") + branch.set_input("salary", JANUARY, np.array([5000.0, 1000.0])) + branch.calculate("disposable_income", JANUARY) + branch.set_input("salary", JANUARY, np.array([7000.0, 1000.0])) + + fresh = _with_salary(tax_benefit_system, [7000.0, 1000.0]) + assert np.array_equal( + branch.calculate("disposable_income", JANUARY), + fresh.calculate("disposable_income", JANUARY), + ) + + +def test_set_input_on_a_simulation_that_is_not_a_branch_drops_nothing( + tax_benefit_system, +): + """Unchanged: a root simulation's later inputs do not reach its cache.""" + simulation = _build(tax_benefit_system) + cached = simulation.calculate("income_tax", JANUARY).copy() + simulation.set_input("salary", JANUARY, np.array([9000.0, 9000.0])) + + assert np.array_equal(simulation.calculate("income_tax", JANUARY), cached) + + +def test_existing_child_branches_keep_their_values(tax_benefit_system): + simulation = _build(tax_benefit_system) + branch = simulation.get_branch("branch") + branch.calculate("income_tax", JANUARY) + child = branch.get_branch("child") + child_value = child.calculate("income_tax", JANUARY).copy() + + branch.set_input("salary", JANUARY, np.array([0.0, 0.0])) + + assert np.array_equal(child.calculate("income_tax", JANUARY), child_value) + assert np.array_equal(branch.calculate("income_tax", JANUARY), [0.0, 0.0]) + + +def test_disk_storage_drops_calculated_values_and_keeps_inputs(tax_benefit_system): + situation = { + "persons": {"a": {"birth": {"ETERNITY": "1980-01-01"}}}, + "households": {"h": {"parents": ["a"]}}, + } + simulation = _build(tax_benefit_system, situation) + simulation.memory_config = MemoryConfig(max_memory_occupation=0) + for variable in ("rent", "housing_allowance"): + holder = simulation.get_holder(variable) + holder._disk_storage = holder.create_disk_storage() + holder._on_disk_storable = True + simulation.set_input("rent", JANUARY, np.array([700.0])) + simulation.calculate("housing_allowance", JANUARY) + assert simulation.get_holder("housing_allowance")._disk_storage._files + + branch = simulation.get_branch("cheaper") + branch.set_input("rent", JANUARY, np.array([300.0])) + + dropped = branch.get_holder("housing_allowance")._disk_storage._files == {} + fresh = _build(tax_benefit_system, situation) + fresh.set_input("rent", JANUARY, np.array([300.0])) + assert np.array_equal( + branch.calculate("housing_allowance", JANUARY), + fresh.calculate("housing_allowance", JANUARY), + ) + assert dropped + rent_disk = branch.get_holder("rent")._disk_storage + assert set(rent_disk._files) == {"default_2017-01", "cheaper_2017-01"} + assert rent_disk._input_keys == {"default_2017-01", "cheaper_2017-01"} + + +def test_value_served_from_the_macro_cache_is_tracked(tax_benefit_system, monkeypatch): + """A macro-cache read is not stored, but what reads it is.""" + import policyengine_core.simulations.simulation as simulation_module + + class CachedIncomeTax: + def __init__(self, tax_benefit_system): + pass + + def set_cache_path(self, *args): + pass + + def get_cache_path(self): + return type("Path", (), {"exists": lambda self: True})() + + def get_cache_value(self, path): + return np.array([100.0, 100.0], dtype=np.float32) + + monkeypatch.setattr(simulation_module, "SimulationMacroCache", CachedIncomeTax) + simulation = _build(tax_benefit_system) + simulation.macro_cache_read = True + simulation.dataset = type( + "Dataset", (), {"file_path": simulation_module.Path("cache"), "name": "d"} + )() + monkeypatch.setattr( + type(simulation), + "check_macro_cache", + lambda self, variable, period: variable == "income_tax", + ) + simulation.calculate("disposable_income", JANUARY) + assert not _stored_keys(simulation, "income_tax") + + branch = simulation.get_branch("branch") + branch.set_input("income_tax", JANUARY, np.array([0.0, 0.0])) + + assert not _stored_keys(branch, "disposable_income") + fresh = _build(tax_benefit_system) + fresh.set_input("income_tax", JANUARY, np.array([0.0, 0.0])) + assert np.array_equal( + branch.calculate("disposable_income", JANUARY), + fresh.calculate("disposable_income", JANUARY), + ) + + +# ----- drop_computed_arrays ----- # + + +def test_drop_computed_arrays_keeps_only_inputs(tax_benefit_system): + simulation = _build(tax_benefit_system) + simulation.calculate("total_taxes", JANUARY) + branch = simulation.get_branch("branch") + branch.set_input("rent", JANUARY, np.array([100.0])) + branch.calculate("total_benefits", JANUARY) + + dropped = branch.drop_computed_arrays() + + assert dropped > 0 + remaining = { + (name, key) + for population in branch.populations.values() + for name, holder in population._holders.items() + for key in holder._memory_storage._arrays + } + inputs = { + (name, key) + for population in branch.populations.values() + for name, holder in population._holders.items() + for key in holder._memory_storage._input_keys + } + assert remaining == inputs + assert {("rent", "branch:2017-01"), ("salary", "default:2017-01")} <= remaining + assert not _stored_keys(branch, "total_taxes") + assert not _stored_keys(branch, "total_benefits") + assert np.array_equal(branch.calculate("rent", JANUARY), [100.0]) + # The parent is untouched. + assert _stored_keys(simulation, "total_taxes") == {"default:2017-01"} + + +def test_drop_computed_arrays_after_changing_a_branch_parameter(tax_benefit_system): + """Parameters are not inputs: a branch with other parameters drops its values.""" + simulation = _build(tax_benefit_system) + simulation.calculate("income_tax", JANUARY) + + branch = simulation.get_branch("higher_rate", clone_system=True) + branch.tax_benefit_system.parameters.taxes.income_tax_rate.update( + period="2017", value=0.5 + ) + stale = branch.calculate("income_tax", JANUARY).copy() + branch.drop_computed_arrays() + + assert np.array_equal(branch.calculate("income_tax", JANUARY), [2000.0, 500.0]) + assert not np.array_equal(stale, [2000.0, 500.0]) + + +def test_apply_reform_keeps_branch_inputs(tax_benefit_system): + simulation = _build(tax_benefit_system) + branch = simulation.get_branch("branch") + branch.set_input("salary", JANUARY, np.array([10.0, 20.0])) + branch.calculate("income_tax", JANUARY) + + class NoOp(Reform): + def apply(self): + pass + + simulation.apply_reform(NoOp) + + assert _stored_keys(branch, "income_tax") == set() + assert np.array_equal(branch.calculate("salary", JANUARY), [10.0, 20.0]) + + +def test_subsample_starts_a_new_store_history(): + """``subsample`` replaces every value, so earlier stores no longer count.""" + from policyengine_core.country_template import Microsimulation + + simulation = Microsimulation() + simulation.calculate("income_tax", JANUARY_2022) + last_before = next_sequence_number() + simulation.subsample(n=3, seed="store-history", time_period="2022") + + recorded = [ + number + for stored in simulation._store_history._first_stored.values() + for number in stored.values() + ] + assert recorded and min(recorded) > last_before + + # A branch's input still reaches what the parent calculated first. + salary = np.asarray(simulation.calculate("salary", JANUARY_2022)) + 1000 + before = simulation.get_branch("before") + before.set_input("salary", JANUARY_2022, salary) + expected = np.asarray( + before.calculate("social_security_contribution", JANUARY_2022) + ) + simulation.calculate("social_security_contribution", JANUARY_2022) + after = simulation.get_branch("after") + after.set_input("salary", JANUARY_2022, salary) + assert np.array_equal( + np.asarray(after.calculate("social_security_contribution", JANUARY_2022)), + expected, + ) + + +# ----- Storage and history ----- # + + +def test_storage_records_sequence_numbers_and_inputs_through_clone(): + storage = InMemoryStorage(is_eternal=False) + storage.put(np.array([1.0]), periods.period("2020"), is_input=True) + first = storage._sequence_numbers["default:2020"] + storage.put(np.array([2.0]), periods.period("2021")) + clone = storage.clone() + + assert clone._sequence_numbers == storage._sequence_numbers + assert clone._sequence_numbers["default:2021"] > first + assert clone._input_keys == {"default:2020"} + + assert clone.drop_computed(since=first) == 1 + assert set(clone._arrays) == {"default:2020"} + assert set(storage._arrays) == {"default:2020", "default:2021"} + + +def test_storage_drops_only_values_at_or_after_since(): + storage = InMemoryStorage(is_eternal=False) + numbers = [] + for year in (2020, 2021, 2022): + number = next_sequence_number() + storage.put(np.array([0.0]), periods.period(year), sequence_number=number) + numbers.append(number) + + assert storage.drop_computed(since=numbers[1]) == 2 + assert set(storage._arrays) == {"default:2020"} + storage.delete(periods.period(2020)) + assert storage._sequence_numbers == {} + + +def test_history_matches_overlapping_periods_and_derived_values(): + history = StoreHistory() + history.record_store("v", periods.period("2020-03"), 10) + history.record_store("v", periods.period("2020-03"), 99) + history.record_store("v", periods.period("2021"), 20) + + assert history.earliest_dependency("v", periods.period("2020")) == 10 + assert history.earliest_dependency("v", periods.period("2020-04")) is None + assert history.earliest_dependency("v", periods.period("2021-06")) == 20 + assert history.earliest_dependency("w", periods.period("2020")) is None + + history.record_derived("v", 15) + assert history.earliest_dependency("v", periods.period("2020-04")) == 15 + history.record_store("e", periods.period(periods.ETERNITY), 5) + assert history.earliest_dependency("e", periods.period("1999-01")) == 5 + + +# ----- Derived values and branches inside formulas, on a synthetic system ----- # + + +def test_value_carried_over_from_an_earlier_input_is_dropped(): + simulation = synthetic_simulation(ROOT_INPUTS, system=CARRY_OVER_SYSTEM) + simulation.calculate("p_cond", "2015") # p_c 2015 carried over from 2012 + + branch = simulation.get_branch("branch") + branch.set_input("p_c", "2014", np.array([70.0, 80.0, 90.0])) + + fresh = synthetic_simulation( + {**ROOT_INPUTS, ("p_c", "2014"): (70.0, 80.0, 90.0)}, + system=CARRY_OVER_SYSTEM, + ) + assert np.array_equal(fresh.calculate("p_c", "2015"), [70.0, 80.0, 90.0]) + for variable in ("p_c", "p_cond"): + assert np.array_equal( + branch.calculate(variable, "2015"), fresh.calculate(variable, "2015") + ), variable + + +def test_uprated_value_is_dropped(): + simulation = synthetic_simulation(ROOT_INPUTS) + simulation.calculate("p_prod", "2015") # p_up 2015 uprated from 2012 + + branch = simulation.get_branch("branch") + branch.set_input("p_up", "2014", np.array([1.0, 2.0, 3.0])) + + fresh = synthetic_simulation({**ROOT_INPUTS, ("p_up", "2014"): (1.0, 2.0, 3.0)}) + assert np.allclose( + branch.calculate("p_prod", "2015"), fresh.calculate("p_prod", "2015") + ) + + +def test_value_calculated_only_inside_another_branch_is_tracked(): + """A formula's own branch stores the overridden variable; its result is dropped. + + ``p_inner`` calculates ``p_inner_only`` in a branch it deletes when done, + so the parent never stores ``p_inner_only`` itself. It takes in that + branch's history when its calculation returns, and the parent's later + branches copy it. + """ + simulation = synthetic_simulation(ROOT_INPUTS) + simulation.calculate("p_inner", "2013") + assert not simulation.get_holder("p_inner_only")._memory_storage._arrays + + branch = simulation.get_branch("branch") + branch.set_input("p_inner_only", "2013", np.array([1.0, 1.0, 1.0])) + + fresh = synthetic_simulation( + {**ROOT_INPUTS, ("p_inner_only", "2013"): (1.0, 1.0, 1.0)} + ) + assert np.array_equal( + branch.calculate("p_inner", "2013"), fresh.calculate("p_inner", "2013") + ) + + +def test_value_a_holder_does_not_keep_is_tracked(): + """``variables_to_drop`` values are not stored, but what reads them is.""" + config = MemoryConfig(max_memory_occupation=1, variables_to_drop=["p_sum"]) + simulation = synthetic_simulation(ROOT_INPUTS, memory_config=config) + simulation.calculate("p_prod", "2013") + assert not simulation.get_holder("p_sum")._memory_storage._arrays + + branch = simulation.get_branch("branch") + branch.set_input("p_sum", "2013", np.array([0.0, 0.0, 0.0])) + + assert np.array_equal(branch.calculate("p_prod", "2013"), [0.0, 0.0, 0.0]) + + +def test_value_a_spiral_defaults_is_tracked(): + """The default a spiral returns is not stored, but what reads it is.""" + simulation = synthetic_simulation(ROOT_INPUTS) + branch = simulation.get_branch("branch") + branch.calculate("p_spiral", "2015") + deepest = min( + periods.period(key.split(":", 1)[1]) + for key in branch.get_holder("p_spiral")._memory_storage._arrays + ).last_year + + branch.set_input("p_spiral", deepest, np.array([100.0, 100.0, 100.0])) + + fresh = synthetic_simulation( + {**ROOT_INPUTS, ("p_spiral", str(deepest)): (100.0,) * 3} + ) + assert np.array_equal( + branch.calculate("p_spiral", "2015"), fresh.calculate("p_spiral", "2015") + ) + + +# ----- Values that reach a simulation without being stored there ----- # + + +def test_macro_cache_files_a_branch_wrote_are_not_read_after_its_input_changes( + tmp_path, +): + """Macro-cache files are keyed by branch and period, not by inputs.""" + from policyengine_core.data.dataset import Dataset + from policyengine_core.entities import build_entity + from policyengine_core.simulations import Simulation + from policyengine_core.taxbenefitsystems import TaxBenefitSystem + from policyengine_core.variables import Variable + + person = build_entity("person", "persons", "Person", is_person=True) + + class source(Variable): + label = "Source" + value_type = float + entity = person + definition_period = periods.YEAR + + class result(Variable): + label = "Result" + value_type = float + entity = person + definition_period = periods.YEAR + exhaustive_parameter_dependencies = [] + + def formula(person, period): + return 2 * person("source", period) + + class OnePerson(Dataset): + name = "one_person" + label = "One person" + file_path = tmp_path / "one_person.h5" + data_format = Dataset.ARRAYS + time_period = "2022" + + def generate(self): + self.save_dataset({"person_id": np.array([0]), "source": np.array([1.0])}) + + system = TaxBenefitSystem([person]) + system.add_variables(source, result) + parent = Simulation(tax_benefit_system=system, dataset=OnePerson) + parent.macro_cache_read = True + branch = parent.get_branch("test") + assert branch.calculate("result", "2022").tolist() == [2.0] # writes the file + + branch.set_input("source", "2022", np.array([5.0])) + + assert branch.calculate("result", "2022").tolist() == [10.0] + + +def test_blacklisted_value_is_tracked(): + """With ``opt_out_cache``, blacklisted values are not stored; what reads them is.""" + simulation = synthetic_simulation(ROOT_INPUTS, opt_out_cache=True) + simulation.calculate("p_prod", "2013") + assert not simulation.get_holder("p_sum")._memory_storage._arrays + + branch = simulation.get_branch("branch") + branch.set_input("p_sum", "2013", np.array([0.0, 0.0, 0.0])) + + assert np.array_equal(branch.calculate("p_prod", "2013"), [0.0, 0.0, 0.0]) + + +def test_value_written_straight_into_storage_counts_as_a_dependency( + tax_benefit_system, +): + """A value stored without ``put`` has no number, so anything may depend on it.""" + simulation = _build(tax_benefit_system) + holder = simulation.get_holder("salary") + holder._memory_storage._arrays["default:2017-02"] = np.array([4000.0, 0.0]) + simulation.calculate("income_tax", "2017-02") + + branch = simulation.get_branch("branch") + branch.set_input("salary", "2017-02", np.array([0.0, 0.0])) + + assert np.array_equal(branch.calculate("income_tax", "2017-02"), [0.0, 0.0]) + + +def test_cached_value_another_simulation_returns_is_tracked(): + """A formula that reads another simulation's cached value takes in its history.""" + simulation = branching_simulation() + # Cached in the persistent branch before any formula of the parent reads it. + simulation.get_branch("persistent").calculate("tax", "2020") + simulation.calculate("from_persistent_branch", "2020") # a cache hit there + + branch = simulation.get_branch("branch") + branch.set_input("tax", "2020", np.zeros(3)) + + # A new simulation given the input first creates its persistent branch + # from itself, so the branch's input reaches it. + assert np.array_equal( + branch.calculate("from_persistent_branch", "2020"), np.zeros(3) + ) + + +def test_history_is_taken_in_again_after_a_drop_forgets_it(): + """A drop forgets records; reading the other simulation again restores them.""" + simulation = branching_simulation() + branch = simulation.get_branch("branch") + branch.calculate("from_persistent_branch", "2020") + branch.drop_computed_arrays() # forgets what the persistent branch stored + branch.calculate("from_persistent_branch", "2020") # a cache hit there + + child = branch.get_branch("child") + child.set_input("tax", "2020", np.zeros(3)) + + assert np.array_equal( + child.calculate("from_persistent_branch", "2020"), np.zeros(3) + ) + + +def _one_person_system(*variables): + from policyengine_core.country_template import CountryTaxBenefitSystem + + system = CountryTaxBenefitSystem() + system.add_variables(*variables) + return system + + +def _yearly_variable(name, formula=None): + from policyengine_core.country_template import entities + from policyengine_core.variables import Variable + + namespace = dict( + value_type=float, + entity=entities.Person, + definition_period=periods.YEAR, + label=name, + ) + if formula is not None: + namespace["formula"] = formula + return type(name, (Variable,), namespace) + + +@pytest.mark.parametrize("drop", ["drop_computed_arrays", "set_input"]) +def test_drop_while_a_formula_runs_keeps_what_it_read_recorded(drop): + """A formula still holding a value it read keeps that value's record.""" + + def result(person, period): + source = person("source", period) + if drop == "drop_computed_arrays": + person.simulation.drop_computed_arrays() + else: + person.simulation.set_input("trigger", period, np.ones(person.count)) + return source * 2 + + system = _one_person_system( + _yearly_variable("source", lambda person, period: np.ones(person.count)), + _yearly_variable("trigger", lambda person, period: np.zeros(person.count)), + _yearly_variable("result", result), + ) + simulation = SimulationBuilder().build_default_simulation(system) + simulation.calculate("trigger", "2020") # stored before ``source`` + branch = simulation.get_branch("branch") + assert branch.calculate("result", "2020").tolist() == [2.0] + + branch.set_input("source", "2020", np.array([3.0])) + + assert branch.calculate("result", "2020").tolist() == [6.0] + + +def test_disk_restore_reads_the_latest_file_of_each_key(tmp_path, monkeypatch): + """Sequence numbers restart in each process; restore goes by write time.""" + import itertools + import os + + import policyengine_core.data_storage.on_disk_storage as on_disk_storage + import policyengine_core.data_storage.store_history as store_history + from policyengine_core.data_storage import OnDiskStorage + + monkeypatch.setattr(on_disk_storage, "_PROCESS_TOKEN", "aaaaaaaaaaaa") + writer = OnDiskStorage(str(tmp_path), preserve_storage_dir=True) + for _ in range(5): + writer.put(np.array([1.0]), periods.period("2020")) + # A later process: its counter restarts, and it has its own token. + monkeypatch.setattr(store_history, "_sequence", itertools.count(1)) + monkeypatch.setattr(on_disk_storage, "_PROCESS_TOKEN", "bbbbbbbbbbbb") + later = OnDiskStorage(str(tmp_path), preserve_storage_dir=True) + later.put(np.array([2.0]), periods.period("2020")) + # The writer's files clearly earlier, even on a coarse file clock. + for path in tmp_path.glob("default_2020.aaaaaaaaaaaa.*.npy"): + stat = os.stat(path) + os.utime(path, ns=(stat.st_atime_ns, stat.st_mtime_ns - 10**9)) + + reader = OnDiskStorage(str(tmp_path), preserve_storage_dir=True) + reader.restore() + + assert reader.get(periods.period("2020")).tolist() == [2.0] + + +def test_uprated_value_another_simulation_returns_is_tracked(): + """Taking in another simulation's history includes its uprated values.""" + inputs = {("p_up", "2012"): (100.0, 100.0, 100.0)} + simulation = synthetic_simulation(inputs) + simulation.calculate("p_imported_up", "2015") # uprated in a temporary branch + + branch = simulation.get_branch("branch") + branch.set_input("p_up", "2014", np.full(3, 300.0)) + + fresh = synthetic_simulation({**inputs, ("p_up", "2014"): (300.0,) * 3}) + assert np.allclose( + branch.calculate("p_imported_up", "2015"), + fresh.calculate("p_imported_up", "2015"), + ) + + +def test_value_a_formula_calculates_in_a_worker_thread_is_tracked(): + """A thread started without the formula's context still hands its history up.""" + from concurrent.futures import ThreadPoolExecutor + + def result(person, period): + simulation = person.simulation + child = simulation.get_branch("worker") + try: + with ThreadPoolExecutor(max_workers=1) as pool: + return pool.submit(child.calculate, "value", period).result() * 2 + finally: + del simulation.branches["worker"] + + system = _one_person_system( + _yearly_variable("value", lambda person, period: np.zeros(person.count)), + _yearly_variable("result", result), + ) + simulation = SimulationBuilder().build_default_simulation(system) + simulation.calculate("result", "2020") + + branch = simulation.get_branch("branch") + branch.set_input("value", "2020", np.array([10.0])) + + assert branch.calculate("result", "2020").tolist() == [20.0] + + +def test_failed_prerequisite_request_does_not_satisfy_the_gate(): + def prerequisite(person, period): + raise RuntimeError("the prerequisite failed") + + dependent = _yearly_variable( + "dependent", lambda person, period: np.full(person.count, 42.0) + ) + dependent.requires_computation_after = "prerequisite" + system = _one_person_system( + _yearly_variable("prerequisite", prerequisite), dependent + ) + simulation = SimulationBuilder().build_default_simulation(system) + with pytest.raises(RuntimeError): + simulation.calculate("prerequisite", "2020") + + with pytest.raises(ValueError, match="requires prerequisite"): + simulation.calculate("dependent", "2020") + + +def test_disk_files_from_another_process_are_not_overwritten(tmp_path, monkeypatch): + """Processes write their own files even when their sequence numbers repeat.""" + import itertools + + import policyengine_core.data_storage.on_disk_storage as on_disk_storage + import policyengine_core.data_storage.store_history as store_history + from policyengine_core.data_storage import OnDiskStorage + + monkeypatch.setattr(store_history, "_sequence", itertools.count(1)) + monkeypatch.setattr(on_disk_storage, "_PROCESS_TOKEN", "aaaaaaaaaaaa") + writer = OnDiskStorage(str(tmp_path), preserve_storage_dir=True) + writer.put(np.array([1.0]), periods.period("2020")) + snapshot = writer.clone() + + # A later process: its counter restarts. + monkeypatch.setattr(store_history, "_sequence", itertools.count(1)) + monkeypatch.setattr(on_disk_storage, "_PROCESS_TOKEN", "bbbbbbbbbbbb") + reader = OnDiskStorage(str(tmp_path), preserve_storage_dir=True) + reader.restore() + reader.put(np.array([99.0]), periods.period("2020")) + + assert snapshot.get(periods.period("2020")).tolist() == [1.0] + reader.restore() + assert reader.get(periods.period("2020")).tolist() == [99.0] + + +def test_result_calculated_across_its_own_input_change_is_calculated_again(): + """A formula that read an input and then replaced it runs again from the new one.""" + + def result(person, period): + source = person("source", period) + person.simulation.set_input("source", period, np.full(person.count, 3.0)) + return source * 2 + + system = _one_person_system( + _yearly_variable("source", lambda person, period: np.ones(person.count)), + _yearly_variable("result", result), + ) + branch = SimulationBuilder().build_default_simulation(system).get_branch("b") + assert branch.calculate("result", "2020").tolist() == [6.0] + + assert branch.calculate("result", "2020").tolist() == [6.0] + + +def test_result_calculated_again_is_kept_for_carry_over(): + """A result that was not kept would leave its period unknown to carry-over.""" + from policyengine_core.country_template import entities + from policyengine_core.variables import Variable + + class flag(Variable): + value_type = float + entity = entities.Person + definition_period = periods.YEAR + label = "flag" + + class source(Variable): + value_type = float + entity = entities.Person + definition_period = periods.YEAR + label = "source" + + def formula_2021(person, period): + if np.any(person("flag", "2021") == 0): + person.simulation.set_input("flag", "2021", np.ones(person.count)) + return np.full(person.count, 3.0) + + def formula_2022(person, period): + return None + + system = _one_person_system(flag, source) + system.auto_carry_over_input_variables = True + root = SimulationBuilder().build_default_simulation(system) + root.set_input("flag", "2021", np.array([0.0])) + branch = root.get_branch("branch") + + assert branch.calculate("source", "2021").tolist() == [3.0] + assert branch.calculate("source", "2022").tolist() == [3.0] # carried over + + +def test_result_is_calculated_again_until_no_input_changes(): + """Each run may change another input it read; the kept result follows the last.""" + from policyengine_core.country_template import entities + from policyengine_core.variables import Variable + + def flag(name): + return type( + name, + (Variable,), + dict( + value_type=float, + entity=entities.Person, + definition_period=periods.YEAR, + label=name, + ), + ) + + class source(Variable): + value_type = float + entity = entities.Person + definition_period = periods.YEAR + label = "source" + + def formula_2021(person, period): + for name in ("first_flag", "second_flag"): + if np.any(person(name, "2021") == 0): + person.simulation.set_input(name, "2021", np.ones(person.count)) + break + return np.full(person.count, 3.0) + + def formula_2022(person, period): + return None + + system = _one_person_system(flag("first_flag"), flag("second_flag"), source) + system.auto_carry_over_input_variables = True + root = SimulationBuilder().build_default_simulation(system) + root.set_input("first_flag", "2021", np.array([0.0])) + root.set_input("second_flag", "2021", np.array([0.0])) + branch = root.get_branch("branch") + + assert branch.calculate("source", "2021").tolist() == [3.0] + assert branch.calculate("source", "2022").tolist() == [3.0] # carried over + + +def test_direct_sum_is_taken_again_when_a_term_changes_an_earlier_input(): + from policyengine_core.country_template import entities + from policyengine_core.variables import Variable + + class source(Variable): + value_type = float + entity = entities.Person + definition_period = periods.MONTH + label = "source" + + class result(Variable): + value_type = float + entity = entities.Person + definition_period = periods.MONTH + label = "result" + + def formula(person, period): + if str(period) == "2020-02": + january = person("source", "2020-01") + if np.any(january == 1): + person.simulation.set_input( + "source", "2020-01", np.full(person.count, 3.0) + ) + return person("source", period) + + system = _one_person_system(source, result) + root = SimulationBuilder().build_default_simulation(system) + for month in periods.period("2020").get_subperiods(periods.MONTH): + root.set_input("source", month, np.array([1.0])) + branch = root.get_branch("branch") + + assert branch.calculate_add("result", "2020").tolist() == [14.0] + + +def test_input_a_formula_sets_wins_over_a_carried_over_default(): + """A formula returning None after setting its own period's input.""" + from policyengine_core.country_template import entities + from policyengine_core.variables import Variable + + class source(Variable): + value_type = float + entity = entities.Person + definition_period = periods.YEAR + label = "source" + + def formula_2020(person, period): + person.simulation.set_input("source", period, np.full(person.count, 3.0)) + return None + + system = _one_person_system(source) + system.auto_carry_over_input_variables = True + root = SimulationBuilder().build_default_simulation(system) + root.set_input("source", "2021", np.array([9.0])) # later: no carry-over + branch = root.get_branch("branch") + + assert branch.calculate("source", "2020").tolist() == [3.0] + + +def test_input_set_under_the_default_key_during_a_branch_calculation_wins(): + """Holder.set_input stores under "default", which the branch reads.""" + + def source(person, period): + person.simulation.get_holder("source").set_input( + period, np.full(person.count, 3.0) + ) + return np.ones(person.count) + + system = _one_person_system(_yearly_variable("source", source)) + branch = SimulationBuilder().build_default_simulation(system).get_branch("b") + + assert branch.calculate("source", "2020").tolist() == [3.0] + assert branch.calculate("source", "2020").tolist() == [3.0] + + +def test_direct_divide_keeps_ignoring_an_earlier_monthly_input(): + """Only an input set during the calculation wins; one stored before does not.""" + from policyengine_core.country_template import entities + from policyengine_core.variables import Variable + + def store_whole_period(holder, period, array): + holder._set(period, array) + + class flag(Variable): + value_type = float + entity = entities.Person + definition_period = periods.YEAR + label = "flag" + + class source(Variable): + value_type = float + entity = entities.Person + definition_period = periods.YEAR + label = "source" + set_input = store_whole_period + + def formula(person, period): + if np.any(person("flag", "2020") == 0): + person.simulation.set_input("flag", "2020", np.ones(person.count)) + return np.full(person.count, 120.0) + + system = _one_person_system(flag, source) + + def branch(final_inputs): + root = SimulationBuilder().build_default_simulation(system) + root.set_input("flag", "2020", np.array([0.0])) + root.set_input("source", "2020-01", np.array([99.0])) + child = root.get_branch("b") + if final_inputs: + child.set_input("flag", "2020", np.array([1.0])) + return child + + fresh = branch(final_inputs=True).calculate_divide("source", "2020-01") + assert fresh.tolist() == [10.0] + assert branch(final_inputs=False).calculate_divide( + "source", "2020-01" + ).tolist() == [10.0] + + +def test_reruns_do_not_multiply_through_nested_calculations(): + """Only the outermost calculation runs again; the inner ones run with it.""" + from collections import Counter + + calls = Counter() + + def leaf(person, period): + calls["leaf"] += 1 + value = person("source", period) + person.simulation.set_input("source", period, value + 1) # every run + return value + + def level(name, inner): + def formula(person, period): + calls[name] += 1 + return person(inner, period) + + return formula + + system = _one_person_system( + _yearly_variable("source"), + _yearly_variable("level1", leaf), + _yearly_variable("level2", level("level2", "level1")), + _yearly_variable("level3", level("level3", "level2")), + ) + root = SimulationBuilder().build_default_simulation(system) + root.set_input("source", "2020", np.array([0.0])) + branch = root.get_branch("branch") + branch.calculate("level3", "2020") + + reruns = simulation_module._RERUNS_AFTER_INPUT_CHANGE + assert calls == {"leaf": reruns + 1, "level2": reruns + 1, "level3": reruns + 1} + assert branch._calculations_in_flight == 0 + + +def test_input_handler_calling_calculate_directly_still_drops(): + """A private _calculate (here a carry-over, no formula) also counts as calculating.""" + from policyengine_core.country_template import entities + from policyengine_core.variables import Variable + + def two_months(holder, period, array): + holder._set(period.first_month, array) + holder.simulation._calculate("source", periods.period("2020-03")) + holder._set(period.first_month.offset(1), array) + + class source(Variable): + value_type = float + entity = entities.Person + definition_period = periods.MONTH + label = "source" + set_input = two_months + + system = _one_person_system(source) + system.auto_carry_over_input_variables = True + + def branch(handler_input): + root = SimulationBuilder().build_default_simulation(system) + root.set_input("source", "2020-01", np.array([1.0])) + root.set_input("source", "2020-02", np.array([1.0])) + child = root.get_branch("branch") + if handler_input: + child.set_input("source", "2020", np.array([3.0])) + else: + child.set_input("source", "2020-01", np.array([3.0])) + child.set_input("source", "2020-02", np.array([3.0])) + return child + + assert branch(handler_input=False).calculate("source", "2020-03").tolist() == [3.0] + assert branch(handler_input=True).calculate("source", "2020-03").tolist() == [3.0] + + +@pytest.mark.parametrize("mode", ["calculate", "add", "divide"]) +def test_result_refused_in_one_simulation_is_not_kept_by_its_caller_in_another(mode): + """A parent formula calling into a branch whose formula changes its own input.""" + from policyengine_core.country_template import entities + from policyengine_core.variables import Variable + + def variable(name, formula=None, definition_period=periods.YEAR): + namespace = dict( + value_type=float, + entity=entities.Person, + definition_period=definition_period, + label=name, + ) + if formula is not None: + namespace["formula"] = formula + return type(name, (Variable,), namespace) + + def leaf(person, period): + previous = person("source", period.this_year) + if np.any(previous != 3.0): + person.simulation.set_input( + "source", period.this_year, np.full(person.count, 3.0) + ) + return previous * 2 + + def from_child(person, period): + worker = person.simulation.get_branch("worker") + if mode == "add": + return worker.calculate_add("leaf", period) + if mode == "divide": + return worker.calculate_divide("leaf", period.first_month) + return worker.calculate("leaf", period) + + def child_outer(person, period): + return person.simulation.parent_branch.calculate("from_child", period) + + def result_with(input_first): + def result(person, period): + simulation = person.simulation + worker = simulation.get_branch("worker") + try: + if input_first: + worker.set_input("source", period, np.full(person.count, 3.0)) + return worker.calculate("child_outer", period) + finally: + del simulation.branches["worker"] + + return result + + def run(input_first): + system = _one_person_system( + variable("source"), + variable("leaf", leaf, periods.MONTH if mode == "add" else periods.YEAR), + variable("from_child", from_child), + variable("child_outer", child_outer), + variable("result", result_with(input_first)), + ) + root = SimulationBuilder().build_default_simulation(system) + root.set_input("source", "2020", np.array([1.0])) + return ( + root.calculate("result", "2020").tolist(), + root.calculate("from_child", "2020").tolist(), + ) + + assert run(input_first=False) == run(input_first=True) + + +def test_input_stored_with_an_earlier_number_still_wins(): + """Inputs are numbered when stored, whatever number a caller passes.""" + from policyengine_core.country_template import entities + from policyengine_core.variables import Variable + + def january_only(holder, period, array): + holder._set(period.first_month, array, sequence_number=1) + + class source(Variable): + value_type = float + entity = entities.Person + definition_period = periods.MONTH + label = "source" + set_input = january_only + + def formula(person, period): + person.simulation.set_input("source", "2020", np.full(person.count, 42.0)) + return np.ones(person.count) + + system = _one_person_system(source) + branch = SimulationBuilder().build_default_simulation(system).get_branch("b") + + assert branch.calculate("source", "2020-01").tolist() == [42.0] + assert branch.get_holder("source")._is_input(periods.period("2020-01"), "b") + + +def test_direct_sum_runs_its_terms_again_only_as_a_whole(): + from collections import Counter + + from policyengine_core.country_template import entities + from policyengine_core.variables import Variable + + calls = Counter() + + class source(Variable): + value_type = float + entity = entities.Person + definition_period = periods.MONTH + label = "source" + + class counted(Variable): + value_type = float + entity = entities.Person + definition_period = periods.MONTH + label = "counted" + + def formula(person, period): + calls["counted"] += 1 + value = person("source", period) + person.simulation.set_input("source", period, value + 1) # every run + return value + + system = _one_person_system(source, counted) + root = SimulationBuilder().build_default_simulation(system) + root.set_input("source", "2020-01", np.array([0.0])) + branch = root.get_branch("branch") + branch.calculate_add("counted", "2020-01") + + assert calls["counted"] == simulation_module._RERUNS_AFTER_INPUT_CHANGE + 1 + + +def test_input_change_in_an_unrelated_simulation_leaves_this_one_caching(): + """A settled result from another family's simulation is kept, so carry-over finds it.""" + from policyengine_core.country_template import entities + from policyengine_core.variables import Variable + + def leaf(person, period): + if np.any(person("source", period) != 3): + person.simulation.set_input("source", period, np.full(person.count, 3.0)) + return person("source", period) * 2 + + other_root = SimulationBuilder().build_default_simulation( + _one_person_system(_yearly_variable("source"), _yearly_variable("leaf", leaf)) + ) + other_root.set_input("source", "2020", np.array([1.0])) + other = other_root.get_branch("worker") + + class result(Variable): + value_type = float + entity = entities.Person + definition_period = periods.YEAR + label = "result" + + def formula_2020(person, period): + return other.calculate("leaf", "2020") + + def formula_2021(person, period): + return None + + system = _one_person_system(result) + system.auto_carry_over_input_variables = True + branch = SimulationBuilder().build_default_simulation(system).get_branch("consumer") + + assert branch.calculate("result", "2020").tolist() == [6.0] + assert branch.calculate("result", "2021").tolist() == [6.0] # carried over + + +def test_refused_result_is_not_kept_by_a_caller_in_another_family(): + """A worker (its own family) calls back into the consumer, which calls the worker again.""" + from policyengine_core.country_template import entities + from policyengine_core.variables import Variable + + def variable(name, formula=None): + namespace = dict( + value_type=float, + entity=entities.Person, + definition_period=periods.YEAR, + label=name, + ) + if formula is not None: + namespace["formula"] = formula + return type(name, (Variable,), namespace) + + def run(input_first): + holders = {} + + def leaf(person, period): + previous = person("source", period) + if np.any(previous != 3.0): + person.simulation.set_input( + "source", period, np.full(person.count, 3.0) + ) + return previous * 2 + + def worker_outer(person, period): + return holders["consumer"].calculate("from_worker", period) + + worker_root = SimulationBuilder().build_default_simulation( + _one_person_system( + variable("source"), + variable("leaf", leaf), + variable("worker_outer", worker_outer), + ) + ) + worker_root.set_input("source", "2020", np.array([1.0])) + worker = worker_root.get_branch("worker") + if input_first: + worker.set_input("source", "2020", np.array([3.0])) + + consumer = ( + SimulationBuilder() + .build_default_simulation( + _one_person_system( + variable( + "from_worker", + lambda person, period: worker.calculate("leaf", period), + ), + variable( + "result", + lambda person, period: worker.calculate("worker_outer", period), + ), + ) + ) + .get_branch("consumer") + ) + holders["consumer"] = consumer + return ( + consumer.calculate("result", "2020").tolist(), + consumer.calculate("from_worker", "2020").tolist(), + ) + + assert run(input_first=False) == run(input_first=True) == ([6.0], [6.0]) + + +def test_caller_in_the_same_family_does_not_keep_what_it_read_before_the_change(): + """A parent read its branch, then called a branch formula that changed the branch's input.""" + + def run(input_first): + def changes_source(person, period): + if np.any(person("source", period) != 3.0): + person.simulation.set_input( + "source", period, np.full(person.count, 3.0) + ) + return np.zeros(person.count) + + def result(person, period): + branch = person.simulation.branches["kept"] + doubled = branch.calculate("doubled", period) # read before the change + branch.calculate("changes_source", period) + return doubled + + system = _one_person_system( + _yearly_variable("source"), + _yearly_variable( + "doubled", lambda person, period: person("source", period) * 2 + ), + _yearly_variable("changes_source", changes_source), + _yearly_variable("result", result), + ) + root = SimulationBuilder().build_default_simulation(system) + root.set_input("source", "2020", np.array([1.0])) + kept = root.get_branch("kept") + if input_first: + kept.set_input("source", "2020", np.array([3.0])) + first = root.calculate("result", "2020").tolist() + return first, root.calculate("result", "2020").tolist() + + live_first, live_next = run(input_first=False) + assert live_first == [2.0] # documented: what it read stays read + assert live_next == run(input_first=True)[1] == [6.0] # but it was not kept + + +def _worker_and_consumer(leaf, auto_carry_over=False, worker_inputs=()): + """A worker family whose ``worker_outer`` calls back into a consumer's ``from_worker``.""" + from policyengine_core.country_template import entities + from policyengine_core.variables import Variable + + def variable(name, formula=None): + namespace = dict( + value_type=float, + entity=entities.Person, + definition_period=periods.YEAR, + label=name, + ) + if formula is not None: + namespace["formula"] = formula + return type(name, (Variable,), namespace) + + def run(input_first): + holders = {} + worker_system = _one_person_system( + variable("source"), + variable("leaf", leaf), + variable( + "worker_outer", + lambda person, period: holders["consumer"].calculate( + "from_worker", period + ), + ), + ) + worker_system.auto_carry_over_input_variables = auto_carry_over + worker_root = SimulationBuilder().build_default_simulation(worker_system) + worker_root.set_input("source", "2020", np.array([1.0])) + for variable_name, period, value in worker_inputs: + worker_root.set_input(variable_name, period, np.array([value])) + worker = worker_root.get_branch("worker") + if input_first: + worker.set_input("source", "2020", np.array([3.0])) + consumer = ( + SimulationBuilder() + .build_default_simulation( + _one_person_system( + variable( + "from_worker", + lambda person, period: worker.calculate("leaf", period), + ), + variable( + "result", + lambda person, period: worker.calculate("worker_outer", period), + ), + ) + ) + .get_branch("consumer") + ) + holders["consumer"] = consumer + first = consumer.calculate("result", "2020").tolist() + return first, consumer.calculate("result", "2020").tolist() + + return run + + +def test_refused_result_that_returns_early_still_taints_its_caller(): + """A carry-over default returned before any cache decision, after an input change.""" + + def leaf(person, period): + if np.any(person("source", period) != 3): + person.simulation.set_input("source", period, np.full(person.count, 3.0)) + return None # carry-over finds only 2021: the default, early + return np.full(person.count, 6.0) + + run = _worker_and_consumer( + leaf, auto_carry_over=True, worker_inputs=[("leaf", "2021", 9.0)] + ) + + assert run(input_first=False) == run(input_first=True) == ([6.0], [6.0]) + + +def test_caller_that_catches_an_error_after_an_input_change_does_not_keep_its_result(): + from policyengine_core.country_template import entities + from policyengine_core.variables import Variable + + def leaf(person, period): + if np.any(person("source", period) != 3): + person.simulation.set_input("source", period, np.full(person.count, 3.0)) + raise ValueError("changed") + return np.full(person.count, 6.0) + + worker_root = SimulationBuilder().build_default_simulation( + _one_person_system(_yearly_variable("source"), _yearly_variable("leaf", leaf)) + ) + worker_root.set_input("source", "2020", np.array([1.0])) + worker = worker_root.get_branch("worker") + + def result(person, period): + try: + return worker.calculate("leaf", period) + except ValueError: + return np.zeros(person.count) + + consumer = ( + SimulationBuilder() + .build_default_simulation( + _one_person_system(_yearly_variable("result", result)) + ) + .get_branch("consumer") + ) + + assert consumer.calculate("result", "2020").tolist() == [0.0] # its own fallback + assert consumer.calculate("result", "2020").tolist() == [6.0] # not kept + + +@pytest.mark.parametrize("worker_kind", ["other_family", "own_branch_in_a_thread"]) +def test_earlier_read_is_not_kept_after_the_worker_changes_its_input(worker_kind): + """What a formula read before stays read, but its result is not kept.""" + import threading + + def changes_source(person, period): + if np.any(person("source", period) != 3.0): + person.simulation.set_input("source", period, np.full(person.count, 3.0)) + return np.zeros(person.count) + + worker_variables = ( + _yearly_variable("source"), + _yearly_variable( + "doubled", lambda person, period: person("source", period) * 2 + ), + _yearly_variable("changes_source", changes_source), + ) + + if worker_kind == "other_family": + worker_root = SimulationBuilder().build_default_simulation( + _one_person_system(*worker_variables) + ) + worker_root.set_input("source", "2020", np.array([1.0])) + worker = worker_root.get_branch("worker") + + def result(person, period): + doubled = worker.calculate("doubled", period) + worker.calculate("changes_source", period) + return doubled + + simulation = SimulationBuilder().build_default_simulation( + _one_person_system(_yearly_variable("result", result)) + ) + else: + + def result(person, period): + branch = person.simulation.branches["kept"] + doubled = branch.calculate("doubled", period) + thread = threading.Thread( + target=branch.calculate, args=("changes_source", period) + ) + thread.start() + thread.join() + return doubled + + simulation = SimulationBuilder().build_default_simulation( + _one_person_system(*worker_variables, _yearly_variable("result", result)) + ) + simulation.set_input("source", "2020", np.array([1.0])) + simulation.get_branch("kept") + + assert simulation.calculate("result", "2020").tolist() == [2.0] # read before + assert simulation.calculate("result", "2020").tolist() == [6.0] # not kept + + +def test_settled_call_into_own_branch_keeps_the_result_for_carry_over(): + """The branch's change settles before returning: nothing the caller holds is old.""" + from policyengine_core.country_template import entities + from policyengine_core.variables import Variable + + def leaf(person, period): + if np.any(person("source", period) != 3): + person.simulation.set_input("source", period, np.full(person.count, 3.0)) + return person("source", period) * 2 + + class result(Variable): + value_type = float + entity = entities.Person + definition_period = periods.YEAR + label = "result" + + def formula_2020(person, period): + return person.simulation.branches["worker"].calculate("leaf", "2020") + + def formula_2021(person, period): + return None + + system = _one_person_system( + _yearly_variable("source"), _yearly_variable("leaf", leaf), result + ) + system.auto_carry_over_input_variables = True + root = SimulationBuilder().build_default_simulation(system) + root.set_input("source", "2020", np.array([1.0])) + root.get_branch("worker") + + assert root.calculate("result", "2020").tolist() == [6.0] + assert root.calculate("result", "2021").tolist() == [6.0] # carried over + + +def test_input_a_formula_sets_for_its_own_period_wins(): + """As it would had the input been set before the calculation.""" + + def source(person, period): + person.simulation.set_input("source", period, np.full(person.count, 3.0)) + return np.ones(person.count) + + system = _one_person_system(_yearly_variable("source", source)) + branch = SimulationBuilder().build_default_simulation(system).get_branch("b") + + assert branch.calculate("source", "2020").tolist() == [3.0] + assert branch.calculate("source", "2020").tolist() == [3.0] + assert branch.get_holder("source")._is_input(periods.period("2020"), "b") + + +def test_input_handler_that_calculates_between_its_stores(): + """Values a handler calculates before its last store are dropped after it.""" + from policyengine_core.country_template import entities + from policyengine_core.variables import Variable + + def two_months(holder, period, array): + holder._set(period.first_month, array) + holder.simulation.calculate("total", period) + holder._set(period.first_month.offset(1), array) + + class source(Variable): + value_type = float + entity = entities.Person + definition_period = periods.MONTH + label = "source" + set_input = two_months + + class total(Variable): + value_type = float + entity = entities.Person + definition_period = periods.YEAR + label = "total" + + def formula(person, period): + first = period.first_month + return person("source", first) + person("source", first.offset(1)) + + system = _one_person_system(source, total) + root = SimulationBuilder().build_default_simulation(system) + root.set_input("source", "2020-01", np.array([1.0])) + root.set_input("source", "2020-02", np.array([1.0])) + branch = root.get_branch("branch") + branch.set_input("source", "2020", np.array([3.0])) + + assert branch.calculate("total", "2020").tolist() == [6.0] + + +def test_input_handler_that_fails_after_calculating(): + """The inputs it stored stay, so what it calculated before them drops.""" + from policyengine_core.country_template import entities + from policyengine_core.variables import Variable + + def two_months_then_fail(holder, period, array): + holder._set(period.first_month, array) + holder.simulation.calculate("total", period) + holder._set(period.first_month.offset(1), array) + raise RuntimeError("handler failed") + + class source(Variable): + value_type = float + entity = entities.Person + definition_period = periods.MONTH + label = "source" + set_input = two_months_then_fail + + class total(Variable): + value_type = float + entity = entities.Person + definition_period = periods.YEAR + label = "total" + + def formula(person, period): + first = period.first_month + return person("source", first) + person("source", first.offset(1)) + + system = _one_person_system(source, total) + root = SimulationBuilder().build_default_simulation(system) + root.set_input("source", "2020-01", np.array([1.0])) + root.set_input("source", "2020-02", np.array([1.0])) + branch = root.get_branch("branch") + with pytest.raises(RuntimeError, match="handler failed"): + branch.set_input("source", "2020", np.array([3.0])) + + assert branch.calculate("total", "2020").tolist() == [6.0] + + +def test_branch_input_stops_macro_reads_before_any_read(monkeypatch): + """A macro file for the branch's name may come from a branch without the input.""" + import policyengine_core.simulations.simulation as simulation_module + + class CachedResult: + def __init__(self, tax_benefit_system): + pass + + def set_cache_path(self, *args): + pass + + def get_cache_path(self): + return type("Path", (), {"exists": lambda self: True})() + + def get_cache_value(self, path): + return np.array([2.0]) # written by another simulation's "override" + + def set_cache_value(self, path, value): + pass + + monkeypatch.setattr(simulation_module, "SimulationMacroCache", CachedResult) + system = _one_person_system( + _yearly_variable("source", lambda person, period: np.ones(person.count)), + _yearly_variable("result", lambda person, period: person("source", period) * 2), + _yearly_variable("total", lambda person, period: person("result", period) * 2), + ) + simulation = SimulationBuilder().build_default_simulation(system) + system.data_modified = False + simulation.macro_cache_read = True + simulation.dataset = type( + "Dataset", (), {"file_path": simulation_module.Path("cache"), "name": "d"} + )() + monkeypatch.setattr( + type(simulation), + "check_macro_cache", + lambda self, variable, period: variable == "result", + ) + + branch = simulation.get_branch("override") + branch.set_input("source", "2020", np.array([3.0])) # before any calculation + + assert branch.calculate("total", "2020").tolist() == [12.0] + + +def test_result_obsolete_after_its_own_input_change_is_not_written_to_the_macro_cache( + monkeypatch, +): + import policyengine_core.simulations.simulation as simulation_module + + written = [] + + class RecordingCache: + def __init__(self, tax_benefit_system): + pass + + def set_cache_path(self, *args): + pass + + def get_cache_path(self): + return type("Path", (), {"exists": lambda self: False})() + + def get_cache_value(self, path): + return None + + def set_cache_value(self, path, value): + written.append(np.array(value).tolist()) + + def result(person, period): + source = person("source", period) + if np.any(source == 1): + person.simulation.set_input("source", period, np.full(person.count, 3.0)) + return source * 2 + + monkeypatch.setattr(simulation_module, "SimulationMacroCache", RecordingCache) + system = _one_person_system( + _yearly_variable("source", lambda person, period: np.ones(person.count)), + _yearly_variable("result", result), + ) + simulation = SimulationBuilder().build_default_simulation(system) + simulation.dataset = type( + "Dataset", (), {"file_path": simulation_module.Path("cache"), "name": "d"} + )() + monkeypatch.setattr( + type(simulation), + "check_macro_cache", + lambda self, variable, period: variable == "result", + ) + branch = simulation.get_branch("b") + + assert branch.calculate("result", "2020").tolist() == [6.0] + assert [2.0] not in written # the first run's result, from the old input + + +def test_macro_cache_read_counts_as_depending_on_every_input(monkeypatch): + """What a macro-cache value was calculated from was never calculated here.""" + import policyengine_core.simulations.simulation as simulation_module + + class CachedResult: + def __init__(self, tax_benefit_system): + pass + + def set_cache_path(self, *args): + pass + + def get_cache_path(self): + return type("Path", (), {"exists": lambda self: True})() + + def get_cache_value(self, path): + return np.array([2.0]) + + def set_cache_value(self, path, value): + pass + + monkeypatch.setattr(simulation_module, "SimulationMacroCache", CachedResult) + system = _one_person_system( + _yearly_variable("source", lambda person, period: np.ones(person.count)), + _yearly_variable("result", lambda person, period: person("source", period) * 2), + _yearly_variable("total", lambda person, period: person("result", period) * 2), + ) + simulation = SimulationBuilder().build_default_simulation(system) + system.data_modified = False # adding variables marks the system modified + simulation.macro_cache_read = True + simulation.dataset = type( + "Dataset", (), {"file_path": simulation_module.Path("cache"), "name": "d"} + )() + monkeypatch.setattr( + type(simulation), + "check_macro_cache", + lambda self, variable, period: variable == "result", + ) + assert simulation.calculate("total", "2020").tolist() == [4.0] + assert not simulation.get_holder("source").get_known_periods() # never calculated + + branch = simulation.get_branch("branch") + branch.set_input("source", "2020", np.array([3.0])) # never calculated here + + assert branch.calculate("total", "2020").tolist() == [12.0] + + +def test_history_hand_back_failure_still_ends_the_calculation(monkeypatch): + def fail(self): + raise RuntimeError("hand-back failed") + + system = _one_person_system( + _yearly_variable("source", lambda person, period: np.ones(person.count)) + ) + simulation = SimulationBuilder().build_default_simulation(system) + monkeypatch.setattr(type(simulation), "_share_store_history_with_caller", fail) + + with pytest.raises(RuntimeError, match="hand-back failed"): + simulation.calculate("source", "2020") + + assert simulation._calculations_in_flight == 0 + assert not simulation.tracer.stack + + +def test_failed_calculation_in_another_simulation_hands_its_history_back(): + """Whether a calculation fails can depend on what it read.""" + + def checked(person, period): + source = person("source", period) + if (source == 0).any(): + raise ValueError("source is zero") + return source * 2 + + def result(person, period): + simulation = person.simulation + child = simulation.get_branch("child") + try: + return child.calculate("checked", period) + except ValueError: + return np.full(person.count, 7.0) + finally: + del simulation.branches["child"] + + system = _one_person_system( + _yearly_variable("source", lambda person, period: np.zeros(person.count)), + _yearly_variable("checked", checked), + _yearly_variable("result", result), + ) + simulation = SimulationBuilder().build_default_simulation(system) + assert simulation.calculate("result", "2020").tolist() == [7.0] + + branch = simulation.get_branch("branch") + branch.set_input("source", "2020", np.array([3.0])) + + assert branch.calculate("result", "2020").tolist() == [6.0] + + +def test_restored_simulation_drops_restored_values_an_input_may_have_fed(tmp_path): + from policyengine_core.tools.simulation_dumper import ( + dump_simulation, + restore_simulation, + ) + + system = _one_person_system( + _yearly_variable("source"), + _yearly_variable("result", lambda person, period: person("source", period) * 2), + ) + simulation = SimulationBuilder().build_default_simulation(system) + simulation.set_input("source", "2020", np.array([1.0])) + simulation.calculate("result", "2020") + dump_simulation(simulation, str(tmp_path / "dump")) + + restored = restore_simulation(str(tmp_path / "dump"), system) + branch = restored.get_branch("branch") + branch.set_input("source", "2020", np.array([3.0])) + + assert branch.calculate("result", "2020").tolist() == [6.0] + assert restored.calculate("source", "2020").tolist() == [1.0] # still an input + + +@pytest.mark.parametrize("read", ["carried_over", "not_kept"]) +def test_restored_values_drop_without_records_of_what_they_read(tmp_path, read): + """A dump keeps values, not what each was calculated from.""" + from policyengine_core.tools.simulation_dumper import ( + dump_simulation, + restore_simulation, + ) + + system = _one_person_system( + _yearly_variable("source"), + _yearly_variable( + "result", lambda person, period: person("source", period) * 2 + 7 + ), + ) + simulation = SimulationBuilder().build_default_simulation(system) + if read == "carried_over": + system.auto_carry_over_input_variables = True + simulation.set_input("source", "2020", np.array([1.0])) + override, fresh_inputs = "2021", {"2020": 1.0, "2021": 3.0} + else: + system.cache_blacklist = {"source"} + simulation.opt_out_cache = True + override, fresh_inputs = "2022", {"2022": 3.0} + simulation.calculate("result", "2022") + dump_simulation(simulation, str(tmp_path / "dump")) + + restored = restore_simulation(str(tmp_path / "dump"), system) + branch = restored.get_branch("branch") + branch.set_input("source", override, np.array([3.0])) + + fresh = SimulationBuilder().build_default_simulation(system) + for period, value in fresh_inputs.items(): + fresh.set_input("source", period, np.array([value])) + assert ( + branch.calculate("result", "2022").tolist() + == fresh.calculate("result", "2022").tolist() + ) + + +def test_dump_of_a_branch_holds_the_branch_values(tmp_path): + from policyengine_core.tools.simulation_dumper import ( + dump_simulation, + restore_simulation, + ) + + system = _one_person_system( + _yearly_variable("source"), + _yearly_variable("result", lambda person, period: person("source", period) * 2), + ) + simulation = SimulationBuilder().build_default_simulation(system) + simulation.set_input("source", "2020", np.array([1.0])) + simulation.set_input("source", "2021", np.array([5.0])) + branch = simulation.get_branch("branch") + branch.set_input("source", "2020", np.array([3.0])) + branch.calculate("result", "2020") + dump_simulation(branch, str(tmp_path / "dump")) + + restored = restore_simulation(str(tmp_path / "dump"), system) + + assert restored.calculate("source", "2020").tolist() == [3.0] + assert restored.calculate("source", "2021").tolist() == [5.0] # inherited + assert restored.calculate("result", "2020").tolist() == [6.0] + holder = restored.get_holder("source") + assert holder._is_input(periods.period("2020")) + assert holder._is_input(periods.period("2021")) + + +def test_merge_reads_again_what_was_recorded_while_it_merged(): + """A record added to the source during a merge (another thread) is not skipped.""" + source = StoreHistory() + source.record_store("a", periods.period("2020"), 1) + destination = StoreHistory() + original = destination.record_store + calls = [] + + def record_store(variable_name, period, sequence_number): + if not calls: # the other thread records while this merge runs + source.record_store("v", periods.period("2020"), 2) + calls.append(variable_name) + original(variable_name, period, sequence_number) + + destination.record_store = record_store + destination.merge(source) + destination.merge(source) + + assert destination.earliest_dependency("v", periods.period("2020")) == 2 + + +def test_history_pickled_before_journals_still_records(): + import pickle + + history = StoreHistory() + history.record_store("v", periods.period("2020"), 1) + del history._journal, history._generation # as pickled before journals + restored = pickle.loads(pickle.dumps(history)) + + restored.record_store("w", periods.period("2020"), 2) + + assert restored.earliest_dependency("w", periods.period("2020")) == 2 + + +def test_disk_restore_reads_older_file_names_and_breaks_time_ties(tmp_path): + import os + + import policyengine_core.data_storage.on_disk_storage as on_disk_storage + from policyengine_core.data_storage import OnDiskStorage + + # A file named before process tokens, and one from another process. + np.save(tmp_path / "default_2020.4.npy", np.array([1.0])) + np.save(tmp_path / "default_2020.cccccccccccc.999999999999.npy", np.array([2.0])) + storage = OnDiskStorage(str(tmp_path), preserve_storage_dir=True) + storage.put(np.array([3.0]), periods.period("2020")) # this process, number lower + for path in tmp_path.glob("*.npy"): + os.utime(path, ns=(10**18, 10**18)) # a coarse clock: every time ties + + reader = OnDiskStorage(str(tmp_path), preserve_storage_dir=True) + reader.restore() + + assert set(reader._files) == {"default_2020"} # one key, old name included + assert reader.get(periods.period("2020")).tolist() == [3.0] + assert on_disk_storage._PROCESS_TOKEN in reader._files["default_2020"] + + +def test_forked_process_gets_its_own_disk_file_token(): + import os + import warnings + + import policyengine_core.data_storage.on_disk_storage as on_disk_storage + + if not hasattr(os, "fork"): + pytest.skip("no fork on this platform") + # A bare fork whose child only writes to a pipe: a process pool forked + # from a multi-threaded test process can deadlock. + read, write = os.pipe() + with warnings.catch_warnings(): + warnings.simplefilter("ignore", DeprecationWarning) + pid = os.fork() + if pid == 0: + try: + os.write(write, on_disk_storage._PROCESS_TOKEN.encode()) + finally: + os._exit(0) + os.close(write) + child_token = os.read(read, 64).decode() + os.close(read) + os.waitpid(pid, 0) + + assert child_token and child_token != on_disk_storage._PROCESS_TOKEN + + +# ----- Only what may depend on the input is dropped ----- # + + +def test_input_on_one_branch_does_not_make_another_drop(tax_benefit_system): + simulation = _build(tax_benefit_system) + sibling = simulation.get_branch("sibling") + sibling.set_input("rent", "2017-02", np.array([500.0])) + sibling.calculate("housing_allowance", "2017-02") + simulation.calculate("disposable_income", JANUARY) + held = { + (name, key) + for name in ("disposable_income", "income_tax") + for key in _stored_keys(simulation, name) + } + + branch = simulation.get_branch("branch") + dropped = branch._drop_values_that_may_depend_on("rent", periods.period("2017-02")) + + assert dropped == 0 + assert held <= { + (name, key) + for name in ("disposable_income", "income_tax") + for key in _stored_keys(branch, name) + } + + +def test_branches_that_choose_between_overrides_calculate_shared_values_once(): + """Values the overridden variable cannot reach are calculated once per arm. + + ``choose`` compares ``tax`` in two branches that set it, as + policyengine-us itemization does; ``agi`` does not depend on it. + """ + # A marginal-rate branch after the parent has calculated everything. + simulation = branching_simulation() + simulation.calculate("net", "2020") + FORMULA_RUNS["agi"] = 0 + rate = simulation.calculate("marginal_rate", "2020") + assert FORMULA_RUNS["agi"] == 1 + assert np.allclose(rate, 0.85) + + # A baseline branch calculating after the reform arm has. + FORMULA_RUNS["agi"] = 0 + simulation = branching_simulation() + baseline = simulation.get_branch("baseline") + simulation.calculate("net", "2020") + baseline.calculate("net", "2020") + assert FORMULA_RUNS["agi"] == 2 + + fresh = branching_simulation() + assert np.array_equal( + baseline.calculate("net", "2020"), fresh.calculate("net", "2020") + ) + + +# ----- Other paths ----- # + + +def test_prerequisite_requested_before_a_branch_input_still_counts(): + """``requires_computation_after`` holds if the prerequisite was requested, even if dropped.""" + from policyengine_core.entities import build_entity + from policyengine_core.taxbenefitsystems import TaxBenefitSystem + from policyengine_core.variables import Variable + + person = build_entity("person", "persons", "Person", is_person=True) + + class source(Variable): + label = "Source" + value_type = float + entity = person + definition_period = periods.YEAR + + class prerequisite(Variable): + label = "Prerequisite" + value_type = float + entity = person + definition_period = periods.YEAR + + def formula(person, period): + return person("source", period) * 0 + 1 + + class result(Variable): + label = "Result" + value_type = float + entity = person + definition_period = periods.YEAR + requires_computation_after = "prerequisite" + + def formula(person, period): + return person("source", period) * 2 + + system = TaxBenefitSystem([person]) + system.add_variables(source, prerequisite, result) + simulation = SimulationBuilder().build_default_simulation(system) + simulation.set_input("source", "2022", np.array([1.0])) + simulation.calculate("prerequisite", "2022") + simulation.calculate("result", "2022") + + branch = simulation.get_branch("branch") + branch.set_input("source", "2022", np.array([2.0])) + + assert branch.calculate("result", "2022").tolist() == [4.0] + + +def test_values_a_custom_input_handler_calculates_are_not_inputs(): + from policyengine_core.entities import build_entity + from policyengine_core.taxbenefitsystems import TaxBenefitSystem + from policyengine_core.variables import Variable + + person = build_entity("person", "persons", "Person", is_person=True) + + def store_first_month_then_calculate(holder, period, array): + holder._set(period.first_month, array) + holder.simulation.calculate("dependent", period) + + class dispatched(Variable): + label = "Dispatched" + value_type = float + entity = person + definition_period = periods.MONTH + set_input = store_first_month_then_calculate + + class dependent(Variable): + label = "Dependent" + value_type = float + entity = person + definition_period = periods.YEAR + + def formula(person, period): + return 2 * person("dispatched", period.first_month) + + system = TaxBenefitSystem([person]) + system.add_variables(dispatched, dependent) + branch = SimulationBuilder().build_default_simulation(system).get_branch("branch") + branch.set_input("dispatched", "2020", np.array([3.0])) + assert branch.calculate("dependent", "2020").tolist() == [6.0] + + branch.set_input("dispatched", "2020-01", np.array([4.0])) + + assert branch.calculate("dependent", "2020").tolist() == [8.0] + assert not branch.get_holder("dependent")._memory_storage._input_keys + + +def test_year_input_after_calculating_one_of_its_months(): + """The drop comes before a yearly input is divided, so only input months count.""" + simulation = synthetic_simulation(ROOT_INPUTS) + branch = simulation.get_branch("branch") + branch.calculate("p_month", "2015-03") # 2015 has no p_m input: default + + branch.set_input("p_m", "2015", np.array([120.0, 120.0, 120.0])) + + fresh = synthetic_simulation( + {**ROOT_INPUTS, **{("p_m", f"2015-{m:02d}"): (10.0,) * 3 for m in range(1, 13)}} + ) + assert np.allclose( + branch.calculate("p_month", "2015-03"), fresh.calculate("p_month", "2015-03") + ) + + +def test_derivative_keeps_inputs_set_after_construction(): + """``derivative`` keeps every input of the simulation it differentiates.""" + simulation = synthetic_simulation({("p_a", "2013"): (1.0, 1.0, 1.0)}) + branch = simulation.get_branch("branch") + branch.set_input("p_b", "2013", np.array([9.0, 9.0, 9.0])) + + assert np.allclose(branch.derivative("p_sum", "p_b", "2013"), 2.0) + assert np.array_equal(branch.calculate("p_b", "2013"), [9.0, 9.0, 9.0]) + + +def test_disk_branch_keeps_its_values_when_its_parent_recalculates(): + """Each disk store writes its own file, so a recalculation leaves children's files.""" + config = MemoryConfig(max_memory_occupation=0) + simulation = synthetic_simulation(ROOT_INPUTS, memory_config=config) + branch = simulation.get_branch("branch") + before = branch.calculate("p_sum", "2013").copy() + child = branch.get_branch("child") + + branch.set_input("p_a", "2013", np.zeros(3)) + branch.calculate("p_sum", "2013") # recalculated: a new file + + assert np.array_equal(child.calculate("p_sum", "2013"), before) + + +def test_disk_branches_with_the_same_name_keep_their_own_values(): + config = MemoryConfig(max_memory_occupation=0) + simulation = synthetic_simulation(ROOT_INPUTS, memory_config=config) + first = simulation.get_branch("a").get_branch("leaf") + second = simulation.get_branch("b").get_branch("leaf") + first.set_input("p_a", "2013", np.full(3, 3.0)) + second.set_input("p_a", "2013", np.full(3, 4.0)) + # A name reused after its branch is forgotten. + old = simulation.get_branch("reused") + old.set_input("p_a", "2013", np.full(3, 7.0)) + del simulation.branches["reused"] + new = simulation.get_branch("reused") + new.set_input("p_a", "2013", np.full(3, 8.0)) + + assert np.array_equal(first.calculate("p_a", "2013"), [3.0] * 3) + assert np.array_equal(second.calculate("p_a", "2013"), [4.0] * 3) + assert np.array_equal(old.calculate("p_a", "2013"), [7.0] * 3) + assert np.array_equal(new.calculate("p_a", "2013"), [8.0] * 3) + + +def test_values_unpickled_from_another_process_stay_earlier(monkeypatch): + """Sequence numbers restart in each process; unpickling moves this one past them.""" + import itertools + import pickle + + import policyengine_core.data_storage.store_history as store_history + + storage = InMemoryStorage(is_eternal=False) + history = StoreHistory() + for _ in range(100): + next_sequence_number() + storage.put(np.array([1.0]), periods.period("2020")) + history.record_store("v", periods.period("2020"), next_sequence_number()) + source = StoreHistory() # kept alive: its weak read position must not stop pickling + history.merge(source) + assert len(history._merged) == 1 + payloads = pickle.dumps(storage), pickle.dumps(history) + + monkeypatch.setattr(store_history, "_sequence", itertools.count(1)) # a new process + restored_storage, restored_history = (pickle.loads(p) for p in payloads) + restored_storage.put(np.array([2.0]), periods.period("2021")) + + numbers = restored_storage._sequence_numbers + assert numbers["default:2021"] > numbers["default:2020"] + assert next_sequence_number() > restored_history.earliest_dependency( + "v", periods.period("2020") + ) + + +def test_storage_pickled_before_stores_were_numbered_still_stores(): + import pickle + + storage = InMemoryStorage(is_eternal=False) + storage.put(np.array([1.0]), periods.period("2020")) + del storage._sequence_numbers, storage._input_keys # as pickled before + restored = pickle.loads(pickle.dumps(storage)) + + restored.put(np.array([2.0]), periods.period("2021"), is_input=True) + + assert restored._input_keys == {"default:2021"} diff --git a/tests/core/test_branch_input_invalidation_property.py b/tests/core/test_branch_input_invalidation_property.py new file mode 100644 index 00000000..c76d7c44 --- /dev/null +++ b/tests/core/test_branch_input_invalidation_property.py @@ -0,0 +1,303 @@ +"""Branch results do not depend on what was calculated before. + +Random sequences of calculations, branches (including reused and forgotten +names), inputs, ``drop_computed_arrays`` and dumps restored as new +simulations run on a synthetic system, with +values held in memory, on disk, not kept, or blacklisted. Every calculation +in a branch must equal the same calculation in a new simulation given the +branch's inputs, its own and those it inherited when it was created, before +calculating anything. The inputs a branch should hold are modelled here, not +read back from the branch. ``test_branch_input_invalidation.py`` pins the +same behaviour with examples. +""" + +from __future__ import annotations + +import os +import shutil +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.experimental import MemoryConfig +from policyengine_core.tools.simulation_dumper import ( + dump_simulation, + restore_simulation, +) +from tests.fixtures.branch_input_invalidation import ( + PEOPLE, + ROOT_INPUTS, + SYNTHETIC_SYSTEM, + YEARS, + synthetic_simulation, +) + +SETTABLE = [ + "p_a", + "p_b", + "p_up", + "p_c", + "p_m", + "p_sum", + "p_prod", + "p_switch", + "p_inner_only", +] +CALCULABLE = [ + "p_sum", + "p_prod", + "p_cond", + "p_lag", + "p_months", + "p_inner", + "p_inner_only", + "p_imported_up", + "p_up", + "p_c", +] +MONTHS = ["2013-01", "2013-07", "2014-12", "2015-03"] +# How the root simulation (and so every branch) stores values. +MODES = { + "memory": dict(), + "disk": dict(memory_config=lambda: MemoryConfig(max_memory_occupation=0)), + "not_kept": dict( + memory_config=lambda: MemoryConfig( + max_memory_occupation=1, variables_to_drop=["p_sum", "p_inner_only"] + ) + ), + # The synthetic system blacklists p_sum and p_inner_only. + "blacklist": dict(opt_out_cache=True), +} + +# Indices pick a simulation modulo how many exist (-1: the newest); small ones +# keep most operations on the root and the first few branches, where they +# interact. A branch is named "x" or "y" (so two lineages can share a name, +# and a forgotten name can be reused) or given a name of its own (None). +simulation_index = st.integers(0, 3) +branch_name = st.sampled_from(["x", "y", None, None]) +values = st.tuples(*[st.integers(0, 100).map(float)] * PEOPLE) +calculate = st.one_of( + st.tuples( + st.just("calculate"), + simulation_index, + st.sampled_from(CALCULABLE), + st.sampled_from(YEARS), + ), + st.tuples( + st.just("calculate"), + simulation_index, + st.just("p_month"), + st.sampled_from(MONTHS), + ), +) +set_input = st.one_of( + st.tuples( + st.just("set"), + simulation_index, + st.sampled_from(SETTABLE), + st.sampled_from(YEARS), + values, + ), + st.tuples( + st.just("set"), + simulation_index, + st.sampled_from(["p_m", "p_month"]), + st.sampled_from(MONTHS), + values, + ), +) +single_operation = st.one_of( + calculate, + calculate, + set_input, + set_input, + st.tuples(st.just("branch"), simulation_index, branch_name), + st.tuples(st.just("forget"), simulation_index), + st.tuples(st.just("drop"), simulation_index), + # Dump a simulation (a branch, say) and restore it as a new one. + st.tuples(st.just("dump"), simulation_index), +) +# Pairs of a calculated value and an input it depends on, each reaching it a +# different way: through a formula's own branch, uprating from an earlier +# year (directly, and through a formula's own branch), a month of a year, a +# value the holder may not keep, a lagged year, and a year summed from months. +DEPENDENCIES = [ + ("p_inner", "2013", "p_inner_only", "2013"), + ("p_prod", "2015", "p_up", "2014"), + ("p_imported_up", "2015", "p_up", "2014"), + ("p_month", "2013-07", "p_m", "2013"), + ("p_month", "2015-03", "p_m", "2015"), + ("p_prod", "2013", "p_sum", "2013"), + ("p_lag", "2014", "p_a", "2013"), + ("p_cond", "2014", "p_b", "2014"), + ("p_months", "2014", "p_m", "2014-12"), +] +# Calculate a value, branch, override an input it depends on in the new +# branch (index -1), and calculate the value again there; or branch first and +# calculate the value in the branch before and after overriding the input. +override_after_calculating = st.tuples( + simulation_index, st.sampled_from(DEPENDENCIES), values, st.booleans() +).map( + lambda drawn: ( + [ + ("calculate", drawn[0], drawn[1][0], drawn[1][1]), + ("branch", drawn[0], None), + ("set", -1, drawn[1][2], drawn[1][3], drawn[2]), + ("calculate", -1, drawn[1][0], drawn[1][1]), + ] + if drawn[3] + else [ + ("branch", drawn[0], None), + ("calculate", -1, drawn[1][0], drawn[1][1]), + ("set", -1, drawn[1][2], drawn[1][3], drawn[2]), + ("calculate", -1, drawn[1][0], drawn[1][1]), + ] + ) +) +# The same through a dump: calculate the value, dump that simulation, restore +# it, branch from the restored one, override the input and calculate again. +override_after_restoring = st.tuples( + simulation_index, st.sampled_from(DEPENDENCIES), values +).map( + lambda drawn: [ + ("calculate", drawn[0], drawn[1][0], drawn[1][1]), + ("dump", drawn[0]), + ("branch", -1, None), + ("set", -1, drawn[1][2], drawn[1][3], drawn[2]), + ("calculate", -1, drawn[1][0], drawn[1][1]), + ] +) +operations = st.lists( + st.one_of( + single_operation.map(lambda operation: [operation]), + override_after_calculating, + override_after_restoring, + ), + min_size=1, + max_size=15, +).map(lambda chunks: [operation for chunk in chunks for operation in chunk]) + + +def _stored_inputs(own_inputs, variable, period, value): + """What ``set_input(variable, period, value)`` stores on a branch, by its own rule. + + Monthly ``p_m`` divides a yearly value between the months the branch has + not set as inputs itself (what is left after the months it has set). + Returns ``None`` when every month is already set (an error unless the + totals match, so the program skips it). + """ + period = periods.period(period) + variable_period = SYNTHETIC_SYSTEM.get_variable(variable).definition_period + if variable_period != periods.MONTH or period.unit == periods.MONTH: + return {(variable, str(period)): tuple(value)} + months = [str(month) for month in period.get_subperiods(periods.MONTH)] + unset = [month for month in months if (variable, month) not in own_inputs] + if not unset: + return None + remaining = np.asarray(value, dtype=float) - sum( + ( + np.asarray(own_inputs[(variable, month)]) + for month in months + if month not in unset + ), + np.zeros(PEOPLE), + ) + return {(variable, month): tuple(remaining / len(unset)) for month in unset} + + +def _run(program, mode): + """Run ``program``; return each calculation with the inputs it should reflect. + + A branch's inputs are those of the simulation it was created from, as + they were then, and those set on it since. A simulation restored from a + dump has the inputs the dumped one had. + """ + options = MODES[mode] + root = synthetic_simulation( + ROOT_INPUTS, + memory_config=options.get("memory_config", lambda: None)(), + opt_out_cache=options.get("opt_out_cache", False), + ) + simulations = [root] + inputs = [dict(ROOT_INPUTS)] + own_inputs = [{}] + children = {} # (parent index, name) -> index, while the parent keeps it + roots = {0} # the root and simulations restored from dumps + results = [] + for operation in program: + kind, index = operation[0], operation[1] % len(simulations) + if operation[1] == -1: + index = len(simulations) - 1 # The newest simulation. + simulation = simulations[index] + if kind == "branch": + name = operation[2] or f"b{len(simulations)}" + if (index, name) in children or name == simulation.branch_name: + continue # get_branch would return an existing simulation. + simulations.append(simulation.get_branch(name)) + inputs.append(dict(inputs[index])) + own_inputs.append({}) + children[(index, name)] = len(simulations) - 1 + elif kind == "forget": + # The parent forgets a branch (formulas delete theirs); the branch + # object keeps working, and its name can be reused. + names = sorted(name for parent, name in children if parent == index) + if names: + del simulation.branches[names[0]] + del children[(index, names[0])] + elif kind == "set": + if index in roots: + continue # Inputs on a root after calculating are out of scope. + _, _, variable, period, value = operation + stored = _stored_inputs(own_inputs[index], variable, period, value) + if stored is None: + continue + simulation.set_input(variable, period, np.asarray(value)) + inputs[index].update(stored) + own_inputs[index].update(stored) + elif kind == "calculate": + _, _, variable, period = operation + value = np.array(simulation.calculate(variable, period), copy=True) + results.append((dict(inputs[index]), variable, period, value)) + elif kind == "dump": + directory = tempfile.mkdtemp(prefix="policyengine-branch-dump-") + try: + dump_simulation(simulation, directory) + restored = restore_simulation(directory, SYNTHETIC_SYSTEM) + finally: + shutil.rmtree(directory, ignore_errors=True) + simulations.append(restored) + inputs.append(dict(inputs[index])) + own_inputs.append({}) + roots.add(len(simulations) - 1) + else: + simulation.drop_computed_arrays() + return results + + +# CI runs a fixed set of programs; set POLICYENGINE_BRANCH_PROPERTY_EXAMPLES +# to explore new ones (e.g. 5000 for a soak). +_SOAK_EXAMPLES = os.environ.get("POLICYENGINE_BRANCH_PROPERTY_EXAMPLES") + + +@hypothesis.settings( + max_examples=int(_SOAK_EXAMPLES or 200), + derandomize=not _SOAK_EXAMPLES, + database=None, + deadline=None, + suppress_health_check=[hypothesis.HealthCheck.too_slow], +) +@hypothesis.given(program=operations, mode=st.sampled_from(sorted(MODES))) +def test_branch_calculations_match_a_simulation_given_its_inputs_first(program, mode): + for branch_inputs, variable, period, value in _run(program, mode): + expected = synthetic_simulation(branch_inputs).calculate(variable, period) + # Uprating chains, and months divided from a year, may round + # differently in float32 depending on what was calculated first. + np.testing.assert_allclose( + value, expected, rtol=1e-5, atol=1e-4, err_msg=f"{variable} {period}" + ) diff --git a/tests/core/test_branch_scoped_delete_arrays.py b/tests/core/test_branch_scoped_delete_arrays.py index 5014bc62..2dcb7793 100644 --- a/tests/core/test_branch_scoped_delete_arrays.py +++ b/tests/core/test_branch_scoped_delete_arrays.py @@ -4,6 +4,8 @@ import gc +from pathlib import Path + import numpy as np from policyengine_core.country_template import CountryTaxBenefitSystem @@ -149,7 +151,9 @@ def test_child_disk_delete_keeps_parent_view_intact(tmp_path): np.testing.assert_array_equal(inherited_value, [3_000.0]) assert disk_key in parent_storage._files assert disk_key not in child_storage._files - assert (tmp_path / "salary" / f"{disk_key}.npy").is_file() + # Each store writes its own file; the parent's is still there. + assert Path(parent_storage._files[disk_key]).is_file() + assert Path(parent_storage._files[disk_key]).parent == tmp_path / "salary" np.testing.assert_array_equal( parent_holder._disk_storage.get(PERIOD, "default"), [3_000.0], diff --git a/tests/core/test_branch_shared_arrays.py b/tests/core/test_branch_shared_arrays.py index 61aa398c..9da7a933 100644 --- a/tests/core/test_branch_shared_arrays.py +++ b/tests/core/test_branch_shared_arrays.py @@ -128,11 +128,16 @@ def test_branch_copies_only_what_it_reads(tax_benefit_system): branch.calculate("income_tax", JANUARY) still_shared = shared_keys(branch) + held = set(stored_arrays(branch)) assert still_shared < inherited - # ``income_tax`` reads ``salary`` only, so nothing else was copied. + # ``income_tax`` reads ``salary`` only, so nothing else was copied: each + # other array is still shared, or was dropped by the input (values + # calculated after ``salary`` was first stored). for variable in ("rent", "accommodation_size", "housing_tax", "basic_income"): keys = {key for key in inherited if key[0] == variable} - assert keys and keys <= still_shared, variable + assert keys and all(key in still_shared or key not in held for key in keys), ( + variable + ) branch_arrays = stored_arrays(branch) for key in still_shared: assert np.shares_memory(branch_arrays[key], parent_arrays[key]), key @@ -439,11 +444,9 @@ def test_set_input_on_branch_leaves_parent(tax_benefit_system): branch.set_input("salary", JANUARY, np.array([10_000.0, 0.0, 0.0, 0.0])) branch_tax = branch.calculate("income_tax", JANUARY) - # ``income_tax`` was cached before branching, so the branch keeps the - # parent's value, exactly as when the branch held a copy of it. - assert np.array_equal(branch_tax, parent_tax) - branch.delete_arrays("income_tax", JANUARY) - assert branch.calculate("income_tax", JANUARY)[0] > parent_tax[0] + # The input drops the ``income_tax`` the branch inherited, so the branch + # calculates it from the new salary. + assert branch_tax[0] > parent_tax[0] _assert_unchanged(simulation, snapshot) assert np.array_equal(simulation.calculate("income_tax", JANUARY), parent_tax) diff --git a/tests/fixtures/branch_input_invalidation.py b/tests/fixtures/branch_input_invalidation.py new file mode 100644 index 00000000..06f737cc --- /dev/null +++ b/tests/fixtures/branch_input_invalidation.py @@ -0,0 +1,263 @@ +"""A synthetic system for derived values and branches inside formulas.""" + +import tempfile + +import numpy as np + +from policyengine_core import periods +from policyengine_core.country_template import CountryTaxBenefitSystem, entities +from policyengine_core.holders import set_input_divide_by_period +from policyengine_core.simulations import SimulationBuilder +from policyengine_core.variables import Variable + + +def _yearly(name, formula=None, **attributes): + namespace = dict( + value_type=float, + entity=entities.Person, + definition_period=periods.YEAR, + label=name, + **attributes, + ) + if formula is not None: + namespace["formula"] = formula + return type(name, (Variable,), namespace) + + +def _inner_branch_formula(person, period): + """Calculate a value in a branch that exists only while this formula runs.""" + simulation = person.simulation + inner = simulation.get_branch("inner") + try: + inner.set_input("p_switch", period, np.ones(person.count)) + inner_value = inner.calculate("p_inner_only", period) + finally: + del simulation.branches["inner"] + return inner_value + person("p_sum", period) + + +def _imported_uprating_formula(person, period): + """An uprated value, calculated in a branch that exists only while this runs.""" + simulation = person.simulation + temporary = simulation.get_branch("temporary_uprating") + try: + return temporary.calculate("p_up", period) * 2 + finally: + del simulation.branches["temporary_uprating"] + + +def _monthly_formula(person, period): + return ( + person("p_a", period) + + person("p_b", period.this_year) / 12 + + person("p_m", period) + ) + + +def _monthly(name, formula=None, **attributes): + namespace = dict( + value_type=float, + entity=entities.Person, + definition_period=periods.MONTH, + label=name, + **attributes, + ) + if formula is not None: + namespace["formula"] = formula + return type(name, (Variable,), namespace) + + +SYNTHETIC_VARIABLES = [ + # Inputs. p_up is uprated by a parameter that changes every year from + # 2012 to 2015; p_c is given for 2012 only, and carried over where the + # system carries inputs over. + _yearly("p_a"), + _yearly("p_b"), + _yearly("p_up", uprating="taxes.income_tax_rate"), + _yearly("p_c"), + _yearly( + "p_sum", + lambda person, period: person("p_a", period) + 2 * person("p_b", period), + ), + _yearly( + "p_prod", + lambda person, period: ( + person("p_sum", period) * (1 + person("p_up", period) / 100) + ), + ), + _yearly( + "p_cond", + lambda person, period: np.where( + person("p_sum", period) > 60, + person("p_prod", period), + person("p_c", period), + ), + ), + _yearly( + "p_lag", + lambda person, period: ( + person("p_sum", period.last_year) + person("p_a", period) + ), + ), + # A monthly input; an input for a year is divided between its months. + _monthly("p_m", set_input=set_input_divide_by_period), + _monthly("p_month", _monthly_formula), + _yearly("p_months", lambda person, period: person("p_month", period)), + _yearly("p_switch", lambda person, period: np.zeros(person.count)), + _yearly( + "p_inner_only", + lambda person, period: 3 * person("p_a", period) + person("p_switch", period), + ), + _yearly("p_inner", _inner_branch_formula), + _yearly("p_imported_up", _imported_uprating_formula), + # Reads itself a year earlier, back until core's spiral detection gives + # up and returns the default. + _yearly( + "p_spiral", lambda person, period: person("p_spiral", period.last_year) + 1 + ), +] + + +def _synthetic_system(carry_over): + class System(CountryTaxBenefitSystem): + auto_carry_over_input_variables = carry_over + + system = System() + system.add_variables(*SYNTHETIC_VARIABLES) + # Not stored by simulations that opt out of the cache (``opt_out_cache``). + system.cache_blacklist = {"p_sum", "p_inner_only"} + return system + + +# Carry-over is left out of the property test's random programs: core's +# carry-over itself depends on calculation order (a later period calculated +# first stops an earlier input being carried into the years between), in any +# simulation. +SYNTHETIC_SYSTEM = _synthetic_system(carry_over=False) +CARRY_OVER_SYSTEM = _synthetic_system(carry_over=True) + +PEOPLE = 3 +YEARS = ["2012", "2013", "2014", "2015"] + + +def synthetic_simulation( + inputs, memory_config=None, system=SYNTHETIC_SYSTEM, opt_out_cache=False +): + """A simulation of ``PEOPLE`` people given ``inputs`` before anything else.""" + simulation = SimulationBuilder().build_default_simulation(system, count=PEOPLE) + simulation.opt_out_cache = opt_out_cache + if memory_config is not None: + # Each holder's disk storage removes its directory, and this one + # once empty, when garbage-collected. + simulation._data_storage_dir = tempfile.mkdtemp( + prefix="policyengine-branch-tests-" + ) + simulation.memory_config = memory_config + # Holders read the memory configuration when created; nothing is + # stored yet, so create them again. Create them all here, so every + # branch's holders are copies of these: a holder first created in a + # branch would own (and remove on deletion) a directory the whole + # family's disk storage shares. + for population in simulation.populations.values(): + population._holders = {} + for variable in system.variables: + simulation.get_holder(variable) + for (variable, period), values in sorted( + inputs.items(), + key=lambda item: periods.key_period_size(periods.period(item[0][1])), + ): + simulation.set_input(variable, period, np.asarray(values, dtype=float)) + return simulation + + +ROOT_INPUTS = { + **{("p_a", year): (10.0 * i, 20.0, 35.0 + i) for i, year in enumerate(YEARS)}, + **{("p_b", year): (5.0, 15.0 + i, 25.0) for i, year in enumerate(YEARS)}, + ("p_up", "2012"): (100.0, 200.0, 300.0), + ("p_c", "2012"): (7.0, 8.0, 9.0), + **{ + ("p_m", f"{year}-{month:02d}"): (1.0 * month, 2.0, 0.5 * month) + for year in (2013, 2014) + for month in range(1, 13) + }, +} + + +# ----- A system whose formulas choose between branches, as itemization does ----- # + +# How many times each counted formula ran. +FORMULA_RUNS = {"agi": 0} + + +def _agi(person, period): + FORMULA_RUNS["agi"] += 1 + return person("earn", period) * 1.0 + + +def _tax_if(choice): + """Tax with ``choose`` set to ``choice``, from a branch deleted afterwards.""" + + def formula(person, period): + simulation = person.simulation + name = f"choose_{choice}" + branch = simulation.get_branch(name) + try: + branch.set_input("choose", period, np.full(person.count, float(choice))) + return branch.calculate("tax", period) + finally: + del simulation.branches[name] + + return formula + + +def _marginal_rate(person, period): + """Net income with one more unit of earnings, from a branch, as MTRs are.""" + simulation = person.simulation + branch = simulation.get_branch("raise") + try: + branch.set_input("earn", period, person("earn", period) + 1) + return branch.calculate("net", period) - person("net", period) + finally: + del simulation.branches["raise"] + + +def _from_persistent_branch(person, period): + """Read a value from a branch kept between calls.""" + return person.simulation.get_branch("persistent").calculate("tax", period) + + +BRANCHING_VARIABLES = [ + _yearly("earn"), + _yearly("agi", _agi), + _yearly( + "tax", + lambda person, period: ( + person("agi", period) * (0.2 - 0.05 * person("choose", period)) + ), + ), + _yearly("tax_if_chosen", _tax_if(1)), + _yearly("tax_if_not_chosen", _tax_if(0)), + _yearly( + "choose", + lambda person, period: ( + (person("tax_if_chosen", period) < person("tax_if_not_chosen", period)) + * 1.0 + ), + ), + _yearly( + "net", lambda person, period: person("agi", period) - person("tax", period) + ), + _yearly("marginal_rate", _marginal_rate), + _yearly("from_persistent_branch", _from_persistent_branch), +] + +BRANCHING_SYSTEM = CountryTaxBenefitSystem() +BRANCHING_SYSTEM.add_variables(*BRANCHING_VARIABLES) + + +def branching_simulation(earn=(10.0, 20.0, 30.0), period="2020"): + simulation = SimulationBuilder().build_default_simulation( + BRANCHING_SYSTEM, count=len(earn) + ) + simulation.set_input("earn", period, np.asarray(earn)) + return simulation