This is an automated email from the ASF dual-hosted git repository.

kaxil pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git


The following commit(s) were added to refs/heads/main by this push:
     new 73b07c5aedb Name the tools a durable agent retry will run again 
(#73873)
73b07c5aedb is described below

commit 73b07c5aedbf6e5e0b78aabc290bca9b1965d77f
Author: Kaxil Naik <[email protected]>
AuthorDate: Wed Sep 30 00:43:31 2026 +0100

    Name the tools a durable agent retry will run again (#73873)
    
    * Name the tools a durable retry will run again
    
    With durable=True, a step whose cache write is skipped (a tool result that 
is
    not JSON-serializable, a task state store write that fails) ran live but is
    not cached, so a retry runs it again. The caching wrappers still counted it
    as cached, and the only record was a backend warning naming a step key, not
    the tool. For a tool with side effects, nothing said the side effect would
    repeat.
    
    The save methods now return whether they wrote, the step counter records
    skipped model responses and the names of skipped tools, and each skipped
    write logs a warning naming the tool. AgentOperator logs its durable summary
    from a finally block, so the failed attempt that Airflow retries gets it 
too,
    and the summary lists every tool that was not cached.
    
    * Run the durable skip-count tests on Airflow 3.0 too
    
    The cap_structlog fixture needs airflow._shared, which lands in Airflow 3.1,
    so the three tests asserting the skip warning errored on the 3.0.6 compat
    run and took their counting assertions with them. Split each: the counting
    check runs on every version, the log check only on 3.1+.
---
 providers/common/ai/docs/durable_execution.rst     |  14 ++-
 .../airflow/providers/common/ai/durable/base.py    |   9 +-
 .../providers/common/ai/durable/caching_model.py   |  14 ++-
 .../providers/common/ai/durable/caching_toolset.py |  16 ++-
 .../providers/common/ai/durable/step_counter.py    |   4 +
 .../airflow/providers/common/ai/durable/storage.py |  16 ++-
 .../common/ai/durable/task_state_store.py          |  16 ++-
 .../airflow/providers/common/ai/operators/agent.py |  55 +++++++---
 .../unit/common/ai/durable/test_caching_model.py   |  41 ++++++-
 .../unit/common/ai/durable/test_caching_toolset.py |  36 ++++++-
 .../tests/unit/common/ai/durable/test_storage.py   |   9 ++
 .../common/ai/durable/test_task_state_store.py     |  22 +++-
 .../tests/unit/common/ai/operators/test_agent.py   | 118 ++++++++++++++++++++-
 13 files changed, 324 insertions(+), 46 deletions(-)

diff --git a/providers/common/ai/docs/durable_execution.rst 
b/providers/common/ai/docs/durable_execution.rst
index 262f48f6aa5..ab6647f2492 100644
--- a/providers/common/ai/docs/durable_execution.rst
+++ b/providers/common/ai/docs/durable_execution.rst
@@ -113,8 +113,9 @@ not invalidate an already-cached result for an identical 
call, and pointing
 ``llm_conn_id`` at a different endpoint serving the same model name does not
 invalidate cached responses -- clear the cache to force a fully fresh run.
 
-After the run, a single INFO summary line reports how many steps were
-replayed vs executed fresh. Per-step detail is available at DEBUG level.
+When the run ends, successfully or not, an INFO line reports how many steps
+were replayed from the cache and how many new steps were cached. Per-step
+detail is available at DEBUG level.
 
 The cache is scoped to a single task instance (Dag id, run id, task id, and
 map index), so each run replays only its own steps. On Airflow >= 3.3 the cache
@@ -152,9 +153,12 @@ example, check whether the operation already completed 
before acting, or
 use database constraints to prevent duplicate writes.
 
 Tool results must be JSON-serializable to be cached. If a tool returns a
-non-serializable value (e.g. ``BinaryContent`` from MCP tools), that step is
-skipped with a warning and will re-execute on retry instead of replaying from
-cache. The task itself still succeeds.
+non-serializable value (e.g. ``BinaryContent`` from MCP tools), or a write to
+the task state store fails, the step is not cached and runs again on retry
+instead of replaying. The step itself still succeeds. Each such step logs a
+WARNING naming the tool, and the end-of-run summary lists every tool that was
+not cached. A step that re-runs can change what the model sees next, so the
+steps after it may re-run too.
 
 See also
 --------
diff --git 
a/providers/common/ai/src/airflow/providers/common/ai/durable/base.py 
b/providers/common/ai/src/airflow/providers/common/ai/durable/base.py
index 9fff491012e..56cd17b0e93 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/durable/base.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/durable/base.py
@@ -46,13 +46,18 @@ class DurableStorageProtocol(Protocol):
     
:class:`~airflow.providers.common.ai.durable.task_state_store.TaskStateStoreDurableStorage`
     (AIP-103 task state store, Airflow >= 3.3). ``CachingModel`` and
     ``CachingToolset`` depend on this interface, not a concrete backend.
+
+    Both ``save_*`` methods return whether the entry was written. A backend may
+    skip a write (a tool result that is not JSON-serializable, a store write 
that
+    fails) without failing the step; the step then re-runs live on retry, and 
the
+    caller counts it as skipped rather than cached.
     """
 
-    def save_model_response(self, key: str, response: ModelResponse, *, 
fingerprint: str | None) -> None: ...
+    def save_model_response(self, key: str, response: ModelResponse, *, 
fingerprint: str | None) -> bool: ...
 
     def load_model_response(self, key: str) -> tuple[ModelResponse | None, str 
| None]: ...
 
-    def save_tool_result(self, key: str, result: Any, *, fingerprint: str | 
None) -> None: ...
+    def save_tool_result(self, key: str, result: Any, *, fingerprint: str | 
None) -> bool: ...
 
     def load_tool_result(self, key: str) -> tuple[bool, Any, str | None]: ...
 
diff --git 
a/providers/common/ai/src/airflow/providers/common/ai/durable/caching_model.py 
b/providers/common/ai/src/airflow/providers/common/ai/durable/caching_model.py
index 6118311f70c..9c702aae50d 100644
--- 
a/providers/common/ai/src/airflow/providers/common/ai/durable/caching_model.py
+++ 
b/providers/common/ai/src/airflow/providers/common/ai/durable/caching_model.py
@@ -112,7 +112,15 @@ class CachingModel(WrapperModel):
             )
 
         response = await self.wrapped.request(messages, model_settings, 
model_request_parameters)
-        self.storage.save_model_response(key, response, 
fingerprint=fingerprint)
-        self.counter.cached_model += 1
-        log.debug("Durable: cached model response", step=step)
+        if self.storage.save_model_response(key, response, 
fingerprint=fingerprint):
+            self.counter.cached_model += 1
+            log.debug("Durable: cached model response", step=step)
+        else:
+            self.counter.skipped_model += 1
+            # A re-run model step returns fresh tool call ids, so every later 
step's
+            # fingerprint changes and re-runs too.
+            log.warning(
+                "Durable: model response not cached; a retry re-runs this step 
and every step after it",
+                step=step,
+            )
         return response
diff --git 
a/providers/common/ai/src/airflow/providers/common/ai/durable/caching_toolset.py
 
b/providers/common/ai/src/airflow/providers/common/ai/durable/caching_toolset.py
index b411c57ef69..19ced8b5e27 100644
--- 
a/providers/common/ai/src/airflow/providers/common/ai/durable/caching_toolset.py
+++ 
b/providers/common/ai/src/airflow/providers/common/ai/durable/caching_toolset.py
@@ -89,7 +89,17 @@ class CachingToolset(WrapperToolset[Any]):
             )
 
         result = await self.wrapped.call_tool(name, tool_args, ctx, tool)
-        self.storage.save_tool_result(key, result, fingerprint=fingerprint)
-        self.counter.cached_tool += 1
-        log.debug("Durable: cached tool result", step=step, tool=name)
+        if self.storage.save_tool_result(key, result, fingerprint=fingerprint):
+            self.counter.cached_tool += 1
+            log.debug("Durable: cached tool result", step=step, tool=name)
+        else:
+            self.counter.skipped_tools.append(name)
+            # Named here rather than only in the end-of-run summary: this 
warning is
+            # logged on every path, including the failed attempt that Airflow 
retries.
+            log.warning(
+                "Durable: tool result not cached; a retry runs this tool 
again, "
+                "and may re-run the steps after it",
+                step=step,
+                tool=name,
+            )
         return result
diff --git 
a/providers/common/ai/src/airflow/providers/common/ai/durable/step_counter.py 
b/providers/common/ai/src/airflow/providers/common/ai/durable/step_counter.py
index a32b4a1ee3d..314195cbefb 100644
--- 
a/providers/common/ai/src/airflow/providers/common/ai/durable/step_counter.py
+++ 
b/providers/common/ai/src/airflow/providers/common/ai/durable/step_counter.py
@@ -35,6 +35,10 @@ class DurableStepCounter:
         self.replayed_tool: int = 0
         self.cached_model: int = 0
         self.cached_tool: int = 0
+        # Steps that ran live but whose cache write the backend refused. They
+        # re-run on retry, so they must not count as cached.
+        self.skipped_model: int = 0
+        self.skipped_tools: list[str] = []
 
     def next_step(self) -> int:
         """Return the current step and advance the counter."""
diff --git 
a/providers/common/ai/src/airflow/providers/common/ai/durable/storage.py 
b/providers/common/ai/src/airflow/providers/common/ai/durable/storage.py
index 0ff1c84afbe..a3f974bcb42 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/durable/storage.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/durable/storage.py
@@ -108,8 +108,12 @@ class DurableStorage:
         path.parent.mkdir(parents=True, exist_ok=True)
         path.write_text(json.dumps(self._cache))
 
-    def save_model_response(self, key: str, response: ModelResponse, *, 
fingerprint: str | None) -> None:
-        """Serialize and store a ModelResponse with the request fingerprint 
that produced it."""
+    def save_model_response(self, key: str, response: ModelResponse, *, 
fingerprint: str | None) -> bool:
+        """
+        Serialize and store a ModelResponse with the request fingerprint that 
produced it.
+
+        :return: Always ``True``. Unlike the task state store, this backend 
never skips a model response.
+        """
         cache = self._load_cache()
         # Store the dumped messages as native JSON-compatible objects, not a
         # pre-encoded string: the whole cache is JSON-encoded once in
@@ -120,6 +124,7 @@ class DurableStorage:
             "data": ModelMessagesTypeAdapter.dump_python([response], 
mode="json"),
         }
         self._save_cache()
+        return True
 
     def load_model_response(self, key: str) -> tuple[ModelResponse | None, str 
| None]:
         """
@@ -149,13 +154,15 @@ class DurableStorage:
             return None, None
         return messages[0], fingerprint  # type: ignore[return-value]
 
-    def save_tool_result(self, key: str, result: Any, *, fingerprint: str | 
None) -> None:
+    def save_tool_result(self, key: str, result: Any, *, fingerprint: str | 
None) -> bool:
         """
         Store a tool call result with the call fingerprint that produced it.
 
         Non-serializable results (e.g. BinaryContent from MCP tools) are
         skipped with a warning -- the tool call still succeeds, but won't
         be replayed on retry.
+
+        :return: ``True`` if the entry was written, ``False`` if it was 
skipped.
         """
         cache = self._load_cache()
         try:
@@ -170,9 +177,10 @@ class DurableStorage:
                 key=key,
                 type=type(result).__name__,
             )
-            return
+            return False
         cache[key] = {_SENTINEL: True, "value": result, "fingerprint": 
fingerprint}
         self._save_cache()
+        return True
 
     def load_tool_result(self, key: str) -> tuple[bool, Any, str | None]:
         """
diff --git 
a/providers/common/ai/src/airflow/providers/common/ai/durable/task_state_store.py
 
b/providers/common/ai/src/airflow/providers/common/ai/durable/task_state_store.py
index acbada88b51..ecf1c0e066d 100644
--- 
a/providers/common/ai/src/airflow/providers/common/ai/durable/task_state_store.py
+++ 
b/providers/common/ai/src/airflow/providers/common/ai/durable/task_state_store.py
@@ -75,13 +75,15 @@ class TaskStateStoreDurableStorage:
         # attempt; those are reclaimed by the DAG-run cascade, not here.
         self._keys: set[str] = set()
 
-    def save_model_response(self, key: str, response: ModelResponse, *, 
fingerprint: str | None) -> None:
+    def save_model_response(self, key: str, response: ModelResponse, *, 
fingerprint: str | None) -> bool:
         """
         Serialize and store a ModelResponse with the request fingerprint that 
produced it.
 
         Best-effort: the save runs *after* the live model call already 
succeeded, so a
         failed write (e.g. a value over the backend's size limit) must not 
fail the step.
         It is skipped with a warning and simply re-runs live on the next retry.
+
+        :return: ``True`` if the entry was written, ``False`` if it was 
skipped.
         """
         try:
             self._store.set(
@@ -94,8 +96,9 @@ class TaskStateStoreDurableStorage:
             )
         except Exception:
             log.warning("Durable: skipping cache for model response", key=key, 
exc_info=True)
-            return
+            return False
         self._keys.add(key)
+        return True
 
     def load_model_response(self, key: str) -> tuple[ModelResponse | None, str 
| None]:
         """
@@ -119,13 +122,15 @@ class TaskStateStoreDurableStorage:
         fingerprint = raw.get("fingerprint")
         return messages[0], fingerprint if isinstance(fingerprint, str) else 
None  # type: ignore[return-value]
 
-    def save_tool_result(self, key: str, result: Any, *, fingerprint: str | 
None) -> None:
+    def save_tool_result(self, key: str, result: Any, *, fingerprint: str | 
None) -> bool:
         """
         Store a tool call result with the call fingerprint that produced it.
 
         Non-serializable results (e.g. BinaryContent from MCP tools) are 
skipped
         with a warning -- the tool call still succeeds, but won't be replayed 
on
         retry.
+
+        :return: ``True`` if the entry was written, ``False`` if it was 
skipped.
         """
         try:
             # The store validates against pydantic ``JsonValue``, which is 
stricter than
@@ -141,7 +146,7 @@ class TaskStateStoreDurableStorage:
                 key=key,
                 type=type(result).__name__,
             )
-            return
+            return False
         try:
             # Best-effort like the model-response save: a write that fails 
after the tool
             # already ran (e.g. an oversize value) must not fail the step -- 
skip and re-run
@@ -153,8 +158,9 @@ class TaskStateStoreDurableStorage:
             )
         except Exception:
             log.warning("Durable: skipping cache for tool result", key=key, 
exc_info=True)
-            return
+            return False
         self._keys.add(key)
