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()

Reply via email to