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 d886882ac14 Let Common AI toolsets use the agent's tool retry budget 
(#73957)
d886882ac14 is described below

commit d886882ac14eed7674ca456038ec6d01efa059c0
Author: Kaxil Naik <[email protected]>
AuthorDate: Thu Oct 1 10:25:41 2026 +0100

    Let Common AI toolsets use the agent's tool retry budget (#73957)
    
    SQLToolset, HookToolset, DataFusionToolset and ObjectStorageToolset pinned
    max_retries=1 on every tool, overriding the agent's retries. An agent given
    retries=3 still ended its run on the second failed query in a row, so a 
model
    that needed a couple of attempts to get a column name right could not.
    
    Each toolset now takes max_retries, defaulting to the run's own tool retry
    budget, the way pydantic-ai's FunctionToolset does. Outside a pydantic-ai 
run
    (the LangChain, Strands and ADK bridges) there is no agent budget, so tools
    get pydantic-ai's default of one correction, and each bridged call now gets 
a
    context with its own retry count, so ctx.last_attempt means what it does
    inside a pydantic-ai run.
---
 providers/common/ai/docs/toolsets/datafusion.rst   |  2 +
 providers/common/ai/docs/toolsets/hook.rst         |  3 +
 providers/common/ai/docs/toolsets/index.rst        | 42 ++++++++++
 .../common/ai/docs/toolsets/object_storage.rst     |  3 +
 providers/common/ai/docs/toolsets/sql.rst          |  4 +-
 .../providers/common/ai/tools/_from_toolset.py     | 18 +++-
 .../providers/common/ai/toolsets/datafusion.py     | 10 ++-
 .../airflow/providers/common/ai/toolsets/hook.py   | 11 ++-
 .../common/ai/toolsets/langchain_bridge.py         |  5 +-
 .../providers/common/ai/toolsets/managed_agent.py  | 11 ++-
 .../providers/common/ai/toolsets/object_storage.py | 11 ++-
 .../airflow/providers/common/ai/toolsets/sql.py    | 10 ++-
 .../providers/common/ai/utils/toolset_base.py      | 19 +++++
 .../unit/common/ai/tools/test__from_toolset.py     | 33 +++++++-
 .../unit/common/ai/toolsets/test_retry_budget.py   | 98 ++++++++++++++++++++++
 15 files changed, 259 insertions(+), 21 deletions(-)

diff --git a/providers/common/ai/docs/toolsets/datafusion.rst 
b/providers/common/ai/docs/toolsets/datafusion.rst
index 27db2d8a109..55570ab32f5 100644
--- a/providers/common/ai/docs/toolsets/datafusion.rst
+++ b/providers/common/ai/docs/toolsets/datafusion.rst
@@ -86,6 +86,8 @@ Parameters
 - ``max_rows``: Maximum rows returned from the ``query`` tool. Default ``50``.
 - ``max_result_bytes``: Budget for the serialized ``query`` result. Default 64 
KiB.
   See :ref:`bounded-query-results`.
+- ``max_retries``: How many times the model may correct a failed call to these
+  tools. Default ``None``, the agent's ``retries``. See 
:ref:`toolset-retry-budget`.
 
 When to choose it
 -----------------
diff --git a/providers/common/ai/docs/toolsets/hook.rst 
b/providers/common/ai/docs/toolsets/hook.rst
index 9acf6ed2a7b..15a6ad0d2b8 100644
--- a/providers/common/ai/docs/toolsets/hook.rst
+++ b/providers/common/ai/docs/toolsets/hook.rst
@@ -133,6 +133,9 @@ Parameters
   (e.g. ``"s3_"`` produces ``"s3_list_keys"``).
 - ``pinned_arguments``: Arguments fixed by the Dag author rather than chosen 
by the
   model. See above.
+- ``max_retries``: How many times the model may correct a call with invalid 
arguments,
+  or one that changes a pinned argument. Default ``None``, the agent's 
``retries``. See
+  :ref:`toolset-retry-budget`.
 
 When to choose it
 -----------------
diff --git a/providers/common/ai/docs/toolsets/index.rst 
b/providers/common/ai/docs/toolsets/index.rst
index c2e91ec63fd..fa9db38d70a 100644
--- a/providers/common/ai/docs/toolsets/index.rst
+++ b/providers/common/ai/docs/toolsets/index.rst
@@ -222,6 +222,48 @@ the upstream toolset declares, so the setting is not 
theirs to make. Do not read
 this as a reason to choose one route over another; read it as something to 
expect
 from all four.
 
+.. _toolset-retry-budget:
+
+How often the model may correct a failed call
+---------------------------------------------
+
+When the model calls a tool with arguments that fail its schema, or the tool 
asks the
+model to try again (``ModelRetry``), the error goes back to the model so it 
can correct
+the call. What counts differs by toolset:
+
+- ``SQLToolset`` and ``DataFusionToolset`` turn every query error into 
``ModelRetry``,
+  so a misspelled column and a dropped connection both count.
+- ``HookToolset`` counts invalid arguments and a call that tries to change a 
pinned
+  argument. An exception from the hook itself fails the run straight away.
+- ``ObjectStorageToolset`` counts invalid arguments only. A path that does not 
exist or
+  cannot be read goes back to the model as a failed result without using the 
budget;
+  bound repeated failed reads with ``usage_limits``.
+
+These toolsets allow as many corrections as the agent's tool retry budget, 
pydantic-ai's
+``retries`` (one by default), the same way pydantic-ai's own toolsets do. Pass
+``max_retries`` to a toolset to give its tools a budget of their own. Once the 
budget is
+used up the run fails, and Airflow's task retries take over.
+
+.. code-block:: python
+
+    AgentOperator(
+        task_id="revenue_agent",
+        prompt="What was last week's revenue?",
+        llm_conn_id="pydanticai_default",
+        toolsets=[SQLToolset(db_conn_id="warehouse")],
+        agent_params={"retries": {"tools": 3}},
+    )
+
+An integer ``retries`` sets both the tool budget and the output-validation 
budget; a
+dict such as ``{"tools": 3}`` or ``{"output": 3}`` raises only one of them. 
These toolsets
+used to allow exactly one correction whatever ``retries`` said, so an agent 
that sets
+``retries`` now applies it to them too: ``retries=0`` fails the run on the 
first bad
+query, and a large integer ``retries`` meant for output validation also lets a 
failing
+database be queried that many times. Pass ``max_retries=1`` to a toolset to 
keep the old
+behaviour. Outside a pydantic-ai agent (the LangChain, Strands and Google ADK 
bridges)
+there is no agent budget, so each tool gets one correction unless its toolset 
sets
+``max_retries``.
+
 Layering
 --------
 
diff --git a/providers/common/ai/docs/toolsets/object_storage.rst 
b/providers/common/ai/docs/toolsets/object_storage.rst
index fdc069bc41d..67984d575e7 100644
--- a/providers/common/ai/docs/toolsets/object_storage.rst
+++ b/providers/common/ai/docs/toolsets/object_storage.rst
@@ -90,6 +90,9 @@ Parameters
     A prefix for the tool names, needed when the agent has another toolset 
with the same
     tool names: a second ``ObjectStorageToolset``, or a ``SandboxToolset``, 
which has a
     ``read_file`` of its own.
+``max_retries``
+    How many times the model may correct a call with invalid arguments. A 
failed read does
+    not count. Default ``None``, the agent's ``retries``. See 
:ref:`toolset-retry-budget`.
 
 When to choose it
 -----------------
diff --git a/providers/common/ai/docs/toolsets/sql.rst 
b/providers/common/ai/docs/toolsets/sql.rst
index b34433ef8bd..1c8481ef705 100644
--- a/providers/common/ai/docs/toolsets/sql.rst
+++ b/providers/common/ai/docs/toolsets/sql.rst
@@ -205,6 +205,8 @@ Parameters
   transferred is its own call. See :ref:`bounded-query-results`.
 - ``max_result_bytes``: Budget for the serialized ``query`` result. Default 64 
KiB.
   See :ref:`bounded-query-results`.
+- ``max_retries``: How many times the model may correct a failed call to these
+  tools. Default ``None``, the agent's ``retries``. See 
:ref:`toolset-retry-budget`.
 
 .. _bounded-query-results:
 
@@ -296,7 +298,7 @@ subqueries and joins.
 - It does not classify failures. A connection error or a typo in a column
   name reaching ``list_tables``, ``get_schema`` or ``query`` becomes one
   ``ModelRetry``, so the two are treated the same way until the retry budget
-  runs out and the task fails for Airflow to retry. Two paths do not raise:
+  (:ref:`toolset-retry-budget`) runs out and the task fails for Airflow to 
retry. Two paths do not raise:
   ``check_query`` catches its own errors and reports them back as a normal
   ``{"valid": false, ...}`` result, and ``get_schema`` returns a normal
   ``{"error": ...}`` result instead of raising when the requested table is
diff --git 
a/providers/common/ai/src/airflow/providers/common/ai/tools/_from_toolset.py 
b/providers/common/ai/src/airflow/providers/common/ai/tools/_from_toolset.py
index cb30d7401a4..21bd028f338 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/tools/_from_toolset.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/tools/_from_toolset.py
@@ -26,6 +26,7 @@ import time
 from collections.abc import Hashable, Iterator
 from contextlib import contextmanager
 from contextvars import ContextVar
+from dataclasses import replace
 from typing import TYPE_CHECKING, Any
 
 from pydantic import ValidationError
@@ -79,12 +80,15 @@ def airflow_tools_from_toolset(toolset: 
AbstractToolset[Any], *, deps: Any = Non
 
     Outside a pydantic-ai run there is no live ``RunContext``, so an inert one 
with a
     placeholder model is passed. That is enough for toolsets whose 
``call_tool`` ignores
-    the context, which is true of every toolset this provider exposes this way.
+    the context, which is true of every toolset this provider exposes this 
way. Its
+    ``max_retries`` is the pydantic-ai agent's default of one, the budget a 
tool gets
+    when its toolset takes it from the run. Each call then gets its own copy 
carrying
+    that tool's ``retry`` count and ``max_retries``, as inside a pydantic-ai 
run.
 
     :param toolset: The pydantic-ai toolset to expose.
     :param deps: Exposed to the toolset as ``ctx.deps``.
     """
-    ctx: RunContext[Any] = RunContext(deps=deps, model=TestModel(), 
usage=RunUsage())
+    ctx: RunContext[Any] = RunContext(deps=deps, model=TestModel(), 
usage=RunUsage(), max_retries=1)
     toolset_tools = run_coroutine_sync(toolset.get_tools(ctx))
     in_order = _InOrder()
     return [_as_airflow_tool(toolset, name, tool, ctx, in_order) for name, 
tool in toolset_tools.items()]
@@ -162,6 +166,11 @@ class _RetryBudget:
         self._failed_turn: Hashable | None = None
         self._counted_until = float("-inf")
 
+    @property
+    def failures(self) -> int:
+        """Consecutive failures counted so far, what pydantic-ai passes a tool 
as ``ctx.retry``."""
+        return self._failures
+
     def start(self, run: Hashable | None) -> None:
         """Start the count over when a call belongs to a new run."""
         if run is not None and run != self._run:
@@ -209,8 +218,11 @@ def _as_airflow_tool(
             return correctable(started, turn, str(e))
         # A ValidationError raised by the tool itself is not the model's to 
fix: the call
         # may already have had a side effect, so it propagates rather than 
inviting a retry.
+        # Per call, as pydantic-ai's ToolManager does, so ``ctx.retry`` and 
``ctx.last_attempt``
+        # mean the same here as inside a pydantic-ai run.
+        call_ctx = replace(ctx, retry=budget.failures, 
max_retries=tool.max_retries)
         try:
-            result = await toolset.call_tool(name, validated, ctx, tool)
+            result = await toolset.call_tool(name, validated, call_ctx, tool)
         except ModelRetry as e:
             return correctable(started, turn, e.message)
         except ToolFailed as e:
diff --git 
a/providers/common/ai/src/airflow/providers/common/ai/toolsets/datafusion.py 
b/providers/common/ai/src/airflow/providers/common/ai/toolsets/datafusion.py
index 805a5f4eea3..e338cf9bb83 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/toolsets/datafusion.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/toolsets/datafusion.py
@@ -42,7 +42,7 @@ from airflow.providers.common.ai.utils.query_results import (
     build_query_result,
 )
 from airflow.providers.common.ai.utils.tool_definition import 
build_args_validator
-from airflow.providers.common.ai.utils.toolset_base import AirflowToolset
+from airflow.providers.common.ai.utils.toolset_base import AirflowToolset, 
validate_max_retries
 
 if TYPE_CHECKING:
     from pydantic_ai._run_context import RunContext
@@ -121,6 +121,9 @@ class DataFusionToolset(AirflowToolset):
         rather than skipping it and packing later ones, so one wide row early 
in the
         result ends it. The result reports which limit it hit so the agent can 
narrow
         its projection rather than page through the table.
+    :param max_retries: How many times the model may correct a failed call to 
one of these
+        tools before the run fails. ``None`` (the default) uses the agent's 
tool retry
+        budget, its ``retries``, as pydantic-ai's own toolsets do.
     """
 
     def __init__(
@@ -130,7 +133,9 @@ class DataFusionToolset(AirflowToolset):
         allow_writes: bool = False,
         max_rows: int = 50,
         max_result_bytes: int = DEFAULT_MAX_RESULT_BYTES,
+        max_retries: int | None = None,
     ) -> None:
+        self._max_retries = validate_max_retries(max_retries)
         if not datasource_configs:
             raise ValueError("datasource_configs must contain at least one 
DataSourceConfig")
         self._datasource_configs = datasource_configs
@@ -154,6 +159,7 @@ class DataFusionToolset(AirflowToolset):
         return self._engine
 
     async def get_tools(self, ctx: RunContext[Any]) -> dict[str, 
ToolsetTool[Any]]:
+        max_retries = self._get_tool_max_retries(ctx)
         tools: dict[str, ToolsetTool[Any]] = {}
 
         for name, description, schema in (
@@ -170,7 +176,7 @@ class DataFusionToolset(AirflowToolset):
             tools[name] = ToolsetTool(
                 toolset=self,
                 tool_def=tool_def,
-                max_retries=1,
+                max_retries=max_retries,
                 args_validator=build_args_validator(schema),
             )
         return tools
diff --git 
a/providers/common/ai/src/airflow/providers/common/ai/toolsets/hook.py 
b/providers/common/ai/src/airflow/providers/common/ai/toolsets/hook.py
index 382c9940eb2..0dac4f5b080 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/toolsets/hook.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/toolsets/hook.py
@@ -33,7 +33,7 @@ from airflow.providers.common.ai.utils.tool_definition import 
(
     return_schema_kwargs,
     serialize_for_llm,
 )
-from airflow.providers.common.ai.utils.toolset_base import AirflowToolset
+from airflow.providers.common.ai.utils.toolset_base import AirflowToolset, 
validate_max_retries
 
 if TYPE_CHECKING:
     from collections.abc import Callable, Iterable, Sequence
@@ -80,6 +80,10 @@ class HookToolset(AirflowToolset):
         as a method taking ``bucket`` or only ``**kwargs``, raises 
``ValueError``, because
         the model could still choose the value through it. Expose such a 
method from a
         second ``HookToolset``.
+    :param max_retries: How many times the model may correct a call with 
invalid arguments,
+        or one that changes a pinned argument, before the run fails. An 
exception from the
+        hook itself fails the run straight away. ``None`` (the default) uses 
the agent's
+        tool retry budget, its ``retries``, as pydantic-ai's own toolsets do.
     """
 
     # Rendered, on a copy, by AgentOperator. Deliberately not 
``template_fields``, which
@@ -93,7 +97,9 @@ class HookToolset(AirflowToolset):
         allowed_methods: list[str],
         tool_name_prefix: str = "",
         pinned_arguments: dict[str, Any] | None = None,
+        max_retries: int | None = None,
     ) -> None:
+        self._max_retries = validate_max_retries(max_retries)
         if not allowed_methods:
             raise ValueError("allowed_methods must be a non-empty list.")
 
@@ -160,6 +166,7 @@ class HookToolset(AirflowToolset):
         return f"hook-{name}-{self.conn_id}" if self.conn_id else 
f"hook-{name}"
 
     async def get_tools(self, ctx: RunContext[Any]) -> dict[str, 
ToolsetTool[Any]]:
+        max_retries = self._get_tool_max_retries(ctx)
         tools: dict[str, ToolsetTool[Any]] = {}
         for method_name in self._allowed_methods:
             method = getattr(self._hook, method_name)
@@ -192,7 +199,7 @@ class HookToolset(AirflowToolset):
             tools[tool_name] = ToolsetTool(
                 toolset=self,
                 tool_def=tool_def,
-                max_retries=1,
+                max_retries=max_retries,
                 args_validator=build_args_validator(json_schema),
             )
         return tools
diff --git 
a/providers/common/ai/src/airflow/providers/common/ai/toolsets/langchain_bridge.py
 
b/providers/common/ai/src/airflow/providers/common/ai/toolsets/langchain_bridge.py
index c6dad5d8102..b2d821367f1 100644
--- 
a/providers/common/ai/src/airflow/providers/common/ai/toolsets/langchain_bridge.py
+++ 
b/providers/common/ai/src/airflow/providers/common/ai/toolsets/langchain_bridge.py
@@ -89,8 +89,9 @@ def airflow_toolset_to_langchain_tools(
         live :class:`~pydantic_ai.RunContext` carries the model, usage, and
         message history. Outside an agent run there is no such context, so this
         bridge builds a minimal one with an inert placeholder model. The 
curated
-        common.ai toolsets (``SQLToolset``, ``HookToolset``, ``MCPToolset``)
-        ignore the context, so this works for them. A custom toolset that reads
+        common.ai toolsets (``SQLToolset``, ``HookToolset``, ``MCPToolset``) 
read
+        only its retry budget, which the bridge sets, so this works for them. A
+        custom toolset that reads
         live run state (``ctx.model``, ``ctx.messages``, ``ctx.usage``) will 
not
         behave correctly when bridged standalone.
 
diff --git 
a/providers/common/ai/src/airflow/providers/common/ai/toolsets/managed_agent.py 
b/providers/common/ai/src/airflow/providers/common/ai/toolsets/managed_agent.py
index 700da193ef3..d48729b0b92 100644
--- 
a/providers/common/ai/src/airflow/providers/common/ai/toolsets/managed_agent.py
+++ 
b/providers/common/ai/src/airflow/providers/common/ai/toolsets/managed_agent.py
@@ -30,7 +30,7 @@ from airflow.providers.common.ai.utils.tool_definition import 
(
     return_schema_kwargs,
     serialize_for_llm,
 )
-from airflow.providers.common.ai.utils.toolset_base import AirflowToolset
+from airflow.providers.common.ai.utils.toolset_base import AirflowToolset, 
validate_max_retries
 from airflow.providers.common.compat.sdk import Stats
 
 if TYPE_CHECKING:
@@ -108,8 +108,7 @@ class BaseManagedAgentToolset(AirflowToolset):
     ) -> None:
         if not tool_name:
             raise ValueError("tool_name must be a non-empty string.")
-        if max_retries < 0:
-            raise ValueError(f"max_retries must not be negative, got 
{max_retries}.")
+        validate_max_retries(max_retries)
         cls = type(self)
         if (
             cls.invoke is BaseManagedAgentToolset.invoke
@@ -221,12 +220,12 @@ class BaseManagedAgentToolset(AirflowToolset):
                 toolset=self,
                 tool_def=tool_def,
                 # How many times the calling model may rephrase after 
``invoke``
-                # raises ``ModelRetry``. One by default, matching HookToolset: 
a
-                # managed agent invocation is expensive, so the budget is 
small.
+                # raises ``ModelRetry``. One by default rather than the 
agent's own
+                # budget: a managed agent invocation is expensive, so it stays 
small.
                 # Zero disables the ``ModelRetry`` path entirely -- the first 
one
                 # becomes a hard error -- so raise it only when the remote 
agent's
                 # rejections are genuinely worth re-prompting.
-                max_retries=self._max_retries,
+                max_retries=self._get_tool_max_retries(ctx),
                 args_validator=build_args_validator(_PROMPT_SCHEMA),
             )
         }
diff --git 
a/providers/common/ai/src/airflow/providers/common/ai/toolsets/object_storage.py
 
b/providers/common/ai/src/airflow/providers/common/ai/toolsets/object_storage.py
index 5a7dde59b06..6e19b6fe4a0 100644
--- 
a/providers/common/ai/src/airflow/providers/common/ai/toolsets/object_storage.py
+++ 
b/providers/common/ai/src/airflow/providers/common/ai/toolsets/object_storage.py
@@ -44,7 +44,7 @@ from airflow.providers.common.ai.utils.tool_definition import 
(
     return_schema_kwargs,
     serialize_for_llm,
 )
-from airflow.providers.common.ai.utils.toolset_base import AirflowToolset
+from airflow.providers.common.ai.utils.toolset_base import AirflowToolset, 
validate_max_retries
 from airflow.providers.common.compat.sdk import 
AirflowOptionalProviderFeatureException, ObjectStoragePath
 
 if TYPE_CHECKING:
@@ -147,6 +147,10 @@ class ObjectStorageToolset(AirflowToolset):
         ``reports_read_file``. Set this when one agent has another toolset 
with the same tool
         names, such as a second ``ObjectStorageToolset`` or a 
``SandboxToolset``, whose
         ``read_file`` would collide, since duplicate tool names are rejected.
+    :param max_retries: How many times the model may correct a call with 
invalid arguments
+        before the run fails. A failed read goes back to the model without 
using it.
+        ``None`` (the default) uses the agent's tool retry budget, its 
``retries``, as
+        pydantic-ai's own toolsets do.
     """
 
     # Rendered, on a copy, by AgentOperator. Deliberately not 
``template_fields``, which
@@ -162,7 +166,9 @@ class ObjectStorageToolset(AirflowToolset):
         max_read_bytes: int = 10 * 1024 * 1024,
         max_output_bytes: int = 50 * 1024,
         tool_prefix: str = "",
+        max_retries: int | None = None,
     ) -> None:
+        self._max_retries = validate_max_retries(max_retries)
         for name, value in (
             ("max_files", max_files),
             ("max_read_bytes", max_read_bytes),
@@ -188,6 +194,7 @@ class ObjectStorageToolset(AirflowToolset):
         return f"{self._tool_prefix}_{base}" if self._tool_prefix else base
 
     async def get_tools(self, ctx: RunContext[Any]) -> dict[str, 
ToolsetTool[Any]]:
+        max_retries = self._get_tool_max_retries(ctx)
         tools: dict[str, ToolsetTool[Any]] = {}
         for base, schema in _SCHEMAS.items():
             name = self._tool_name(base)
@@ -199,7 +206,7 @@ class ObjectStorageToolset(AirflowToolset):
                     parameters_json_schema=schema,
                     **return_schema_kwargs({"type": "string"}),
                 ),
-                max_retries=1,
+                max_retries=max_retries,
                 args_validator=build_args_validator(schema),
             )
         return tools
diff --git 
a/providers/common/ai/src/airflow/providers/common/ai/toolsets/sql.py 
b/providers/common/ai/src/airflow/providers/common/ai/toolsets/sql.py
index 19384e2c669..d3271315f18 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/toolsets/sql.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/toolsets/sql.py
@@ -47,7 +47,7 @@ from airflow.providers.common.ai.utils.query_results import (
     build_query_result,
 )
 from airflow.providers.common.ai.utils.tool_definition import 
build_args_validator, return_schema_kwargs
-from airflow.providers.common.ai.utils.toolset_base import AirflowToolset
+from airflow.providers.common.ai.utils.toolset_base import AirflowToolset, 
validate_max_retries
 from airflow.providers.common.compat.sdk import BaseHook
 
 if TYPE_CHECKING:
@@ -250,6 +250,9 @@ class SQLToolset(AirflowToolset):
         rather than skipping it and packing later ones, so one wide row early 
in the
         result ends it. The result reports which limit it hit so the agent can 
narrow
         its projection rather than page through the table.
+    :param max_retries: How many times the model may correct a failed call to 
one of these
+        tools before the run fails. ``None`` (the default) uses the agent's 
tool retry
+        budget, its ``retries``, as pydantic-ai's own toolsets do.
     """
 
     # Rendered, on a copy, by AgentOperator. Deliberately not 
``template_fields``, which
@@ -266,7 +269,9 @@ class SQLToolset(AirflowToolset):
         allow_writes: bool = False,
         max_rows: int = 50,
         max_result_bytes: int = DEFAULT_MAX_RESULT_BYTES,
+        max_retries: int | None = None,
     ) -> None:
+        self._max_retries = validate_max_retries(max_retries)
         self._allowed_tables: frozenset[str] | None
         if allowed_tables is _UNSET:
             self._allowed_tables = None
@@ -369,6 +374,7 @@ class SQLToolset(AirflowToolset):
     # ------------------------------------------------------------------
 
     async def get_tools(self, ctx: RunContext[Any]) -> dict[str, 
ToolsetTool[Any]]:
+        max_retries = self._get_tool_max_retries(ctx)
         tools: dict[str, ToolsetTool[Any]] = {}
 
         for name, description, schema in (
@@ -392,7 +398,7 @@ class SQLToolset(AirflowToolset):
             tools[name] = ToolsetTool(
                 toolset=self,
                 tool_def=tool_def,
-                max_retries=1,
+                max_retries=max_retries,
                 args_validator=build_args_validator(schema),
             )
         return tools
diff --git 
a/providers/common/ai/src/airflow/providers/common/ai/utils/toolset_base.py 
b/providers/common/ai/src/airflow/providers/common/ai/utils/toolset_base.py
index 65c85557945..0b2da926221 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/utils/toolset_base.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/utils/toolset_base.py
@@ -151,6 +151,13 @@ def _mask_attributes(error: Exception) -> Exception:
     return error
 
 
+def validate_max_retries(max_retries: int | None) -> int | None:
+    """Return ``max_retries`` unchanged, or raise ``ValueError`` if it is 
negative."""
+    if max_retries is not None and max_retries < 0:
+        raise ValueError(f"max_retries must not be negative, got 
{max_retries}.")
+    return max_retries
+
+
 class AirflowToolset(AbstractToolset[Any]):
     """
     A toolset whose tool results are safe to hand to a model.
@@ -181,6 +188,18 @@ class AirflowToolset(AbstractToolset[Any]):
             name, self.execute_tool(name, tool_args, ctx=ctx, tool=tool), 
count_as=type(self).__name__
         )
 
+    # A subclass that takes ``max_retries`` stores it here; ``None`` follows 
the run.
+    _max_retries: int | None = None
+
+    def _get_tool_max_retries(self, ctx: RunContext[Any]) -> int:
+        """
+        Return how many times the model may correct a failed call to this 
toolset's tools.
+
+        The toolset's own ``max_retries`` if set, else the run's tool retry 
budget (the
+        agent's ``retries``), the same order pydantic-ai's ``FunctionToolset`` 
uses.
+        """
+        return ctx.max_retries if self._max_retries is None else 
self._max_retries
+
     @abstractmethod
     async def execute_tool(
         self,
diff --git 
a/providers/common/ai/tests/unit/common/ai/tools/test__from_toolset.py 
b/providers/common/ai/tests/unit/common/ai/tools/test__from_toolset.py
index 22040f67b08..e4d8dd0f2e7 100644
--- a/providers/common/ai/tests/unit/common/ai/tools/test__from_toolset.py
+++ b/providers/common/ai/tests/unit/common/ai/tools/test__from_toolset.py
@@ -44,7 +44,7 @@ def _by_name(tools) -> dict:
     return {tool.name: tool for tool in tools}
 
 
-def _scripted(*outcomes, max_retries: int = 1):
+def _scripted(*outcomes, max_retries: int | None = 1):
     """A toolset with one tool, ``step``, that returns or raises each outcome 
in turn."""
     remaining = iter(outcomes)
 
@@ -208,6 +208,37 @@ class TestRetryBudget:
         with pytest.raises(ToolCallError, match="kept failing"):
             asyncio.run(turn())
 
+    def test_a_toolset_without_its_own_budget_gets_one_correction(self):
+        """With no agent to take a budget from, a tool whose toolset follows 
the run gets
+        pydantic-ai's default of one correction."""
+        step = _scripted(ModelRetry("bad"), ModelRetry("bad"), 
max_retries=None)
+
+        async def call_in(turn: str) -> ToolResult:
+            with tool_call_scope(run="run", turn=turn):
+                return await step.call({})
+
+        assert asyncio.run(call_in("turn-1")).is_error
+        with pytest.raises(ToolCallError, match="after 1 correction"):
+            asyncio.run(call_in("turn-2"))
+
+    def test_last_attempt_means_what_it_does_in_a_pydantic_ai_run(self):
+        """Each call sees its own retry count, so a tool can fall back on its 
last attempt."""
+
+        def lookup(ctx: RunContext[None]) -> str:
+            """Look the answer up."""
+            if ctx.last_attempt:
+                return "fallback"
+            raise ModelRetry("try again")
+
+        tool = airflow_tools_from_toolset(FunctionToolset([lookup], 
max_retries=1))[0]
+
+        async def call_in(turn: str) -> ToolResult:
+            with tool_call_scope(run="run", turn=turn):
+                return await tool.call({})
+
+        assert asyncio.run(call_in("turn-1")).is_error
+        assert asyncio.run(call_in("turn-2")).content == "fallback"
+
     def test_a_new_run_starts_with_a_fresh_budget(self):
         """An agent reused for a second run gets its full budget again, as in 
pydantic-ai."""
         step = _scripted(ModelRetry("bad"), ModelRetry("bad"), max_retries=1)
diff --git 
a/providers/common/ai/tests/unit/common/ai/toolsets/test_retry_budget.py 
b/providers/common/ai/tests/unit/common/ai/toolsets/test_retry_budget.py
new file mode 100644
index 00000000000..263fcf0bbdd
--- /dev/null
+++ b/providers/common/ai/tests/unit/common/ai/toolsets/test_retry_budget.py
@@ -0,0 +1,98 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements.  See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership.  The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License.  You may obtain a copy of the License at
+#
+#   http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied.  See the License for the
+# specific language governing permissions and limitations
+# under the License.
+"""How many times the model may correct a failed call to a Common AI toolset's 
tools."""
+
+from __future__ import annotations
+
+import asyncio
+
+import pytest
+from pydantic_ai import Agent, RunContext
+from pydantic_ai.exceptions import UnexpectedModelBehavior
+from pydantic_ai.messages import ModelResponse, TextPart, ToolCallPart, 
ToolReturnPart
+from pydantic_ai.models.function import FunctionModel
+from pydantic_ai.models.test import TestModel
+from pydantic_ai.usage import RunUsage
+
+from airflow.providers.common.ai.toolsets.datafusion import DataFusionToolset
+from airflow.providers.common.ai.toolsets.hook import HookToolset
+from airflow.providers.common.ai.toolsets.object_storage import 
ObjectStorageToolset
+from airflow.providers.common.ai.toolsets.sql import SQLToolset
+from airflow.providers.common.compat.sdk import BaseHook
+
+from unit.common.ai.toolsets.test_datafusion import 
_make_mock_datasource_config
+from unit.common.ai.toolsets.test_sql import _make_mock_db_hook
+
+
+class _ListKeysHook(BaseHook):
+    def list_keys(self, bucket: str) -> list[str]:
+        """List the keys in a bucket."""
+        return []
+
+
+TOOLSETS = [
+    pytest.param(lambda **kw: SQLToolset("pg_default", **kw), id="sql"),
+    pytest.param(lambda **kw: HookToolset(_ListKeysHook(), 
allowed_methods=["list_keys"], **kw), id="hook"),
+    pytest.param(lambda **kw: 
DataFusionToolset([_make_mock_datasource_config()], **kw), id="datafusion"),
+    pytest.param(lambda **kw: ObjectStorageToolset("memory://bucket/reports", 
**kw), id="object-storage"),
+]
+
+
[email protected]("make_toolset", TOOLSETS)
+class TestToolsetRetryBudget:
+    @pytest.mark.parametrize(
+        ("kwargs", "expected"),
+        [pytest.param({}, 3, id="agents-budget"), pytest.param({"max_retries": 
0}, 0, id="own-budget")],
+    )
+    def test_tools_get_the_toolsets_budget_else_the_agents(self, make_toolset, 
kwargs, expected):
+        ctx = RunContext(deps=None, model=TestModel(), usage=RunUsage(), 
max_retries=3)
+
+        tools = asyncio.run(make_toolset(**kwargs).get_tools(ctx))
+
+        assert {tool.max_retries for tool in tools.values()} == {expected}
+
+    def test_negative_budget_is_rejected(self, make_toolset):
+        with pytest.raises(ValueError, match="max_retries must not be 
negative"):
+            make_toolset(max_retries=-1)
+
+
+class TestSQLToolsetInAnAgent:
+    @staticmethod
+    def _run(agent_retries: int | None, failures: int) -> str:
+        hook = _make_mock_db_hook()
+        hook.run.side_effect = [RuntimeError('column "totl" does not exist')] 
* failures + [[(42,)]]
+        toolset = SQLToolset("pg_default")
+        toolset._hook = hook
+
+        def model_fn(messages, info):
+            if any(isinstance(p, ToolReturnPart) for m in messages for p in 
m.parts):
+                return ModelResponse(parts=[TextPart(content="done")])
+            call = ToolCallPart(tool_name="query", args={"sql": "SELECT 1"}, 
tool_call_id=f"c{len(messages)}")
+            return ModelResponse(parts=[call])
+
+        return (
+            Agent(FunctionModel(model_fn), toolsets=[toolset], 
retries=agent_retries)
+            .run_sync("total?")
+            .output
+        )
+
+    def 
test_the_agents_retries_let_the_model_correct_its_sql_more_than_once(self):
+        assert self._run(agent_retries=3, failures=2) == "done"
+
+    def test_the_default_budget_still_ends_the_run(self):
+        with pytest.raises(UnexpectedModelBehavior, match="exceeded max 
retries count of 1"):
+            self._run(agent_retries=None, failures=2)

Reply via email to