From 513a43552b4613e466bc2f9ed8028bdaf416d67c Mon Sep 17 00:00:00 2001 From: Thor Whalen <1906276+thorwhalen@users.noreply.github.com> Date: Tue, 22 Sep 2026 17:57:33 +0000 Subject: [PATCH 1/5] fix(CachedDag): cache the inputs a call was given (#34) - Persist explicitly given var-node inputs in CachedDag.cache after a successful computation. - Re-passing an input with its cached value is allowed; a different value raises ValueError. - Cached values take precedence over the dag's defaults. - A user-supplied Mapping cache is now actually used (was left unset). - Fix the always-true assertion in cached_dag_test; add meshed/tests/test_cached_dag.py. Co-Authored-By: Claude Opus 5 --- meshed/scrap/cached_dag.py | 35 ++++++++++++++--- meshed/tests/test_cached_dag.py | 70 +++++++++++++++++++++++++++++++++ 2 files changed, 99 insertions(+), 6 deletions(-) create mode 100644 meshed/tests/test_cached_dag.py diff --git a/meshed/scrap/cached_dag.py b/meshed/scrap/cached_dag.py index f6d4d74b..3f8b8b74 100644 --- a/meshed/scrap/cached_dag.py +++ b/meshed/scrap/cached_dag.py @@ -227,7 +227,11 @@ def __init__(self, dag, cache=True, name=None): "This type of cache is not implemented (must resolve to a Mapping): " f"{cache=}" ) - self._cache = ChainMap(self.defaults, self.cache) + else: + self.cache = cache + # Cached values (computed outputs and inputs the dag was called with) take + # precedence over the dag's defaults. + self._cache = ChainMap(self.cache, self.defaults) @property def __name__(self): @@ -246,14 +250,25 @@ def func_node_id(self, k): def __call__(self, k, /, **input_kwargs): # print(f"Calling ({k=},{input_kwargs=})\t{self.cache=}") input_kwargs = dict(input_kwargs) - if intersection := (input_kwargs.keys() & self.cache.keys()): + if conflicts := { + name + for name in input_kwargs.keys() & self.cache.keys() + if input_kwargs[name] != self.cache[name] + }: # TODO: Can give the user a more informative/correct message, since the # user has more options than just the root nodes: They some combination of # intermediates would also satisfy requirements. raise ValueError( - f"input_kwargs can't contain any keys that are already in cache! " - f"These names were in both: {intersection}" + f"input_kwargs can't contain keys that are already in cache with a " + f"different value! These names were in both: {conflicts}" ) + output = self._compute(k, input_kwargs) + # Only persist inputs once they led to a successful computation, so that a + # failed (e.g. mistaken) call doesn't pin values in the cache. + self._cache_inputs(input_kwargs) + return output + + def _compute(self, k, input_kwargs): _cache = ChainMap(input_kwargs, self._cache) if k in _cache: return _cache[k] @@ -287,6 +302,14 @@ def __call__(self, k, /, **input_kwargs): f"argument: '{k}'" ) + def _cache_inputs(self, input_kwargs): + """Persist the (explicitly given) values of the dag's var nodes in the cache, + so later calls can reuse them (see https://github.com/i2mint/meshed/issues/34). + Keys that are not var nodes of the dag are not cached.""" + for name, value in input_kwargs.items(): + if name in self.var_nodes: + self.cache[name] = value + def _call(self, k, /, **kwargs): return self(k, **kwargs) @@ -359,9 +382,9 @@ def g(a, y=2): dag = DAG([f, g]) c = CachedDag(dag) - c("g", a=1) + assert c("g", a=1) == 2 assert c.cache == {"g": 2, "a": 1} - assert c("f" == 2) + assert c("f") == 2 def add(a, b=1): diff --git a/meshed/tests/test_cached_dag.py b/meshed/tests/test_cached_dag.py new file mode 100644 index 00000000..e385bddc --- /dev/null +++ b/meshed/tests/test_cached_dag.py @@ -0,0 +1,70 @@ +"""Tests for ``meshed.scrap.cached_dag.CachedDag`` (see i2mint/meshed#34).""" + +import pytest + +from meshed import DAG +from meshed.scrap.cached_dag import CachedDag, cached_dag_test + + +def f(a, x=1): + return a + x + + +def g(a, y=2): + return a * y + + +def _cached_dag(**kwargs): + return CachedDag(DAG([f, g]), **kwargs) + + +def test_cached_dag_test_function(): + cached_dag_test() + + +def test_inputs_are_cached_and_reused(): + c = _cached_dag() + assert c("g", a=1) == 2 + assert c.cache == {"g": 2, "a": 1} + assert c("f") == 2 # uses the cached ``a`` + assert c.cache == {"g": 2, "a": 1, "f": 2} + + +def test_repeating_an_input_with_the_same_value_is_allowed(): + c = _cached_dag() + c("g", a=1) + assert c("f", a=1) == 2 + + +def test_conflicting_input_raises(): + c = _cached_dag() + c("g", a=1) + with pytest.raises(ValueError): + c("f", a=10) + + +def test_cached_input_takes_precedence_over_default(): + c = _cached_dag() + assert c("g", a=1, y=5) == 5 + assert c.cache["y"] == 5 + assert c("y") == 5 # not the dag's default (2) + + +def test_non_var_node_inputs_are_not_cached(): + c = _cached_dag() + c("g", a=1, not_a_node=3) + assert "not_a_node" not in c.cache + + +def test_failed_call_does_not_cache_inputs(): + c = _cached_dag() + with pytest.raises(TypeError): + c("g", y=5) # missing ``a`` + assert c.cache == {} + + +def test_custom_mapping_cache(): + cache = {} + c = _cached_dag(cache=cache) + assert c("g", a=1) == 2 + assert cache == {"g": 2, "a": 1} From 0b815e062e61f06e3ab4c686aec01087f4e3c54d Mon Sep 17 00:00:00 2001 From: Thor Whalen <1906276+thorwhalen@users.noreply.github.com> Date: Tue, 22 Sep 2026 18:00:48 +0000 Subject: [PATCH 2/5] fix(CachedDag): validate and cache inputs once per top-level call Address review of #83: recurse through _compute (not __call__) so the cache check and input caching run once per call; compare values without raising (identity first, then guarded ==); reject inputs that contradict cached downstream outputs; more tests (multi-level, array-like, nan, arg order, intermediate inputs). Co-Authored-By: Claude Opus 5 --- meshed/scrap/cached_dag.py | 76 +++++++++++++++++++++++++-------- meshed/tests/test_cached_dag.py | 70 ++++++++++++++++++++++++++++++ 2 files changed, 129 insertions(+), 17 deletions(-) diff --git a/meshed/scrap/cached_dag.py b/meshed/scrap/cached_dag.py index 3f8b8b74..56a20d59 100644 --- a/meshed/scrap/cached_dag.py +++ b/meshed/scrap/cached_dag.py @@ -70,6 +70,27 @@ def __setitem__(self, key, value): NoSuchKey = type("NoSuchKey", (), {}) +def _is_same_value(a, b): + """Whether ``a`` and ``b`` are the same value, never raising. + + Identity counts as sameness (so a ``nan`` object equals itself). Values whose + ``==`` doesn't give a plain truth value (e.g. numpy arrays) are only considered + the same if they are identical. + + >>> _is_same_value(1, 1), _is_same_value(1, 2) + (True, False) + >>> nan = float('nan') + >>> _is_same_value(nan, nan), _is_same_value(nan, float('nan')) + (True, False) + """ + if a is b: + return True + try: + return bool(a == b) + except Exception: + return False + + # TODO: Cache validation and invalidation # TODO: Continue constructing uppward towards lazyprop-using class (instances are # varnodes) @@ -248,13 +269,29 @@ def func_node_id(self, k): # TODO: Consider having args and kwargs instead of just input_kwargs. # or making it (k, /, *args, **kwargs) def __call__(self, k, /, **input_kwargs): - # print(f"Calling ({k=},{input_kwargs=})\t{self.cache=}") input_kwargs = dict(input_kwargs) - if conflicts := { + self._validate_inputs_against_cache(input_kwargs) + output = self._compute(k, input_kwargs) + # Only persist inputs once they led to a successful computation, so that a + # failed (e.g. mistaken) call doesn't pin values in the cache. + self._cache_inputs(input_kwargs) + return output + + def _validate_inputs_against_cache(self, input_kwargs): + """Raise a ``ValueError`` if ``input_kwargs`` contradicts the cache. + + An input contradicts the cache if: + + - it is already cached with a different value, or + - it is not cached, but cached outputs downstream of it were computed using + its default (or it has no default), and the given value differs from it. + """ + conflicts = { name for name in input_kwargs.keys() & self.cache.keys() - if input_kwargs[name] != self.cache[name] - }: + if not _is_same_value(input_kwargs[name], self.cache[name]) + } + if conflicts: # TODO: Can give the user a more informative/correct message, since the # user has more options than just the root nodes: They some combination of # intermediates would also satisfy requirements. @@ -262,34 +299,38 @@ def __call__(self, k, /, **input_kwargs): f"input_kwargs can't contain keys that are already in cache with a " f"different value! These names were in both: {conflicts}" ) - output = self._compute(k, input_kwargs) - # Only persist inputs once they led to a successful computation, so that a - # failed (e.g. mistaken) call doesn't pin values in the cache. - self._cache_inputs(input_kwargs) - return output + for name in input_kwargs.keys() - self.cache.keys(): + if name not in self.var_nodes: + continue + cached_downstream = descendants(self.dag.graph_ids, [name]) & set( + self.cache + ) + if cached_downstream and not _is_same_value( + input_kwargs[name], self.defaults.get(name, NoSuchKey) + ): + raise ValueError( + f"The value given for {name!r} contradicts the cache: " + f"{sorted(cached_downstream)} were already computed without it." + ) def _compute(self, k, input_kwargs): _cache = ChainMap(input_kwargs, self._cache) if k in _cache: return _cache[k] - input_kwargs = dict(input_kwargs) func_node_id = self.func_node_id(k) - # print(f"{func_node_id=}") if func_node_id: if (output := self.cache.get(func_node_id)) is not None: return output else: func_node = self.func_node_of_id[func_node_id] input_sources = { - src: self(src, **input_kwargs) for src in func_node.bind.values() + src: self._compute(src, input_kwargs) + for src in func_node.bind.values() } - # inputs = dict(input_sources, **input_kwargs) # # TODO: do we need to include **self.defaults in the middle? inputs = ChainMap(_cache, input_sources) - # print(f"Computing {func_node_id}: ", end=" ") output = func_node.call_on_scope(inputs, write_output_into_scope=False) self.cache[func_node_id] = output - # print(f"result -> {output}") return output else: # k is a root node assert k in self.roots, f"Was expecting this to be a root node: {k}" @@ -305,9 +346,10 @@ def _compute(self, k, input_kwargs): def _cache_inputs(self, input_kwargs): """Persist the (explicitly given) values of the dag's var nodes in the cache, so later calls can reuse them (see https://github.com/i2mint/meshed/issues/34). - Keys that are not var nodes of the dag are not cached.""" + Keys that are not var nodes of the dag, or that are already cached (and were + validated to hold the same value), are not written.""" for name, value in input_kwargs.items(): - if name in self.var_nodes: + if name in self.var_nodes and name not in self.cache: self.cache[name] = value def _call(self, k, /, **kwargs): diff --git a/meshed/tests/test_cached_dag.py b/meshed/tests/test_cached_dag.py index e385bddc..5dba8f6e 100644 --- a/meshed/tests/test_cached_dag.py +++ b/meshed/tests/test_cached_dag.py @@ -68,3 +68,73 @@ def test_custom_mapping_cache(): c = _cached_dag(cache=cache) assert c("g", a=1) == 2 assert cache == {"g": 2, "a": 1} + + +def k2(y, a): + return y - a + + +def test_failed_call_does_not_cache_inputs_regardless_of_arg_order(): + c = CachedDag(DAG([k2])) + with pytest.raises(TypeError): + c("k2", y=5) # missing ``a``, which comes *after* ``y`` + assert c.cache == {} + + +def h(f, g): + return f + g + + +def test_multi_level_dag_with_values_without_plain_equality(): + class Arr: + """Stand-in for a numpy array: ``==`` has an ambiguous truth value.""" + + def __init__(self, v): + self.v = v + + def __add__(self, other): + return Arr(self.v + (other.v if isinstance(other, Arr) else other)) + + __radd__ = __add__ + + def __mul__(self, other): + return Arr(self.v * other) + + def __eq__(self, other): + class Ambiguous: + def __bool__(self): + raise ValueError("ambiguous") + + return Ambiguous() + + __ne__ = __eq__ + + c = CachedDag(DAG([f, g, h])) + a = Arr(1) + assert c("h", a=a).v == 4 # (1 + 1) + (1 * 2) + assert c.cache["a"] is a + assert c("h", a=a).v == 4 # same object again: allowed + + +def test_nan_input(): + nan = float("nan") + c = _cached_dag() + out = c("f", a=nan) + assert out != out # nan + c("g", a=nan) # the same nan object is accepted again + + +def test_input_contradicting_cached_outputs_raises(): + c = _cached_dag() + c("f", a=1) # computed with the default x=1 + assert c("f", a=1, x=1) == 2 # same as the default used: fine + with pytest.raises(ValueError): + c("f", a=1, x=5) # f was computed with x=1 + + +def test_intermediate_node_as_input(): + c = CachedDag(DAG([f, g, h])) + assert c("h", f=100, a=1) == 102 + assert c.cache["f"] == 100 + with pytest.raises(ValueError): + c("h", f=5) From 931cc59833a3cebc9c01bb81b62952db61559fdf Mon Sep 17 00:00:00 2001 From: Thor Whalen <1906276+thorwhalen@users.noreply.github.com> Date: Tue, 22 Sep 2026 18:02:55 +0000 Subject: [PATCH 3/5] fix(CachedDag): roll back partial results when a call fails Second review round of #83: a failed call left its intermediate outputs in the cache while (deliberately) not caching its inputs, so the retry was rejected by the downstream-contradiction check. Now the keys a failed call added are removed, and retries are tested. Co-Authored-By: Claude Opus 5 --- meshed/scrap/cached_dag.py | 14 +++++++++++++- meshed/tests/test_cached_dag.py | 31 +++++++++++++++++++++++++++++++ 2 files changed, 44 insertions(+), 1 deletion(-) diff --git a/meshed/scrap/cached_dag.py b/meshed/scrap/cached_dag.py index 56a20d59..0badb5eb 100644 --- a/meshed/scrap/cached_dag.py +++ b/meshed/scrap/cached_dag.py @@ -271,7 +271,15 @@ def func_node_id(self, k): def __call__(self, k, /, **input_kwargs): input_kwargs = dict(input_kwargs) self._validate_inputs_against_cache(input_kwargs) - output = self._compute(k, input_kwargs) + keys_before = set(self.cache) + try: + output = self._compute(k, input_kwargs) + except BaseException: + # Roll back the outputs computed during this failed call, so the cache + # stays consistent with its (uncached) inputs and the call can be retried. + for key in set(self.cache) - keys_before: + del self.cache[key] + raise # Only persist inputs once they led to a successful computation, so that a # failed (e.g. mistaken) call doesn't pin values in the cache. self._cache_inputs(input_kwargs) @@ -285,6 +293,10 @@ def _validate_inputs_against_cache(self, input_kwargs): - it is already cached with a different value, or - it is not cached, but cached outputs downstream of it were computed using its default (or it has no default), and the given value differs from it. + + Values are compared with ``_is_same_value``: for values without a plain + ``==`` truth value (e.g. numpy arrays), only the very same object counts as + the same value. """ conflicts = { name diff --git a/meshed/tests/test_cached_dag.py b/meshed/tests/test_cached_dag.py index 5dba8f6e..66bc8fa4 100644 --- a/meshed/tests/test_cached_dag.py +++ b/meshed/tests/test_cached_dag.py @@ -138,3 +138,34 @@ def test_intermediate_node_as_input(): assert c.cache["f"] == 100 with pytest.raises(ValueError): c("h", f=5) + + +def k(f, b): + return f + b + + +def test_failed_call_can_be_retried(): + c = CachedDag(DAG([f, k])) + with pytest.raises(TypeError): + c("k", a=1) # missing ``b`` (after ``f`` was computed) + assert c.cache == {} # partial results were rolled back + assert c("k", a=1, b=3) == 5 + assert c.cache == {"f": 2, "k": 5, "a": 1, "b": 3} + + +def test_exception_in_function_can_be_retried(): + calls = [] + + def flaky(a): + calls.append(a) + if len(calls) == 1: + raise RuntimeError("transient") + return a + + def total(f, flaky): + return f + flaky + + c = CachedDag(DAG([f, flaky, total])) + with pytest.raises(RuntimeError): + c("total", a=1) + assert c("total", a=1) == 3 From fe6a6c82015ee4a622a6f42e887e887cb1cb5391 Mon Sep 17 00:00:00 2001 From: Thor Whalen <1906276+thorwhalen@users.noreply.github.com> Date: Tue, 22 Sep 2026 18:05:30 +0000 Subject: [PATCH 4/5] fix(CachedDag): reject inputs already determined by the cache; safe rollback Third review round of #83: an intermediate var node given as an input was never checked against cached values it would be derived from, which silently poisoned the cache. It is now rejected. The rollback of a failed call no longer masks the original exception if the cache doesn't support deletion. Documented that inputs of a single call are not validated against each other. Co-Authored-By: Claude Opus 5 --- meshed/scrap/cached_dag.py | 21 ++++++++++++++++++--- meshed/tests/test_cached_dag.py | 7 +++++++ 2 files changed, 25 insertions(+), 3 deletions(-) diff --git a/meshed/scrap/cached_dag.py b/meshed/scrap/cached_dag.py index 0badb5eb..78eb1d2e 100644 --- a/meshed/scrap/cached_dag.py +++ b/meshed/scrap/cached_dag.py @@ -277,8 +277,11 @@ def __call__(self, k, /, **input_kwargs): except BaseException: # Roll back the outputs computed during this failed call, so the cache # stays consistent with its (uncached) inputs and the call can be retried. - for key in set(self.cache) - keys_before: - del self.cache[key] + try: + for key in set(self.cache) - keys_before: + del self.cache[key] + except Exception: # pragma: no cover - e.g. a cache without __delitem__ + pass # never mask the original error with a rollback error raise # Only persist inputs once they led to a successful computation, so that a # failed (e.g. mistaken) call doesn't pin values in the cache. @@ -292,7 +295,13 @@ def _validate_inputs_against_cache(self, input_kwargs): - it is already cached with a different value, or - it is not cached, but cached outputs downstream of it were computed using - its default (or it has no default), and the given value differs from it. + its default (or it has no default), and the given value differs from it, or + - it is not cached and not a root, but values it would be computed from are + cached (so the cache already determines it). + + Note that inputs are only validated against the *cache*: values given in the + same call are not checked against each other (``c('h', f=100, a=1)`` is + accepted even if ``f`` wouldn't be computed as ``100`` from ``a=1``). Values are compared with ``_is_same_value``: for values without a plain ``==`` truth value (e.g. numpy arrays), only the very same object counts as @@ -324,6 +333,12 @@ def _validate_inputs_against_cache(self, input_kwargs): f"The value given for {name!r} contradicts the cache: " f"{sorted(cached_downstream)} were already computed without it." ) + cached_upstream = descendants(self.reversed_graph, [name]) & set(self.cache) + if cached_upstream: + raise ValueError( + f"{name!r} is determined by values that are already cached " + f"({sorted(cached_upstream)}), so it can't be given as an input." + ) def _compute(self, k, input_kwargs): _cache = ChainMap(input_kwargs, self._cache) diff --git a/meshed/tests/test_cached_dag.py b/meshed/tests/test_cached_dag.py index 66bc8fa4..d7403804 100644 --- a/meshed/tests/test_cached_dag.py +++ b/meshed/tests/test_cached_dag.py @@ -169,3 +169,10 @@ def total(f, flaky): with pytest.raises(RuntimeError): c("total", a=1) assert c("total", a=1) == 3 + + +def test_input_determined_by_cached_upstream_values_raises(): + c = CachedDag(DAG([f, g, h])) + c("f", a=1) # caches a=1, which determines g (and therefore h) + with pytest.raises(ValueError): + c("h", g=99) From 1878ce3ca1518eb4d8bd4675511f7cb0774a9020 Mon Sep 17 00:00:00 2001 From: Thor Whalen <1906276+thorwhalen@users.noreply.github.com> Date: Tue, 22 Sep 2026 18:08:11 +0000 Subject: [PATCH 5/5] fix(CachedDag): only reject an intermediate input the cache really determines Fourth review round of #83: the upstream check fired whenever any ancestor was cached; now it only fires when every source of the node is settled by the cache or the dag's defaults, so giving an intermediate whose other inputs are unknown still works. Co-Authored-By: Claude Opus 5 --- meshed/scrap/cached_dag.py | 25 ++++++++++++++++++++----- meshed/tests/test_cached_dag.py | 12 ++++++++++++ 2 files changed, 32 insertions(+), 5 deletions(-) diff --git a/meshed/scrap/cached_dag.py b/meshed/scrap/cached_dag.py index 78eb1d2e..73a5cb99 100644 --- a/meshed/scrap/cached_dag.py +++ b/meshed/scrap/cached_dag.py @@ -296,8 +296,8 @@ def _validate_inputs_against_cache(self, input_kwargs): - it is already cached with a different value, or - it is not cached, but cached outputs downstream of it were computed using its default (or it has no default), and the given value differs from it, or - - it is not cached and not a root, but values it would be computed from are - cached (so the cache already determines it). + - it is not cached and not a root, but the cache (and the dag's defaults) + already determine everything it would be computed from. Note that inputs are only validated against the *cache*: values given in the same call are not checked against each other (``c('h', f=100, a=1)`` is @@ -333,13 +333,28 @@ def _validate_inputs_against_cache(self, input_kwargs): f"The value given for {name!r} contradicts the cache: " f"{sorted(cached_downstream)} were already computed without it." ) - cached_upstream = descendants(self.reversed_graph, [name]) & set(self.cache) - if cached_upstream: + if name not in self.roots and self._is_determined_by_cache(name): + cached_upstream = descendants(self.reversed_graph, [name]) & set( + self.cache + ) raise ValueError( - f"{name!r} is determined by values that are already cached " + f"{name!r} is already determined by the cache " f"({sorted(cached_upstream)}), so it can't be given as an input." ) + def _is_determined_by_cache(self, k): + """Whether ``k``'s value is already fixed by the cache and the dag's defaults + (that is, whether ``self(k)``, with no inputs, would give a value).""" + if k in self.cache or k in self.defaults: + return True + func_node_id = self.func_node_id(k) + if func_node_id is None: # a root node with no value in sight + return False + func_node = self.func_node_of_id[func_node_id] + return all( + self._is_determined_by_cache(src) for src in func_node.bind.values() + ) + def _compute(self, k, input_kwargs): _cache = ChainMap(input_kwargs, self._cache) if k in _cache: diff --git a/meshed/tests/test_cached_dag.py b/meshed/tests/test_cached_dag.py index d7403804..33e5d2eb 100644 --- a/meshed/tests/test_cached_dag.py +++ b/meshed/tests/test_cached_dag.py @@ -176,3 +176,15 @@ def test_input_determined_by_cached_upstream_values_raises(): c("f", a=1) # caches a=1, which determines g (and therefore h) with pytest.raises(ValueError): c("h", g=99) + + +def test_input_not_determined_by_cache_is_allowed(): + def u(t): + return t + 1 + + def v(u, s): + return u + s + + c = CachedDag(DAG([u, v])) + c("u", t=1) # caches t, u -- but ``v`` also needs ``s``, so it isn't determined + assert c("v", v=10) == 10