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 dcc9ce6246f Fix Common AI durable retries not replaying tools from a
capability without an id (#74314)
dcc9ce6246f is described below
commit dcc9ce6246f49a028eca964b5ec791b9ca99f892
Author: Kaxil Naik <[email protected]>
AuthorDate: Tue Oct 6 08:09:18 2026 +0100
Fix Common AI durable retries not replaying tools from a capability without
an id (#74314)
* Fix Common AI durable retries not replaying tools from a capability
without an id
pydantic-ai 2.40+ gives a capability without an explicit id a random one
per run and stamps it on each of its tools. The durable fingerprint hashed
the whole request parameters, so a retry never matched its cached model
response and every step re-ran live, including tools with side effects.
The run-local capability ids are now left out of the fingerprint.
* Sort revealed tool names in the durable fingerprint and tighten its tests
revealed_tool_names is a set and dumped in iteration order, which differs
between processes, so a durable retry with two or more revealed tools also
missed the cache.
---
.../providers/common/ai/durable/fingerprint.py | 36 ++++++---
.../unit/common/ai/durable/test_fingerprint.py | 87 ++++++++++++++++++++++
.../tests/unit/common/ai/operators/test_agent.py | 57 ++++++++++++++
3 files changed, 171 insertions(+), 9 deletions(-)
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/durable/fingerprint.py
b/providers/common/ai/src/airflow/providers/common/ai/durable/fingerprint.py
index be3c2b6f1c2..5ff860f04bf 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/durable/fingerprint.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/durable/fingerprint.py
@@ -32,11 +32,11 @@ fingerprint, so stale tool results recorded under the old
conversation no
longer match.
Fields that pydantic-ai regenerates on every attempt (message-level
-``timestamp``/``run_id``/``conversation_id`` and part-level ``timestamp``)
-are excluded from the fingerprint. Requests that cannot be serialized to
-JSON fingerprint as ``None``, which degrades that step to unverified
-positional replay (the pre-fingerprint behavior) rather than disabling
-caching.
+``timestamp``/``run_id``/``conversation_id``, part-level ``timestamp``) and
+capability ids are excluded from the fingerprint, and set-valued request
+parameters are sorted. Requests that cannot be serialized to JSON
+fingerprint as ``None``, which degrades that step to unverified positional
+replay (the pre-fingerprint behavior) rather than disabling caching.
"""
from __future__ import annotations
@@ -102,6 +102,24 @@ def _strip_volatile(messages_dump: list[dict[str, Any]])
-> list[dict[str, Any]]
return stripped
+def _normalize_params(params_dump: dict[str, Any]) -> dict[str, Any]:
+ """
+ Drop capability ids and sort set-valued fields from dumped request
parameters.
+
+ A capability without an explicit ``id`` gets a random one per run
(``<toolset:d0d75e>``),
+ stamped on its tools' ``capability_id``, so hashing it would make every
retry miss the
+ cache. What the model sees of capabilities (the deferred-capability
catalog in the
+ instructions, tool visibility, revealed tool names) is hashed through
other fields, so
+ ``deferred_capability_ids`` is dropped too. A set dumps in iteration
order, which differs
+ between processes (``PYTHONHASHSEED``), and a retry runs in a new process.
+ """
+ cleaned = {k: v for k, v in params_dump.items() if k !=
"deferred_capability_ids"}
+ cleaned["revealed_tool_names"] = sorted(cleaned["revealed_tool_names"])
+ for key in ("function_tools", "output_tools"):
+ cleaned[key] = [{k: v for k, v in tool.items() if k !=
"capability_id"} for tool in cleaned[key]]
+ return cleaned
+
+
def _digest(payload: Any) -> str:
# No ``default=`` fallback: a non-JSON-serializable value must raise so the
# callers degrade to an unverifiable (None) fingerprint instead of hashing
@@ -119,9 +137,9 @@ def fingerprint_model_request(
"""
Fingerprint a model request: model identity, message history, settings,
and request parameters.
- The full ``ModelRequestParameters`` object is hashed (tool definitions,
- output mode and schema, native tools, ...) so any change to what is sent
- to the model invalidates the cached response.
+ The ``ModelRequestParameters`` object is hashed (tool definitions, output
+ mode and schema, native tools, ...) so any change to what is sent to the
+ model invalidates the cached response; only capability ids are left out.
Returns ``None`` when the request cannot be serialized; ``None`` compares
equal to ``None``, so requests that cannot be fingerprinted degrade to
@@ -135,7 +153,7 @@ def fingerprint_model_request(
"model": model_identifier,
"messages": _strip_volatile(dumped),
"settings": _content_settings(model_settings),
- "params": params,
+ "params": _normalize_params(params),
}
)
except (TypeError, ValueError):
diff --git
a/providers/common/ai/tests/unit/common/ai/durable/test_fingerprint.py
b/providers/common/ai/tests/unit/common/ai/durable/test_fingerprint.py
index d555336d955..aa14d45274a 100644
--- a/providers/common/ai/tests/unit/common/ai/durable/test_fingerprint.py
+++ b/providers/common/ai/tests/unit/common/ai/durable/test_fingerprint.py
@@ -19,10 +19,12 @@ from __future__ import annotations
import datetime
import httpx
+import pytest
from pydantic_ai.messages import (
ModelRequest,
ModelResponse,
SystemPromptPart,
+ TextPart,
ToolCallPart,
UserPromptPart,
)
@@ -30,6 +32,7 @@ from pydantic_ai.models import ModelRequestParameters
from pydantic_ai.tools import ToolDefinition
from airflow.providers.common.ai.durable.fingerprint import (
+ _normalize_params,
fingerprint_model_request,
fingerprint_tool_call,
)
@@ -118,6 +121,90 @@ class TestModelRequestFingerprint:
assert fp1 != fp2
+ def test_stable_across_capability_ids(self):
+ """A capability without an ``id`` gets a random one per run; it never
reaches the model."""
+
+ def params(capability_id):
+ tool = ToolDefinition(
+ name="t", parameters_json_schema={"type": "object"},
capability_id=capability_id
+ )
+ return ModelRequestParameters(function_tools=[tool],
output_tools=[tool])
+
+ fp1 = fingerprint_model_request("m", make_messages(), None,
params("<toolset:d0d75e>"))
+ fp2 = fingerprint_model_request("m", make_messages(), None,
params("<toolset:78ba70>"))
+
+ assert fp1 is not None
+ assert fp1 == fp2
+
+ @pytest.mark.parametrize("tools_param", ["function_tools", "output_tools"])
+ @pytest.mark.parametrize(
+ "change",
+ [
+ pytest.param({"name": "other"}, id="name"),
+ pytest.param({"description": "other"}, id="description"),
+ pytest.param(
+ {"parameters_json_schema": {"type": "object", "properties":
{"q": {"type": "string"}}}},
+ id="schema",
+ ),
+ ],
+ )
+ def test_changes_with_tool_content_next_to_the_capability_id(self,
tools_param, change):
+ """Only ``capability_id`` is dropped from a tool definition; what the
model sees still counts."""
+ base = {"name": "t", "parameters_json_schema": {"type": "object"},
"capability_id": "lookup"}
+ fp1 = fingerprint_model_request(
+ "m", make_messages(), None, ModelRequestParameters(**{tools_param:
[ToolDefinition(**base)]})
+ )
+ fp2 = fingerprint_model_request(
+ "m",
+ make_messages(),
+ None,
+ ModelRequestParameters(**{tools_param: [ToolDefinition(**{**base,
**change})]}),
+ )
+
+ assert fp1 != fp2
+
+ def test_changes_with_revealed_tool_names(self):
+ fp1 = fingerprint_model_request(
+ "m", make_messages(), None,
ModelRequestParameters(revealed_tool_names={"search"})
+ )
+ fp2 = fingerprint_model_request(
+ "m", make_messages(), None,
ModelRequestParameters(revealed_tool_names={"search", "fetch"})
+ )
+
+ assert fp1 != fp2
+
+ def test_revealed_tool_names_hash_in_any_order(self):
+ """A set dumps in iteration order, which differs between processes,
and a retry is a new one."""
+ names = ["search", "fetch", "summarize", "rank"]
+ dumps = [
+ {"revealed_tool_names": order, "function_tools": [],
"output_tools": []}
+ for order in (names, list(reversed(names)))
+ ]
+
+ assert _normalize_params(dumps[0]) == _normalize_params(dumps[1])
+
+ def test_changes_with_message_metadata(self):
+ """
+ Message ``metadata`` is not sent to the model, but pydantic-ai keeps
routing state in it
+ (``FallbackModel``'s continuation pin under ``__pydantic_ai__``), so
it stays in the hash.
+ """
+
+ def messages(pinned_model):
+ return [
+ ModelRequest(parts=[UserPromptPart(content="q")]),
+ ModelResponse(
+ parts=[TextPart(content="partial")],
+ metadata={"__pydantic_ai__": {"fallback_model_id":
pinned_model}},
+ ),
+ ]
+
+ fp1 = fingerprint_model_request("m", messages("openai:gpt-5"), None,
ModelRequestParameters())
+ fp2 = fingerprint_model_request(
+ "m", messages("anthropic:claude-sonnet-4-5"), None,
ModelRequestParameters()
+ )
+
+ assert fp1 != fp2
+
def test_volatile_keys_inside_user_data_are_not_stripped(self):
"""Only pydantic-ai's own message/part fields are volatile; a tool
argument
legitimately named run_id must still affect the fingerprint."""
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 318f5efe6fd..3823eacee0e 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
@@ -1639,6 +1639,63 @@ class TestAgentOperatorDurable:
assert calls["n"] == 1
+ @pytest.mark.parametrize(
+ "capability",
+ [
+ pytest.param(lambda tool: Toolset(FunctionToolset([tool])),
id="anonymous"),
+ pytest.param(lambda tool: Toolset(FunctionToolset([tool]),
id="lookup"), id="with-id"),
+ ],
+ )
+ def test_retry_replays_steps_of_a_toolset_capability(self, capability):
+ """
+ pydantic-ai gives a capability without an ``id`` a random one per run
and stamps it
+ on its tools. A retry is a new run, so the model request differs only
in that id;
+ it must still replay; ``with-id`` is the control. The model issues
fresh tool call ids,
+ as a real provider does.
+ """
+ storage = _InMemoryDurableStorage()
+ live = {"model": 0, "tool": 0}
+ fail_after_tool = [True]
+
+ def my_tool() -> str:
+ live["tool"] += 1
+ return "tool-result"
+
+ def model_fn(messages, info):
+ live["model"] += 1
+ if any(isinstance(p, ToolReturnPart) for m in messages for p in
m.parts):
+ if fail_after_tool[0]:
+ fail_after_tool[0] = False
+ raise RuntimeError("transient model failure")
+ return ModelResponse(parts=[TextPart(content="done")])
+ return ModelResponse(parts=[ToolCallPart(tool_name="my_tool",
args={})])
+
+ for try_number in (1, 2):
+ live.update(model=0, tool=0)
+ op = AgentOperator(
+ task_id="t",
+ prompt="hi",
+ llm_conn_id="c",
+ durable=True,
+ enable_tool_logging=False,
+ capabilities=[capability(my_tool)],
+ )
+ op.llm_hook = MagicMock(spec=["create_agent"])
+ op.llm_hook.create_agent.side_effect = lambda **kw:
Agent(FunctionModel(model_fn), **kw)
+ context = _make_context(ti=_make_ti(id=f"ti-{try_number}",
try_number=try_number))
+ with (
+ patch.object(AgentOperator, "_build_durable_storage",
autospec=True, return_value=storage),
+ pytest.raises(RuntimeError, match="transient") if try_number
== 1 else nullcontext(),
+ ):
+ op.execute(context=context)
+ if try_number == 1:
+ # Verified replay, not positional replay of unfingerprintable
(None) steps.
+ assert storage.models
+ assert all(fingerprint is not None for _, fingerprint in
storage.models.values())
+
+ # Attempt 2 replays model step 0 and the tool call; only the step that
failed runs live.
+ assert live == {"model": 1, "tool": 0}
+
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."""