This is an automated email from the ASF dual-hosted git repository.
eladkal 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 771c09c5f20 Fix common.ai durable execution skipping
Toolset-capability tools (#69881)
771c09c5f20 is described below
commit 771c09c5f204a3d790d7a5b25be4cff73f316f42
Author: Kaxil Naik <[email protected]>
AuthorDate: Fri Jul 17 12:52:50 2026 +0100
Fix common.ai durable execution skipping Toolset-capability tools (#69881)
When durable=True, tools supplied via a pydantic-ai Toolset capability
(capabilities=[Toolset(ts)]) bypassed the CachingToolset applied to
toolsets=, so their results re-executed on every retry instead of
replaying. Wrap the inner toolset of concrete Toolset capabilities with
the same CachingToolset. A Toolset capability backed by a callable
factory can't be wrapped (it resolves per run) and now logs a warning.
Also fixes two latent durable-execution bugs found in review:
- cleanup() ran before the message-history XCom push and output
serialization. A failure in either wiped the cache, so the retry
re-ran every already-completed step. Move cleanup to after them.
- the pre-3.3 ObjectStorage cache filename joined dag/task/run with "_",
aliasing distinct tasks (dag "etl"/task "load_data" and dag
"etl_load"/task "data" both mapped to one file, so one task could
read, overwrite, or delete another's cache). Hash the identity
components instead.
---
.../airflow/providers/common/ai/durable/storage.py | 16 ++-
.../airflow/providers/common/ai/operators/agent.py | 85 +++++++++++--
.../tests/unit/common/ai/durable/test_storage.py | 29 +++--
.../tests/unit/common/ai/operators/test_agent.py | 132 ++++++++++++++++++++-
4 files changed, 239 insertions(+), 23 deletions(-)
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 d50107631a4..0ff1c84afbe 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
@@ -19,6 +19,7 @@
from __future__ import annotations
import contextlib
+import hashlib
import json
from functools import lru_cache
from typing import Any
@@ -54,8 +55,9 @@ class DurableStorage:
Stores step-level caches in a single JSON file on ObjectStorage.
All step caches (model responses and tool results) are stored as entries
- in a single JSON blob, written to a file named after the task execution:
- ``{base_path}/{dag_id}_{task_id}_{run_id}[_{map_index}].json``.
+ in a single JSON blob, written to ``{base_path}/{cache_id}.json`` where
+ ``cache_id`` is a hash of the task instance's identity (dag, task, run,
+ map index) so distinct task instances never share a file.
The file survives Airflow task retries since it lives outside the
XCom system. It is deleted on successful task completion.
@@ -74,8 +76,14 @@ class DurableStorage:
run_id: str,
map_index: int = -1,
) -> None:
- suffix = f"_{map_index}" if map_index >= 0 else ""
- self._cache_id = f"{dag_id}_{task_id}_{run_id}{suffix}"
+ # Hash the identity components with a separator that cannot appear in
+ # them, so distinct task instances can never alias to the same cache
+ # file. A plain ``_``-joined string collides -- e.g. dag ``etl`` + task
+ # ``load_data`` and dag ``etl_load`` + task ``data`` both yield
+ # ``etl_load_data`` -- letting one task read, overwrite, or delete
+ # another's durable cache.
+ identity = "\x00".join([dag_id, task_id, run_id, str(map_index)])
+ self._cache_id = hashlib.sha256(identity.encode()).hexdigest()
self._cache: dict[str, Any] | None = None
def _get_path(self):
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 5d61588f4d0..f83a58f0e2a 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
@@ -20,6 +20,7 @@ from __future__ import annotations
import json
from collections.abc import Sequence
+from dataclasses import replace
from datetime import timedelta
from functools import cached_property
from typing import TYPE_CHECKING, Any, ClassVar
@@ -157,6 +158,14 @@ class AgentOperator(BaseOperator, HITLReviewMixin):
On Airflow >= 3.3 the cache is kept in the AIP-103 task state store, so
no extra configuration is needed. On older cores it is persisted to
ObjectStorage and requires ``[common.ai] durable_cache_path`` to be
set.
+ Tools are durably cached when provided via ``toolsets=`` or via a
+ concrete pydantic-ai ``Toolset`` capability. Tools reaching the agent
+ through any *other* capability -- ``MCP``, ``PrefixTools``,
+ ``CombinedCapability``, a ``Toolset`` backed by a callable factory, or
+ capabilities loaded from a ``spec_file`` -- are not cached and re-run
on
+ retry; put tools you need replayed in ``toolsets=``. Provider-native
+ capabilities such as ``WebSearch`` and ``Thinking`` execute inside the
+ model call and are covered by model-response caching.
:param code_mode: When ``True``, wraps the agent's tools in a single
``run_code`` tool powered by the Monty sandbox (pydantic-ai-harness
``CodeMode``). Instead of one model round-trip per tool call, the model
@@ -266,6 +275,12 @@ class AgentOperator(BaseOperator, HITLReviewMixin):
self.durable = durable
self.code_mode = code_mode
+ # Populated per run in ``execute`` when durable=True. Declared here so
+ # ``_build_agent`` -- also reached via ``regenerate_with_feedback``
+ # outside ``execute`` -- can read them unconditionally.
+ self._durable_storage: DurableStorageProtocol | None = None
+ self._durable_counter: DurableStepCounter | None = None
+
if durable and enable_hitl_review:
raise ValueError("durable=True and enable_hitl_review=True cannot
be used together.")
@@ -306,18 +321,24 @@ class AgentOperator(BaseOperator, HITLReviewMixin):
def _build_agent(self) -> Agent[object, Any]:
"""Build and return a pydantic-ai Agent from the operator's config."""
extra_kwargs = dict(self.agent_params)
+ storage = self._durable_storage
+ counter = self._durable_counter
if self.toolsets:
toolsets = self.toolsets
- if self.durable and self._durable_storage is not None and
self._durable_counter is not None:
- toolsets = self._build_durable_toolsets(
- toolsets, self._durable_storage, self._durable_counter
- )
+ if self.durable and storage is not None and counter is not None:
+ toolsets = self._build_durable_toolsets(toolsets, storage,
counter)
if self.enable_tool_logging:
toolsets = wrap_toolsets_for_logging(toolsets, self.log)
extra_kwargs["toolsets"] = toolsets
+ capabilities = list(extra_kwargs.get("capabilities") or [])
+ if self.durable and storage is not None and counter is not None:
+ # Tools supplied through a ``Toolset`` capability bypass the
+ # ``toolsets=`` wrapping above, so their results would re-execute
on
+ # every retry instead of replaying; wrap their inner toolset too.
+ capabilities = self._build_durable_capabilities(capabilities,
storage, counter)
if self.code_mode:
- capabilities = list(extra_kwargs.get("capabilities") or [])
capabilities.append(_build_code_mode())
+ if capabilities:
extra_kwargs["capabilities"] = capabilities
return self.llm_hook.create_agent(
output_type=self.output_type,
@@ -333,6 +354,49 @@ class AgentOperator(BaseOperator, HITLReviewMixin):
return [CachingToolset(wrapped=ts, storage=storage, counter=counter)
for ts in toolsets]
+ def _build_durable_capabilities(
+ self, capabilities: list[Any], storage: DurableStorageProtocol,
counter: DurableStepCounter
+ ) -> list[Any]:
+ """
+ Wrap toolsets provided via a pydantic-ai ``Toolset`` capability for
durable replay.
+
+ Tools reaching the agent through ``capabilities=[Toolset(ts)]`` bypass
the
+ operator's ``toolsets=`` list, so the ``CachingToolset`` applied in
+ :meth:`_build_durable_toolsets` never sees them and their results
+ re-execute on every retry instead of replaying. Wrap each ``Toolset``
+ capability's inner toolset with the same ``CachingToolset``, preserving
+ the capability's other fields. Non-``Toolset`` capabilities pass
through
+ unchanged, as does a ``Toolset`` holding a callable factory rather
than a
+ concrete toolset (only a concrete toolset can be wrapped here).
+ """
+ # pydantic-ai (and the pydantic-ai-importing CachingToolset) are
imported
+ # lazily to keep them out of DAG-parse-time imports, matching
+ # ``_build_durable_toolsets`` and the rest of this module.
+ from pydantic_ai.capabilities import Toolset
+ from pydantic_ai.toolsets.abstract import AbstractToolset
+
+ from airflow.providers.common.ai.durable.caching_toolset import
CachingToolset
+
+ rewrapped: list[Any] = []
+ for capability in capabilities:
+ # ``Toolset.toolset`` can be a concrete toolset or a callable
factory
+ # resolved per run; only a concrete toolset can be wrapped here.
+ if isinstance(capability, Toolset) and
isinstance(capability.toolset, AbstractToolset):
+ cached = CachingToolset(wrapped=capability.toolset,
storage=storage, counter=counter)
+ rewrapped.append(replace(capability, toolset=cached))
+ continue
+ if isinstance(capability, Toolset):
+ # The toolset is a callable factory resolved per run, so there
is
+ # no concrete toolset to wrap; its results won't be cached for
+ # replay. Warn so durable users aren't silently surprised on
retry.
+ self.log.warning(
+ "durable=True: tools from a Toolset capability backed by a
callable "
+ "factory are not cached for replay; pass the toolset via
`toolsets=` "
+ "for durability."
+ )
+ rewrapped.append(capability)
+ return rewrapped
+
def _build_durable_storage(self, context: Context) ->
DurableStorageProtocol:
"""
Return the durable storage backend for the current task instance.
@@ -418,9 +482,6 @@ class AgentOperator(BaseOperator, HITLReviewMixin):
c.cached_tool,
)
- if self._durable_storage is not None:
- self._durable_storage.cleanup()
-
if self.message_history is not None:
self._emit_message_history(context, result)
@@ -445,6 +506,14 @@ class AgentOperator(BaseOperator, HITLReviewMixin):
if self._serialize_model_output and isinstance(output, BaseModel):
output = output.model_dump()
+
+ # Clean up the durable cache only after the run and every post-run step
+ # that can still fail (the message-history XCom push above and output
+ # serialization) has succeeded. Cleaning up earlier and then raising
+ # would leave the Airflow retry with an empty cache, re-executing every
+ # already-completed model and tool step.
+ if self._durable_storage is not None:
+ self._durable_storage.cleanup()
return output
def _resolve_message_history(self) -> list[ModelMessage] | None:
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 03f85be30e1..d4faf580244 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
@@ -50,16 +50,25 @@ def sample_response():
class TestDurableStorageInit:
- def test_cache_id_format(self, storage):
- assert storage._cache_id == "test_dag_my_task_run_1"
-
- def test_cache_id_with_map_index(self):
- s = DurableStorage(dag_id="d", task_id="t", run_id="r", map_index=3)
- assert s._cache_id == "d_t_r_3"
-
- def test_cache_id_without_map_index(self):
- s = DurableStorage(dag_id="d", task_id="t", run_id="r", map_index=-1)
- assert "_-1" not in s._cache_id
+ def test_cache_id_is_deterministic(self):
+ """The same task identity always maps to the same cache file (so
retries resume)."""
+ a = DurableStorage(dag_id="d", task_id="t", run_id="r", map_index=-1)
+ b = DurableStorage(dag_id="d", task_id="t", run_id="r", map_index=-1)
+ assert a._cache_id == b._cache_id
+
+ def test_cache_id_differs_by_map_index(self):
+ base = DurableStorage(dag_id="d", task_id="t", run_id="r",
map_index=-1)
+ mapped = DurableStorage(dag_id="d", task_id="t", run_id="r",
map_index=3)
+ assert base._cache_id != mapped._cache_id
+
+ def test_cache_id_no_collision_across_tasks(self):
+ """Distinct (dag, task) pairs that concatenate to the same string must
not
+ share a cache file -- e.g. dag ``etl`` + task ``load_data`` vs dag
+ ``etl_load`` + task ``data``. A plain ``_``-join aliased them, letting
one
+ task read, overwrite, or delete another task's durable cache."""
+ a = DurableStorage(dag_id="etl", task_id="load_data", run_id="r")
+ b = DurableStorage(dag_id="etl_load", task_id="data", run_id="r")
+ assert a._cache_id != b._cache_id
class TestSaveLoadModelResponse:
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 67c28ebfafe..2631c544138 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
@@ -22,15 +22,24 @@ from unittest.mock import MagicMock, patch
import pytest
from pydantic import BaseModel
+from pydantic_ai import Agent
+from pydantic_ai.capabilities import Toolset
from pydantic_ai.messages import (
ModelMessagesTypeAdapter,
ModelRequest,
ModelResponse,
TextPart,
+ ToolCallPart,
+ ToolReturnPart,
UserPromptPart,
)
+from pydantic_ai.models.function import FunctionModel
+from pydantic_ai.toolsets.function import FunctionToolset
from pydantic_ai.usage import UsageLimits
+from airflow.providers.common.ai.durable.base import DurableStorageProtocol
+from airflow.providers.common.ai.durable.caching_toolset import CachingToolset
+from airflow.providers.common.ai.durable.step_counter import DurableStepCounter
from airflow.providers.common.ai.durable.storage import DurableStorage
from airflow.providers.common.ai.operators.agent import AgentOperator,
HITLReviewLink, _build_code_mode
from airflow.providers.common.ai.toolsets.logging import LoggingToolset
@@ -71,6 +80,33 @@ def _make_mock_agent(output):
return mock_agent
+class _InMemoryDurableStorage:
+ """In-memory DurableStorageProtocol backend for exercising real replay in
tests."""
+
+ def __init__(self):
+ self.models: dict = {}
+ self.tools: dict = {}
+
+ def save_model_response(self, key, response, *, fingerprint):
+ self.models[key] = (response, fingerprint)
+
+ def load_model_response(self, key):
+ return self.models.get(key, (None, None))
+
+ def save_tool_result(self, key, result, *, fingerprint):
+ self.tools[key] = (result, fingerprint)
+
+ def load_tool_result(self, key):
+ if key in self.tools:
+ value, fingerprint = self.tools[key]
+ return True, value, fingerprint
+ return False, None, None
+
+ def cleanup(self):
+ self.models.clear()
+ self.tools.clear()
+
+
class TestAgentOperatorValidation:
def test_requires_llm_conn_id(self):
with pytest.raises(TypeError):
@@ -584,7 +620,10 @@ class TestAgentOperatorDurable:
storage = op._build_durable_storage({"task_instance": ti})
assert isinstance(storage, DurableStorage)
- assert storage._cache_id == "d_t_r"
+ # cache_id is a stable hash of the identity components, not a raw
concat.
+ assert (
+ storage._cache_id == DurableStorage(dag_id="d", task_id="t",
run_id="r", map_index=-1)._cache_id
+ )
@patch("pydantic_ai.models.wrapper.infer_model", side_effect=lambda m: m)
@patch("pydantic_ai.models.infer_model", autospec=True)
@@ -629,6 +668,97 @@ class TestAgentOperatorDurable:
# run_sync called directly, no override
mock_agent.run_sync.assert_called_once_with("test", usage_limits=None)
+ def test_build_durable_capabilities_wraps_toolset_capability(self):
+ """A ``Toolset`` capability's inner toolset is wrapped with
CachingToolset;
+ capabilities that are not ``Toolset`` pass through unchanged."""
+ inner = FunctionToolset()
+ passthrough = object()
+ op = AgentOperator(task_id="t", prompt="p", llm_conn_id="c",
durable=True)
+
+ result = op._build_durable_capabilities(
+ [Toolset(inner), passthrough],
MagicMock(spec=DurableStorageProtocol), DurableStepCounter()
+ )
+
+ assert isinstance(result[0], Toolset)
+ assert isinstance(result[0].toolset, CachingToolset)
+ assert result[0].toolset.wrapped is inner
+ assert result[1] is passthrough
+
+ def test_build_durable_capabilities_skips_callable_toolset_factory(self):
+ """A ``Toolset`` holding a callable factory (resolved per run with
+ RunContext) cannot be wrapped with CachingToolset, so it passes
through."""
+
+ def factory(ctx):
+ return FunctionToolset()
+
+ cap = Toolset(factory)
+ op = AgentOperator(task_id="t", prompt="p", llm_conn_id="c",
durable=True)
+
+ result = op._build_durable_capabilities(
+ [cap], MagicMock(spec=DurableStorageProtocol), DurableStepCounter()
+ )
+
+ assert result[0] is cap
+
+ def test_toolset_capability_tool_replayed_on_retry(self):
+ """A tool supplied via a ``Toolset`` capability is cached and replayed
on a
+ retry instead of re-executing. Regression: such tools bypassed the
+ ``CachingToolset`` because they did not arrive via the ``toolsets=``
list."""
+ 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")])
+
+ # Shared storage across two attempts; the second (a retry) must replay
the
+ # cached tool result rather than executing the tool a second time.
+ storage = _InMemoryDurableStorage()
+ for _ in range(2):
+ op = AgentOperator(
+ task_id="t",
+ prompt="hi",
+ llm_conn_id="c",
+ durable=True,
+ enable_tool_logging=False,
+ agent_params={"capabilities":
[Toolset(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")
+
+ assert calls["n"] == 1
+
+ @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")
+ @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook",
autospec=True)
+ def test_cleanup_skipped_when_post_run_step_fails(self, mock_hook_cls,
mock_build_storage, mock_infer, _):
+ """Durable cleanup must not run if a post-run step (the
message-history XCom
+ push) fails, so the Airflow retry can still replay the cached steps."""
+ storage = MagicMock(spec=DurableStorageProtocol)
+ mock_build_storage.return_value = storage
+
+ mock_agent = MagicMock(spec=["run_sync", "model", "override"])
+ mock_agent.run_sync.return_value = _make_mock_run_result("ok")
+ mock_agent.model = "test-model"
+ mock_hook_cls.get_hook.return_value.create_agent.return_value =
mock_agent
+
+ op = AgentOperator(task_id="t", prompt="p", llm_conn_id="c",
durable=True, message_history="[]")
+ with patch.object(op, "_emit_message_history",
side_effect=RuntimeError("xcom down")):
+ with pytest.raises(RuntimeError, match="xcom down"):
+ op.execute(context={})
+
+ storage.cleanup.assert_not_called()
+
@pytest.mark.skipif(
not AIRFLOW_V_3_1_PLUS, reason="Human in the loop is only compatible with
Airflow >= 3.1.0"