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 873bd644b51 Let a person approve an agent's tool calls before they run
(#73586)
873bd644b51 is described below
commit 873bd644b518b71fda452f244d95ce486dcf5d20
Author: Kaxil Naik <[email protected]>
AuthorDate: Mon Sep 28 14:49:40 2026 +0100
Let a person approve an agent's tool calls before they run (#73586)
Tools marked with pydantic-ai's approval API (toolset.approval_required()
or Tool(requires_approval=True)) now pause AgentOperator / @task.agent in
the awaiting_input state and ask on the Required Actions page. Approve
runs the call and the agent carries on; Reject tells the agent why and it
carries on without the call. A task instance asks once per Dag run.
Needs Airflow 3.3+.
---
providers/common/ai/docs/examples.rst | 3 +
providers/common/ai/docs/features.rst | 9 +-
providers/common/ai/docs/observability.rst | 4 +-
providers/common/ai/docs/operators/agent.rst | 4 +-
providers/common/ai/docs/tool_approval.rst | 144 ++++++
.../ai/example_dags/example_agent_tool_approval.py | 58 +++
.../src/airflow/providers/common/ai/exceptions.py | 32 +-
.../airflow/providers/common/ai/mixins/approval.py | 30 ++
.../airflow/providers/common/ai/operators/agent.py | 334 ++++++++++++-
.../airflow/providers/common/ai/operators/llm.py | 24 +-
.../providers/common/ai/toolsets/logging.py | 6 +
.../tests/unit/common/ai/operators/test_agent.py | 6 +-
.../ai/operators/test_agent_tool_approval.py | 556 +++++++++++++++++++++
.../tests/unit/common/ai/toolsets/test_logging.py | 18 +
14 files changed, 1183 insertions(+), 45 deletions(-)
diff --git a/providers/common/ai/docs/examples.rst
b/providers/common/ai/docs/examples.rst
index c9a86288dd0..592000c9938 100644
--- a/providers/common/ai/docs/examples.rst
+++ b/providers/common/ai/docs/examples.rst
@@ -93,6 +93,9 @@ Agents & tools
- Connecting an agent to an MCP server through an Airflow connection.
* - :doc:`hitl_review`
- Adding a human-in-the-loop review gate to agent output.
+ * - :doc:`tool_approval`
+ - Pausing an agent before a tool call a person must approve
+ (`example_agent_tool_approval.py
<https://github.com/apache/airflow/blob/providers-common-ai/|version|/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_agent_tool_approval.py>`__).
* - :doc:`use_cases/research_agent_with_review`
- A LangChain ReAct agent that decides its own tool calls, composed with
``LLMOperator`` for
report formatting and AIP-90 HITL review
diff --git a/providers/common/ai/docs/features.rst
b/providers/common/ai/docs/features.rst
index 50c58806b4c..f8c7a7cfbf0 100644
--- a/providers/common/ai/docs/features.rst
+++ b/providers/common/ai/docs/features.rst
@@ -34,10 +34,12 @@ you use. Each is a parameter on the operator or decorator.
approves, edits or rejects the output.
- :doc:`hitl_review`: ``enable_hitl_review=True`` opens an iterative review
loop on an agent,
with a chat UI and REST API for the reviewer.
+- :doc:`tool_approval`: a tool marked with pydantic-ai's approval API pauses
the agent before
+ the call runs, until a person approves or rejects it.
-The last two are different tools for different jobs: an approval gate is a
one-shot decision on
-one output, a HITL review is a conversation with a running agent. Each page
opens with the
-other in a *see also* note.
+The last three are different tools for different jobs: an approval gate is a
one-shot decision on
+one output, a HITL review is a conversation with a running agent, and a tool
approval is a
+decision on one action before it happens.
Making retries cheap with ``durable=True`` is a reliability feature and lives
under
:doc:`operations`.
@@ -52,3 +54,4 @@ Making retries cheap with ``durable=True`` is a reliability
feature and lives un
Code mode <code_mode>
Approve outputs <approval_gates>
Review agent sessions <hitl_review>
+ Approve tool calls <tool_approval>
diff --git a/providers/common/ai/docs/observability.rst
b/providers/common/ai/docs/observability.rst
index a8987a61461..1282df14f0f 100644
--- a/providers/common/ai/docs/observability.rst
+++ b/providers/common/ai/docs/observability.rst
@@ -65,7 +65,9 @@ How it works
(``ti.xcom_pull(task_ids="my_agent", key="run_id")``) and a trace backend can
join a task's output to its agent trace without parsing logs. With
``enable_hitl_review`` the ``run_id`` and ``usage`` reflect the initial model
- run, not the human-feedback regenerations.
+ run, not the human-feedback regenerations. A run that resumes after a tool
+ approval (see :doc:`tool_approval`) continues as ``<task-instance
id>-resumed``,
+ which is the ``run_id`` the operator pushes; ``usage`` covers both parts.
* **Scope.** The ``airflow.*`` identity attributes and the ``run_id`` /
``usage``
XComs come only from ``AgentOperator`` and ``@task.agent``. The other LLM
operators still emit GenAI spans correlated to the task span by nesting, but
diff --git a/providers/common/ai/docs/operators/agent.rst
b/providers/common/ai/docs/operators/agent.rst
index f443150b5ac..ffccc626ef0 100644
--- a/providers/common/ai/docs/operators/agent.rst
+++ b/providers/common/ai/docs/operators/agent.rst
@@ -244,7 +244,7 @@ replayed on retry; they run again. Pass tools you need
replayed in ``toolsets=``
Agent features
--------------
-Four features have pages of their own:
+Five features have pages of their own:
- :doc:`../message_history`: pass ``message_history`` to carry a conversation
across runs.
- :doc:`../durable_execution`: set ``durable=True`` to replay completed model
and tool steps
@@ -253,6 +253,8 @@ Four features have pages of their own:
through ``agent_params``.
- :doc:`../code_mode`: set ``code_mode=True`` to collapse the agent's tools
into a single
``run_code`` tool the model drives by writing Python.
+- :doc:`../tool_approval`: mark tools that need a person's approval, and the
task pauses before
+ a marked call runs.
.. _agent-durable-execution:
diff --git a/providers/common/ai/docs/tool_approval.rst
b/providers/common/ai/docs/tool_approval.rst
new file mode 100644
index 00000000000..347d67e4519
--- /dev/null
+++ b/providers/common/ai/docs/tool_approval.rst
@@ -0,0 +1,144 @@
+ .. 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.
+
+.. _howto/tool_approval:
+
+Approve an agent's tool calls
+=============================
+
+.. seealso::
+ To approve, edit or reject an LLM operator's output instead, see
:doc:`approval_gates`;
+ to review an agent's final answer over several rounds, see
:doc:`hitl_review`.
+
+An agent that can only read is easy to trust. The moment one of its tools does
+something you cannot undo -- refunds an order, sends an email, writes to a
+production table -- you want a person to see the call before it runs, without
+giving up the agent's freedom to look things up on its own.
+
+Mark those tools with pydantic-ai's approval API. ``AgentOperator`` and
+``@task.agent`` then pause the task in front of a marked call and ask on the
+**Required Actions** page:
+
+.. exampleinclude::
/../../ai/src/airflow/providers/common/ai/example_dags/example_agent_tool_approval.py
+ :language: python
+ :start-after: [START howto_agent_tool_approval]
+ :end-before: [END howto_agent_tool_approval]
+
+``lookup_order`` runs whenever the agent calls it. When the agent calls
+``refund_order``, the task stops and the reviewer sees the tool name and its
+arguments:
+
+.. code-block:: text
+
+ Approve tool call for task `handle_ticket`
+
+ The agent wants to run:
+
+ refund_order
+ {
+ "order_id": 7
+ }
+
+ [Approve] [Reject] reason: ______
+
+``.approval_required()`` works on any toolset, including ``SQLToolset`` and
+``MCPToolset``, and the function decides per call, so you can gate on the
+arguments as well as the tool name. For a single function tool,
+``Tool(refund_order, requires_approval=True)`` does the same.
+
+Under ``airflow dags test`` the task waits until someone answers from Required
+Actions in the UI of an api-server on the same metadata database;
+``airflow standalone`` gives you one.
+
+What each decision does
+-----------------------
+
+**Approve** runs the call and the agent carries on from where it stopped.
+
+**Reject** does not fail the task. The agent is told the call was denied --
with
+the reviewer's reason when one is given, such as "Refunds over $40 need a
support
+ticket number." -- and carries on without it, so it can explain the refusal,
+ask for what is missing, or take another route.
+
+While it waits, the task is in the ``awaiting_input`` state and holds no worker
+slot. When the model calls several tools in one step, the ones that need
+approval are decided together and the others run straight away, once.
+
+``usage_limits`` covers both sides of the pause, so a ``cost_limit`` is not
reset
+by it.
+
+One approval per task instance
+------------------------------
+
+A task instance asks for approval at most once per Dag run, and that includes
its
+retries and clears. Airflow keeps a single approval request per task instance,
+and a second request would show the reviewer the first one's tool call while
+asking about the new one. So when the agent asks again -- later in the same
run,
+or in a retry after the first request -- the task fails with
+``ToolApprovalAlreadyRequestedError`` and does not retry.
+
+Have the agent request the gated calls in one step, or give each irreversible
+action its own task. A task whose approved call ran and that then failed cannot
+ask again in the same Dag run; trigger a new run, and keep gated tools
+idempotent, since the approved call has already happened once.
+
+Timeouts and who decides
+------------------------
+
+``tool_approval_timeout`` bounds the pause; ``None`` (the default) waits for as
+long as it takes. What a timeout does is set by ``on_tool_approval_timeout``:
+
+- ``"fail"`` (default) fails the task.
+- ``"deny"`` rejects the pending calls, and the agent carries on without them.
+ The agent is told nobody answered in time, not that a person refused.
``"deny"``
+ needs a ``tool_approval_timeout``.
+
+There is no approve-on-timeout: a pause that approves itself when nobody
answers
+guards nothing.
+
+``tool_approval_assigned_users`` limits who may decide, with the same user list
+``require_approval`` on the LLM operators takes (see :doc:`approval_gates`).
+
+The arguments shown to the reviewer pass through Airflow's secrets masker
first,
+so a value under a key such as ``api_key`` or ``password``, or a value already
+registered as a secret, is shown masked.
+
+Requirements and limits
+-----------------------
+
+- Airflow 3.3 or later. On older versions a tool marked for approval fails the
+ task, as it did before.
+- Not together with ``durable=True``, ``enable_hitl_review=True``,
+ ``code_mode=True``, or a ``SandboxToolset``. Each assumes the run finishes
in one
+ go; a sandbox, for one, is destroyed when the run pauses. With any of them, a
+ marked tool fails the task.
+- Tools that hand work to an external system (pydantic-ai's ``CallDeferred``)
are
+ not supported; the task fails without retrying.
+- The conversation so far, including tool results, is kept in the task's
+ :doc:`task state store <apache-airflow:core-concepts/task-state-store>` while
+ the task waits. It is deleted when the task resumes, whether the resumed run
+ succeeds or fails, and a later try deletes one left behind by a try that
ended
+ while waiting. With a ``tool_approval_timeout``, it also expires a day after
the
+ timeout.
+- If a templated connection id (see :ref:`sql-toolset-templated-connection`)
renders
+ differently when the task resumes, the task fails rather than run an approved
+ call against a connection the reviewer did not see. The check compares
toolset
+ ids, which name the connection for ``SQLToolset``, ``MCPToolset`` and
+ ``HookToolset``; it does not notice an edit to the connection itself.
+- The resumed run gets its own pydantic-ai ``run_id``, ``<task-instance
id>-resumed``.
+ The ``run_id`` XCom holds that id, and the ``usage`` XCom covers both sides
of
+ the pause.
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_agent_tool_approval.py
b/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_agent_tool_approval.py
new file mode 100644
index 00000000000..4edcef5a9ca
--- /dev/null
+++
b/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_agent_tool_approval.py
@@ -0,0 +1,58 @@
+# 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.
+"""Example Dag: an agent that looks orders up freely but needs a person before
it refunds one."""
+
+from __future__ import annotations
+
+from datetime import timedelta
+
+from pydantic_ai.toolsets.function import FunctionToolset
+
+from airflow.providers.common.ai.operators.agent import AgentOperator
+from airflow.providers.common.compat.sdk import dag
+
+
+# [START howto_agent_tool_approval]
+def lookup_order(order_id: int) -> str:
+ """Look up an order."""
+ return f"order {order_id}: $42, delivered"
+
+
+def refund_order(order_id: int) -> str:
+ """Refund an order. Irreversible."""
+ return f"refunded order {order_id}"
+
+
+shop = FunctionToolset(tools=[lookup_order, refund_order]).approval_required(
+ lambda ctx, tool_def, args: tool_def.name == "refund_order"
+)
+
+
+@dag(tags=["example"])
+def example_agent_tool_approval():
+ AgentOperator(
+ task_id="handle_ticket",
+ llm_conn_id="pydanticai_default",
+ prompt="The customer says order 7 never arrived. Resolve it.",
+ toolsets=[shop],
+ tool_approval_timeout=timedelta(hours=4),
+ )
+
+
+# [END howto_agent_tool_approval]
+
+example_agent_tool_approval()
diff --git a/providers/common/ai/src/airflow/providers/common/ai/exceptions.py
b/providers/common/ai/src/airflow/providers/common/ai/exceptions.py
index 3a5a80178b8..38ddab555a5 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/exceptions.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/exceptions.py
@@ -16,13 +16,43 @@
# under the License.
from __future__ import annotations
-from airflow.providers.common.compat.sdk import AirflowException
+from airflow.providers.common.compat.sdk import AirflowException,
AirflowFailException
class HITLMaxIterationsError(AirflowException):
"""Raised when the HITL review loop exhausts max iterations without
approval or rejection."""
+class ToolApprovalError(RuntimeError):
+ """
+ Raised when a run paused for tool approval cannot resume safely.
+
+ The saved transcript, or the agent's rendered toolset ids, no longer match
what the
+ reviewer saw. A retry starts a fresh run.
+ """
+
+
+class ToolApprovalAlreadyRequestedError(AirflowFailException):
+ """
+ Raised when a task instance asks for a second tool approval.
+
+ Airflow keeps one approval request per task instance, across retries and
clears, and a
+ second request would show the reviewer the first one's details. The task
fails without
+ retrying, since a retry would ask again.
+ """
+
+
+class UnsupportedToolDeferralError(AirflowFailException):
+ """
+ Raised when an agent defers a tool call that ``AgentOperator`` cannot
resolve.
+
+ Either the tool hands its work to an external system, or it needs approval
where
+ approval is not available (before Airflow 3.3, or with ``durable``,
+ ``enable_hitl_review``, ``code_mode`` or a ``SandboxToolset``). A retry
would repeat
+ the same call, so the task fails without retrying.
+ """
+
+
class LLMFileAnalysisError(ValueError):
"""Base class for file-analysis validation errors."""
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/mixins/approval.py
b/providers/common/ai/src/airflow/providers/common/ai/mixins/approval.py
index 01b779d3553..49e02d2f3b5 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/mixins/approval.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/mixins/approval.py
@@ -19,6 +19,7 @@ from __future__ import annotations
import json
import logging
+from collections.abc import Iterable
from datetime import timedelta
from typing import TYPE_CHECKING, Any, ClassVar, Literal, Protocol
@@ -42,6 +43,35 @@ if TYPE_CHECKING:
from airflow.sdk.execution_time.hitl import HITLUser
+def normalize_assigned_users(value: Any, *, param: str) -> list[HITLUser]:
+ """
+ Return *value* as a list of ``{'id': str, 'name': str}`` users.
+
+ Accepts ``None``, a single user dict, or an iterable of them, and raises
``TypeError``
+ naming *param* for anything else, so a malformed list fails when the Dag
is parsed
+ rather than when the task first asks for a review.
+ """
+ users: list[Any]
+ if value is None:
+ users = []
+ elif isinstance(value, dict):
+ users = [value]
+ elif isinstance(value, str) or not isinstance(value, Iterable):
+ raise TypeError(
+ f"{param} must be a {{'id': str, 'name': str}} dict or an iterable
of them, got {value!r}"
+ )
+ else:
+ users = list(value)
+ for user in users:
+ if (
+ not isinstance(user, dict)
+ or not isinstance(user.get("id"), str)
+ or not isinstance(user.get("name"), str)
+ ):
+ raise TypeError(f"{param} entries must be {{'id': str, 'name':
str}} dicts, got {user!r}")
+ return users
+
+
class DeferForApprovalProtocol(Protocol):
"""Protocol for defer for approval mixin."""
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 b7b64d1a42a..08bdf7b7aa1 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
@@ -19,18 +19,28 @@
from __future__ import annotations
import copy
+import hashlib
import json
from collections.abc import Iterable, Sequence
from dataclasses import replace
from datetime import timedelta
from functools import cached_property
-from typing import TYPE_CHECKING, Any, ClassVar
+from typing import TYPE_CHECKING, Any, ClassVar, Literal, NoReturn
-from pydantic import BaseModel
+from pydantic import BaseModel, TypeAdapter
+from pydantic_ai import DeferredToolRequests, DeferredToolResults, ToolDenied
from pydantic_ai.capabilities import Toolset
+from pydantic_ai.messages import ModelMessagesTypeAdapter
from pydantic_ai.toolsets.abstract import AbstractToolset
+from pydantic_ai.usage import RunUsage
+from airflow.providers.common.ai.exceptions import (
+ ToolApprovalAlreadyRequestedError,
+ ToolApprovalError,
+ UnsupportedToolDeferralError,
+)
from airflow.providers.common.ai.hooks.pydantic_ai import PydanticAIHook
+from airflow.providers.common.ai.mixins.approval import LLMApprovalMixin,
normalize_assigned_users
from airflow.providers.common.ai.mixins.cancellable_run import
CancellableAgentRunMixin
from airflow.providers.common.ai.mixins.hitl_review import HITLReviewMixin
from airflow.providers.common.ai.observability import (
@@ -47,8 +57,16 @@ from airflow.providers.common.compat.sdk import (
BaseOperator,
BaseOperatorLink,
conf,
+ redact,
)
from airflow.providers.common.compat.version_compat import AIRFLOW_V_3_1_PLUS,
AIRFLOW_V_3_3_PLUS
+from airflow.providers.standard.exceptions import HITLTimeoutError,
HITLTriggerEventError
+
+if AIRFLOW_V_3_3_PLUS:
+ # Per-tool approval parks the task in AWAITING_INPUT, which older cores do
not have.
+ from airflow.sdk.exceptions import TaskAwaitingInput
+ from airflow.sdk.execution_time.context import NEVER_EXPIRE
+ from airflow.sdk.execution_time.hitl import upsert_hitl_detail
try:
# See LLMOperator: new enough cores register declared ``output_type``
classes
@@ -68,6 +86,17 @@ if TYPE_CHECKING:
from airflow.providers.common.ai.durable.step_counter import
DurableStepCounter
from airflow.providers.common.compat.sdk import TaskInstanceKey
from airflow.sdk import Context
+ from airflow.sdk.execution_time.context import TaskStateStoreAccessor
+ from airflow.sdk.execution_time.hitl import HITLUser
+
+# Task state store keys: the transcript of a run paused for tool approval, and
a marker that
+# this task instance has asked once. The store is keyed by Dag run, task and
map index, so the
+# marker survives retries and clears, as the task instance's single approval
request does.
+_TOOL_APPROVAL_TRANSCRIPT_KEY = "common_ai_tool_approval_transcript"
+_TOOL_APPROVAL_REQUESTED_KEY = "common_ai_tool_approval_requested"
+# How long the transcript outlives a timed pause, so a resume that runs late
still finds it.
+_TRANSCRIPT_RETENTION_MARGIN = timedelta(days=1)
+_RUN_USAGE_ADAPTER: TypeAdapter[RunUsage] = TypeAdapter(RunUsage)
class HITLReviewLink(BaseOperatorLink):
@@ -280,6 +309,31 @@ class AgentOperator(CancellableAgentRunMixin,
BaseOperator, HITLReviewMixin):
operator blocks until a terminal action).
:param hitl_poll_interval: Seconds between XCom polls
while waiting for a human response. Default ``10``.
+
+ **Per-tool approval** (Airflow 3.3+):
+
+ Mark the tools a human must approve with pydantic-ai's own API --
+ ``toolset.approval_required(...)``, or ``requires_approval=True`` on a
function
+ tool -- and the task pauses before running them. The pending calls, with
their
+ arguments, appear on the **Required Actions** page; the task waits in the
+ ``awaiting_input`` state without holding a worker slot. On **Approve** the
calls
+ run and the agent carries on. On **Reject** the agent is told the call was
denied
+ (with the reviewer's reason, when given) and carries on without it. A task
+ instance asks at most once per Dag run, across retries and clears; a second
+ request fails the task. ``usage_limits`` applies to both sides of the
pause.
+ Not available together with ``durable``, ``enable_hitl_review``,
``code_mode``,
+ or a ``SandboxToolset``; there, a tool that requires approval fails the
task as
+ before.
+
+ :param tool_approval_timeout: How long the pause waits for a decision.
+ ``None`` (default) waits indefinitely. Must be positive.
+ :param on_tool_approval_timeout: What a timed-out pause does: ``"fail"``
+ (default) fails the task, ``"deny"`` rejects the pending calls so the
agent
+ carries on without them, and needs a ``tool_approval_timeout``. There
is no
+ approve-on-timeout.
+ :param tool_approval_assigned_users: Users allowed to decide. ``None``
(default)
+ leaves it to anyone who can act on the task's Required Actions.
+
:param serialize_output: If ``True`` and ``output_type`` is a Pydantic
``BaseModel`` subclass, the model instance is dumped to a ``dict`` via
``model_dump()`` before being pushed to XCom. Default ``False`` --
@@ -328,6 +382,9 @@ class AgentOperator(CancellableAgentRunMixin, BaseOperator,
HITLReviewMixin):
hitl_timeout: timedelta | None = None,
hitl_poll_interval: float = 10.0,
serialize_output: bool = False,
+ tool_approval_timeout: timedelta | None = None,
+ on_tool_approval_timeout: Literal["fail", "deny"] = "fail",
+ tool_approval_assigned_users: HITLUser | Iterable[HITLUser] | None =
None,
**kwargs: Any,
) -> None:
super().__init__(**kwargs)
@@ -393,6 +450,20 @@ class AgentOperator(CancellableAgentRunMixin,
BaseOperator, HITLReviewMixin):
self.hitl_timeout = hitl_timeout
self.hitl_poll_interval = hitl_poll_interval
+ if on_tool_approval_timeout not in ("fail", "deny"):
+ raise ValueError(
+ f"on_tool_approval_timeout must be 'fail' or 'deny', got
{on_tool_approval_timeout!r}."
+ )
+ if tool_approval_timeout is not None and tool_approval_timeout <=
timedelta(0):
+ raise ValueError(f"tool_approval_timeout must be positive, got
{tool_approval_timeout!r}.")
+ if on_tool_approval_timeout == "deny" and tool_approval_timeout is
None:
+ raise ValueError("on_tool_approval_timeout='deny' needs a
tool_approval_timeout to fire.")
+ self.tool_approval_timeout = tool_approval_timeout
+ self.on_tool_approval_timeout = on_tool_approval_timeout
+ self.tool_approval_assigned_users = normalize_assigned_users(
+ tool_approval_assigned_users, param="tool_approval_assigned_users"
+ )
+
def _reject_sandbox_without_continuity(self, *, durable: bool,
enable_hitl_review: bool) -> None:
"""
Refuse a ``SandboxToolset`` under a feature that assumes the sandbox
outlives the run.
@@ -413,11 +484,7 @@ class AgentOperator(CancellableAgentRunMixin,
BaseOperator, HITLReviewMixin):
capabilities, since those are the compositions the documentation
recommends.
A toolset resolved per run from a callable cannot be inspected here.
"""
- candidates = list(self.toolsets or [])
- for capability in self.agent_params.get("capabilities") or ():
- if _is_concrete_toolset_capability(capability):
- candidates.append(capability.toolset)
- if find_toolset(candidates, SandboxToolset) is None:
+ if find_toolset(self._declared_toolsets(), SandboxToolset) is None:
return
flag = "durable=True" if durable else "enable_hitl_review=True"
why = (
@@ -535,11 +602,58 @@ class AgentOperator(CancellableAgentRunMixin,
BaseOperator, HITLReviewMixin):
if capabilities:
extra_kwargs["capabilities"] = capabilities
return self.llm_hook.create_agent(
- output_type=self.output_type,
+ output_type=self._agent_output_type(),
instructions=self.system_prompt,
**extra_kwargs,
)
+ def _supports_tool_approval(self) -> bool:
+ """
+ Whether a tool that requires approval pauses the task instead of
failing it.
+
+ Each excluded feature assumes the run finishes in one go: durable
replay counts
+ steps across a single run, HITL review and code mode wrap the run, and
a
+ sandbox is destroyed when the run ends, so its files would be gone on
resume.
+ """
+ if not AIRFLOW_V_3_3_PLUS or self.durable or self.enable_hitl_review
or self.code_mode:
+ return False
+ return find_toolset(self._declared_toolsets(), SandboxToolset) is None
+
+ def _agent_output_type(self) -> Any:
+ """
+ Return ``output_type``, plus ``DeferredToolRequests`` when tool
approval is supported.
+
+ pydantic-ai drops ``DeferredToolRequests`` from the output schema the
model sees;
+ it only lets the run end on a tool call awaiting approval.
``self.output_type`` is
+ left alone because it is part of the serialized Dag.
+ """
+ if not self._supports_tool_approval():
+ return self.output_type
+ declared = self.output_type if isinstance(self.output_type, (list,
tuple)) else [self.output_type]
+ return [*declared, DeferredToolRequests]
+
+ def _declared_toolsets(self) -> list[AbstractToolset[Any]]:
+ """Toolsets passed via ``toolsets=``, ``agent_params["toolsets"]`` and
concrete ``Toolset`` capabilities."""
+ candidates = [
+ toolset
+ for toolset in (*(self.toolsets or []),
*(self.agent_params.get("toolsets") or []))
+ if isinstance(toolset, AbstractToolset)
+ ]
+ for capability in self.agent_params.get("capabilities") or ():
+ if _is_concrete_toolset_capability(capability):
+ candidates.append(capability.toolset)
+ return candidates
+
+ def _toolset_ids(self) -> list[str]:
+ """Ids of every leaf toolset, which for SQL and MCP toolsets name the
connection."""
+ # Declared order, not sorted: two toolsets that swapped connections
must not compare equal.
+ return [
+ leaf.id
+ for toolset in self._declared_toolsets()
+ for leaf in iter_toolsets(toolset)
+ if leaf.id is not None
+ ]
+
def _build_durable_toolsets(
self, toolsets: list[AbstractToolset], storage:
DurableStorageProtocol, counter: DurableStepCounter
) -> list[AbstractToolset]:
@@ -623,6 +737,12 @@ class AgentOperator(CancellableAgentRunMixin,
BaseOperator, HITLReviewMixin):
# Coerced first so a bad rendered value fails before the expensive
setup below.
usage_limits = coerce_usage_limits(self.usage_limits)
+ # A try that paused and then ended some other way (marked failed while
waiting, a failed
+ # request) leaves its transcript behind; a fresh run never reads it.
``.get``: a context
+ # built by hand in a unit test has no task state store, and nothing to
clean up.
+ if self._supports_tool_approval() and (store :=
context.get("task_state_store")) is not None:
+ self._delete_approval_transcript(store)
+
self._durable_storage = None
self._durable_counter = None
@@ -664,7 +784,13 @@ class AgentOperator(CancellableAgentRunMixin,
BaseOperator, HITLReviewMixin):
else:
result = self.run_agent_sync(agent, self.prompt, **run_kwargs)
+ return self._complete_run(context, result)
+
+ def _complete_run(self, context: Context, result: Any) -> Any:
+ """Finish a run, or pause it when the agent is waiting on a tool call
to be approved."""
log_run_summary(self.log, result)
+ if isinstance(result.output, DeferredToolRequests):
+ self._pause_for_tool_approval(context, result)
self._emit_run_metadata(context, result)
if self._durable_counter is not None:
@@ -717,6 +843,191 @@ class AgentOperator(CancellableAgentRunMixin,
BaseOperator, HITLReviewMixin):
self._durable_storage.cleanup()
return output
+ def _pause_for_tool_approval(self, context: Context, result: Any) ->
NoReturn:
+ """
+ Park the task until a human approves or rejects the tool calls the
agent is waiting on.
+
+ The transcript goes to the task state store, not the continuation
kwargs: it holds tool
+ results such as query rows, which do not belong on the task instance
row. The
+ continuation carries its hash, so the resume refuses a transcript that
changed while
+ the task waited, and the rendered toolset ids, so it refuses to run
approved calls
+ against a connection id the reviewer did not see.
+
+ A task instance asks at most once. Airflow keeps one approval request
per task
+ instance, across retries and clears, and a repeat request keeps the
first one's
+ subject and body, so the reviewer would approve a new call while
reading the old one.
+ """
+ requests: DeferredToolRequests = result.output
+ if requests.calls:
+ raise UnsupportedToolDeferralError(
+ "The agent called tools that need external execution "
+ f"({', '.join(call.tool_name for call in requests.calls)});
AgentOperator only "
+ "supports tools that need approval."
+ )
+ pending_names = ", ".join(call.tool_name for call in
requests.approvals)
+ if not self._supports_tool_approval():
+ # DeferredToolRequests in a user-set output_type reaches here
where approval is off.
+ raise UnsupportedToolDeferralError(
+ f"The agent called tools that need approval ({pending_names}),
but tool approval "
+ "needs Airflow 3.3+ and is not available with durable,
enable_hitl_review, "
+ "code_mode or a SandboxToolset."
+ )
+ store = context["task_state_store"]
+ if store.get(_TOOL_APPROVAL_REQUESTED_KEY):
+ raise ToolApprovalAlreadyRequestedError(
+ f"The agent asked for a second tool approval
({pending_names}), but this task "
+ "instance already asked once in this Dag run. Airflow keeps
one approval request per "
+ "task instance and would show the reviewer the earlier
request's details, so the task "
+ "fails instead of pausing. Have the agent ask for the gated
calls in one step, give "
+ "each irreversible action its own task, or trigger a new Dag
run."
+ )
+ transcript =
ModelMessagesTypeAdapter.dump_json(result.all_messages()).decode()
+ retention = (
+ self.tool_approval_timeout + _TRANSCRIPT_RETENTION_MARGIN
+ if self.tool_approval_timeout is not None
+ else NEVER_EXPIRE
+ )
+ store.set(_TOOL_APPROVAL_TRANSCRIPT_KEY, transcript,
retention=retention)
+
+ # Tool arguments can carry credentials (an HTTP header, an MCP token);
mask them first.
+ pending = "\n\n".join(
+ f"**{call.tool_name}**\n\n```json\n"
+ f"{json.dumps(redact(call.args_as_dict()), indent=2,
default=str)}\n```"
+ for call in requests.approvals
+ )
+ upsert_hitl_detail(
+ ti_id=context["task_instance"].id,
+ options=[LLMApprovalMixin.APPROVE, LLMApprovalMixin.REJECT],
+ subject=f"Approve tool call for task `{self.task_id}`",
+ body=f"The agent wants to run:\n\n{pending}",
+ defaults=[LLMApprovalMixin.REJECT] if
self.on_tool_approval_timeout == "deny" else None,
+ multiple=False,
+ params={
+ "reason": {
+ # "null" in the type is what makes the field optional in
the review form (a
+ # plain "string" forces a reason before Approve can be
clicked); the empty
+ # default renders as an empty box, where some UI versions
show a null one
+ # as "[object Object]".
+ "value": "",
+ "description": "Sent to the agent when you reject the call
(optional).",
+ "schema": {"type": ["string", "null"]},
+ },
+ },
+ assigned_users=self.tool_approval_assigned_users,
+ )
+ # Only once the request exists: a failed request must not block the
retry from asking.
+ store.set(_TOOL_APPROVAL_REQUESTED_KEY, True, retention=NEVER_EXPIRE)
+ self.log.info("Waiting for approval of %s", pending_names)
+ raise TaskAwaitingInput(
+ method_name="resume_after_tool_approval",
+ kwargs={
+ "tool_call_ids": [call.tool_call_id for call in
requests.approvals],
+ "usage": _RUN_USAGE_ADAPTER.dump_python(result.usage,
mode="json"),
+ "transcript_sha256":
hashlib.sha256(transcript.encode()).hexdigest(),
+ "toolset_ids": self._toolset_ids(),
+ },
+ timeout=self.tool_approval_timeout,
+ )
+
+ def _delete_approval_transcript(self, store: TaskStateStoreAccessor) ->
None:
+ # Best-effort: the transcript holds tool results, but failing the task
over its cleanup
+ # would be worse, and the row goes with the Dag run anyway.
+ try:
+ store.delete(_TOOL_APPROVAL_TRANSCRIPT_KEY)
+ except Exception:
+ self.log.warning("Could not delete the tool approval transcript",
exc_info=True)
+
+ def resume_after_tool_approval(
+ self,
+ context: Context,
+ tool_call_ids: list[str],
+ usage: dict[str, Any],
+ transcript_sha256: str,
+ toolset_ids: list[str],
+ event: dict[str, Any],
+ ) -> Any:
+ """Continue a run paused by :meth:`_pause_for_tool_approval` with the
reviewer's decision."""
+ store = context["task_state_store"]
+ try:
+ return self._resume_after_tool_approval(
+ context, store, tool_call_ids, usage, transcript_sha256,
toolset_ids, event
+ )
+ finally:
+ # Whatever the outcome. A second pause fails closed before writing
a new transcript.
+ self._delete_approval_transcript(store)
+
+ def _resume_after_tool_approval(
+ self,
+ context: Context,
+ store: TaskStateStoreAccessor,
+ tool_call_ids: list[str],
+ usage: dict[str, Any],
+ transcript_sha256: str,
+ toolset_ids: list[str],
+ event: dict[str, Any],
+ ) -> Any:
+ if "error" in event:
+ if event.get("error_type") == "timeout":
+ raise HITLTimeoutError(f"Tool approval timed out:
{event['error']}")
+ raise HITLTriggerEventError(event)
+ if (current := self._toolset_ids()) != toolset_ids:
+ raise ToolApprovalError(
+ f"The agent's toolsets changed while it waited for approval
(paused with {toolset_ids}, "
+ f"resumed with {current}), so the approved calls would reach a
connection the reviewer "
+ "did not see."
+ )
+ transcript = store.get(_TOOL_APPROVAL_TRANSCRIPT_KEY)
+ if not isinstance(transcript, str) or (
+ hashlib.sha256(transcript.encode()).hexdigest() !=
transcript_sha256
+ ):
+ raise ToolApprovalError(
+ "The transcript saved when the task paused for tool approval
is missing or was modified."
+ )
+
+ approval: bool | ToolDenied
+ if event.get("timedout"):
+ # on_tool_approval_timeout="deny": nobody refused, so do not tell
the agent a person did.
+ approval = ToolDenied(
+ "No reviewer answered within the approval timeout, so this
call was not run."
+ )
+ self.log.info(
+ "Tool calls denied: nobody answered within
tool_approval_timeout=%s",
+ self.tool_approval_timeout,
+ )
+ elif LLMApprovalMixin.APPROVE in event["chosen_options"]:
+ approval = True
+ self.log.info("Tool calls approved by %s",
LLMApprovalMixin._describe_responder(event))
+ else:
+ reason = (event.get("params_input") or {}).get("reason")
+ # Only a typed reason reaches the agent; an untouched field can
come back as None, "",
+ # or, from some UI versions, the whole parameter spec.
+ if not isinstance(reason, str) or not reason.strip():
+ reason = "A reviewer denied this tool call."
+ approval = ToolDenied(reason)
+ self.log.info("Tool calls rejected by %s",
LLMApprovalMixin._describe_responder(event))
+
+ agent = self._build_agent()
+ ti = context["task_instance"]
+ self._run_identity_attrs = build_run_identity_attributes(ti)
+ stamp_identity_on_agent_spans(agent, self._run_identity_attrs)
+ # The full transcript, not a trimmed one: pydantic-ai reads its last
request to skip
+ # the calls that already ran in the paused step. No new prompt: it
would land after
+ # the tool results as a second user turn.
+ result = self.run_agent_sync(
+ agent,
+ None,
+ message_history=ModelMessagesTypeAdapter.validate_json(transcript),
+ deferred_tool_results=DeferredToolResults(
+ approvals={tool_call_id: approval for tool_call_id in
tool_call_ids}
+ ),
+ usage=_RUN_USAGE_ADAPTER.validate_python(usage),
+ usage_limits=coerce_usage_limits(self.usage_limits),
+ # pydantic-ai refuses a run_id already in the history; the
task-instance id stays
+ # the prefix, so the resumed run still joins back to the task.
+ run_id=f"{ti.id}-resumed",
+ )
+ return self._complete_run(context, result)
+
def _resolve_message_history(self) -> list[ModelMessage] | None:
"""
Deserialize :attr:`message_history` into a list of pydantic-ai
messages.
@@ -731,19 +1042,12 @@ class AgentOperator(CancellableAgentRunMixin,
BaseOperator, HITLReviewMixin):
if isinstance(raw, (str, bytes)) and not raw.strip():
# A template that renders to empty (no prior XCom) starts a fresh
session.
return []
- # pydantic-ai is imported lazily here to match this module's pattern of
- # keeping pydantic-ai out of DAG-parse-time imports.
- from pydantic_ai.messages import ModelMessagesTypeAdapter
-
if isinstance(raw, (str, bytes)):
return ModelMessagesTypeAdapter.validate_json(raw)
return ModelMessagesTypeAdapter.validate_python(raw)
def _emit_message_history(self, context: Context, result: Any) -> None:
"""Push the full post-run transcript to XCom for the next turn to
resume."""
- # Lazy import: see _resolve_message_history.
- from pydantic_ai.messages import ModelMessagesTypeAdapter
-
transcript =
ModelMessagesTypeAdapter.dump_json(result.all_messages()).decode()
context["task_instance"].xcom_push(key="message_history",
value=transcript)
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/operators/llm.py
b/providers/common/ai/src/airflow/providers/common/ai/operators/llm.py
index df3d0e8d515..a06f27f04a1 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/operators/llm.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/operators/llm.py
@@ -26,7 +26,7 @@ from typing import TYPE_CHECKING, Any, ClassVar, Literal
from pydantic import BaseModel
from airflow.providers.common.ai.hooks.pydantic_ai import PydanticAIHook
-from airflow.providers.common.ai.mixins.approval import LLMApprovalMixin
+from airflow.providers.common.ai.mixins.approval import LLMApprovalMixin,
normalize_assigned_users
from airflow.providers.common.ai.mixins.cancellable_run import
CancellableAgentRunMixin
from airflow.providers.common.ai.policies.decision import DecisionPolicy
from airflow.providers.common.ai.utils.decision import (
@@ -272,27 +272,7 @@ class LLMOperator(CancellableAgentRunMixin, BaseOperator,
LLMApprovalMixin):
for notifier in self.approval_notifiers:
if not isinstance(notifier, BaseNotifier):
raise TypeError(f"approval_notifiers must contain BaseNotifier
instances, got {notifier!r}")
- assigned_users: list[Any]
- if approval_assigned_users is None:
- assigned_users = []
- elif isinstance(approval_assigned_users, dict):
- assigned_users = [approval_assigned_users]
- elif isinstance(approval_assigned_users, str) or not
isinstance(approval_assigned_users, Iterable):
- raise TypeError(
- "approval_assigned_users must be a {'id': str, 'name': str}
dict or an iterable of them, "
- f"got {approval_assigned_users!r}"
- )
- else:
- assigned_users = list(approval_assigned_users)
- for user in assigned_users:
- if (
- not isinstance(user, dict)
- or not isinstance(user.get("id"), str)
- or not isinstance(user.get("name"), str)
- ):
- raise TypeError(
- f"approval_assigned_users entries must be {{'id': str,
'name': str}} dicts, got {user!r}"
- )
+ assigned_users = normalize_assigned_users(approval_assigned_users,
param="approval_assigned_users")
if assigned_users and not AIRFLOW_V_3_1_PLUS:
raise
AirflowOptionalProviderFeatureException("approval_assigned_users needs Airflow
3.1+.")
self.approval_assigned_users: list[HITLUser] = assigned_users
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/toolsets/logging.py
b/providers/common/ai/src/airflow/providers/common/ai/toolsets/logging.py
index 0ce234f7550..3c325e4d66e 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/toolsets/logging.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/toolsets/logging.py
@@ -24,6 +24,7 @@ import time
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any
+from pydantic_ai.exceptions import ApprovalRequired
from pydantic_ai.toolsets.wrapper import WrapperToolset
if TYPE_CHECKING:
@@ -55,6 +56,11 @@ class LoggingToolset(WrapperToolset[Any]):
self.logger.info("Tool %s returned in %.2fs", name, elapsed)
self.logger.info("::endgroup::")
return result
+ except ApprovalRequired:
+ # Not a failure: the run pauses here until a person approves or
rejects the call.
+ self.logger.info("Tool %s is waiting to be approved", name)
+ self.logger.info("::endgroup::")
+ raise
except Exception:
elapsed = time.monotonic() - start
self.logger.exception("Tool %s failed after %.2fs", name, elapsed)
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 f6d65279ef2..6b47a176af0 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
@@ -24,7 +24,7 @@ from unittest.mock import ANY, MagicMock, patch
import pytest
from pydantic import BaseModel
-from pydantic_ai import Agent
+from pydantic_ai import Agent, DeferredToolRequests
from pydantic_ai.capabilities import Toolset
from pydantic_ai.exceptions import UsageLimitExceeded
from pydantic_ai.messages import (
@@ -636,8 +636,10 @@ class TestAgentOperatorExecute:
mock_hook_cls.get_hook.assert_called_once_with(
"my_llm", hook_params={"model_id": None, "fallback_conn_ids": None}
)
+ # On 3.3+ the agent may also end on a tool call awaiting approval.
+ expected_output_type = [str, DeferredToolRequests] if
AIRFLOW_V_3_3_PLUS else str
mock_hook_cls.get_hook.return_value.create_agent.assert_called_once_with(
- output_type=str, instructions="You are helpful."
+ output_type=expected_output_type, instructions="You are helpful."
)
mock_agent.run_sync.assert_called_once_with(
"What is the answer?", usage_limits=None, run_id="ti-1",
cancellation_token=ANY
diff --git
a/providers/common/ai/tests/unit/common/ai/operators/test_agent_tool_approval.py
b/providers/common/ai/tests/unit/common/ai/operators/test_agent_tool_approval.py
new file mode 100644
index 00000000000..45866e4a391
--- /dev/null
+++
b/providers/common/ai/tests/unit/common/ai/operators/test_agent_tool_approval.py
@@ -0,0 +1,556 @@
+# 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.
+"""Per-tool approval: AgentOperator pauses before a tool that requires
approval and resumes."""
+
+from __future__ import annotations
+
+from datetime import timedelta
+from typing import Any
+from unittest.mock import MagicMock, patch
+
+import pytest
+from pydantic_ai import Agent, CallDeferred, DeferredToolRequests, Tool
+from pydantic_ai.exceptions import UsageLimitExceeded
+from pydantic_ai.messages import ModelResponse, TextPart, ToolCallPart,
ToolReturnPart
+from pydantic_ai.models.function import AgentInfo, FunctionModel
+from pydantic_ai.toolsets.function import FunctionToolset
+
+from airflow.providers.common.ai.exceptions import (
+ ToolApprovalAlreadyRequestedError,
+ ToolApprovalError,
+ UnsupportedToolDeferralError,
+)
+from airflow.providers.common.ai.operators.agent import (
+ _TOOL_APPROVAL_REQUESTED_KEY,
+ _TOOL_APPROVAL_TRANSCRIPT_KEY,
+ AgentOperator,
+)
+from airflow.providers.common.ai.sandbox.base import SandboxBackend
+from airflow.providers.common.ai.toolsets.sandbox import SandboxToolset
+from airflow.providers.common.ai.toolsets.sql import SQLToolset
+from airflow.providers.standard.exceptions import HITLTimeoutError
+
+from tests_common.test_utils.version_compat import AIRFLOW_V_3_3_PLUS
+
+pytestmark = pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="Per-tool
approval needs Airflow 3.3+.")
+
+if AIRFLOW_V_3_3_PLUS:
+ from airflow.sdk.exceptions import TaskAwaitingInput
+
+UPSERT = "airflow.providers.common.ai.operators.agent.upsert_hitl_detail"
+
+
+class _FakeTaskStateStore:
+ """Dict-backed stand-in for the task state store accessor."""
+
+ def __init__(self):
+ self.data: dict[str, Any] = {}
+
+ def get(self, key, default=None):
+ return self.data.get(key, default)
+
+ def set(self, key, value, *, retention=None):
+ self.data[key] = value
+
+ def delete(self, key):
+ self.data.pop(key, None)
+
+
+class _Shop:
+ """Two tools: ``lookup`` runs freely, ``refund`` needs a human."""
+
+ def __init__(self):
+ self.calls: list[tuple[str, int]] = []
+
+ def toolset(self):
+ def lookup(order_id: int) -> str:
+ self.calls.append(("lookup", order_id))
+ return f"order {order_id} costs $10"
+
+ def refund(order_id: int) -> str:
+ self.calls.append(("refund", order_id))
+ return f"refunded order {order_id}"
+
+ return FunctionToolset(tools=[lookup, refund]).approval_required(
+ lambda ctx, tool_def, args: tool_def.name == "refund"
+ )
+
+
+def _returns(messages) -> list[str]:
+ return [str(p.content) for m in messages for p in m.parts if isinstance(p,
ToolReturnPart)]
+
+
+def _lookup_and_refund(messages, info: AgentInfo) -> ModelResponse:
+ """Call both tools in parallel, then report every tool result."""
+ if not _returns(messages):
+ return ModelResponse(
+ parts=[
+ ToolCallPart("lookup", {"order_id": 1},
tool_call_id="c-lookup"),
+ ToolCallPart("refund", {"order_id": 1},
tool_call_id="c-refund"),
+ ]
+ )
+ return ModelResponse(parts=[TextPart(" | ".join(_returns(messages)))])
+
+
+def _two_refunds(messages, info: AgentInfo) -> ModelResponse:
+ """Refund order 1, then order 2, one call per step: two separate
approvals."""
+ done = _returns(messages)
+ if len(done) < 2:
+ n = len(done) + 1
+ return ModelResponse(parts=[ToolCallPart("refund", {"order_id": n},
tool_call_id=f"c-{n}")])
+ return ModelResponse(parts=[TextPart(" | ".join(done))])
+
+
+def _operator(model_fn, toolsets, **kwargs) -> AgentOperator:
+ op = AgentOperator(task_id="t", prompt="Refund order 1",
llm_conn_id="llm", toolsets=toolsets, **kwargs)
+ hook = MagicMock(spec=["create_agent"])
+ hook.create_agent.side_effect = lambda **kw:
Agent(FunctionModel(model_fn), **kw)
+ op.llm_hook = hook
+ return op
+
+
+def _context(store: _FakeTaskStateStore, *, try_number: int = 1) -> Any: # a
Context stand-in
+ ti = MagicMock(spec=["id", "dag_id", "task_id", "run_id", "map_index",
"try_number", "xcom_push"])
+ ti.configure_mock(id="ti-1", dag_id="d", task_id="t", run_id="r",
map_index=-1, try_number=try_number)
+ return {"task_instance": ti, "task_state_store": store}
+
+
+APPROVE = {"chosen_options": ["Approve"], "params_input": {},
"responded_by_user": {"name": "alice"}}
+
+
+def _reject(reason: str = "") -> dict[str, Any]:
+ return {
+ "chosen_options": ["Reject"],
+ "params_input": {"reason": reason},
+ "responded_by_user": {"name": "bob"},
+ }
+
+
+def _pause(op: AgentOperator, ctx: Any) -> TaskAwaitingInput:
+ with patch(UPSERT, autospec=True):
+ with pytest.raises(TaskAwaitingInput) as exc:
+ op.execute(ctx)
+ return exc.value
+
+
+class TestPause:
+ def test_pauses_before_the_gated_tool_and_runs_the_ungated_one(self):
+ shop, store = _Shop(), _FakeTaskStateStore()
+ op = _operator(
+ _lookup_and_refund,
+ [shop.toolset()],
+ tool_approval_timeout=timedelta(hours=1),
+ tool_approval_assigned_users={"id": "u1", "name": "alice"},
+ )
+
+ with patch(UPSERT, autospec=True) as upsert:
+ with pytest.raises(TaskAwaitingInput) as exc:
+ op.execute(_context(store))
+
+ assert shop.calls == [("lookup", 1)]
+ assert exc.value.method_name == "resume_after_tool_approval"
+ assert exc.value.kwargs["tool_call_ids"] == ["c-refund"]
+ assert _TOOL_APPROVAL_TRANSCRIPT_KEY in store.data
+ body = upsert.call_args.kwargs["body"]
+ assert "**refund**" in body
+ assert '"order_id": 1' in body
+ assert "lookup" not in body
+ assert upsert.call_args.kwargs["options"] == ["Approve", "Reject"]
+ assert upsert.call_args.kwargs["defaults"] is None
+ assert upsert.call_args.kwargs["subject"] == "Approve tool call for
task `t`"
+ assert upsert.call_args.kwargs["multiple"] is False
+ # Optional in the review form: without "null" the UI requires a reason
even to approve.
+ assert upsert.call_args.kwargs["params"]["reason"]["schema"] ==
{"type": ["string", "null"]}
+ assert upsert.call_args.kwargs["assigned_users"] == [{"id": "u1",
"name": "alice"}]
+ assert exc.value.timeout == timedelta(hours=1)
+ assert store.data[_TOOL_APPROVAL_REQUESTED_KEY] is True
+
+ def test_a_function_tool_marked_requires_approval_pauses_too(self):
+ refunds = []
+
+ def refund(order_id: int) -> str:
+ refunds.append(order_id)
+ return "refunded"
+
+ def model(messages, info):
+ if _returns(messages):
+ return ModelResponse(parts=[TextPart("done")])
+ return ModelResponse(parts=[ToolCallPart("refund", {"order_id":
3})])
+
+ toolset = FunctionToolset(tools=[Tool(refund, requires_approval=True)])
+ with patch(UPSERT, autospec=True):
+ with pytest.raises(TaskAwaitingInput):
+ _operator(model,
[toolset]).execute(_context(_FakeTaskStateStore()))
+
+ assert refunds == []
+
+ @pytest.mark.enable_redact
+ def test_tool_arguments_are_masked_in_the_review_body(self):
+ def call_api(url: str, api_key: str) -> str:
+ return "ok"
+
+ toolset = FunctionToolset(tools=[call_api]).approval_required()
+
+ def model(messages, info):
+ return ModelResponse(parts=[ToolCallPart("call_api", {"url":
"https://x", "api_key": "s3cr3t"})])
+
+ with patch(UPSERT, autospec=True) as upsert:
+ with pytest.raises(TaskAwaitingInput):
+ _operator(model,
[toolset]).execute(_context(_FakeTaskStateStore()))
+
+ assert "s3cr3t" not in upsert.call_args.kwargs["body"]
+
+ def test_a_later_try_fails_closed_once_an_approval_was_requested(self):
+ """Core keeps one approval request per task instance across retries
and clears, with the
+ first request's subject and body, so a retry's request would show the
earlier call."""
+ shop, store = _Shop(), _FakeTaskStateStore()
+ _pause(_operator(_lookup_and_refund, [shop.toolset()]),
_context(store))
+
+ with patch(UPSERT, autospec=True) as upsert:
+ with pytest.raises(ToolApprovalAlreadyRequestedError,
match="already asked once"):
+ _operator(_lookup_and_refund,
[shop.toolset()]).execute(_context(store, try_number=2))
+
+ upsert.assert_not_called()
+
+ def test_a_failed_request_does_not_block_the_retry(self):
+ shop, store = _Shop(), _FakeTaskStateStore()
+ with patch(UPSERT, autospec=True, side_effect=RuntimeError("api
down")):
+ with pytest.raises(RuntimeError, match="api down"):
+ _operator(_lookup_and_refund,
[shop.toolset()]).execute(_context(store))
+
+ assert _TOOL_APPROVAL_REQUESTED_KEY not in store.data
+
+ def test_a_fresh_run_deletes_a_stale_transcript(self):
+ store = _FakeTaskStateStore()
+ store.data[_TOOL_APPROVAL_TRANSCRIPT_KEY] = "left over from a try that
ended while waiting"
+
+ def answer(messages, info):
+ return ModelResponse(parts=[TextPart("done")])
+
+ assert _operator(answer, [_Shop().toolset()]).execute(_context(store))
== "done"
+ assert _TOOL_APPROVAL_TRANSCRIPT_KEY not in store.data
+
+ def test_deny_on_timeout_sets_reject_as_the_timeout_default(self):
+ with patch(UPSERT, autospec=True) as upsert:
+ with pytest.raises(TaskAwaitingInput):
+ _operator(
+ _lookup_and_refund,
+ [_Shop().toolset()],
+ on_tool_approval_timeout="deny",
+ tool_approval_timeout=timedelta(hours=1),
+ ).execute(_context(_FakeTaskStateStore()))
+
+ assert upsert.call_args.kwargs["defaults"] == ["Reject"]
+
+ def
test_a_user_set_deferred_output_type_fails_where_approval_is_unsupported(self):
+ """Adding DeferredToolRequests by hand must not open a pause that
durable, HITL review,
+ code mode or a sandbox cannot survive."""
+ op = _operator(
+ _lookup_and_refund,
+ [_Shop().toolset()],
+ output_type=[str, DeferredToolRequests],
+ enable_hitl_review=True,
+ )
+
+ with patch(UPSERT, autospec=True) as upsert:
+ with pytest.raises(UnsupportedToolDeferralError, match="not
available with durable"):
+ op.execute(_context(_FakeTaskStateStore()))
+
+ upsert.assert_not_called()
+
+ def test_external_execution_calls_are_refused(self):
+ def slow_job() -> str:
+ raise CallDeferred
+
+ def model(messages, info):
+ return ModelResponse(parts=[ToolCallPart("slow_job", {})])
+
+ with patch(UPSERT, autospec=True):
+ with pytest.raises(UnsupportedToolDeferralError, match="need
external execution"):
+ _operator(model,
[FunctionToolset(tools=[slow_job])]).execute(_context(_FakeTaskStateStore()))
+
+
+class TestResume:
+ def test_approve_runs_the_call_once_and_finishes(self):
+ shop, store = _Shop(), _FakeTaskStateStore()
+ paused = _pause(_operator(_lookup_and_refund, [shop.toolset()]),
_context(store))
+
+ # A resume is a fresh process: a new operator, same task state store.
+ output = _operator(_lookup_and_refund,
[shop.toolset()]).resume_after_tool_approval(
+ _context(store), **paused.kwargs, event=APPROVE
+ )
+
+ assert shop.calls == [("lookup", 1), ("refund", 1)]
+ assert "refunded order 1" in output
+ assert "order 1 costs $10" in output
+ assert _TOOL_APPROVAL_TRANSCRIPT_KEY not in store.data
+
+ def test_the_resumed_run_can_be_cancelled_by_on_kill(self):
+ shop, store = _Shop(), _FakeTaskStateStore()
+ paused = _pause(_operator(_lookup_and_refund, [shop.toolset()]),
_context(store))
+ tokens = []
+
+ def recording_model(messages, info: AgentInfo) -> ModelResponse:
+ tokens.append(resumed._cancellation_token)
+ return _lookup_and_refund(messages, info)
+
+ resumed = _operator(recording_model, [shop.toolset()])
+ resumed.resume_after_tool_approval(_context(store), **paused.kwargs,
event=APPROVE)
+
+ assert tokens
+ assert all(token is not None for token in tokens)
+ assert resumed._cancellation_token is None
+
+ def test_reject_tells_the_agent_why_and_skips_the_call(self):
+ shop, store = _Shop(), _FakeTaskStateStore()
+ paused = _pause(_operator(_lookup_and_refund, [shop.toolset()]),
_context(store))
+
+ output = _operator(_lookup_and_refund,
[shop.toolset()]).resume_after_tool_approval(
+ _context(store), **paused.kwargs, event=_reject("refunds need a
ticket")
+ )
+
+ assert shop.calls == [("lookup", 1)]
+ assert "refunds need a ticket" in output
+
+ @pytest.mark.parametrize(
+ "untouched",
+ [
+ pytest.param("", id="empty"),
+ pytest.param(None, id="null"),
+ pytest.param({"value": None, "schema": {"type": ["string",
"null"]}}, id="param-spec"),
+ ],
+ )
+ def test_reject_without_a_typed_reason_sends_a_default_message(self,
untouched):
+ """An untouched reason field comes back as "", None, or (from some UI
versions) the spec."""
+ shop, store = _Shop(), _FakeTaskStateStore()
+ paused = _pause(_operator(_lookup_and_refund, [shop.toolset()]),
_context(store))
+ event = {
+ "chosen_options": ["Reject"],
+ "params_input": {"reason": untouched},
+ "responded_by_user": None,
+ }
+
+ output = _operator(_lookup_and_refund,
[shop.toolset()]).resume_after_tool_approval(
+ _context(store), **paused.kwargs, event=event
+ )
+
+ assert "A reviewer denied this tool call." in output
+ assert "schema" not in output
+
+ def test_timeout_fails_the_task_by_default(self):
+ shop, store = _Shop(), _FakeTaskStateStore()
+ paused = _pause(_operator(_lookup_and_refund, [shop.toolset()]),
_context(store))
+
+ with pytest.raises(HITLTimeoutError):
+ _operator(_lookup_and_refund,
[shop.toolset()]).resume_after_tool_approval(
+ _context(store), **paused.kwargs, event={"error": "expired",
"error_type": "timeout"}
+ )
+ assert shop.calls == [("lookup", 1)]
+
+ def test_a_timed_out_deny_does_not_claim_a_reviewer_refused(self):
+ shop, store = _Shop(), _FakeTaskStateStore()
+ paused = _pause(_operator(_lookup_and_refund, [shop.toolset()]),
_context(store))
+ timed_out = {
+ "chosen_options": ["Reject"],
+ "params_input": {},
+ "responded_by_user": None,
+ "timedout": True,
+ }
+
+ output = _operator(_lookup_and_refund,
[shop.toolset()]).resume_after_tool_approval(
+ _context(store), **paused.kwargs, event=timed_out
+ )
+
+ assert "No reviewer answered within the approval timeout" in output
+ assert shop.calls == [("lookup", 1)]
+
+ def test_a_second_approval_in_the_same_try_fails_closed(self):
+ """Airflow keeps one approval request per task instance and a second
one would show the
+ first one's details, so the reviewer would approve refund 2 while
reading refund 1."""
+ shop, store = _Shop(), _FakeTaskStateStore()
+ first = _pause(_operator(_two_refunds, [shop.toolset()]),
_context(store))
+
+ with patch(UPSERT, autospec=True) as upsert:
+ with pytest.raises(ToolApprovalAlreadyRequestedError,
match="already asked once"):
+ _operator(_two_refunds,
[shop.toolset()]).resume_after_tool_approval(
+ _context(store), **first.kwargs, event=APPROVE
+ )
+
+ assert shop.calls == [("refund", 1)]
+ upsert.assert_not_called()
+ assert _TOOL_APPROVAL_TRANSCRIPT_KEY not in store.data
+
+ @pytest.mark.parametrize(
+ ("kwargs_override", "event", "error"),
+ [
+ pytest.param({}, {"error": "expired", "error_type": "timeout"},
HITLTimeoutError, id="timeout"),
+ pytest.param({"toolset_ids": ["sql-other"]}, APPROVE,
ToolApprovalError, id="toolsets-changed"),
+ ],
+ )
+ def test_the_transcript_is_deleted_when_the_resume_fails(self,
kwargs_override, event, error):
+ """The transcript holds tool results, so it must not outlive a failed
resume."""
+ shop, store = _Shop(), _FakeTaskStateStore()
+ paused = _pause(_operator(_lookup_and_refund, [shop.toolset()]),
_context(store))
+
+ with pytest.raises(error):
+ _operator(_lookup_and_refund,
[shop.toolset()]).resume_after_tool_approval(
+ _context(store), **{**paused.kwargs, **kwargs_override},
event=event
+ )
+
+ assert _TOOL_APPROVAL_TRANSCRIPT_KEY not in store.data
+
+ def test_usage_limits_span_the_pause(self):
+ """Each side of the pause makes one model request; a limit of one must
stop the second."""
+ shop, store = _Shop(), _FakeTaskStateStore()
+ limits = {"request_limit": 1}
+ paused = _pause(_operator(_lookup_and_refund, [shop.toolset()],
usage_limits=limits), _context(store))
+
+ with pytest.raises(UsageLimitExceeded):
+ _operator(_lookup_and_refund, [shop.toolset()],
usage_limits=limits).resume_after_tool_approval(
+ _context(store), **paused.kwargs, event=APPROVE
+ )
+
+ def test_changed_toolsets_refuse_to_run_the_approved_call(self):
+ shop, store = _Shop(), _FakeTaskStateStore()
+ paused = _pause(_operator(_lookup_and_refund, [shop.toolset()]),
_context(store))
+ kwargs = {**paused.kwargs, "toolset_ids": ["sql-tenant_acme"]}
+
+ with pytest.raises(ToolApprovalError, match="toolsets changed"):
+ _operator(_lookup_and_refund,
[shop.toolset()]).resume_after_tool_approval(
+ _context(store), **kwargs, event=APPROVE
+ )
+ assert shop.calls == [("lookup", 1)]
+
+ def
test_a_rendered_connection_that_changed_refuses_to_run_the_approved_call(self):
+ shop, store = _Shop(), _FakeTaskStateStore()
+
+ def operator():
+ toolsets = [shop.toolset(), SQLToolset(db_conn_id="tenant_{{
params.customer }}")]
+ return _operator(_lookup_and_refund, toolsets)
+
+ before = operator()
+ before.render_template_fields({"params": {"customer": "acme"}})
+ paused = _pause(before, _context(store))
+ after = operator()
+ after.render_template_fields({"params": {"customer": "globex"}})
+
+ assert paused.kwargs["toolset_ids"] == ["sql-tenant_acme"]
+ with pytest.raises(ToolApprovalError, match="toolsets changed"):
+ after.resume_after_tool_approval(_context(store), **paused.kwargs,
event=APPROVE)
+ assert shop.calls == [("lookup", 1)]
+
+ def test_modified_transcript_is_refused(self):
+ shop, store = _Shop(), _FakeTaskStateStore()
+ paused = _pause(_operator(_lookup_and_refund, [shop.toolset()]),
_context(store))
+ store.data[_TOOL_APPROVAL_TRANSCRIPT_KEY] += " "
+
+ with pytest.raises(ToolApprovalError, match="missing or was modified"):
+ _operator(_lookup_and_refund,
[shop.toolset()]).resume_after_tool_approval(
+ _context(store), **paused.kwargs, event=APPROVE
+ )
+
+ def test_message_history_output_holds_the_whole_run(self):
+ shop, store = _Shop(), _FakeTaskStateStore()
+ paused = _pause(_operator(_lookup_and_refund, [shop.toolset()],
message_history=[]), _context(store))
+ ctx = _context(store)
+
+ _operator(_lookup_and_refund, [shop.toolset()],
message_history=[]).resume_after_tool_approval(
+ ctx, **paused.kwargs, event=APPROVE
+ )
+
+ pushed = {c.kwargs["key"]: c.kwargs["value"] for c in
ctx["task_instance"].xcom_push.call_args_list}
+ assert "order 1 costs $10" in pushed["message_history"]
+ assert "refunded order 1" in pushed["message_history"]
+
+
+class _NoopBackend(SandboxBackend):
+ """A backend that is never reached: these tests stop before any run."""
+
+ name = "noop"
+
+ def create(self, *, spec=None):
+ raise NotImplementedError
+
+ def run_command(self, sandbox, command, *, timeout, max_output_bytes):
+ raise NotImplementedError
+
+ def destroy(self, sandbox):
+ pass
+
+
+class TestWhenApprovalApplies:
+ def
test_output_type_gains_deferred_requests_without_changing_the_attribute(self):
+ op = AgentOperator(task_id="t", prompt="p", llm_conn_id="llm")
+
+ assert op._agent_output_type() == [str, DeferredToolRequests]
+ assert op.output_type is str
+
+ def test_a_list_output_type_is_extended_not_nested(self):
+ op = AgentOperator(task_id="t", prompt="p", llm_conn_id="llm",
output_type=[str, int])
+
+ assert op._agent_output_type() == [str, int, DeferredToolRequests]
+
+ @pytest.mark.parametrize(
+ "kwargs",
+ [
+ pytest.param({"durable": True}, id="durable"),
+ pytest.param({"code_mode": True}, id="code_mode"),
+ pytest.param({"enable_hitl_review": True}, id="hitl_review"),
+ pytest.param({"toolsets": [SandboxToolset(_NoopBackend())]},
id="sandbox"),
+ pytest.param(
+ {"agent_params": {"toolsets":
[SandboxToolset(_NoopBackend())]}}, id="sandbox-in-agent-params"
+ ),
+ ],
+ )
+ def
test_features_that_assume_one_uninterrupted_run_keep_the_output_type(self,
kwargs):
+ op = AgentOperator(task_id="t", prompt="p", llm_conn_id="llm",
**kwargs)
+
+ assert op._agent_output_type() is str
+
+ def test_on_tool_approval_timeout_rejects_unknown_values(self):
+ with pytest.raises(ValueError, match="on_tool_approval_timeout"):
+ AgentOperator(task_id="t", prompt="p", llm_conn_id="llm",
on_tool_approval_timeout="approve")
+
+ @pytest.mark.parametrize("timeout", [timedelta(0), timedelta(seconds=-1)])
+ def test_tool_approval_timeout_must_be_positive(self, timeout):
+ with pytest.raises(ValueError, match="must be positive"):
+ AgentOperator(task_id="t", prompt="p", llm_conn_id="llm",
tool_approval_timeout=timeout)
+
+ def test_deny_needs_a_timeout_to_fire(self):
+ with pytest.raises(ValueError, match="needs a tool_approval_timeout"):
+ AgentOperator(task_id="t", prompt="p", llm_conn_id="llm",
on_tool_approval_timeout="deny")
+
+ def test_a_single_assigned_user_is_accepted(self):
+ op = AgentOperator(
+ task_id="t", prompt="p", llm_conn_id="llm",
tool_approval_assigned_users={"id": "u1", "name": "a"}
+ )
+
+ assert op.tool_approval_assigned_users == [{"id": "u1", "name": "a"}]
+
+ def test_malformed_assigned_users_fail_at_parse_time(self):
+ with pytest.raises(TypeError, match="tool_approval_assigned_users
entries must be"):
+ AgentOperator(task_id="t", prompt="p", llm_conn_id="llm",
tool_approval_assigned_users=["alice"])
+
+ def
test_a_sandbox_passed_through_agent_params_is_refused_with_durable(self):
+ with pytest.raises(ValueError, match="SandboxToolset"):
+ AgentOperator(
+ task_id="t",
+ prompt="p",
+ llm_conn_id="llm",
+ durable=True,
+ agent_params={"toolsets": [SandboxToolset(_NoopBackend())]},
+ )
diff --git a/providers/common/ai/tests/unit/common/ai/toolsets/test_logging.py
b/providers/common/ai/tests/unit/common/ai/toolsets/test_logging.py
index 2bf88a4cd25..6ffecbf4926 100644
--- a/providers/common/ai/tests/unit/common/ai/toolsets/test_logging.py
+++ b/providers/common/ai/tests/unit/common/ai/toolsets/test_logging.py
@@ -20,6 +20,9 @@ import logging
from unittest.mock import AsyncMock, MagicMock
import pytest
+from pydantic_ai import RunContext
+from pydantic_ai.exceptions import ApprovalRequired
+from pydantic_ai.toolsets.abstract import AbstractToolset, ToolsetTool
from airflow.providers.common.ai.toolsets.logging import LoggingToolset
@@ -91,6 +94,21 @@ class TestLoggingToolset:
assert any("Tool bad_tool failed after" in r.message for r in
caplog.records)
assert any("::endgroup::" in r.message for r in caplog.records)
+ @pytest.mark.asyncio
+ async def test_a_call_waiting_for_approval_is_not_logged_as_a_failure(
+ self, logging_toolset, wrapped_toolset, caplog
+ ):
+ wrapped_toolset.call_tool = AsyncMock(spec=AbstractToolset.call_tool,
side_effect=ApprovalRequired())
+
+ with caplog.at_level(logging.INFO, logger="test.logging_toolset"):
+ with pytest.raises(ApprovalRequired):
+ await logging_toolset.call_tool(
+ "refund", {}, MagicMock(spec=RunContext),
MagicMock(spec=ToolsetTool)
+ )
+
+ assert any("Tool refund is waiting to be approved" in r.message for r
in caplog.records)
+ assert not any(r.levelno >= logging.ERROR for r in caplog.records)
+
@pytest.mark.asyncio
async def test_delegates_get_tools(self, logging_toolset, wrapped_toolset):
ctx = MagicMock()