+        return True
 
     def load_tool_result(self, key: str) -> tuple[bool, Any, str | None]:
         """
diff --git 
a/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py 
b/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py
index d6ef9ad8100..8eb0dd85fcb 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py
@@ -18,6 +18,7 @@
 
 from __future__ import annotations
 
+import collections
 import copy
 import hashlib
 import json
@@ -728,6 +729,38 @@ class AgentOperator(CancellableAgentRunMixin, 
BaseOperator, HITLReviewMixin):
             rewrapped.append(capability)
         return rewrapped
 
+    def _log_durable_summary(self, counter: DurableStepCounter) -> None:
+        """
+        Log what this attempt replayed and cached, and which steps it could 
not cache.
+
+        A step whose cache write was skipped (a tool result that is not
+        JSON-serializable, a store write that fails) ran live but is not 
cached,
+        so a retry runs it again. For a tool with side effects the side effect
+        repeats, so those tools are named rather than counted as cached.
+        """
+        self.log.info(
+            "Durable: replayed %d cached steps (%d model, %d tool), cached %d 
new steps (%d model, %d tool)",
+            counter.replayed_model + counter.replayed_tool,
+            counter.replayed_model,
+            counter.replayed_tool,
+            counter.cached_model + counter.cached_tool,
+            counter.cached_model,
+            counter.cached_tool,
+        )
+        if counter.skipped_tools:
+            calls = collections.Counter(counter.skipped_tools)
+            self.log.warning(
+                "Durable: %d tool results were not cached, and a retry runs 
them again: %s",
+                len(counter.skipped_tools),
+                ", ".join(name if n == 1 else f"{name} (x{n})" for name, n in 
calls.items()),
+            )
+        if counter.skipped_model:
+            self.log.warning(
+                "Durable: %d model responses were not cached, and a retry 
re-runs them "
+                "and every step after the first of them",
+                counter.skipped_model,
+            )
+
     def _build_durable_storage(self, context: Context) -> 
DurableStorageProtocol:
         """
         Return the durable storage backend for the current task instance.
