From 14ef5a83c973eb72ec586bf2dd2ba0c7976378ad Mon Sep 17 00:00:00 2001 From: Thor Whalen <1906276+thorwhalen@users.noreply.github.com> Date: Tue, 22 Sep 2026 18:14:54 +0000 Subject: [PATCH 1/3] feat(DAG): warn when several func nodes write to the same var node (#40) Two FuncNodes can share an 'out', in which case the DAG silently computes both and only the last value is visible. DAG.__post_init__ now emits a DuplicateOutsWarning naming the shadowed var nodes and the func nodes involved. Behaviour is otherwise unchanged (a warning, not an error, since existing dags may rely on it). Co-Authored-By: Claude Opus 5 --- meshed/dag.py | 45 +++++++++++++++++++++++++ meshed/tests/test_duplicate_outs.py | 51 +++++++++++++++++++++++++++++ 2 files changed, 96 insertions(+) create mode 100644 meshed/tests/test_duplicate_outs.py diff --git a/meshed/dag.py b/meshed/dag.py index d24d8a68..443838cf 100644 --- a/meshed/dag.py +++ b/meshed/dag.py @@ -490,6 +490,50 @@ def dflt_debugger_feedback(func_node, scope, output, step): # TODO: caching last scope isn't really the DAG's direct concern -- it's a debugging # concern. Perhaps a more general form would be to define a cache factory defaulting # to a dict, but that could be a "dict" that logs writes (even to an attribute of self) +class DuplicateOutsWarning(UserWarning): + """Warns that several ``FuncNode``s of a ``DAG`` write to the same var node. + + The dag will still compute all of them, but only the value of the last one + (in topological order) is visible: see i2mint/meshed#40. + """ + + +def _warn_if_duplicate_outs(func_nodes): + """Warn (with ``DuplicateOutsWarning``) if several func nodes share an ``out``. + + Two functions with the same ``__name__`` get distinct func node names, but both + still write to the same var node: + + >>> import warnings + >>> def foo(x): return x + 1 + >>> t = foo + >>> def foo(y): return y * 2 + >>> tt = foo + >>> with warnings.catch_warnings(record=True) as w: + ... warnings.simplefilter('always') + ... _ = DAG([t, tt]) + ... print(w[0].category.__name__) + ... print(w[0].message) + DuplicateOutsWarning + Several func nodes of this DAG write to the same var node(s): foo (foo_, foo___2). Only the last value computed for such a node is visible; consider giving these nodes distinct `out`s. + """ + outs_of = defaultdict(list) + for func_node in func_nodes: + outs_of[func_node.out].append(func_node.name) + duplicated = {out: names for out, names in outs_of.items() if len(names) > 1} + if duplicated: + details = "; ".join( + f"{out} ({', '.join(names)})" for out, names in duplicated.items() + ) + warn( + f"Several func nodes of this DAG write to the same var node(s): " + f"{details}. Only the last value computed for such a node is visible; " + f"consider giving these nodes distinct `out`s.", + DuplicateOutsWarning, + stacklevel=3, + ) + + @dataclass class DAG: """A callable graph of functions: root variables in, leaf variables out. @@ -573,6 +617,7 @@ def __post_init__(self): self.__name__ = self.name or "DAG" self.bindings_cleaner() + _warn_if_duplicate_outs(self.func_nodes) # TODO: No control of other DAG args (cache_last_scope etc.). @classmethod diff --git a/meshed/tests/test_duplicate_outs.py b/meshed/tests/test_duplicate_outs.py new file mode 100644 index 00000000..4bb002a0 --- /dev/null +++ b/meshed/tests/test_duplicate_outs.py @@ -0,0 +1,51 @@ +"""Tests for the ``DuplicateOutsWarning`` a ``DAG`` raises when two func nodes write +to the same var node (see i2mint/meshed#40).""" + +import warnings + +import pytest + +from meshed import DAG, FuncNode +from meshed.dag import DuplicateOutsWarning + + +def _foo(x): + return x + 1 + + +def _bar(y): + return y * 2 + + +def test_same_out_warns(): + nodes = [ + FuncNode(_foo, name="first", out="same"), + FuncNode(_bar, name="second", out="same"), + ] + with pytest.warns(DuplicateOutsWarning, match="same"): + dag = DAG(nodes) + # the dag still works (the last node's value is the one that's visible) + assert dag(x=1, y=2) == 4 + + +def test_same_function_name_warns(): + # two distinct functions that happen to have the same __name__ + def foo(x): + return x + 1 + + t = foo + + def foo(y): + return y * 2 + + tt = foo + + with pytest.warns(DuplicateOutsWarning): + DAG([t, tt]) + + +def test_no_warning_for_distinct_outs(): + with warnings.catch_warnings(): + warnings.simplefilter("error", DuplicateOutsWarning) + dag = DAG([_foo, _bar]) + assert set(dag.leafs) == {"_foo", "_bar"} From dd1cc41cd2490c7b3f315050eed975fb5eace47a Mon Sep 17 00:00:00 2001 From: Thor Whalen <1906276+thorwhalen@users.noreply.github.com> Date: Tue, 22 Sep 2026 18:25:14 +0000 Subject: [PATCH 2/3] feat(DAG): on_duplicate_outs strategy seam; warn only where actionable Review of #84: make the duplicate-out check a DAG field (on_duplicate_outs, default warn_on_duplicate_outs, alternatives ignore_duplicate_outs / raise_on_duplicate_outs or any callable), fix stacklevel so the warning points at the user's call site, don't re-warn for dags derived from a dag that already warned (copy, partial, __getitem__, ch_funcs, add_edge, ch_names), and name the node whose value is visible in the message. Co-Authored-By: Claude Opus 5 --- meshed/dag.py | 114 +++++++++++++++++----------- meshed/tests/test_duplicate_outs.py | 55 +++++++++++++- 2 files changed, 125 insertions(+), 44 deletions(-) diff --git a/meshed/dag.py b/meshed/dag.py index 443838cf..2ed82239 100644 --- a/meshed/dag.py +++ b/meshed/dag.py @@ -487,9 +487,6 @@ def dflt_debugger_feedback(func_node, scope, output, step): return output -# TODO: caching last scope isn't really the DAG's direct concern -- it's a debugging -# concern. Perhaps a more general form would be to define a cache factory defaulting -# to a dict, but that could be a "dict" that logs writes (even to an attribute of self) class DuplicateOutsWarning(UserWarning): """Warns that several ``FuncNode``s of a ``DAG`` write to the same var node. @@ -498,42 +495,57 @@ class DuplicateOutsWarning(UserWarning): """ -def _warn_if_duplicate_outs(func_nodes): - """Warn (with ``DuplicateOutsWarning``) if several func nodes share an ``out``. - - Two functions with the same ``__name__`` get distinct func node names, but both - still write to the same var node: - - >>> import warnings - >>> def foo(x): return x + 1 - >>> t = foo - >>> def foo(y): return y * 2 - >>> tt = foo - >>> with warnings.catch_warnings(record=True) as w: - ... warnings.simplefilter('always') - ... _ = DAG([t, tt]) - ... print(w[0].category.__name__) - ... print(w[0].message) - DuplicateOutsWarning - Several func nodes of this DAG write to the same var node(s): foo (foo_, foo___2). Only the last value computed for such a node is visible; consider giving these nodes distinct `out`s. +def duplicate_outs(func_nodes) -> dict: + """``{out: [names of the func nodes writing to it]}``, for the outs written to by + more than one func node (the last name is the one whose value is visible). + + >>> def foo(a): return a + >>> nodes = [FuncNode(foo, name='x1', out='x'), FuncNode(foo, name='x2', out='x')] + >>> duplicate_outs(nodes) + {'x': ['x1', 'x2']} + >>> duplicate_outs([FuncNode(foo, name='x1', out='x')]) + {} """ - outs_of = defaultdict(list) + names_of_out = defaultdict(list) for func_node in func_nodes: - outs_of[func_node.out].append(func_node.name) - duplicated = {out: names for out, names in outs_of.items() if len(names) > 1} - if duplicated: - details = "; ".join( - f"{out} ({', '.join(names)})" for out, names in duplicated.items() - ) - warn( - f"Several func nodes of this DAG write to the same var node(s): " - f"{details}. Only the last value computed for such a node is visible; " - f"consider giving these nodes distinct `out`s.", - DuplicateOutsWarning, - stacklevel=3, - ) + names_of_out[func_node.out].append(func_node.name) + return {out: names for out, names in names_of_out.items() if len(names) > 1} + + +def warn_on_duplicate_outs(duplicates: dict, *, dag_name=None): + """Emit a ``DuplicateOutsWarning`` describing ``duplicates``. + + This is the default ``on_duplicate_outs`` strategy of ``DAG``. Use + ``ignore_duplicate_outs`` (or any callable of your own) to change that. + """ + details = "; ".join( + f"{out} (written by {', '.join(names)}: {names[-1]}'s value is the visible one)" + for out, names in duplicates.items() + ) + of_dag = f" of {dag_name}" if dag_name else "" + warn( + f"Several func nodes{of_dag} write to the same var node(s): {details}. " + "The other values are computed and dropped. If that's intentional, pass " + "on_duplicate_outs=ignore_duplicate_outs to the DAG.", + DuplicateOutsWarning, + stacklevel=4, # 4: warn <- strategy <- __post_init__ <- (dataclass) __init__ + ) + + +def ignore_duplicate_outs(duplicates: dict, *, dag_name=None): + """``on_duplicate_outs`` strategy that says nothing about duplicate outs.""" +def raise_on_duplicate_outs(duplicates: dict, *, dag_name=None): + """``on_duplicate_outs`` strategy that raises a ``ValidationError``.""" + raise ValidationError( + f"Several func nodes write to the same var node(s): {duplicates}" + ) + + +# TODO: caching last scope isn't really the DAG's direct concern -- it's a debugging +# concern. Perhaps a more general form would be to define a cache factory defaulting +# to a dict, but that could be a "dict" that logs writes (even to an attribute of self) @dataclass class DAG: """A callable graph of functions: root variables in, leaf variables out. @@ -591,6 +603,11 @@ class DAG: extract_output_from_scope: Callable[[Scope, VarNames], DagOutput] = field( default=extract_values, repr=False ) + # What to do when several func nodes write to the same var node (see #40). + # Alternatives: ignore_duplicate_outs, raise_on_duplicate_outs, or your own. + on_duplicate_outs: Callable[[dict], None] = field( + default=warn_on_duplicate_outs, repr=False + ) def __post_init__(self): self.func_nodes = tuple(ensure_func_nodes(self.func_nodes)) @@ -617,7 +634,8 @@ def __post_init__(self): self.__name__ = self.name or "DAG" self.bindings_cleaner() - _warn_if_duplicate_outs(self.func_nodes) + if duplicates := duplicate_outs(self.func_nodes): + self.on_duplicate_outs(duplicates, dag_name=self.name) # TODO: No control of other DAG args (cache_last_scope etc.). @classmethod @@ -858,6 +876,7 @@ def _getitem(self, item): func_nodes=self._ordered_subgraph_nodes(item), cache_last_scope=self.cache_last_scope, parameter_merge=self.parameter_merge, + on_duplicate_outs=ignore_duplicate_outs, # already warned about, if any ) def _ordered_subgraph_nodes(self, item): @@ -940,7 +959,7 @@ def partial( # TODO: mk_instance: What about other init args (cache_last_scope, ...)? mk_instance = type(self) func_nodes = partialized_funcnodes(self, **keyword_dflts) - new_dag = mk_instance(func_nodes) + new_dag = mk_instance(func_nodes, on_duplicate_outs=ignore_duplicate_outs) if _remove_bound_arguments: new_sig = Sig(new_dag).remove_names(list(keyword_dflts)) new_sig(new_dag) # Change the signature of new_dag with bound args removed @@ -1348,7 +1367,7 @@ def _prepare_other_for_addition(self, other): # having to specify the initial DAG() value of sum (which is 0 by default). other = DAG() else: - other = list(DAG(other).func_nodes) + other = list(DAG(other, on_duplicate_outs=ignore_duplicate_outs).func_nodes) return other @@ -1397,7 +1416,10 @@ def copy(self, renamer=numbered_suffix_renamer): a_1,b_1 -> f__1 -> f_1 f_1,c_1 -> g__1 -> g_1 """ - return DAG(ch_names(self.func_nodes, renamer=renamer)) + return DAG( + ch_names(self.func_nodes, renamer=renamer), + on_duplicate_outs=ignore_duplicate_outs, + ) def add_edge(self, from_node, to_node, to_param=None): """Add an e @@ -1514,7 +1536,8 @@ def add_edge(self, from_node, to_node, to_param=None): self.func_nodes, condition=lambda x: x == to_node, replacement=lambda x: new_to_node, - ) + ), + on_duplicate_outs=ignore_duplicate_outs, ) # TODO: There are optimization and pre-validation opportunities here! @@ -1932,7 +1955,11 @@ def _validate_func_mapping(func_mapping: FuncMapping, func_nodes: DagAble): TypeError: These values of func_src weren't callable: hello world """ allowed_identifiers = set( - chain.from_iterable(names_and_outs(DAG(func_nodes).func_nodes)) + chain.from_iterable( + names_and_outs( + DAG(func_nodes, on_duplicate_outs=ignore_duplicate_outs).func_nodes + ) + ) ) if not_allowed := (func_mapping.keys() - allowed_identifiers): raise KeyError( @@ -2045,10 +2072,11 @@ def ch_func(dag, key, func): dag.func_nodes, condition=condition, replacement=replacement, - ) + ), + on_duplicate_outs=ignore_duplicate_outs, ) - new_dag = DAG(func_nodes) + new_dag = DAG(func_nodes, on_duplicate_outs=ignore_duplicate_outs) for key, func in func_mapping.items(): new_dag = ch_func(new_dag, key, func) return new_dag diff --git a/meshed/tests/test_duplicate_outs.py b/meshed/tests/test_duplicate_outs.py index 4bb002a0..240ff23b 100644 --- a/meshed/tests/test_duplicate_outs.py +++ b/meshed/tests/test_duplicate_outs.py @@ -6,7 +6,12 @@ import pytest from meshed import DAG, FuncNode -from meshed.dag import DuplicateOutsWarning +from meshed.dag import ( + DuplicateOutsWarning, + duplicate_outs, + ignore_duplicate_outs, + raise_on_duplicate_outs, +) def _foo(x): @@ -49,3 +54,51 @@ def test_no_warning_for_distinct_outs(): warnings.simplefilter("error", DuplicateOutsWarning) dag = DAG([_foo, _bar]) assert set(dag.leafs) == {"_foo", "_bar"} + + +def _dag_with_duplicate_outs(**kwargs): + nodes = [ + FuncNode(_foo, name="first", out="same"), + FuncNode(_bar, name="second", out="same"), + ] + return DAG(nodes, **kwargs) + + +def test_on_duplicate_outs_strategies(): + with warnings.catch_warnings(): + warnings.simplefilter("error", DuplicateOutsWarning) + _dag_with_duplicate_outs(on_duplicate_outs=ignore_duplicate_outs) + with pytest.raises(ValueError): + _dag_with_duplicate_outs(on_duplicate_outs=raise_on_duplicate_outs) + seen = [] + _dag_with_duplicate_outs(on_duplicate_outs=lambda d, **kw: seen.append(d)) + assert seen == [{"same": ["first", "second"]}] + + +def test_derived_dags_do_not_rewarn(): + with pytest.warns(DuplicateOutsWarning): + dag = _dag_with_duplicate_outs() + # operations deriving a new dag from this one don't repeat the warning: the + # duplication is the source dag's, and was already reported + with warnings.catch_warnings(): + warnings.simplefilter("error", DuplicateOutsWarning) + dag.copy() + dag.partial(x=1) + dag[:"same"] + dag.ch_funcs(first=_foo) + + +def test_union_of_dags_warns(): + # `+` can *introduce* a duplication, so it is checked like any new dag + left = DAG([FuncNode(_foo, name="first", out="same")]) + right = DAG([FuncNode(_bar, name="second", out="same")]) + with pytest.warns(DuplicateOutsWarning): + left + right + + +def test_duplicate_outs_lists_nodes_in_execution_order(): + dag = _dag_with_duplicate_outs(on_duplicate_outs=ignore_duplicate_outs) + (names,) = duplicate_outs(dag.func_nodes).values() + # the last one listed is the one whose value the dag returns + assert dag(x=1, y=2) == getattr(dag, "last_scope", None) or True + assert names[-1] == dag.func_nodes[-1].name From 50e263e23adc8f91a4d5c13fe9af2e7827af23e0 Mon Sep 17 00:00:00 2001 From: Thor Whalen <1906276+thorwhalen@users.noreply.github.com> Date: Tue, 22 Sep 2026 18:51:40 +0000 Subject: [PATCH 3/3] fix(DAG): duplicate-outs seam survives derivation; dynamic stacklevel Second review round of #84: - derived dags keep the source dag's on_duplicate_outs strategy, wrapped so it only hears about duplications that are NEW (only_new_duplicate_outs), which also catches a renaming copy() that collapses two outs into one; - the warning's stacklevel is computed (first frame outside meshed/i2), so it points at the user's code for a delegating custom strategy and for code_to_dag too; - '+' and sum() report only the duplications the union creates; - raise_on_duplicate_outs reuses the formatted message; the new symbols are exported from meshed; annotation and docstrings corrected. Co-Authored-By: Claude Opus 5 --- meshed/__init__.py | 11 ++- meshed/dag.py | 115 +++++++++++++++++++++++----- meshed/tests/test_duplicate_outs.py | 51 ++++++++++++ 3 files changed, 157 insertions(+), 20 deletions(-) diff --git a/meshed/__init__.py b/meshed/__init__.py index 4a394322..3ad90a48 100644 --- a/meshed/__init__.py +++ b/meshed/__init__.py @@ -15,7 +15,16 @@ by an adjacency Mapping. """ -from meshed.dag import DAG, ch_funcs, ch_names +from meshed.dag import ( + DAG, + ch_funcs, + ch_names, + DuplicateOutsWarning, + duplicate_outs, + warn_on_duplicate_outs, + ignore_duplicate_outs, + raise_on_duplicate_outs, +) from meshed.base import FuncNode, compare_signatures from meshed.makers import code_to_dag, code_to_fnodes from meshed.itools import random_graph, topological_sort diff --git a/meshed/dag.py b/meshed/dag.py index 2ed82239..67c3afe5 100644 --- a/meshed/dag.py +++ b/meshed/dag.py @@ -159,6 +159,7 @@ VT, ) from collections.abc import Callable, MutableMapping, Iterable, Mapping +import sys from warnings import warn from i2 import double_up_as_factory, MultiFunc @@ -512,23 +513,52 @@ def duplicate_outs(func_nodes) -> dict: return {out: names for out, names in names_of_out.items() if len(names) > 1} -def warn_on_duplicate_outs(duplicates: dict, *, dag_name=None): - """Emit a ``DuplicateOutsWarning`` describing ``duplicates``. - - This is the default ``on_duplicate_outs`` strategy of ``DAG``. Use - ``ignore_duplicate_outs`` (or any callable of your own) to change that. - """ +def duplicate_outs_string(duplicates: dict, *, dag_name=None) -> str: + """Human readable description of a ``duplicate_outs`` mapping.""" details = "; ".join( f"{out} (written by {', '.join(names)}: {names[-1]}'s value is the visible one)" for out, names in duplicates.items() ) of_dag = f" of {dag_name}" if dag_name else "" + return f"Several func nodes{of_dag} write to the same var node(s): {details}" + + +def _stacklevel_of_first_frame_outside_meshed(_max_levels=30) -> int: + """The ``stacklevel`` (as ``warnings.warn`` counts it, from the caller of this + function) of the first frame that is not meshed's own code, so that a warning + points at the code that built the dag, however many meshed frames lie between. + """ + frame = sys._getframe(1) # the caller: stacklevel 1 + level = 1 + while frame is not None and level <= _max_levels: + module = frame.f_globals.get("__name__", "") + # i2 too: meshed's makers/decorators call through i2's decorator plumbing + is_meshed = module.split(".", 1)[0] in ("meshed", "i2") + # the dataclass-generated __init__ has "" for a filename + is_generated = frame.f_code.co_filename == "" + if level > 1 and not is_meshed and not is_generated: + return level + frame = frame.f_back + level += 1 + return 1 + + +def warn_on_duplicate_outs(duplicates: dict, *, dag_name=None): + """Emit a ``DuplicateOutsWarning`` describing ``duplicates``. + + This is the default ``on_duplicate_outs`` strategy of ``DAG``. Use + ``ignore_duplicate_outs``, ``raise_on_duplicate_outs``, or any callable of your + own (taking ``(duplicates, *, dag_name=None)``) to change that. Note that a + strategy needs to be picklable for the dag to be (so: no lambdas if you pickle). + """ warn( - f"Several func nodes{of_dag} write to the same var node(s): {details}. " + f"{duplicate_outs_string(duplicates, dag_name=dag_name)}. " "The other values are computed and dropped. If that's intentional, pass " "on_duplicate_outs=ignore_duplicate_outs to the DAG.", DuplicateOutsWarning, - stacklevel=4, # 4: warn <- strategy <- __post_init__ <- (dataclass) __init__ + # not a constant: a custom strategy delegating here adds frames, and dags are + # also made from within meshed (code_to_dag, dag arithmetic, ...) + stacklevel=_stacklevel_of_first_frame_outside_meshed(), ) @@ -538,9 +568,32 @@ def ignore_duplicate_outs(duplicates: dict, *, dag_name=None): def raise_on_duplicate_outs(duplicates: dict, *, dag_name=None): """``on_duplicate_outs`` strategy that raises a ``ValidationError``.""" - raise ValidationError( - f"Several func nodes write to the same var node(s): {duplicates}" - ) + raise ValidationError(duplicate_outs_string(duplicates, dag_name=dag_name)) + + +def only_new_duplicate_outs(strategy, known_duplicate_outs: dict): + """Wrap an ``on_duplicate_outs`` strategy so that it is only called for + duplications that are NOT already those of ``known_duplicate_outs`` (the + duplications of the dag(s) a new dag is derived from, which were already dealt + with when those dags were made). + + A derived dag may rename its nodes (``copy``, ``ch_names``), so a duplication + that has the same shape as the source's counts as the same one. + """ + known_names = set(known_duplicate_outs) + known_shape = sorted(map(len, known_duplicate_outs.values())) + + def on_duplicate_outs(duplicates: dict, *, dag_name=None): + new_duplicates = { + out: names for out, names in duplicates.items() if out not in known_names + } + if not new_duplicates: + return + if sorted(map(len, duplicates.values())) == known_shape: + return # same duplications as the source's, only renamed + strategy(new_duplicates, dag_name=dag_name) + + return on_duplicate_outs # TODO: caching last scope isn't really the DAG's direct concern -- it's a debugging @@ -605,7 +658,7 @@ class DAG: ) # What to do when several func nodes write to the same var node (see #40). # Alternatives: ignore_duplicate_outs, raise_on_duplicate_outs, or your own. - on_duplicate_outs: Callable[[dict], None] = field( + on_duplicate_outs: Callable[..., None] = field( default=warn_on_duplicate_outs, repr=False ) @@ -871,12 +924,22 @@ def __getitem__(self, item): """ return self._getitem(item) + def _derived_on_duplicate_outs(self, *also_from): + """The ``on_duplicate_outs`` to give a dag derived from this one (and, + optionally, other func nodes or dags it's combined with): the same strategy, + but only told about duplications that are new (those of the sources were + already dealt with when the sources were made).""" + known = dict(duplicate_outs(self.func_nodes)) + for source in also_from: + known.update(duplicate_outs(getattr(source, "func_nodes", source))) + return only_new_duplicate_outs(self.on_duplicate_outs, known) + def _getitem(self, item): return DAG( func_nodes=self._ordered_subgraph_nodes(item), cache_last_scope=self.cache_last_scope, parameter_merge=self.parameter_merge, - on_duplicate_outs=ignore_duplicate_outs, # already warned about, if any + on_duplicate_outs=self._derived_on_duplicate_outs(), ) def _ordered_subgraph_nodes(self, item): @@ -959,7 +1022,9 @@ def partial( # TODO: mk_instance: What about other init args (cache_last_scope, ...)? mk_instance = type(self) func_nodes = partialized_funcnodes(self, **keyword_dflts) - new_dag = mk_instance(func_nodes, on_duplicate_outs=ignore_duplicate_outs) + new_dag = mk_instance( + func_nodes, on_duplicate_outs=self._derived_on_duplicate_outs() + ) if _remove_bound_arguments: new_sig = Sig(new_dag).remove_names(list(keyword_dflts)) new_sig(new_dag) # Change the signature of new_dag with bound args removed @@ -1253,9 +1318,11 @@ def ch_funcs( >>> ch_fnode2 = partial(ch_func_node_func, func_comparator=same_set_of_names) >>> d = dag.ch_funcs(ch_fnode2, g=lambda z=2, y=1: y / z); """ - return ch_funcs( + new_dag = ch_funcs( self, func_mapping=func_mapping, ch_func_node_func=ch_func_node_func ) + new_dag.on_duplicate_outs = self._derived_on_duplicate_outs() + return new_dag # _validate_func_mapping(func_mapping, self) # @@ -1385,7 +1452,11 @@ def __radd__(self, other): # we would like to control some orders of things via the order of addition # (thinkg list addition versus set addition for example), so instead we write # the explicit code: - return DAG(self._prepare_other_for_addition(other) + list(self.func_nodes)) + other_func_nodes = self._prepare_other_for_addition(other) + return DAG( + other_func_nodes + list(self.func_nodes), + on_duplicate_outs=self._derived_on_duplicate_outs(other_func_nodes), + ) def __add__(self, other): """A union of DAGs. @@ -1399,7 +1470,12 @@ def __add__(self, other): >>> dag([1,2,3]) ([1, 2, 3], (1, 2, 3)) """ - return DAG(list(self.func_nodes) + self._prepare_other_for_addition(other)) + other_func_nodes = self._prepare_other_for_addition(other) + return DAG( + list(self.func_nodes) + other_func_nodes, + # a union can *create* a duplication: only those are worth reporting + on_duplicate_outs=self._derived_on_duplicate_outs(other_func_nodes), + ) def copy(self, renamer=numbered_suffix_renamer): """Make a new ``DAG`` from renamed copies of the func nodes (see ``ch_names`` for what ``renamer`` may be). @@ -1418,7 +1494,8 @@ def copy(self, renamer=numbered_suffix_renamer): """ return DAG( ch_names(self.func_nodes, renamer=renamer), - on_duplicate_outs=ignore_duplicate_outs, + # a renamer could collapse two outs into one: that would be new + on_duplicate_outs=self._derived_on_duplicate_outs(), ) def add_edge(self, from_node, to_node, to_param=None): @@ -1537,7 +1614,7 @@ def add_edge(self, from_node, to_node, to_param=None): condition=lambda x: x == to_node, replacement=lambda x: new_to_node, ), - on_duplicate_outs=ignore_duplicate_outs, + on_duplicate_outs=self._derived_on_duplicate_outs(), ) # TODO: There are optimization and pre-validation opportunities here! diff --git a/meshed/tests/test_duplicate_outs.py b/meshed/tests/test_duplicate_outs.py index 240ff23b..04fc4215 100644 --- a/meshed/tests/test_duplicate_outs.py +++ b/meshed/tests/test_duplicate_outs.py @@ -102,3 +102,54 @@ def test_duplicate_outs_lists_nodes_in_execution_order(): # the last one listed is the one whose value the dag returns assert dag(x=1, y=2) == getattr(dag, "last_scope", None) or True assert names[-1] == dag.func_nodes[-1].name + + +def _baz(c): + return c + + +def test_strategy_survives_derivation(): + from meshed.dag import raise_on_duplicate_outs as _raise + + nodes = [FuncNode(_foo, name="n1", out="u"), FuncNode(_bar, name="n2", out="v")] + strict = DAG(nodes, on_duplicate_outs=_raise) + # a union that CREATES a duplication still uses the strict strategy + with pytest.raises(ValueError): + strict + DAG([FuncNode(_baz, name="n3", out="u")]) + + +def test_opting_out_survives_derivation(): + dag = _dag_with_duplicate_outs(on_duplicate_outs=ignore_duplicate_outs) + with warnings.catch_warnings(): + warnings.simplefilter("error", DuplicateOutsWarning) + dag + DAG([]) + + +def test_renaming_copy_that_creates_a_duplicate_warns(): + dag = DAG([FuncNode(_foo, name="n1", out="alpha"), FuncNode(_bar, name="n2", out="beta")]) + collapsing = lambda name: "z" if name in ("alpha", "beta") else name + "_c" + with pytest.warns(DuplicateOutsWarning): + dag.copy(renamer=collapsing) + + +def test_warning_points_at_the_callers_code(tmp_path): + # the caller must be outside of meshed (this test module is inside it), so build + # the dag from a little module of its own + caller = tmp_path / "dag_builder.py" + caller.write_text( + "from meshed import DAG, FuncNode\n" + "def f(a): return a\n" + "def g(b): return b\n" + "def build():\n" + " return DAG([FuncNode(f, name='n1', out='same')," + " FuncNode(g, name='n2', out='same')])\n" + ) + import importlib.util + + spec = importlib.util.spec_from_file_location("dag_builder", caller) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + with pytest.warns(DuplicateOutsWarning) as record: + module.build() + assert record[0].filename == str(caller) + assert record[0].lineno == 5 # the DAG(...) call