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 7938c0dbb65 Make `execute_tool` the public method `AirflowToolset` 
subclasses implement (#73938)
7938c0dbb65 is described below

commit 7938c0dbb65a7c83df04ed027bf33eb4a0ba250f
Author: Kaxil Naik <[email protected]>
AuthorDate: Wed Sep 30 12:02:15 2026 +0100

    Make `execute_tool` the public method `AirflowToolset` subclasses implement 
(#73938)
    
    
    A subclass had to implement the private _execute_tool, so a method that was
    free to change was also the one every toolset depends on. It is now the 
public
    execute_tool, with ctx and tool keyword-only so that arguments can be added 
later
    without breaking subclasses. call_tool keeps pydantic-ai's positional 
signature,
    since pydantic-ai calls it that way.
    
    Also renames with_masking to ensure_masked, which says that a toolset that
    already masks its output comes back unchanged, and gives the masking helpers
    names that are easier to tell apart: _masked becomes _mask_call and 
_stripped
    becomes _mask_attributes. MCPToolset._server_once_resolved becomes
    _resolve_server.
---
 .../airflow/providers/common/ai/operators/agent.py |  8 ++---
 .../providers/common/ai/toolsets/datafusion.py     |  3 +-
 .../airflow/providers/common/ai/toolsets/hook.py   |  5 +--
 .../providers/common/ai/toolsets/managed_agent.py  |  3 +-
 .../airflow/providers/common/ai/toolsets/mcp.py    | 11 +++---
 .../providers/common/ai/toolsets/object_storage.py |  3 +-
 .../providers/common/ai/toolsets/sandbox.py        |  3 +-
 .../airflow/providers/common/ai/toolsets/sql.py    |  3 +-
 .../providers/common/ai/utils/toolset_base.py      | 42 +++++++++++++++-------
 .../unit/common/ai/utils/test_tool_metrics.py      |  6 ++--
 .../unit/common/ai/utils/test_toolset_base.py      | 10 +++---
 11 files changed, 60 insertions(+), 37 deletions(-)

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 bbee364a2ef..75a077c16f8 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
@@ -56,7 +56,7 @@ from airflow.providers.common.ai.utils.logging import (
     wrap_toolsets_for_logging,
 )
 from airflow.providers.common.ai.utils.output_type import 
rehydrate_pydantic_output
-from airflow.providers.common.ai.utils.toolset_base import with_masking
+from airflow.providers.common.ai.utils.toolset_base import ensure_masked
 from airflow.providers.common.ai.utils.toolsets import iter_toolsets
 from airflow.providers.common.ai.utils.usage import coerce_usage_limits
 from airflow.providers.common.ai.utils.usage_budget import (
@@ -656,16 +656,16 @@ class AgentOperator(CancellableAgentRunMixin, 
BaseOperator, HITLReviewMixin):
         counter = self._durable_counter
         if self.toolsets:
             # Innermost, so the durable cache only ever stores masked results.
-            toolsets: list[AbstractToolset] = [with_masking(ts) for ts in 
self.toolsets]
+            toolsets: list[AbstractToolset] = [ensure_masked(ts) for ts in 
self.toolsets]
             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
         elif extra_kwargs.get("toolsets"):
-            extra_kwargs["toolsets"] = [with_masking(ts) for ts in 
extra_kwargs["toolsets"]]
+            extra_kwargs["toolsets"] = [ensure_masked(ts) for ts in 
extra_kwargs["toolsets"]]
         capabilities = [
-            replace(capability, toolset=with_masking(capability.toolset))
+            replace(capability, toolset=ensure_masked(capability.toolset))
             if _is_concrete_toolset_capability(capability)
             else capability
             for capability in extra_kwargs.get("capabilities") or []
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 a40cacad529..805a5f4eea3 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
@@ -175,10 +175,11 @@ class DataFusionToolset(AirflowToolset):
             )
         return tools
 
-    async def _execute_tool(
+    async def execute_tool(
         self,
         name: str,
         tool_args: dict[str, Any],
+        *,
         ctx: RunContext[Any],
         tool: ToolsetTool[Any],
     ) -> Any:
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 882bee88eb5..382c9940eb2 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
@@ -178,7 +178,7 @@ class HookToolset(AirflowToolset):
             # sequential=True keeps pydantic-ai from running these calls 
concurrently
             # within a turn; run_blocking's process-wide lock serializes them 
with the
             # blocking calls of the other toolsets that use it.
-            # return_schema is "string": _execute_tool serializes every result 
with
+            # return_schema is "string": execute_tool serializes every result 
with
             # serialize_for_llm, so the tool always returns a (JSON-encoded)
             # string regardless of the method's own return annotation. This 
lets
             # code mode render `-> str` instead of `-> Any`.
@@ -197,10 +197,11 @@ class HookToolset(AirflowToolset):
             )
         return tools
 
-    async def _execute_tool(
+    async def execute_tool(
         self,
         name: str,
         tool_args: dict[str, Any],
+        *,
         ctx: RunContext[Any],
         tool: ToolsetTool[Any],
     ) -> Any:
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 28627c9787e..700da193ef3 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
@@ -231,10 +231,11 @@ class BaseManagedAgentToolset(AirflowToolset):
             )
         }
 
-    async def _execute_tool(
+    async def execute_tool(
         self,
         name: str,
         tool_args: dict[str, Any],
+        *,
         ctx: RunContext[Any],
         tool: ToolsetTool[Any],
     ) -> Any:
diff --git 
a/providers/common/ai/src/airflow/providers/common/ai/toolsets/mcp.py 
b/providers/common/ai/src/airflow/providers/common/ai/toolsets/mcp.py
index f181ee255f6..dfac26562c1 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/toolsets/mcp.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/toolsets/mcp.py
@@ -113,12 +113,12 @@ class MCPToolset(AirflowToolset):
             self._server = hook.get_conn()
         return self._server
 
-    async def _server_once_resolved(self) -> Any:
+    async def _resolve_server(self) -> Any:
         # Resolving the connection talks to the supervisor, so it takes the 
blocking-call lock.
         return self._server if self._server is not None else await 
self.run_blocking(self._get_server)
 
     async def __aenter__(self) -> Self:
-        await (await self._server_once_resolved()).__aenter__()
+        await (await self._resolve_server()).__aenter__()
         return self
 
     async def __aexit__(self, *args: Any) -> bool | None:
@@ -127,16 +127,17 @@ class MCPToolset(AirflowToolset):
         return None
 
     async def get_tools(self, ctx: RunContext[Any]) -> dict[str, 
ToolsetTool[Any]]:
-        return await (await self._server_once_resolved()).get_tools(ctx)
+        return await (await self._resolve_server()).get_tools(ctx)
 
-    async def _execute_tool(
+    async def execute_tool(
         self,
         name: str,
         tool_args: dict[str, Any],
+        *,
         ctx: RunContext[Any],
         tool: ToolsetTool[Any],
     ) -> Any:
-        return await (await self._server_once_resolved()).call_tool(name, 
tool_args, ctx, tool)
+        return await (await self._resolve_server()).call_tool(name, tool_args, 
ctx, tool)
 
     def airflow_tools(self) -> list[AirflowTool]:
         """
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 bd688d14261..5a7dde59b06 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
@@ -204,10 +204,11 @@ class ObjectStorageToolset(AirflowToolset):
             )
         return tools
 
-    async def _execute_tool(
+    async def execute_tool(
         self,
         name: str,
         tool_args: dict[str, Any],
+        *,
         ctx: RunContext[Any],
         tool: ToolsetTool[Any],
     ) -> str:
diff --git 
a/providers/common/ai/src/airflow/providers/common/ai/toolsets/sandbox.py 
b/providers/common/ai/src/airflow/providers/common/ai/toolsets/sandbox.py
index 21cbe12f233..3ae301b8584 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/toolsets/sandbox.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/toolsets/sandbox.py
@@ -680,10 +680,11 @@ class SandboxToolset(AirflowToolset):
             )
         return tools
 
-    async def _execute_tool(
+    async def execute_tool(
         self,
         name: str,
         tool_args: dict[str, Any],
+        *,
         ctx: RunContext[Any],
         tool: ToolsetTool[Any],
     ) -> Any:
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 8b8916e3aec..19384e2c669 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
@@ -397,10 +397,11 @@ class SQLToolset(AirflowToolset):
             )
         return tools
 
-    async def _execute_tool(
+    async def execute_tool(
         self,
         name: str,
         tool_args: dict[str, Any],
+        *,
         ctx: RunContext[Any],
         tool: ToolsetTool[Any],
     ) -> Any:
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 7f4fe9ebbf0..65c85557945 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
@@ -70,7 +70,7 @@ def _call_locked(fn: Callable[P, R], /, *args: P.args, 
**kwargs: P.kwargs) -> R:
         return fn(*args, **kwargs)
 
 
-async def _masked(name: str, call: Awaitable[Any], *, count_as: str | None = 
None) -> Any:
+async def _mask_call(name: str, call: Awaitable[Any], *, count_as: str | None 
= None) -> Any:
     """
     Await a tool call and mask everything it hands on: its result, or the 
exception it raised.
 
@@ -123,7 +123,7 @@ def _strip(error: Exception) -> Exception:
     error.__cause__ = None
     error.__context__ = None
     try:
-        stripped = _stripped(error)
+        stripped = _mask_attributes(error)
         message = str(stripped)
         if (masked := mask_secrets(message)) != message:
             stripped = RuntimeError(f"{type(error).__name__}: {masked}")
@@ -133,7 +133,7 @@ def _strip(error: Exception) -> Exception:
     return stripped
 
 
-def _stripped(error: Exception) -> Exception:
+def _mask_attributes(error: Exception) -> Exception:
     group = getattr(error, "exceptions", None)
     if isinstance(group, tuple) and hasattr(error, "derive"):
         # An exception group's own message and arguments are set when it is 
built.
@@ -155,10 +155,18 @@ class AirflowToolset(AbstractToolset[Any]):
     """
     A toolset whose tool results are safe to hand to a model.
 
-    Subclasses implement :meth:`_execute_tool`. :meth:`call_tool` runs it and 
passes what it
-    returns, and any exception it raises, through Airflow's secret masker, so 
a connection
-    password that ends up in a database error or a hook's return value is 
replaced with
-    ``***`` before the model, the model provider or a trace sees it.
+    Two methods look alike and have different jobs. :meth:`execute_tool` is 
the one a subclass
+    writes: it runs the tool and returns the result as the tool produced it. 
:meth:`call_tool`
+    is pydantic-ai's entry point, implemented here once: it runs 
``execute_tool`` and passes what
+    it returns, and any exception it raises, through Airflow's secret masker, 
so a connection
+    password that ends up in a database error or a hook's return value is 
replaced with ``***``
+    before the model, the model provider or a trace sees it. A subclass that 
overrides
+    ``call_tool`` instead skips that masking, which is why 
:func:`ensure_masked` wraps such a
+    toolset again.
+
+    The two signatures differ on purpose. ``call_tool`` keeps the positional 
shape pydantic-ai
+    invokes it with. ``execute_tool`` takes ``ctx`` and ``tool`` keyword-only, 
so arguments can
+    be added to it later without breaking subclasses.
     """
 
     async def call_tool(
@@ -168,19 +176,27 @@ class AirflowToolset(AbstractToolset[Any]):
         ctx: RunContext[Any],
         tool: ToolsetTool[Any],
     ) -> Any:
-        return await _masked(
-            name, self._execute_tool(name, tool_args, ctx, tool), 
count_as=type(self).__name__
+        # pydantic-ai calls this positionally, so its signature has to stay as 
it defines it.
+        return await _mask_call(
+            name, self.execute_tool(name, tool_args, ctx=ctx, tool=tool), 
count_as=type(self).__name__
         )
 
     @abstractmethod
-    async def _execute_tool(
+    async def execute_tool(
         self,
         name: str,
         tool_args: dict[str, Any],
+        *,
         ctx: RunContext[Any],
         tool: ToolsetTool[Any],
     ) -> Any:
-        """Run tool ``name`` with validated ``tool_args``; :meth:`call_tool` 
masks what it returns."""
+        """
+        Run tool ``name`` with validated ``tool_args`` and return its result 
unmasked.
+
+        This is the method a subclass implements; :meth:`call_tool` runs it 
and masks what it
+        returns. ``ctx`` and ``tool`` are keyword-only so that arguments can 
be added here later
+        without breaking subclasses.
+        """
 
     def airflow_tools(self) -> list[AirflowTool]:
         """
@@ -216,10 +232,10 @@ class MaskingToolset(WrapperToolset[Any]):
         ctx: RunContext[Any],
         tool: ToolsetTool[Any],
     ) -> Any:
-        return await _masked(name, self.wrapped.call_tool(name, tool_args, 
ctx, tool))
+        return await _mask_call(name, self.wrapped.call_tool(name, tool_args, 
ctx, tool))
 
 
-def with_masking(toolset: AbstractToolset[Any] | ToolsetFunc[Any]) -> 
AbstractToolset[Any]:
+def ensure_masked(toolset: AbstractToolset[Any] | ToolsetFunc[Any]) -> 
AbstractToolset[Any]:
     """
     Return ``toolset`` wrapped in :class:`MaskingToolset`, unless it already 
masks its own output.
 
diff --git 
a/providers/common/ai/tests/unit/common/ai/utils/test_tool_metrics.py 
b/providers/common/ai/tests/unit/common/ai/utils/test_tool_metrics.py
index e1a5fd23f57..3b1f26fd966 100644
--- a/providers/common/ai/tests/unit/common/ai/utils/test_tool_metrics.py
+++ b/providers/common/ai/tests/unit/common/ai/utils/test_tool_metrics.py
@@ -34,7 +34,7 @@ from airflow.providers.common.ai.utils.tool_metrics import (
     calling_framework,
     record_tool_call,
 )
-from airflow.providers.common.ai.utils.toolset_base import MaskingToolset, 
with_masking
+from airflow.providers.common.ai.utils.toolset_base import MaskingToolset, 
ensure_masked
 
 from unit.common.ai.operators.test_agent import _InMemoryDurableStorage
 from unit.common.ai.toolsets.test_sql import _make_mock_db_hook
@@ -119,7 +119,7 @@ class TestToolsetsCountTheirCalls:
         storage = _InMemoryDurableStorage()
         for _ in range(2):
             cached = CachingToolset(
-                wrapped=with_masking(_sql_toolset()), storage=storage, 
counter=DurableStepCounter()
+                wrapped=ensure_masked(_sql_toolset()), storage=storage, 
counter=DurableStepCounter()
             )
             ctx = RunContext(deps=None, model=TestModel(), usage=RunUsage(), 
tool_call_id="c1")
 
@@ -151,7 +151,7 @@ class TestOutcomes:
     def test_a_replay_through_a_wrapper_counts_the_toolset_underneath(self, 
stats):
         storage = _InMemoryDurableStorage()
         for _ in range(2):
-            wrapped = with_masking(_sql_toolset().prefixed("wh"))
+            wrapped = ensure_masked(_sql_toolset().prefixed("wh"))
             cached = CachingToolset(wrapped=wrapped, storage=storage, 
counter=DurableStepCounter())
             ctx = RunContext(deps=None, model=TestModel(), usage=RunUsage(), 
tool_call_id="c1")
 
diff --git 
a/providers/common/ai/tests/unit/common/ai/utils/test_toolset_base.py 
b/providers/common/ai/tests/unit/common/ai/utils/test_toolset_base.py
index 8154ce25016..7121cf85175 100644
--- a/providers/common/ai/tests/unit/common/ai/utils/test_toolset_base.py
+++ b/providers/common/ai/tests/unit/common/ai/utils/test_toolset_base.py
@@ -35,7 +35,7 @@ from pydantic_ai.toolsets.abstract import ToolsetTool
 from pydantic_ai.toolsets.function import FunctionToolset
 from pydantic_ai.usage import RunUsage
 
-from airflow.providers.common.ai.utils.toolset_base import AirflowToolset, 
MaskingToolset, with_masking
+from airflow.providers.common.ai.utils.toolset_base import AirflowToolset, 
MaskingToolset, ensure_masked
 
 
 class _ScriptedToolset(AirflowToolset):
@@ -51,7 +51,7 @@ class _ScriptedToolset(AirflowToolset):
     async def get_tools(self, ctx: RunContext[Any]) -> dict[str, 
ToolsetTool[Any]]:
         return {}
 
-    async def _execute_tool(self, name, tool_args, ctx, tool) -> Any:
+    async def execute_tool(self, name, tool_args, *, ctx, tool) -> Any:
         if isinstance(self._outcome, BaseException):
             raise self._outcome
         return self._outcome
@@ -257,14 +257,14 @@ class TestWithMasking:
         def per_run(ctx: RunContext[Any]) -> FunctionToolset:
             return _ApiToolset(registered_secret)
 
-        assert _run_agent(with_masking(per_run)) == "postgres://svc:***@db"
+        assert _run_agent(ensure_masked(per_run)) == "postgres://svc:***@db"
 
     def test_an_airflow_toolset_that_overrides_call_tool_is_wrapped(self, 
registered_secret):
         class Overriding(_ScriptedToolset):
             async def call_tool(self, name, tool_args, ctx, tool) -> Any:
                 return f"token {registered_secret}"
 
-        masked = with_masking(Overriding("unused"))
+        masked = ensure_masked(Overriding("unused"))
 
         assert isinstance(masked, MaskingToolset)
         assert _call(masked) == "token ***"
@@ -272,7 +272,7 @@ class TestWithMasking:
     def test_an_airflow_toolset_that_masks_itself_is_not_wrapped(self):
         toolset = _ScriptedToolset("ok")
 
-        assert with_masking(toolset) is toolset
+        assert ensure_masked(toolset) is toolset
 
     def test_a_failure_masked_by_two_layers_is_logged_once(self, caplog):
         with pytest.raises(ValueError, match="boom"):

Reply via email to