@@ -809,7 +842,11 @@ class AgentOperator(CancellableAgentRunMixin, 
BaseOperator, HITLReviewMixin):
             resolved_model = infer_model(agent.model)
             caching_model = CachingModel(resolved_model, storage=storage, 
counter=counter)
             with agent.override(model=caching_model):
-                result = self.run_agent_sync(agent, self.prompt, **run_kwargs)
+                try:
+                    result = self.run_agent_sync(agent, self.prompt, 
**run_kwargs)
+                finally:
+                    # Also on a raise: the failed attempt is the one Airflow 
retries.
+                    self._log_durable_summary(counter)
         else:
             result = self.run_agent_sync(agent, self.prompt, **run_kwargs)
 
@@ -822,22 +859,6 @@ class AgentOperator(CancellableAgentRunMixin, 
BaseOperator, HITLReviewMixin):
             self._pause_for_tool_approval(context, result)
         self._emit_run_metadata(context, result)
 
-        if self._durable_counter is not None:
-            c = self._durable_counter
-            replayed = c.replayed_model + c.replayed_tool
-            cached = c.cached_model + c.cached_tool
-            if replayed:
-                self.log.info(
-                    "Durable: replayed %d cached steps (%d model, %d tool), "
-                    "executed %d new steps (%d model, %d tool)",
-                    replayed,
-                    c.replayed_model,
-                    c.replayed_tool,
-                    cached,
-                    c.cached_model,
-                    c.cached_tool,
-                )
-
         if self.message_history is not None:
             self._emit_message_history(context, result)
 
