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

Reply via email to