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")