diff --git 
a/providers/common/ai/tests/unit/common/ai/durable/test_caching_model.py 
b/providers/common/ai/tests/unit/common/ai/durable/test_caching_model.py
index 9c5eb288f40..4c6b63064e3 100644
--- a/providers/common/ai/tests/unit/common/ai/durable/test_caching_model.py
+++ b/providers/common/ai/tests/unit/common/ai/durable/test_caching_model.py
@@ -22,16 +22,19 @@ import pytest
 from pydantic_ai.messages import ModelResponse, TextPart
 from pydantic_ai.models import ModelRequestParameters
 
-from airflow.providers.common.ai.durable.base import DURABLE_KEY_PREFIX as P
+from airflow.providers.common.ai.durable.base import DURABLE_KEY_PREFIX as P, 
DurableStorageProtocol
 from airflow.providers.common.ai.durable.caching_model import CachingModel
 from airflow.providers.common.ai.durable.fingerprint import 
fingerprint_model_request
 from airflow.providers.common.ai.durable.step_counter import DurableStepCounter
 
+from tests_common.test_utils.version_compat import AIRFLOW_V_3_1_PLUS
+
 
 @pytest.fixture
 def mock_storage():
-    storage = MagicMock()
+    storage = MagicMock(spec=DurableStorageProtocol)
     storage.load_model_response.return_value = (None, None)
