diff --git a/metaflow/user_configs/config_parameters.py b/metaflow/user_configs/config_parameters.py index ec38d87c161..4d41ac2fdc1 100644 --- a/metaflow/user_configs/config_parameters.py +++ b/metaflow/user_configs/config_parameters.py @@ -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): diff --git a/test/unit/test_delay_evaluator.py b/test/unit/test_delay_evaluator.py new file mode 100644 index 00000000000..0b2a9b1d894 --- /dev/null +++ b/test/unit/test_delay_evaluator.py @@ -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"] + + +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"