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-simulation-copy-pickle-recursion.fixed.md
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Stop `copy.deepcopy(simulation)`, `copy.copy(population)` and unpickling a simulation in the process that pickled it from raising `RecursionError`, keep vectorial parameter nodes as nodes when deep-copied, and keep an `EnumArray`'s `possible_values` through pickling.
38 changes: 37 additions & 1 deletion policyengine_core/enums/enum_array.py
Original file line number Diff line number Diff line change
@@ -1,14 +1,41 @@
from __future__ import annotations

import importlib
import operator
import typing
from typing import Any, NoReturn, Optional, Type
from typing import Any, NoReturn, Optional, Tuple, Type

import numpy

if typing.TYPE_CHECKING:
from policyengine_core.enums import Enum


def _restore_enum_array(
array: numpy.ndarray, enum_name: Optional[Tuple[str, str]]
) -> EnumArray:
"""Rebuild a pickled EnumArray, with its enum if this process can find it.

``enum_name`` is the enum's module and qualified name, or ``None`` for an
array that had no enum.

Tax-benefit systems load variable files under module names that exist only
in the process that loaded them, so an enum defined in one cannot be found
from another process. The array then comes back with ``possible_values``
unset (``None``), where pickling the enum by reference would fail to
unpickle the array at all.
"""
possible_values = None
if enum_name is not None:
module_name, qualified_name = enum_name
try:
module = importlib.import_module(module_name)
possible_values = operator.attrgetter(qualified_name)(module)
except (ImportError, AttributeError):
pass
return EnumArray(array, possible_values)


class EnumArray(numpy.ndarray):
"""
Numpy array subclass representing an array of enum items.
Expand All @@ -35,6 +62,15 @@ def __array_finalize__(self, obj: Optional[numpy.int_]) -> None:

self.possible_values = getattr(obj, "possible_values", None)

def __reduce__(self) -> tuple:
# ndarray's own ``__reduce__`` rebuilds the array without
# ``possible_values``, so an unpickled EnumArray could be neither
# decoded nor compared with an enum item. The enum travels by name
# rather than by reference; see ``_restore_enum_array``.
enum = self.possible_values
name = None if enum is None else (enum.__module__, enum.__qualname__)
return _restore_enum_array, (self.view(numpy.ndarray), name)

def __eq__(self, other: Any) -> bool:
# When comparing to an item of self.possible_values, use the item index
# to speed up the comparison.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -209,6 +209,14 @@ def __init__(self, name: str, vector: ArrayLike, instant_str: str):
self._instant_str = instant_str

def __getattr__(self, attribute: str) -> Any:
# ``vector`` is missing while copy or pickle rebuilds a node, and
# looking it up would recurse here. ``copy.deepcopy`` looks
# ``__deepcopy__`` up on the instance, and the vector's would copy the
# vector alone and hand back a bare ``numpy.recarray``.
if attribute in ("vector", "__deepcopy__"):
raise AttributeError(
f"{type(self).__name__!s} has no attribute {attribute!r}"
)
result = getattr(self.vector, attribute)
if isinstance(result, numpy.recarray):
return VectorialParameterNodeAtInstant(result)
Expand Down
8 changes: 8 additions & 0 deletions policyengine_core/populations/population.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,14 @@ def filled_array(self, value: Any, dtype: Any = None) -> numpy.ndarray:
return numpy.full(self.count, value, dtype)

def __getattr__(self, attribute: str) -> Any:
# The shortcut lookup reads ``self.entity`` and ``self.simulation``.
# They are missing while copy or pickle rebuilds a population (both
# probe the new, empty instance for ``__setstate__`` before restoring
# its ``__dict__``), and looking them up would recurse here.
if attribute in ("entity", "simulation"):
raise AttributeError(
f"{type(self).__name__!s} has no attribute {attribute!r}"
)
projector = projectors.get_projector_from_shortcut(self, attribute)
if not projector:
raise AttributeError(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,10 @@ def __getattr__(
self,
key: str,
) -> Union[TracingParameterNodeAtInstant, Child]:
# ``parameter_node_at_instant`` is missing while copy or pickle
# rebuilds a wrapper, and looking it up would recurse here.
if key == "parameter_node_at_instant":
raise AttributeError(f"{type(self).__name__!s} has no attribute {key!r}")
child = getattr(self.parameter_node_at_instant, key)
return self.get_traced_child(child, key)

Expand Down
Loading
Loading