+    storage.save_model_response.return_value = True
     return storage
 
 
@@ -122,6 +125,40 @@ class TestCachingModelCacheMiss:
         keys = [call[0][0] for call in 
mock_storage.save_model_response.call_args_list]
         assert keys == [f"{P}model_step_0", f"{P}model_step_1"]
 
+    @pytest.mark.asyncio
+    @pytest.mark.parametrize(
+        ("written", "expected"),
+        [pytest.param(True, (1, 0), id="written"), pytest.param(False, (0, 1), 
id="refused")],
+    )
+    async def test_counts_cached_only_when_storage_wrote_it(
+        self, mock_model, mock_storage, counter, sample_response, written, 
expected
+    ):
+        """A response the backend did not store re-runs on retry, so it counts 
as skipped."""
+        mock_model.request = AsyncMock(return_value=sample_response)
+        mock_storage.save_model_response.return_value = written
+        caching = CachingModel(mock_model, storage=mock_storage, 
counter=counter)
+
+        result = await caching.request([], None, ModelRequestParameters())
+
+        assert result is sample_response
+        assert (counter.cached_model, counter.skipped_model) == expected
+
+    @pytest.mark.asyncio
+    @pytest.mark.skipif(
+        not AIRFLOW_V_3_1_PLUS, reason="cap_structlog needs airflow._shared, 
which lands in Airflow 3.1"
+    )
+    @pytest.mark.parametrize("written", [True, False])
+    async def test_warns_only_when_storage_skipped_the_write(
+        self, mock_model, mock_storage, counter, sample_response, written, 
cap_structlog
+    ):
+        mock_model.request = AsyncMock(return_value=sample_response)
+        mock_storage.save_model_response.return_value = written
+        caching = CachingModel(mock_model, storage=mock_storage, 
counter=counter)
+
+        await caching.request([], None, ModelRequestParameters())
+
+        assert ({"step": 0, "log_level": "warning"} in cap_structlog) is not 
written
+
 
 class TestCachingModelReplayVerification:
     @pytest.mark.asyncio
diff --git 
a/providers/common/ai/tests/unit/common/ai/durable/test_caching_toolset.py 
b/providers/common/ai/tests/unit/common/ai/durable/test_caching_toolset.py
index c0570f388a1..d0266e8996c 100644
--- a/providers/common/ai/tests/unit/common/ai/durable/test_caching_toolset.py
+++ b/providers/common/ai/tests/unit/common/ai/durable/test_caching_toolset.py
@@ -23,18 +23,22 @@ import pytest
 from pydantic_ai.messages import ModelResponse, TextPart
 from pydantic_ai.models import ModelRequestParameters
 
