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"):