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
8 changes: 4 additions & 4 deletions metaflow/user_configs/config_parameters.py
Original file line number Diff line number Diff line change
Expand Up @@ -301,17 +301,17 @@ def __init__(self, ex: str, saved_globals: Optional[Dict[str, Any]] = None):
self._cached_expr = None

def __copy__(self):
c = DelayEvaluator(self._config_expr)
# Keep caller globals by reference so config_expr("my_func").project still
# resolves after attribute/item access (which copies this object).
c = DelayEvaluator(self._config_expr, saved_globals=self._globals)
c._access = self._access.copy() if self._access is not None else None
# Globals are not copied -- always kept as a reference
return c

def __deepcopy__(self, memo):
c = DelayEvaluator(self._config_expr)
c = DelayEvaluator(self._config_expr, saved_globals=self._globals)
c._access = (
copy.deepcopy(self._access, memo) if self._access is not None else None
)
# Globals are not copied -- always kept as a reference
return c

def __iter__(self):
Expand Down
73 changes: 73 additions & 0 deletions test/unit/test_delay_evaluator.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,73 @@
import copy

from metaflow.flowspec import FlowStateItems
from metaflow.parameters import current_flow
from metaflow.user_configs.config_parameters import DelayEvaluator


def _sample_globals():
def my_func():
return "hello"

return {"my_func": my_func}


class _DummyFlow:
_flow_state = {FlowStateItems.CONFIGS: {}}


def _with_dummy_flow(fn):
current_flow.flow_cls = _DummyFlow
try:
return fn()
finally:
del current_flow.flow_cls


def test_copy_preserves_saved_globals():
saved = _sample_globals()
evaluator = DelayEvaluator("config", saved_globals=saved)
copied = copy.copy(evaluator)
assert copied._globals is saved
assert copied._globals["my_func"]() == "hello"


def test_deepcopy_preserves_saved_globals():
saved = _sample_globals()
evaluator = DelayEvaluator("config", saved_globals=saved)
copied = copy.deepcopy(evaluator)
assert copied._globals is saved
assert copied._globals["my_func"]() == "hello"


def test_getattr_preserves_saved_globals():
saved = _sample_globals()
evaluator = DelayEvaluator("config", saved_globals=saved)
chained = evaluator.project
assert chained._globals is saved
assert chained._access == ["project"]
Comment thread
greptile-apps[bot] marked this conversation as resolved.


def test_getitem_preserves_saved_globals():
saved = _sample_globals()
evaluator = DelayEvaluator("config", saved_globals=saved)
chained = evaluator["project"]
assert chained._globals is saved
assert chained._access == ["project"]


def test_chained_call_uses_saved_globals():
class _Cfg:
project = "from-globals"

saved = {"my_func": _Cfg}
evaluator = DelayEvaluator("my_func", saved_globals=saved)
chained = evaluator.project
assert _with_dummy_flow(chained) == "from-globals"


def test_copied_call_uses_saved_globals():
saved = _sample_globals()
evaluator = DelayEvaluator("my_func()", saved_globals=saved)
copied = copy.copy(evaluator)
assert _with_dummy_flow(copied) == "hello"