-from airflow.providers.common.ai.durable.base import DURABLE_KEY_PREFIX as P
+from airflow.providers.common.ai.durable.base import DURABLE_KEY_PREFIX as P, 
DurableStorageProtocol
 from airflow.providers.common.ai.durable.caching_model import CachingModel
 from airflow.providers.common.ai.durable.caching_toolset import CachingToolset
 from airflow.providers.common.ai.durable.fingerprint import 
fingerprint_tool_call
 from airflow.providers.common.ai.durable.step_counter import DurableStepCounter
 
+from tests_common.test_utils.version_compat import AIRFLOW_V_3_1_PLUS
+
 
 @pytest.fixture
 def mock_storage():
-    storage = MagicMock()
+    storage = MagicMock(spec=DurableStorageProtocol)
     storage.load_tool_result.return_value = (False, None, None)
     storage.load_model_response.return_value = (None, None)
+    storage.save_tool_result.return_value = True
+    storage.save_model_response.return_value = True
     return storage
 
 
@@ -107,6 +111,34 @@ class TestCachingToolsetCacheMiss:
         keys = [call[0][0] for call in 
mock_storage.save_tool_result.call_args_list]
         assert keys == [f"{P}tool_step_0", f"{P}tool_step_1"]
 
+    @pytest.mark.asyncio
+    async def test_skipped_write_is_recorded_by_tool_name(self, mock_toolset, 
mock_storage, counter):
+        """A result the backend did not store re-runs on retry, so it is not 
counted as cached."""
+        mock_storage.save_tool_result.side_effect = [True, False]
+        caching = CachingToolset(wrapped=mock_toolset, storage=mock_storage, 
counter=counter)
+
+        await caching.call_tool("get_schema", {}, ctx_for("c1"), MagicMock())
+        result = await caching.call_tool("run_query", {}, ctx_for("c2"), 
MagicMock())
+
+        assert result == "fresh result"
+        assert counter.cached_tool == 1
+        assert counter.skipped_tools == ["run_query"]
+
+    @pytest.mark.asyncio
+    @pytest.mark.skipif(
+        not AIRFLOW_V_3_1_PLUS, reason="cap_structlog needs airflow._shared, 
which lands in Airflow 3.1"
+    )
+    async def test_skipped_write_warns_by_tool_name(self, mock_toolset, 
mock_storage, counter, cap_structlog):
+        """The warning names the tool on every path, not only in a successful 
run's summary."""
+        mock_storage.save_tool_result.side_effect = [True, False]
+        caching = CachingToolset(wrapped=mock_toolset, storage=mock_storage, 
counter=counter)
+
+        await caching.call_tool("get_schema", {}, ctx_for("c1"), MagicMock())
+        await caching.call_tool("run_query", {}, ctx_for("c2"), MagicMock())
+
+        assert {"tool": "run_query", "step": 1, "log_level": "warning"} in 
cap_structlog
+        assert {"tool": "get_schema", "log_level": "warning"} not in 
cap_structlog
+
 
 class TestCachingToolsetReplayVerification:
     @pytest.mark.asyncio
diff --git a/providers/common/ai/tests/unit/common/ai/durable/test_storage.py 
b/providers/common/ai/tests/unit/common/ai/durable/test_storage.py
index d4faf580244..c1a4176d164 100644
--- a/providers/common/ai/tests/unit/common/ai/durable/test_storage.py
+++ b/providers/common/ai/tests/unit/common/ai/durable/test_storage.py
@@ -192,6 +192,15 @@ class TestSaveLoadToolResult:
         assert found is False
 
 
