This is an automated email from the ASF dual-hosted git repository. dabla pushed a commit to branch dabla/agent-async/01-task-agent-aexecute in repository https://gitbox.apache.org/repos/asf/airflow.git
commit f55607b5204abdcd1b54367518c2be625b7ca696 Author: David Blain <[email protected]> AuthorDate: Sun Oct 4 19:54:48 2026 +0200 Add native async execution to AgentOperator and @task.agent `@task.agent` could not be used with an `async def` function: the function was called without being awaited, so the task failed on the coroutine it got back as its prompt, while the operator already reported `is_async` through `DecoratedOperator`. An agent run could therefore not await async hooks to build its prompt, nor share the event loop of an async task with other work, although the pydantic-ai agent behind it is async-native. `AgentOperator` is now a `BaseAsyncOperator` with an `aexecute()` that makes no blocking call to the supervisor on the path of a successful run: the hook and its connections are fetched asynchronously, the agent runs with `await agent.run()`, and the usage budget, the tool approval transcript and the XComs go through their async accessors. The usage report of a failed run and the pause for a tool approval are rare and reuse the synchronous code from a worker thread. `durable=True` and `enable_hitl_review=True` do blocking I/O during the run and are rejected with an async function for now. --- providers/common/ai/docs/durable_execution.rst | 3 + providers/common/ai/docs/hitl_review.rst | 3 + providers/common/ai/docs/operators/agent.rst | 24 +++ .../providers/common/ai/decorators/agent.py | 24 +++ .../providers/common/ai/hooks/pydantic_ai.py | 14 ++ .../providers/common/ai/mixins/cancellable_run.py | 16 +- .../airflow/providers/common/ai/operators/agent.py | 174 ++++++++++++++++++--- .../providers/common/ai/utils/usage_budget.py | 34 +++- .../tests/unit/common/ai/decorators/test_agent.py | 88 ++++++++++- .../tests/unit/common/ai/hooks/test_pydantic_ai.py | 11 ++ .../unit/common/ai/mixins/test_cancellable_run.py | 56 ++++++- .../tests/unit/common/ai/operators/test_agent.py | 164 ++++++++++++++++++- .../unit/common/ai/utils/test_usage_budget.py | 55 +++++++ 13 files changed, 628 insertions(+), 38 deletions(-) diff --git a/providers/common/ai/docs/durable_execution.rst b/providers/common/ai/docs/durable_execution.rst index ab6647f2492..3dc6947a17b 100644 --- a/providers/common/ai/docs/durable_execution.rst +++ b/providers/common/ai/docs/durable_execution.rst @@ -38,6 +38,9 @@ after successful completion. Durable execution only helps when the task has retries configured. Without retries there is nothing to replay. +``durable=True`` is not supported with an ``async def`` ``@task.agent`` function +(:ref:`howto/operator:agent-async`). + This page is about making an ``AgentOperator`` retry cheap. Deciding *whether* a task should retry at all is :doc:`retry_policies`; a retried ``LLMBatchOperator`` re-attaches to its running batch instead of resubmitting (:ref:`llm-batch-reattach`). diff --git a/providers/common/ai/docs/hitl_review.rst b/providers/common/ai/docs/hitl_review.rst index 1b40a731dc8..8864a3fe861 100644 --- a/providers/common/ai/docs/hitl_review.rst +++ b/providers/common/ai/docs/hitl_review.rst @@ -27,6 +27,9 @@ terminal action, or until a timeout is reached or max_iterations reached. This document describes the architecture, workflow, API, XCom schema, and usage. +``enable_hitl_review=True`` is not supported with an ``async def`` ``@task.agent`` function +(:ref:`howto/operator:agent-async`). + Overview -------- diff --git a/providers/common/ai/docs/operators/agent.rst b/providers/common/ai/docs/operators/agent.rst index 80c87facfe5..69548014063 100644 --- a/providers/common/ai/docs/operators/agent.rst +++ b/providers/common/ai/docs/operators/agent.rst @@ -83,6 +83,30 @@ the prompt string; all other parameters are passed to the operator. :start-after: [START howto_decorator_agent] :end-before: [END howto_decorator_agent] +.. _howto/operator:agent-async: + +Async callables +^^^^^^^^^^^^^^^ + +On Airflow 3.2+ the decorated function can be an ``async def``. It is awaited for the prompt, and +the agent then runs on the task's event loop: the connection lookup, the model requests and the +bookkeeping of the run make no blocking call. The function can await async hooks, and several such +runs can share one event loop. + +.. code-block:: python + + @task.agent(llm_conn_id="pydanticai_default", system_prompt="You are a support analyst.") + async def summarize_ticket(ticket_id: str) -> str: + ticket = await fetch_ticket(ticket_id) + return f"Summarize this ticket: {ticket}" + +A single run is not faster this way, as its duration is the model's. Tools that call blocking hooks, +such as ``SQLToolset`` and ``HookToolset``, run in a worker thread with either kind of function. + +``durable=True`` and ``enable_hitl_review=True`` do blocking I/O during the run and are not +supported with an ``async def`` function: the Dag fails to parse with a ``ValueError``. Use a regular +function for them. + .. _howto/operator:agent-multimodal: Multimodal prompts diff --git a/providers/common/ai/src/airflow/providers/common/ai/decorators/agent.py b/providers/common/ai/src/airflow/providers/common/ai/decorators/agent.py index d2a92e07668..e966b0ebc52 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/decorators/agent.py +++ b/providers/common/ai/src/airflow/providers/common/ai/decorators/agent.py @@ -24,6 +24,7 @@ and output serialization. from __future__ import annotations +import asyncio from collections.abc import Callable, Collection, Mapping, Sequence from typing import TYPE_CHECKING, Any, ClassVar @@ -40,6 +41,7 @@ from airflow.providers.common.compat.sdk import ( determine_kwargs, task_decorator_factory, ) +from airflow.providers.common.compat.standard.operators import BaseAsyncOperator, is_async_callable if TYPE_CHECKING: from airflow.sdk import Context @@ -83,8 +85,17 @@ class _AgentDecoratedOperator(DecoratedOperator, AgentOperator): prompt=SET_DURING_EXECUTION, **kwargs, ) + if self.is_async: + self._reject_unsupported_on_async_path() + + @property + def is_async(self) -> bool: + return is_async_callable(self.python_callable) def execute(self, context: Context) -> Any: + if self.is_async: + return BaseAsyncOperator.execute(self, context) + context_merge(context, self.op_kwargs) kwargs = determine_kwargs(self.python_callable, self.op_args, context) @@ -101,6 +112,19 @@ class _AgentDecoratedOperator(DecoratedOperator, AgentOperator): self.render_template_fields(context) return AgentOperator.execute(self, context) + async def aexecute(self, context: Context) -> Any: + context_merge(context, self.op_kwargs) + kwargs = determine_kwargs(self.python_callable, self.op_args, context) + + self.prompt = await self.python_callable(*self.op_args, **kwargs) + + validate_prompt(self.prompt, decorator_name="@task.agent") + + # Rendering a template can read a Variable or a Connection with a blocking call to + # the supervisor, which must not happen on the event loop. + await asyncio.to_thread(self.render_template_fields, context) + return await AgentOperator.aexecute(self, context) + def agent_task( python_callable: Callable | None = None, diff --git a/providers/common/ai/src/airflow/providers/common/ai/hooks/pydantic_ai.py b/providers/common/ai/src/airflow/providers/common/ai/hooks/pydantic_ai.py index 750a14c8be2..1718bf1cec2 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/hooks/pydantic_ai.py +++ b/providers/common/ai/src/airflow/providers/common/ai/hooks/pydantic_ai.py @@ -173,6 +173,20 @@ class PydanticAIHook(BaseHook): """ return cls.get_connection(conn_id).get_hook(hook_params=hook_params) + @classmethod + async def aget_hook(cls, conn_id: str, hook_params: dict | None = None): + """ + Return the hook for ``conn_id`` from an async context, built with ``hook_params``. + + The hook is primed with the connection it was dispatched from, so its first + :meth:`aget_conn` or :meth:`acreate_agent` does not fetch that connection again. + """ + conn = await get_async_connection(conn_id) + hook = conn.get_hook(hook_params=hook_params) + if isinstance(hook, PydanticAIHook): + hook._seed_connection(conn, await get_async_extra_dejson(conn)) + return hook + @staticmethod def get_ui_field_behaviour() -> dict[str, Any]: """Return custom field behaviour for the Airflow connection form.""" diff --git a/providers/common/ai/src/airflow/providers/common/ai/mixins/cancellable_run.py b/providers/common/ai/src/airflow/providers/common/ai/mixins/cancellable_run.py index 5e40ce85b50..def69a2bafb 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/mixins/cancellable_run.py +++ b/providers/common/ai/src/airflow/providers/common/ai/mixins/cancellable_run.py @@ -29,10 +29,10 @@ if TYPE_CHECKING: class CancellableAgentRunMixin: """ - Run a pydantic-ai agent synchronously with kill-time cancellation wired in. + Run a pydantic-ai agent with kill-time cancellation wired in. The wrapper holds the in-flight run's ``CancellationToken`` so :meth:`on_kill` can - cancel it. Cancelling makes ``run_sync`` raise ``RunCancelled`` and unwind, giving the + cancel it. Cancelling makes the run raise ``RunCancelled`` and unwind, giving the agent's toolsets a chance to exit (tearing down a provisioned sandbox, for one) before SIGKILL rather than leaving the run to die mid-flight. @@ -57,13 +57,23 @@ class CancellableAgentRunMixin: finally: self._cancellation_token = None + async def run_agent_async( + self, agent: Agent[Any, Any], user_prompt: Any, **run_kwargs: Any + ) -> AgentRunResult[Any]: + """Await ``agent.run`` under a fresh cancellation token held for :meth:`on_kill`.""" + self._cancellation_token = CancellationToken() + try: + return await agent.run(user_prompt, cancellation_token=self._cancellation_token, **run_kwargs) + finally: + self._cancellation_token = None + def on_kill(self) -> None: token = self._cancellation_token if token is None: return self.log.info("Task killed, cancelling in-flight agent run") # Cancel from a separate thread, not inline. on_kill runs in the Task SDK's - # SIGTERM handler on the same thread that drives run_sync, and cancel() only + # SIGTERM handler on the same thread that drives the run, and cancel() only # interrupts a blocked run when issued from a different thread. Called inline it # defers until the in-flight await returns, so the worker is SIGKILLed at the # grace deadline before the run unwinds. 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 11470872e0d..6d1a6c1ff9b 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 @@ -18,6 +18,7 @@ from __future__ import annotations +import asyncio import collections import copy import hashlib @@ -74,6 +75,7 @@ from airflow.providers.common.compat.sdk import ( conf, redact, ) +from airflow.providers.common.compat.standard.operators import BaseAsyncOperator 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 @@ -117,6 +119,14 @@ _TRANSCRIPT_RETENTION_MARGIN = timedelta(days=1) _RUN_USAGE_ADAPTER: TypeAdapter[RunUsage] = TypeAdapter(RunUsage) +async def _axcom_push(ti: Any, *, key: str, value: Any) -> None: + if AIRFLOW_V_3_3_PLUS: + await ti.axcom_push(key=key, value=value) + else: + # Airflow 3.2 runs async tasks but has no ``axcom_push``: push from a worker thread. + await asyncio.to_thread(ti.xcom_push, key=key, value=value) + + class HITLReviewLink(BaseOperatorLink): """ Link that opens the live chat window for a running feedback session. @@ -230,7 +240,7 @@ def _build_code_mode() -> Any: # CancellableAgentRunMixin must precede BaseOperator so its on_kill overrides BaseOperator's # no-op. The other mixins only add methods, so they can trail BaseOperator. See the MRO guard # test in tests/unit/common/ai/mixins/test_cancellable_run.py. -class AgentOperator(CancellableAgentRunMixin, BaseOperator, HITLReviewMixin): +class AgentOperator(CancellableAgentRunMixin, BaseAsyncOperator, HITLReviewMixin): """ Run a pydantic-ai Agent with tools and multi-turn reasoning. @@ -729,17 +739,38 @@ class AgentOperator(CancellableAgentRunMixin, BaseOperator, HITLReviewMixin): agent_params["capabilities"] = render_capabilities(agent_params["capabilities"]) self.agent_params = agent_params - @cached_property - def llm_hook(self) -> PydanticAIHook: - """Return PydanticAIHook for the configured LLM connection.""" - hook_params = { + @property + def is_async(self) -> bool: + """ + Whether the task runs on the worker's event loop, through :meth:`aexecute`. + + ``False`` for the operator itself; ``@task.agent`` returns ``True`` for an ``async def`` callable. + """ + return False + + @property + def _llm_hook_params(self) -> dict[str, Any]: + return { "model_id": self.model_id, "fallback_conn_ids": self.fallback_conn_ids, } - return PydanticAIHook.get_hook(self.llm_conn_id, hook_params=hook_params) + + @cached_property + def llm_hook(self) -> PydanticAIHook: + """Return PydanticAIHook for the configured LLM connection.""" + return PydanticAIHook.get_hook(self.llm_conn_id, hook_params=self._llm_hook_params) def _build_agent(self) -> Agent[object, Any]: """Build and return a pydantic-ai Agent from the operator's config.""" + return self.llm_hook.create_agent(**self._agent_kwargs()) + + async def _abuild_agent(self) -> Agent[object, Any]: + """Async version of :meth:`_build_agent`: the hook and its connections are fetched without blocking.""" + hook = await PydanticAIHook.aget_hook(self.llm_conn_id, hook_params=self._llm_hook_params) + return await hook.acreate_agent(**self._agent_kwargs()) + + def _agent_kwargs(self) -> dict[str, Any]: + """Return the keyword arguments the hook builds the agent from.""" extra_kwargs = dict(self.agent_params) passed_through = extra_kwargs.pop("capabilities", None) if passed_through is not None and self.capabilities is not None: @@ -773,7 +804,7 @@ class AgentOperator(CancellableAgentRunMixin, BaseOperator, HITLReviewMixin): capabilities.append(PromptCaching()) if capabilities: extra_kwargs["capabilities"] = capabilities - return self.llm_hook.create_agent( + return dict( output_type=self._agent_output_type(), instructions=self.system_prompt, **extra_kwargs, @@ -1064,6 +1095,9 @@ class AgentOperator(CancellableAgentRunMixin, BaseOperator, HITLReviewMixin): raise def execute(self, context: Context) -> Any: + if self.is_async: + return BaseAsyncOperator.execute(self, context) + if self.enable_hitl_review and not isinstance(self.prompt, str): raise TypeError( f"{type(self).__name__}: enable_hitl_review=True is not supported " @@ -1105,18 +1139,7 @@ class AgentOperator(CancellableAgentRunMixin, BaseOperator, HITLReviewMixin): self._replay_usage = ReplayUsageLedger(run_usage=run_usage, usage_limits=usage_limits) agent = self._build_agent() - - self._run_identity_attrs = build_run_identity_attributes(ti) - stamp_identity_on_agent_spans(agent, self._run_identity_attrs) - - # A per-attempt key (the task-instance id on Airflow 3, which is regenerated on - # each retry; dag/run/task/map/try on Airflow 2) is a unique, reverse-resolvable - # join key. It lands on result.run_id, the run's messages, and the - # ``gen_ai.agent.call.id`` span attribute. - run_kwargs: dict[str, Any] = {"usage_limits": usage_limits, "run_id": make_task_instance_run_key(ti)} - history = self._resolve_message_history() - if history is not None: - run_kwargs["message_history"] = history + run_kwargs = self._run_kwargs(agent, ti, usage_limits) storage = self._durable_storage counter = self._durable_counter @@ -1145,12 +1168,92 @@ class AgentOperator(CancellableAgentRunMixin, BaseOperator, HITLReviewMixin): self._log_durable_summary(counter) return self._complete_run(context, result, attempt_usage=attempt_usage) - def _complete_run(self, context: Context, result: Any, *, attempt_usage: RunUsage) -> Any: - """Finish a run, or pause it when the agent is waiting on a tool call to be approved.""" + def _run_kwargs( + self, agent: Agent[Any, Any], ti: Any, usage_limits: UsageLimits | None + ) -> dict[str, Any]: + """Stamp the run's identity on the agent's spans and return the keyword arguments of its run.""" + self._run_identity_attrs = build_run_identity_attributes(ti) + stamp_identity_on_agent_spans(agent, self._run_identity_attrs) + + # A per-attempt key (the task-instance id on Airflow 3, which is regenerated on + # each retry; dag/run/task/map/try on Airflow 2) is a unique, reverse-resolvable + # join key. It lands on result.run_id, the run's messages, and the + # ``gen_ai.agent.call.id`` span attribute. + run_kwargs: dict[str, Any] = {"usage_limits": usage_limits, "run_id": make_task_instance_run_key(ti)} + history = self._resolve_message_history() + if history is not None: + run_kwargs["message_history"] = history + return run_kwargs + + def _reject_unsupported_on_async_path(self) -> None: + """Refuse the features that block the event loop: both do synchronous I/O during the run.""" + if self.durable or self.enable_hitl_review: + raise ValueError( + f"{type(self).__name__}: durable=True and enable_hitl_review=True are not supported " + "with an async callable. Use a synchronous callable for them." + ) + + async def aexecute(self, context: Context) -> Any: + """ + Run the agent on the worker's event loop. + + The async version of :meth:`execute`: the connection, the agent run and the bookkeeping of a + successful run (usage budget, run metadata, message history) make no blocking call to the + supervisor, so several of these runs can share one event loop. The rare branches (the usage + report of a failed run, the pause for a tool approval) reuse the synchronous code in a + worker thread. + """ + self._reject_unsupported_on_async_path() + usage_limits = coerce_usage_limits(self.usage_limits) + + if self._supports_tool_approval() and (store := context.get("task_state_store")) is not None: + await self._adelete_approval_transcript(store) + + ti = context["task_instance"] + self._durable_storage = None + self._durable_counter = None + self._replay_usage = None + self._usage_budget = self._build_usage_budget(context, usage_limits, ti=ti) + self._run_usage = await self._usage_budget.aload() if self._usage_budget else RunUsage() + run_usage: RunUsage = self._run_usage + self._run_usage_base = copy_run_usage(run_usage) + + agent = await self._abuild_agent() + run_kwargs = self._run_kwargs(agent, ti, usage_limits) + + try: + try: + result = await self.run_agent_async(agent, self.prompt, usage=run_usage, **run_kwargs) + finally: + if self._usage_budget: + await self._usage_budget.asave(run_usage) + except BaseException: + await asyncio.to_thread(self._report_failed_run, context, run_usage) + raise + attempt_usage = subtract_run_usage(run_usage, self._run_usage_base) + return await self._acomplete_run(context, result, attempt_usage=attempt_usage) + + async def _acomplete_run(self, context: Context, result: Any, *, attempt_usage: RunUsage) -> Any: + """Async version of :meth:`_complete_run`, without the features :meth:`aexecute` refuses.""" log_run_summary(self.log, result, usage=attempt_usage) if isinstance(result.output, DeferredToolRequests): - self._pause_for_tool_approval(context, result, attempt_usage=attempt_usage) - self._emit_run_metadata(context, result, usage=attempt_usage) + await asyncio.to_thread( + self._pause_for_tool_approval, context, result, attempt_usage=attempt_usage + ) + await self._aemit_run_metadata(context, result, usage=attempt_usage) + self._log_cumulative_usage() + + if self.message_history is not None: + await self._aemit_message_history(context, result) + + output = result.output + if self._serialize_model_output and isinstance(output, BaseModel): + output = output.model_dump() + if self._usage_budget: + await self._usage_budget.aclear() + return output + + def _log_cumulative_usage(self) -> None: if self._usage_budget and (run_usage := self._run_usage) is not None: self.log.info( "Cumulative usage across attempts: requests=%s, tool_calls=%s, input_tokens=%s, " @@ -1167,6 +1270,14 @@ class AgentOperator(CancellableAgentRunMixin, BaseOperator, HITLReviewMixin): format(run_usage.cost, "f"), ) + def _complete_run(self, context: Context, result: Any, *, attempt_usage: RunUsage) -> 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, usage=attempt_usage) + if isinstance(result.output, DeferredToolRequests): + self._pause_for_tool_approval(context, result, attempt_usage=attempt_usage) + self._emit_run_metadata(context, result, usage=attempt_usage) + self._log_cumulative_usage() + if self.message_history is not None: self._emit_message_history(context, result) @@ -1296,6 +1407,12 @@ class AgentOperator(CancellableAgentRunMixin, BaseOperator, HITLReviewMixin): except Exception: self.log.warning("Could not delete the tool approval transcript", exc_info=True) + async def _adelete_approval_transcript(self, store: TaskStateStoreAccessor) -> None: + try: + await store.adelete(_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, @@ -1441,6 +1558,17 @@ class AgentOperator(CancellableAgentRunMixin, BaseOperator, HITLReviewMixin): ti.xcom_push(key="run_id", value=result.run_id) ti.xcom_push(key="usage", value=format_usage_for_xcom(usage)) + async def _aemit_message_history(self, context: Context, result: Any) -> None: + transcript = ModelMessagesTypeAdapter.dump_json(result.all_messages()).decode() + await _axcom_push(context["task_instance"], key="message_history", value=transcript) + + async def _aemit_run_metadata(self, context: Context, result: Any, *, usage: RunUsage) -> None: + if not self.do_xcom_push: + return + ti = context["task_instance"] + await _axcom_push(ti, key="run_id", value=result.run_id) + await _axcom_push(ti, key="usage", value=format_usage_for_xcom(usage)) + def regenerate_with_feedback(self, *, feedback: str, message_history: Any) -> tuple[str, Any]: """ Re-run the agent with *feedback* appended to the conversation history. diff --git a/providers/common/ai/src/airflow/providers/common/ai/utils/usage_budget.py b/providers/common/ai/src/airflow/providers/common/ai/utils/usage_budget.py index cb8c43e412e..5c722dff040 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/utils/usage_budget.py +++ b/providers/common/ai/src/airflow/providers/common/ai/utils/usage_budget.py @@ -179,7 +179,13 @@ class TaskStateStoreUsageBudget: def load(self) -> RunUsage: """Return the cumulative usage so far, or a fresh ``RunUsage()`` if none or stale.""" - raw = self._store.get(USAGE_BUDGET_KEY) + return self._parse(self._store.get(USAGE_BUDGET_KEY)) + + async def aload(self) -> RunUsage: + """Async version of :meth:`load`.""" + return self._parse(await self._store.aget(USAGE_BUDGET_KEY)) + + def _parse(self, raw: Any) -> RunUsage: if raw is None: return RunUsage() if not isinstance(raw, dict) or "usage" not in raw: @@ -198,18 +204,32 @@ class TaskStateStoreUsageBudget: # module keeps importing cleanly on older Airflow versions (this module's docstring). from airflow.sdk.execution_time.context import NEVER_EXPIRE - record: dict[str, Any] = { - "version": 1, - "max_tries": self._max_tries, - "usage": dump_run_usage(usage), - } - self._store.set(USAGE_BUDGET_KEY, record, retention=NEVER_EXPIRE) + self._store.set(USAGE_BUDGET_KEY, self._record(usage), retention=NEVER_EXPIRE) + except Exception: + log.warning("Usage budget: failed to persist cumulative usage", exc_info=True) + + async def asave(self, usage: RunUsage) -> None: + """Async version of :meth:`save`.""" + try: + from airflow.sdk.execution_time.context import NEVER_EXPIRE + + await self._store.aset(USAGE_BUDGET_KEY, self._record(usage), retention=NEVER_EXPIRE) except Exception: log.warning("Usage budget: failed to persist cumulative usage", exc_info=True) + def _record(self, usage: RunUsage) -> dict[str, Any]: + return {"version": 1, "max_tries": self._max_tries, "usage": dump_run_usage(usage)} + def clear(self) -> None: """Best-effort delete, called once the whole execute succeeds.""" try: self._store.delete(USAGE_BUDGET_KEY) except Exception: log.warning("Usage budget: failed to delete cumulative usage key", exc_info=True) + + async def aclear(self) -> None: + """Async version of :meth:`clear`.""" + try: + await self._store.adelete(USAGE_BUDGET_KEY) + except Exception: + log.warning("Usage budget: failed to delete cumulative usage key", exc_info=True) diff --git a/providers/common/ai/tests/unit/common/ai/decorators/test_agent.py b/providers/common/ai/tests/unit/common/ai/decorators/test_agent.py index fffdd51a15f..bbd391432ed 100644 --- a/providers/common/ai/tests/unit/common/ai/decorators/test_agent.py +++ b/providers/common/ai/tests/unit/common/ai/decorators/test_agent.py @@ -16,7 +16,8 @@ # under the License. from __future__ import annotations -from unittest.mock import ANY, MagicMock, patch +import threading +from unittest.mock import ANY, AsyncMock, MagicMock, patch import pytest from pydantic import BaseModel @@ -29,6 +30,8 @@ from airflow.providers.common.ai.toolsets.logging import LoggingToolset from airflow.providers.common.ai.utils.prompt_cache import PromptCaching from airflow.providers.common.ai.utils.toolset_base import MaskingToolset +from tests_common.test_utils.version_compat import AIRFLOW_V_3_2_PLUS + try: from airflow.sdk.serde import SUPPORTS_OPERATOR_DESERIALIZATION_WALKER as _CORE_WALKER except ImportError: @@ -236,3 +239,86 @@ class TestAgentDecoratedOperator: llm_conn_id="my_llm", ) assert op.durable is False + + +async def _async_prompt(): + return "Who is our top customer?" + + +def _make_async_context(): + context = _make_context() + context["task_instance"].axcom_push = AsyncMock() + return context + + +def _make_async_mock_agent(mock_hook_cls, make_mock_run_result): + mock_agent = MagicMock(spec=["run", "instrument"]) + mock_agent.run = AsyncMock(return_value=make_mock_run_result("The top customer is Acme Corp.")) + mock_hook_cls.aget_hook = AsyncMock() + mock_hook_cls.aget_hook.return_value.acreate_agent = AsyncMock(return_value=mock_agent) + return mock_agent + + [email protected](not AIRFLOW_V_3_2_PLUS, reason="async tasks need Airflow >= 3.2") +class TestAgentDecoratedOperatorAsync: + def test_is_async_follows_the_callable(self): + def prompt(): + return "Who is our top customer?" + + sync_op = _AgentDecoratedOperator(task_id="sync", python_callable=prompt, llm_conn_id="my_llm") + async_op = _AgentDecoratedOperator( + task_id="async", python_callable=_async_prompt, llm_conn_id="my_llm" + ) + + assert sync_op.is_async is False + assert async_op.is_async is True + + @pytest.mark.asyncio + @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook", autospec=True) + async def test_aexecute_awaits_the_callable_for_the_prompt(self, mock_hook_cls, make_mock_run_result): + mock_agent = _make_async_mock_agent(mock_hook_cls, make_mock_run_result) + + op = _AgentDecoratedOperator(task_id="test", python_callable=_async_prompt, llm_conn_id="my_llm") + result = await op.aexecute(_make_async_context()) + + assert result == "The top customer is Acme Corp." + assert op.prompt == "Who is our top customer?" + mock_agent.run.assert_awaited_once_with( + "Who is our top customer?", usage_limits=None, run_id="ti-1", cancellation_token=ANY, usage=ANY + ) + + @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook", autospec=True) + def test_execute_runs_an_async_callable_on_an_event_loop(self, mock_hook_cls, make_mock_run_result): + """The task runner calls ``execute``: with an ``async def`` callable it drives ``aexecute``.""" + _make_async_mock_agent(mock_hook_cls, make_mock_run_result) + + op = _AgentDecoratedOperator(task_id="test", python_callable=_async_prompt, llm_conn_id="my_llm") + result = op.execute(context=_make_async_context()) + + assert result == "The top customer is Acme Corp." + mock_hook_cls.get_hook.assert_not_called() + + @pytest.mark.asyncio + @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook", autospec=True) + async def test_template_fields_are_rendered_off_the_event_loop(self, mock_hook_cls, make_mock_run_result): + """Rendering can read a Variable or a Connection with a blocking call to the supervisor.""" + _make_async_mock_agent(mock_hook_cls, make_mock_run_result) + rendered_on: list[int] = [] + + op = _AgentDecoratedOperator(task_id="test", python_callable=_async_prompt, llm_conn_id="my_llm") + with patch.object( + _AgentDecoratedOperator, + "render_template_fields", + side_effect=lambda context: rendered_on.append(threading.get_ident()), + ): + await op.aexecute(_make_async_context()) + + assert rendered_on + assert rendered_on[0] != threading.get_ident() + + @pytest.mark.parametrize("feature", ["durable", "enable_hitl_review"]) + def test_async_callable_rejects_features_that_block_the_event_loop(self, feature): + with pytest.raises(ValueError, match="not supported with an async callable"): + _AgentDecoratedOperator( + task_id="test", python_callable=_async_prompt, llm_conn_id="my_llm", **{feature: True} + ) diff --git a/providers/common/ai/tests/unit/common/ai/hooks/test_pydantic_ai.py b/providers/common/ai/tests/unit/common/ai/hooks/test_pydantic_ai.py index cd6bcb72786..79bda339fce 100644 --- a/providers/common/ai/tests/unit/common/ai/hooks/test_pydantic_ai.py +++ b/providers/common/ai/tests/unit/common/ai/hooks/test_pydantic_ai.py @@ -2142,6 +2142,17 @@ class TestPydanticAIHookAsync: assert hook.get_conn() is infer_model_stub.models["openai:gpt-5.6-sol"] async_registry.get_async_connection.assert_awaited_once() + @pytest.mark.asyncio + async def test_aget_hook_primes_the_hook_with_the_connection_it_fetched( + self, async_registry, infer_model_stub + ): + async_registry.add("primary", extra={"model": "openai:gpt-5.6-sol"}) + + hook = await PydanticAIHook.aget_hook("primary", hook_params={"model_id": "openai:gpt-5.6-terra"}) + + assert await hook.aget_conn() is infer_model_stub.models["openai:gpt-5.6-terra"] + async_registry.get_async_connection.assert_awaited_once_with("primary") + @pytest.mark.asyncio @patch("airflow.providers.common.ai.hooks.pydantic_ai.Agent", autospec=True) async def test_acreate_agent(self, mock_agent_cls, async_registry, infer_model_stub): diff --git a/providers/common/ai/tests/unit/common/ai/mixins/test_cancellable_run.py b/providers/common/ai/tests/unit/common/ai/mixins/test_cancellable_run.py index 69cf184dbc8..2592fa971b8 100644 --- a/providers/common/ai/tests/unit/common/ai/mixins/test_cancellable_run.py +++ b/providers/common/ai/tests/unit/common/ai/mixins/test_cancellable_run.py @@ -16,12 +16,15 @@ # under the License. from __future__ import annotations +import asyncio import threading import time -from unittest.mock import DEFAULT, MagicMock +from unittest.mock import DEFAULT, AsyncMock, MagicMock import pytest -from pydantic_ai import CancellationToken +from pydantic_ai import Agent, CancellationToken, RunCancelled +from pydantic_ai.messages import ModelResponse, TextPart +from pydantic_ai.models.function import FunctionModel from airflow.providers.common.ai.mixins.cancellable_run import CancellableAgentRunMixin from airflow.providers.common.ai.operators.agent import AgentOperator @@ -63,6 +66,55 @@ class TestRunAgentSync: assert mixin._cancellation_token is None +class TestRunAgentAsync: + @pytest.mark.asyncio + async def test_forwards_held_cancellation_token_and_clears_after_success(self): + mixin = CancellableAgentRunMixin() + agent = MagicMock(spec=["run"]) + held: dict[str, object] = {} + + async def capture(*args, **kwargs): + held["token"] = mixin._cancellation_token + return DEFAULT + + agent.run = AsyncMock(side_effect=capture) + + result = await mixin.run_agent_async(agent, "prompt", usage_limits=None) + + assert result is agent.run.return_value + passed = agent.run.call_args.kwargs["cancellation_token"] + assert isinstance(passed, CancellationToken) + assert passed is held["token"] + agent.run.assert_awaited_once_with("prompt", cancellation_token=passed, usage_limits=None) + assert mixin._cancellation_token is None + + @pytest.mark.asyncio + async def test_clears_token_when_run_raises(self): + mixin = CancellableAgentRunMixin() + agent = MagicMock(spec=["run"]) + agent.run = AsyncMock(side_effect=RuntimeError("boom")) + + with pytest.raises(RuntimeError): + await mixin.run_agent_async(agent, "prompt") + + assert mixin._cancellation_token is None + + @pytest.mark.asyncio + async def test_on_kill_cancels_an_awaited_run(self): + """on_kill is called on the thread that runs the event loop; the run it awaits must unwind.""" + mixin = CancellableAgentRunMixin() + mixin.log = MagicMock() + + async def hang(messages, info) -> ModelResponse: + await asyncio.sleep(30) + return ModelResponse(parts=[TextPart("too late")]) + + asyncio.get_running_loop().call_later(0.1, mixin.on_kill) + + with pytest.raises(RunCancelled): + await asyncio.wait_for(mixin.run_agent_async(Agent(FunctionModel(hang)), "prompt"), timeout=10) + + class TestOnKill: def test_noop_when_no_run_active(self): mixin = CancellableAgentRunMixin() 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 2d3422f2479..f079e913317 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 @@ -16,6 +16,7 @@ # under the License. from __future__ import annotations +import asyncio import dataclasses import json import logging @@ -25,7 +26,7 @@ from datetime import timedelta from decimal import Decimal from types import ModuleType, SimpleNamespace from typing import Any -from unittest.mock import ANY, MagicMock, PropertyMock, call, patch +from unittest.mock import ANY, AsyncMock, MagicMock, PropertyMock, call, patch import pytest from pydantic import BaseModel @@ -95,7 +96,12 @@ from airflow.providers.common.compat.sdk import ( ) from tests_common.test_utils.compat import OperatorSerialization -from tests_common.test_utils.version_compat import AIRFLOW_V_3_0_PLUS, AIRFLOW_V_3_1_PLUS, AIRFLOW_V_3_3_PLUS +from tests_common.test_utils.version_compat import ( + AIRFLOW_V_3_0_PLUS, + AIRFLOW_V_3_1_PLUS, + AIRFLOW_V_3_2_PLUS, + AIRFLOW_V_3_3_PLUS, +) from unit.common.ai.sandbox.fake_tags import TaggedBackend try: @@ -3418,3 +3424,157 @@ class TestAgentOperatorMasksToolOutput: op._build_agent() assert hook.create_agent.call_args.kwargs["toolsets"] == [sql] + + +def _make_async_task_state_store_accessor(): + """A task state store whose async methods are backed by a dict and whose blocking ones fail the test.""" + from airflow.sdk.execution_time.context import TaskStateStoreAccessor + + store: dict = {} + accessor = MagicMock(spec=TaskStateStoreAccessor) + blocking = AssertionError("blocking task state store call on the event loop") + accessor.get.side_effect = accessor.set.side_effect = accessor.delete.side_effect = blocking + accessor.aget = AsyncMock(side_effect=lambda key, default=None: store.get(key, default)) + accessor.aset = AsyncMock(side_effect=lambda key, value, retention=None: store.__setitem__(key, value)) + accessor.adelete = AsyncMock(side_effect=lambda key: store.pop(key, None)) + return accessor + + +def _make_async_context(task_state_store=None): + """A context for ``aexecute``: XComs go through ``axcom_push`` on Airflow 3.3+, ``xcom_push`` on 3.2.""" + ti = _make_ti() + ti.axcom_push = AsyncMock() + context = {"task_instance": ti} + if task_state_store is not None: + context["task_state_store"] = task_state_store + return context + + +def _make_async_mock_agent(mock_hook_cls, **run): + """Wire a mock agent whose ``run`` is awaitable into the hook's async path, and return it.""" + mock_agent = MagicMock(spec=["run", "instrument"]) + mock_agent.run = AsyncMock(**run) + mock_hook_cls.aget_hook = AsyncMock() + mock_hook_cls.aget_hook.return_value.acreate_agent = AsyncMock(return_value=mock_agent) + return mock_agent + + [email protected](not AIRFLOW_V_3_2_PLUS, reason="async tasks need Airflow >= 3.2") +class TestAgentOperatorAsync: + def test_operator_is_synchronous_unless_a_subclass_says_otherwise(self): + assert AgentOperator(task_id="test", prompt="run", llm_conn_id="my_llm").is_async is False + + @pytest.mark.asyncio + @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook", autospec=True) + async def test_aexecute_awaits_the_run_through_the_async_hook_path( + self, mock_hook_cls, make_mock_run_result + ): + mock_agent = _make_async_mock_agent(mock_hook_cls, return_value=make_mock_run_result("done")) + + op = AgentOperator(task_id="test", prompt="run", llm_conn_id="my_llm") + result = await op.aexecute(_make_async_context()) + + assert result == "done" + mock_hook_cls.aget_hook.assert_awaited_once_with( + "my_llm", hook_params={"model_id": None, "fallback_conn_ids": None} + ) + mock_hook_cls.get_hook.assert_not_called() + mock_agent.run.assert_awaited_once_with( + "run", usage_limits=None, run_id="ti-1", cancellation_token=ANY, usage=ANY + ) + + @pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="task state store needs Airflow >= 3.3") + @pytest.mark.asyncio + @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook", autospec=True) + async def test_successful_run_makes_no_blocking_call_to_the_supervisor(self, mock_hook_cls): + """A real run with a usage budget and a message history: every state store call and XCom + push of the successful path is awaited, so the blocking doubles are never reached.""" + mock_hook_cls.aget_hook = AsyncMock() + mock_hook_cls.aget_hook.return_value.acreate_agent = AsyncMock( + side_effect=lambda **kw: Agent(FunctionModel(_build_priced_response), **kw) + ) + store = _make_async_task_state_store_accessor() + context = _make_async_context(task_state_store=store) + context["task_instance"].xcom_push.side_effect = AssertionError("blocking xcom_push") + + op = AgentOperator( + task_id="test", + prompt="run", + llm_conn_id="my_llm", + usage_limits={"cost_limit": str(PRICED_COST * 2)}, + message_history=[], + ) + result = await op.aexecute(context) + + assert result == "the answer" + pushed = [c.kwargs["key"] for c in context["task_instance"].axcom_push.await_args_list] + assert pushed == ["run_id", "usage", "message_history"] + assert store.aset.await_args.args[0] == USAGE_BUDGET_KEY + store.adelete.assert_awaited_with(USAGE_BUDGET_KEY) + + @pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="task state store needs Airflow >= 3.3") + @pytest.mark.asyncio + @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook", autospec=True) + async def test_failed_run_saves_the_budget_and_reports_its_usage(self, mock_hook_cls): + _make_async_mock_agent(mock_hook_cls, side_effect=RuntimeError("boom")) + store = _make_async_task_state_store_accessor() + context = _make_async_context(task_state_store=store) + + op = AgentOperator( + task_id="test", prompt="run", llm_conn_id="my_llm", usage_limits={"request_limit": 5} + ) + with pytest.raises(RuntimeError, match="boom"): + await op.aexecute(context) + + assert store.aset.await_args.args[0] == USAGE_BUDGET_KEY + # The failure report reuses the synchronous code, from a worker thread. + pushed = {c.kwargs["key"] for c in context["task_instance"].xcom_push.call_args_list} + assert pushed == {"run_id", "usage"} + + @pytest.mark.parametrize("feature", ["durable", "enable_hitl_review"]) + @pytest.mark.asyncio + async def test_aexecute_rejects_features_that_block_the_event_loop(self, feature): + op = AgentOperator(task_id="test", prompt="run", llm_conn_id="my_llm", **{feature: True}) + + with pytest.raises(ValueError, match="not supported with an async callable"): + await op.aexecute(_make_async_context()) + + @pytest.mark.asyncio + @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook", autospec=True) + async def test_run_waiting_on_a_tool_approval_pauses_through_the_synchronous_path( + self, mock_hook_cls, make_mock_run_result + ): + class Paused(Exception): + pass + + _make_async_mock_agent(mock_hook_cls, return_value=make_mock_run_result(DeferredToolRequests())) + context = _make_async_context() + + op = AgentOperator(task_id="test", prompt="run", llm_conn_id="my_llm") + with patch.object(AgentOperator, "_pause_for_tool_approval", side_effect=Paused) as pause: + with pytest.raises(Paused): + await op.aexecute(context) + + pause.assert_called_once() + context["task_instance"].axcom_push.assert_not_awaited() + + @pytest.mark.asyncio + @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook", autospec=True) + async def test_runs_overlap_on_one_event_loop(self, mock_hook_cls): + async def slow(messages: list[ModelMessage], info: AgentInfo) -> ModelResponse: + await asyncio.sleep(0.2) + return ModelResponse(parts=[TextPart(content="the answer")]) + + mock_hook_cls.aget_hook = AsyncMock() + mock_hook_cls.aget_hook.return_value.acreate_agent = AsyncMock( + side_effect=lambda **kw: Agent(FunctionModel(slow), **kw) + ) + ops = [AgentOperator(task_id=f"test_{i}", prompt="run", llm_conn_id="my_llm") for i in range(8)] + + loop = asyncio.get_running_loop() + started = loop.time() + results = await asyncio.gather(*(op.aexecute(_make_async_context()) for op in ops)) + + assert results == ["the answer"] * 8 + # Eight runs of 0.2 s each take 1.6 s one after the other. + assert loop.time() - started < 1.0 diff --git a/providers/common/ai/tests/unit/common/ai/utils/test_usage_budget.py b/providers/common/ai/tests/unit/common/ai/utils/test_usage_budget.py index c95a2ddd4fa..5bb74cfd5eb 100644 --- a/providers/common/ai/tests/unit/common/ai/utils/test_usage_budget.py +++ b/providers/common/ai/tests/unit/common/ai/utils/test_usage_budget.py @@ -149,6 +149,15 @@ class FakeTaskStateStore: def delete(self, key): del self.store[key] + async def aget(self, key, default=None): + return self.get(key, default) + + async def aset(self, key, value, *, retention=None): + self.set(key, value, retention=retention) + + async def adelete(self, key): + self.delete(key) + class RaisingTaskStateStore: """Every method raises -- for the best-effort save()/clear() canaries.""" @@ -162,6 +171,12 @@ class RaisingTaskStateStore: def delete(self, key): raise RuntimeError("store down") + async def aset(self, key, value, *, retention=None): + raise RuntimeError("store down") + + async def adelete(self, key): + raise RuntimeError("store down") + @pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="task state store needs Airflow >= 3.3") class TestTaskStateStoreUsageBudgetLoad: @@ -240,3 +255,43 @@ class TestTaskStateStoreUsageBudgetClear: """Canary: removing the try/except around ``self._store.delete`` in ``TaskStateStoreUsageBudget.clear`` turns this red with ``RuntimeError``.""" TaskStateStoreUsageBudget(RaisingTaskStateStore(), max_tries=1).clear() + + [email protected](not AIRFLOW_V_3_3_PLUS, reason="task state store needs Airflow >= 3.3") +class TestTaskStateStoreUsageBudgetAsync: + @pytest.mark.asyncio + async def test_asave_aload_aclear_round_trip(self): + from airflow.sdk.execution_time.context import NEVER_EXPIRE + + store = FakeTaskStateStore() + budget = TaskStateStoreUsageBudget(store, max_tries=2) + + assert await budget.aload() == RunUsage() + + await budget.asave(RunUsage(requests=1, cost=Decimal("0.10"))) + + assert store.store[USAGE_BUDGET_KEY] == { + "version": 1, + "max_tries": 2, + "usage": dump_run_usage(RunUsage(requests=1, cost=Decimal("0.10"))), + } + assert store.retentions[USAGE_BUDGET_KEY] == NEVER_EXPIRE + assert await budget.aload() == RunUsage(requests=1, cost=Decimal("0.10")) + + await budget.aclear() + + assert USAGE_BUDGET_KEY not in store.store + + @pytest.mark.asyncio + async def test_aload_starts_from_zero_after_a_clear_bumped_max_tries(self): + store = FakeTaskStateStore() + await TaskStateStoreUsageBudget(store, max_tries=1).asave(RunUsage(requests=3)) + + assert await TaskStateStoreUsageBudget(store, max_tries=2).aload() == RunUsage() + + @pytest.mark.asyncio + async def test_failed_asave_and_aclear_are_swallowed_not_raised(self): + budget = TaskStateStoreUsageBudget(RaisingTaskStateStore(), max_tries=1) + + await budget.asave(RunUsage()) + await budget.aclear()
