Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions changelog.d/fix-uprating-branch-visibility.fixed.md
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Uprating and auto-carry-over now use only periods the current branch can read (its own, its parent branches' and `default`), so a period stored only under an unrelated branch no longer raises `TypeError` or caches `NaN`; on-disk storage parses branch names that contain `_`, and `dump_simulation` saves the values the dumped branch reads.
22 changes: 17 additions & 5 deletions policyengine_core/data_storage/on_disk_storage.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,17 @@
from policyengine_core.periods import Period


def _split_key(key: str) -> tuple:
"""Split a ``f"{branch_name}_{period}"`` file key into its two parts.

Branch names often contain ``_`` (policyengine-us uses ``no_salt`` and
``mtr_for_adult_1``) but a period's string form never does, so the key
splits on its last ``_``.
"""
branch_name, period = key.rsplit("_", 1)
return branch_name, period


class OnDiskStorage:
"""
Low-level class responsible for storing and retrieving calculated vectors on disk
Expand Down Expand Up @@ -82,12 +93,13 @@ def delete(self, period: Period = None, branch_name: str = "default") -> None:
if period is None:
# Only wipe files belonging to the requested branch (previously
# this wiped every branch regardless of ``branch_name`` — same
# class of bug as C2 in InMemoryStorage).
branch_prefix = f"{branch_name}_"
# class of bug as C2 in InMemoryStorage). Compare the parsed
# branch name, not a prefix: deleting ``pre_tcja`` must not
# also wipe ``pre_tcja_ctc``.
self._files = {
period_item: value
for period_item, value in self._files.items()
if not period_item.startswith(branch_prefix)
if _split_key(period_item)[0] != branch_name
}
return

Expand All @@ -103,12 +115,12 @@ def delete(self, period: Period = None, branch_name: str = "default") -> None:
}

def get_known_periods(self) -> list:
return list([periods.period(x.split("_")[1]) for x in self._files.keys()])
return [period for _, period in self.get_known_branch_periods()]

def get_known_branch_periods(self) -> list:
return [
(branch_name, periods.period(period))
for branch_name, period in map(lambda x: x.split("_"), self._files.keys())
for branch_name, period in map(_split_key, self._files.keys())
]

def restore(self) -> None:
Expand Down
82 changes: 51 additions & 31 deletions policyengine_core/holders/holder.py
Original file line number Diff line number Diff line change
Expand Up @@ -106,45 +106,53 @@ def _get_array_from_storage(
value = self._disk_storage.get(period, branch_name)
return value

def _readable_branch_names(self, branch_name: str = "default") -> List[str]:
"""
Branches whose stored values ``get_array(period, branch_name)`` can
return, in lookup order.

That is ``branch_name`` itself, then (unless it is ``default``) each
``simulation.parent_branch`` ancestor, then ``default``. Nested
branches inherit values from their parent (e.g. a ``no_salt`` branch
cloned from an ``itemizing`` branch still sees ``tax_unit_itemizes``
set on the ``itemizing`` branch). Previously the fallback returned
the first branch in dict-insertion order (bug C1) — silently swapping
values between unrelated sibling branches (reform vs baseline) and
producing wrong reform deltas. The post-C1 behavior only fell back to
``default``, which broke country-package nested-branch patterns that
relied on the ancestor's input being visible.
"""
if branch_name == "default":
return ["default"]
branch_names = [branch_name]
simulation = getattr(self, "simulation", None)
parent = getattr(simulation, "parent_branch", None) if simulation else None
while parent is not None:
branch_names.append(parent.branch_name)
parent = getattr(parent, "parent_branch", None)
branch_names.append("default")
return list(dict.fromkeys(branch_names))

def get_array(self, period: Period, branch_name: str = "default") -> ArrayLike:
"""
Get the value of the variable for the given period.

If the value is not known, return ``None``.
Values stored under ``branch_name``, its ``parent_branch`` ancestors
and ``default`` are visible, in that order. If the value is not
known on any of them, return ``None``.
"""
if self.variable.is_neutralized:
return self.default_array()
# The branch's own value is the common case: look it up before
# walking the ancestors.
value = self._get_array_from_storage(period, branch_name)
if value is not None:
return value
if value is None and branch_name != "default":
# Walk up ``simulation.parent_branch`` so nested branches inherit
# values from their parent (e.g. a ``no_salt`` branch cloned
# from an ``itemizing`` branch still sees ``tax_unit_itemizes``
# set on the ``itemizing`` branch). Fall back to ``default``
# only if no ancestor branch has a value. Previously the
# fallback returned the first branch in dict-insertion order
# (bug C1) — silently swapping values between unrelated
# sibling branches (reform vs baseline) and producing wrong
# reform deltas. The post-C1 behavior only fell back to
# ``default``, which broke country-package nested-branch
# patterns that relied on the ancestor's input being visible.
parent = (
getattr(self.simulation, "parent_branch", None)
if self.simulation
else None
)
while parent is not None:
ancestor_value = self._get_array_from_storage(
period,
parent.branch_name,
)
if ancestor_value is not None:
return ancestor_value
parent = getattr(parent, "parent_branch", None)
default_value = self._get_array_from_storage(period, "default")
if default_value is not None:
return default_value
for readable_branch_name in self._readable_branch_names(branch_name)[1:]:
value = self._get_array_from_storage(period, readable_branch_name)
if value is not None:
return value
return None

def get_memory_usage(self) -> dict:
"""
Expand Down Expand Up @@ -189,11 +197,23 @@ def get_memory_usage(self) -> dict:

return usage

def get_known_periods(self) -> List[Period]:
def get_known_periods(self, branch_name: str = None) -> List[Period]:
"""
Get the list of periods the variable value is known for.
"""

With ``branch_name``, list only the periods ``get_array(period,
branch_name)`` can read: those stored under that branch, its
``parent_branch`` ancestors or ``default``. Without it, list the
periods stored under every branch in this holder, some of which that
branch may not be able to read.
"""
if branch_name is not None:
readable_branch_names = set(self._readable_branch_names(branch_name))
return [
period
for stored_branch_name, period in self.get_known_branch_periods()
if stored_branch_name in readable_branch_names
]
return list(self._memory_storage.get_known_periods()) + list(
(self._disk_storage.get_known_periods() if self._disk_storage else [])
)
Expand Down
7 changes: 5 additions & 2 deletions policyengine_core/simulations/simulation.py
Original file line number Diff line number Diff line change
Expand Up @@ -864,8 +864,11 @@ def _calculate(self, variable_name: str, period: Period = None) -> ArrayLike:

# If no result, use the default value and cache it
if array is None:
# Check if the variable has a previously defined value
known_periods = holder.get_known_periods()
# Check if the variable has a previously defined value.
# Only periods this branch can read count: a period stored
# only under an unrelated branch would read back as ``None``
# and reach the arithmetic below.
known_periods = holder.get_known_periods(self.branch_name)
earlier_known_periods = [
known_period
for known_period in known_periods
Expand Down
12 changes: 8 additions & 4 deletions policyengine_core/tools/simulation_dumper.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,7 @@ def dump_simulation(simulation, directory):

# Dump variable values
for holder in entity._holders.values():
_dump_holder(holder, directory)
_dump_holder(holder, directory, simulation.branch_name)


def restore_simulation(directory, tax_benefit_system, **kwargs):
Expand Down Expand Up @@ -64,10 +64,14 @@ def restore_simulation(directory, tax_benefit_system, **kwargs):
return simulation


def _dump_holder(holder, directory):
def _dump_holder(holder, directory, branch_name="default"):
# Dump the values the simulation's branch reads. A holder can also store
# periods under branches this one cannot see; reading those back under
# this branch gives ``None``, which would be saved as an object array
# that ``restore_simulation`` cannot load.
disk_storage = holder.create_disk_storage(directory, preserve=True)
for period in holder.get_known_periods():
value = holder.get_array(period)
for period in dict.fromkeys(holder.get_known_periods(branch_name)):
value = holder.get_array(period, branch_name)
disk_storage.put(value, period)


Expand Down
Loading
Loading