+class TestSaveReturnsWhetherWritten:
+    def test_written_entries_return_true(self, storage, sample_response):
+        assert storage.save_model_response("model_step_0", sample_response, 
fingerprint="fp") is True
+        assert storage.save_tool_result("tool_step_1", {"rows": [1]}, 
fingerprint="fp") is True
+
+    def test_non_serializable_tool_result_returns_false(self, storage):
+        assert storage.save_tool_result("tool_step_0", object(), 
fingerprint="fp") is False
+
+
 class TestMalformedEntries:
     def test_empty_data_list_degrades_to_miss(self, storage):
         """A torn entry whose data list is empty loads as a miss, not an 
IndexError."""
diff --git 
a/providers/common/ai/tests/unit/common/ai/durable/test_task_state_store.py 
b/providers/common/ai/tests/unit/common/ai/durable/test_task_state_store.py
index 4b0a810f4c4..9e56354fc59 100644
--- a/providers/common/ai/tests/unit/common/ai/durable/test_task_state_store.py
+++ b/providers/common/ai/tests/unit/common/ai/durable/test_task_state_store.py
@@ -176,7 +176,7 @@ class TestSaveLoadToolResult:
 
     def test_non_serializable_result_is_skipped_not_raised(self, storage, 
accessor):
         """A non-serializable tool result skips caching with a warning; the 
tool step still succeeds."""
-        storage.save_tool_result("tool_step_0", object(), fingerprint="fp")  # 
must not raise
+        assert storage.save_tool_result("tool_step_0", object(), 
fingerprint="fp") is False  # must not raise
 
         assert "tool_step_0" not in accessor.store
         assert storage.load_tool_result("tool_step_0") == (False, None, None)
@@ -210,6 +210,26 @@ class TestSaveLoadToolResult:
         assert value == {"1": "a", "2": "b"}
 
 
+class TestSaveReturnsWhetherWritten:
+    def test_written_entries_return_true(self, storage, sample_response):
+        assert storage.save_model_response("model_step_0", sample_response, 
fingerprint="fp") is True
+        assert storage.save_tool_result("tool_step_1", {"rows": [1]}, 
fingerprint="fp") is True
+
+    @pytest.mark.parametrize("method", ["save_model_response", 
"save_tool_result"])
+    def test_write_rejected_by_the_store_returns_false(self, storage, 
accessor, sample_response, method):
+        """A store write that fails skips the entry and says so."""
+
+        def reject(key, value, *, retention=None):
+            raise ValueError("value exceeds the maximum size")
+
+        accessor.set = reject
+        value = sample_response if method == "save_model_response" else 
"result"
+
+        assert getattr(storage, method)("step_0", value, fingerprint="fp") is 
False
+        # Not tracked for cleanup either: a skipped key was never written.
+        assert "step_0" not in storage._keys
+
+
 class TestCleanup:
     def test_cleanup_deletes_keys_written_this_run(self, storage, accessor, 
sample_response):
         storage.save_model_response("model_step_0", sample_response, 
fingerprint="fp")
diff --git a/providers/common/ai/tests/unit/common/ai/operators/test_agent.py 
b/providers/common/ai/tests/unit/common/ai/operators/test_agent.py
index 1f26bb9751b..f08578cc429 100644
--- a/providers/common/ai/tests/unit/common/ai/operators/test_agent.py
+++ b/providers/common/ai/tests/unit/common/ai/operators/test_agent.py
@@ -118,20 +118,29 @@ def _build_priced_response(messages: list[ModelMessage], 
info: AgentInfo) -> Mod
 
 
 class _InMemoryDurableStorage:
-    """In-memory DurableStorageProtocol backend for exercising real replay in 
tests."""
+    """In-memory DurableStorageProtocol backend for exercising real replay in 
tests.
 
-    def __init__(self):
+    ``refuse_tool_writes`` stands in for a backend that skips a tool result
+    (a store write that fails), so the step is not cached.
+    """
+
+    def __init__(self, *, refuse_tool_writes: bool = False):
         self.models: dict = {}
         self.tools: dict = {}
+        self.refuse_tool_writes = refuse_tool_writes
 
     def save_model_response(self, key, response, *, fingerprint):
         self.models[key] = (response, fingerprint)
+        return True
 
     def load_model_response(self, key):
         return self.models.get(key, (None, None))
 
     def save_tool_result(self, key, result, *, fingerprint):
+        if self.refuse_tool_writes:
+            return False
         self.tools[key] = (result, fingerprint)
+        return True
 
     def load_tool_result(self, key):
         if key in self.tools:
@@ -1290,6 +1299,111 @@ class TestAgentOperatorDurable:
 
         assert calls["n"] == 1
 
+    def 
test_tool_result_refused_by_storage_is_counted_skipped_and_reruns(self):
+        """A tool result the backend refuses to store is not counted as 
cached, and a
+        retry runs the tool again instead of replaying it."""
+        calls = {"n": 0}
+
+        def my_tool() -> str:
+            calls["n"] += 1
+            return "tool-result"
+
+        def model_fn(messages, info):
+            saw_return = any(isinstance(p, ToolReturnPart) for m in messages 
for p in getattr(m, "parts", []))
+            if saw_return:
+                return ModelResponse(parts=[TextPart(content="done")])
+            return ModelResponse(parts=[ToolCallPart(tool_name="my_tool", 
args={}, tool_call_id="c1")])
+
+        storage = _InMemoryDurableStorage(refuse_tool_writes=True)
+        counters = []
+        for _ in range(2):
+            op = AgentOperator(
+                task_id="t",
+                prompt="hi",
+                llm_conn_id="c",
+                durable=True,
+                enable_tool_logging=False,
+                toolsets=[FunctionToolset(tools=[my_tool])],
+            )
+            op._durable_storage = storage
+            op._durable_counter = DurableStepCounter()
+            hook = MagicMock(spec=["create_agent"])
+            hook.create_agent.side_effect = lambda **kw: 
Agent(FunctionModel(model_fn), **kw)
+            op.llm_hook = hook
+            op._build_agent().run_sync("hi")
+            counters.append(op._durable_counter)
+
+        assert calls["n"] == 2
+        first, retry = counters
+        assert (first.cached_tool, first.skipped_tools) == (0, ["my_tool"])
+        assert (retry.replayed_tool, retry.skipped_tools) == (0, ["my_tool"])
+
+    def test_durable_summary_names_tools_that_were_not_cached(self, caplog):
+        counter = DurableStepCounter()
+        counter.cached_model = 2
+        counter.cached_tool = 1
+        counter.skipped_tools = ["run_query", "get_schema", "run_query"]
+        counter.skipped_model = 1
+        op = AgentOperator(task_id="t", prompt="p", llm_conn_id="c", 
durable=True)
+
+        with caplog.at_level("INFO"):
+            op._log_durable_summary(counter)
+
+        assert (
+            "replayed 0 cached steps (0 model, 0 tool), cached 3 new steps (2 
model, 1 tool)" in caplog.text
+        )
+        assert (
+            "3 tool results were not cached, and a retry runs them again: 
run_query (x2), get_schema"
+            in caplog.text
+        )
+        assert "1 model responses were not cached, and a retry re-runs them" 
in caplog.text
+
+    def test_durable_summary_has_no_warning_when_everything_was_cached(self, 
caplog):
+        counter = DurableStepCounter()
+        counter.cached_model = 1
+        op = AgentOperator(task_id="t", prompt="p", llm_conn_id="c", 
durable=True)
+
+        with caplog.at_level("INFO"):
+            op._log_durable_summary(counter)
+
+        assert "cached 1 new steps (1 model, 0 tool)" in caplog.text
+        assert not [r for r in caplog.records if r.levelname == "WARNING"]
+
+    
@patch("airflow.providers.common.ai.operators.agent.AgentOperator._build_durable_storage")
+    def test_failed_run_logs_summary_naming_uncached_tools(self, 
mock_build_storage, caplog):
+        """The attempt that fails is the one Airflow retries, so its summary 
must name
+        the tools that were not cached and will run again."""
+        mock_build_storage.return_value = 
_InMemoryDurableStorage(refuse_tool_writes=True)
+
+        def send_email() -> str:
+            return "sent"
+
+        def explode() -> str:
+            raise RuntimeError("downstream failure")
+
+        def model_fn(messages, info):
+            returned = [p.tool_name for m in messages for p in m.parts if 
isinstance(p, ToolReturnPart)]
+            name = "explode" if "send_email" in returned else "send_email"
+            return ModelResponse(parts=[ToolCallPart(tool_name=name, args={}, 
tool_call_id=name)])
+
+        op = AgentOperator(
+            task_id="t",
+            prompt="hi",
+            llm_conn_id="c",
+            durable=True,
+            enable_tool_logging=False,
+            toolsets=[FunctionToolset(tools=[send_email, explode])],
+        )
+        hook = MagicMock(spec=["create_agent"])
+        hook.create_agent.side_effect = lambda **kw: 
Agent(FunctionModel(model_fn), **kw)
+        op.llm_hook = hook
+
+        with caplog.at_level("INFO"), pytest.raises(RuntimeError, 
match="downstream failure"):
+            op.execute(context=_make_context())
+
+        assert "cached 2 new steps (2 model, 0 tool)" in caplog.text
+        assert "1 tool results were not cached, and a retry runs them again: 
send_email" in caplog.text
+
     @patch("pydantic_ai.models.wrapper.infer_model", side_effect=lambda m: m)
     @patch("pydantic_ai.models.infer_model", autospec=True)
     
@patch("airflow.providers.common.ai.operators.agent.AgentOperator._build_durable_storage")

Reply via email to