This is an automated email from the ASF dual-hosted git repository.

kaxil pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git


The following commit(s) were added to refs/heads/main by this push:
     new a33fe27cace Add async client to `AnthropicHook` (#73967)
a33fe27cace is described below

commit a33fe27cacefd94f561e74cdff868b4c4de554f8
Author: Kaxil Naik <[email protected]>
AuthorDate: Thu Oct 1 12:03:01 2026 +0100

    Add async client to `AnthropicHook` (#73967)
    
    * Add async client to AnthropicHook
    
    AnthropicHook.get_async_conn() returns the async counterpart of the client
    get_conn() builds (AsyncAnthropic, AsyncAnthropicBedrock, 
AsyncAnthropicVertex,
    AsyncAnthropicAWS or AsyncAnthropicFoundry) from the same connection, 
through
    one shared builder. It looks the connection up asynchronously through the 
hook,
    so a subclass's own lookup is honoured, and caches it, so the platform and 
the
    connection's default model are read without a second, blocking lookup. The
    common-compat floor moves to 1.17.0 for get_async_connection's hook 
argument.
    
    * Mark common-compat for the next version instead of raising its floor
---
 providers/anthropic/pyproject.toml                 |   2 +-
 .../airflow/providers/anthropic/hooks/anthropic.py |  99 +++++++++++++--
 .../tests/unit/anthropic/hooks/test_anthropic.py   | 141 ++++++++++++++++++++-
 3 files changed, 227 insertions(+), 15 deletions(-)

diff --git a/providers/anthropic/pyproject.toml 
b/providers/anthropic/pyproject.toml
index 85ce42bab5b..5cfd126d7e2 100644
--- a/providers/anthropic/pyproject.toml
+++ b/providers/anthropic/pyproject.toml
@@ -60,7 +60,7 @@ requires-python = ">=3.10"
 # After you modify the dependencies, and rebuild your Breeze CI image with 
``breeze ci-image build``
 dependencies = [
     "apache-airflow>=3.0.0",
-    "apache-airflow-providers-common-compat>=1.12.0",
+    "apache-airflow-providers-common-compat>=1.12.0",  # use next version
     "anthropic>=1.0.0",
 ]
 
diff --git 
a/providers/anthropic/src/airflow/providers/anthropic/hooks/anthropic.py 
b/providers/anthropic/src/airflow/providers/anthropic/hooks/anthropic.py
index a6d5345e930..fa669186856 100644
--- a/providers/anthropic/src/airflow/providers/anthropic/hooks/anthropic.py
+++ b/providers/anthropic/src/airflow/providers/anthropic/hooks/anthropic.py
@@ -20,10 +20,11 @@ import logging
 import time
 from collections.abc import Mapping
 from copy import deepcopy
+from dataclasses import dataclass
 from decimal import Decimal, InvalidOperation
 from enum import Enum
 from functools import cached_property
-from typing import TYPE_CHECKING, Any, NamedTuple, cast
+from typing import TYPE_CHECKING, Any, Generic, NamedTuple, TypeVar, cast
 
 from anthropic import (
     Anthropic,
@@ -31,6 +32,11 @@ from anthropic import (
     AnthropicBedrock,
     AnthropicFoundry,
     AnthropicVertex,
+    AsyncAnthropic,
+    AsyncAnthropicAWS,
+    AsyncAnthropicBedrock,
+    AsyncAnthropicFoundry,
+    AsyncAnthropicVertex,
     BadRequestError,
     IdentityTokenFile,
     WorkloadIdentityCredentials,
@@ -45,6 +51,7 @@ from airflow.providers.anthropic.exceptions import (
     AnthropicSessionBudgetExceeded,
     AnthropicTriggerEventError,
 )
+from airflow.providers.common.compat.connection import get_async_connection
 from airflow.providers.common.compat.sdk import AirflowSkipException, BaseHook
 
 logger = logging.getLogger(__name__)
@@ -65,6 +72,7 @@ if TYPE_CHECKING:
     from anthropic.types.messages import MessageBatch, 
MessageBatchIndividualResponse
     from anthropic.types.messages.batch_create_params import Request
 
+
 #: Default model used when an operator or hook caller does not specify one.
 #: Prefer configuring the model on the connection so it can be updated without
 #: a provider release when this model ID is retired.
@@ -77,6 +85,29 @@ DEFAULT_MODEL = "claude-opus-4-8"
 FIRST_PARTY_PLATFORMS = frozenset({"anthropic", "aws"})
 
 AnthropicClient = Anthropic | AnthropicBedrock | AnthropicVertex | 
AnthropicAWS | AnthropicFoundry
+AsyncAnthropicClient = (
+    AsyncAnthropic | AsyncAnthropicBedrock | AsyncAnthropicVertex | 
AsyncAnthropicAWS | AsyncAnthropicFoundry
+)
+
+_ClientT = TypeVar("_ClientT")
+
+
+@dataclass(frozen=True)
+class _ClientFactories(Generic[_ClientT]):
+    """
+    The client class to build for each platform.
+
+    :meth:`AnthropicHook.get_conn` and :meth:`AnthropicHook.get_async_conn` 
pass the sync and
+    async SDK classes through the same builder, so the two clients read the 
connection the
+    same way and cannot drift.
+    """
+
+    anthropic: Callable[..., _ClientT]
+    bedrock: Callable[..., _ClientT]
+    vertex: Callable[..., _ClientT]
+    aws: Callable[..., _ClientT]
+    foundry: Callable[..., _ClientT]
+
 
 #: Consecutive failed polls tolerated in the synchronous wait helpers before 
giving up
 #: (transient errors). Mirrors the deferrable triggers' tolerance so a single 
blip does
@@ -310,6 +341,9 @@ class AnthropicHook(BaseHook):
     client is built with no static credential so the SDK resolves them from 
the environment
     — supporting env-driven Workload Identity Federation and ``ant`` profiles.
 
+    :meth:`get_conn` returns the synchronous client. Async callers, such as an 
agent loop or a
+    trigger, ``await`` :meth:`get_async_conn` for its async twin, built from 
the same connection.
+
     .. seealso:: https://docs.claude.com/en/api/client-sdks
 
     :param conn_id: :ref:`Anthropic connection id 
<howto/connection:anthropic>`.
@@ -345,10 +379,10 @@ class AnthropicHook(BaseHook):
 
     def _build_aws_client(
         self,
-        factory: Callable[..., AnthropicClient],
+        factory: Callable[..., _ClientT],
         aws_region: str | None,
         client_kwargs: dict[str, Any],
-    ) -> AnthropicClient:
+    ) -> _ClientT:
         try:
             return factory(aws_region=aws_region, **client_kwargs)
         except ValueError as exc:
@@ -359,27 +393,27 @@ class AnthropicHook(BaseHook):
             raise AnthropicError(
                 f"No AWS region configured for the {self.platform!r} platform. 
Set 'aws_region' in the "
                 f"extra of connection {self.conn_id!r}, or set AWS_REGION / 
AWS_DEFAULT_REGION (or a "
-                "region on the AWS profile) on the worker."
+                "region on the AWS profile) on the worker or triggerer."
             ) from exc
 
-    def get_conn(self) -> AnthropicClient:
-        """Build and return the Anthropic client for the configured 
platform."""
+    def _build_client(self, factories: _ClientFactories[_ClientT]) -> _ClientT:
+        """Build the client for the connection's platform from the given sync 
or async classes."""
         conn = self._connection
         extras = conn.extra_dejson
         client_kwargs = dict(extras.get("anthropic_client_kwargs", {}))
         platform = self.platform
         self.log.debug("Building Anthropic client for platform %r 
(conn_id=%s)", platform, self.conn_id)
         if platform == "bedrock":
-            return self._build_aws_client(AnthropicBedrock, 
extras.get("aws_region"), client_kwargs)
+            return self._build_aws_client(factories.bedrock, 
extras.get("aws_region"), client_kwargs)
         if platform == "vertex":
-            return AnthropicVertex(
+            return factories.vertex(
                 project_id=extras.get("project_id"), 
region=extras.get("region"), **client_kwargs
             )
         if platform == "aws":
-            return self._build_aws_client(AnthropicAWS, 
extras.get("aws_region"), client_kwargs)
+            return self._build_aws_client(factories.aws, 
extras.get("aws_region"), client_kwargs)
         if platform == "foundry":
             api_key = client_kwargs.pop("api_key", None) or conn.password
-            return AnthropicFoundry(api_key=api_key, 
resource=extras.get("resource"), **client_kwargs)
+            return factories.foundry(api_key=api_key, 
resource=extras.get("resource"), **client_kwargs)
         if platform != "anthropic":
             raise AnthropicError(
                 f"Unknown Anthropic platform {platform!r}. "
@@ -388,16 +422,55 @@ class AnthropicHook(BaseHook):
         base_url = client_kwargs.pop("base_url", None) or conn.host or None
         wif = extras.get("workload_identity")
         if wif:
-            return Anthropic(
+            # The async client takes the same credential: the SDK runs a 
synchronous
+            # provider's token exchange in a worker thread, off the event loop.
+            return factories.anthropic(
                 credentials=self._workload_identity_credentials(wif), 
base_url=base_url, **client_kwargs
             )
         api_key = client_kwargs.pop("api_key", None) or conn.password
         if api_key:
-            return Anthropic(api_key=api_key, base_url=base_url, 
**client_kwargs)
+            return factories.anthropic(api_key=api_key, base_url=base_url, 
**client_kwargs)
         # No static key and no explicit federation config: let the SDK resolve 
credentials
         # from the environment, which supports env-driven Workload Identity 
Federation
         # (ANTHROPIC_FEDERATION_RULE_ID etc.) and ``ant`` profiles.
-        return Anthropic(base_url=base_url, **client_kwargs)
+        return factories.anthropic(base_url=base_url, **client_kwargs)
+
+    def get_conn(self) -> AnthropicClient:
+        """Build and return the Anthropic client for the configured 
platform."""
+        factories: _ClientFactories[AnthropicClient] = _ClientFactories(
+            anthropic=Anthropic,
+            bedrock=AnthropicBedrock,
+            vertex=AnthropicVertex,
+            aws=AnthropicAWS,
+            foundry=AnthropicFoundry,
+        )
+        return self._build_client(factories)
+
+    async def get_async_conn(self) -> AsyncAnthropicClient:
+        """
+        Build and return the async Anthropic client for the configured 
platform.
+
+        Reads the same connection fields as :meth:`get_conn` through the same 
builder, and
+        returns the async twin of the client it would build: 
``AsyncAnthropic``,
+        ``AsyncAnthropicBedrock``, ``AsyncAnthropicVertex``, 
``AsyncAnthropicAWS`` or
+        ``AsyncAnthropicFoundry``. The connection is looked up without 
blocking the event
+        loop, so a trigger can call this directly.
+
+        Each call builds a new client. Close it when done, with ``async with`` 
or
+        ``await client.close()``, to release its HTTP connections.
+        """
+        if "_connection" not in self.__dict__:
+            # Fills the cached property, so later reads of ``platform`` and 
``default_model``
+            # use this connection instead of a second, blocking lookup.
+            self._connection = await get_async_connection(self.conn_id, 
hook=self)
+        factories: _ClientFactories[AsyncAnthropicClient] = _ClientFactories(
+            anthropic=AsyncAnthropic,
+            bedrock=AsyncAnthropicBedrock,
+            vertex=AsyncAnthropicVertex,
+            aws=AsyncAnthropicAWS,
+            foundry=AsyncAnthropicFoundry,
+        )
+        return self._build_client(factories)
 
     @staticmethod
     def _workload_identity_credentials(wif: dict[str, Any]) -> 
WorkloadIdentityCredentials:
diff --git a/providers/anthropic/tests/unit/anthropic/hooks/test_anthropic.py 
b/providers/anthropic/tests/unit/anthropic/hooks/test_anthropic.py
index 84b4e37a220..487e32825aa 100644
--- a/providers/anthropic/tests/unit/anthropic/hooks/test_anthropic.py
+++ b/providers/anthropic/tests/unit/anthropic/hooks/test_anthropic.py
@@ -45,7 +45,7 @@ from airflow.providers.anthropic.hooks.anthropic import (
 
 pytest.importorskip("anthropic")
 
-from anthropic import BadRequestError
+from anthropic import AsyncAnthropic, BadRequestError, 
WorkloadIdentityCredentials
 from anthropic.types import BetaMonetaryAmount
 from anthropic.types.beta import BetaManagedAgentsServerToolUsage, 
BetaManagedAgentsSessionUsage
 from anthropic.types.beta.beta_managed_agents_cache_creation_usage import (
@@ -835,6 +835,145 @@ class TestAnthropicHookGetConn:
         mock_anthropic.assert_called_once_with(base_url=None)
 
 
+WIF_EXTRA = {
+    "workload_identity": {
+        "identity_token_file": "/var/run/secrets/anthropic.com/token",
+        "federation_rule_id": "fdrl_x",
+        "organization_id": "org_x",
+        "service_account_id": "svac_x",
+    }
+}
+
+
[email protected]
[email protected](AnthropicHook, "get_connection", autospec=True)
[email protected](f"{HOOK_PATH}.get_async_connection", autospec=True)
+class TestAnthropicHookGetAsyncConn:
+    @pytest.mark.parametrize(
+        ("password", "host", "extra", "async_name", "sync_name", 
"expected_kwargs"),
+        [
+            pytest.param(
+                "sk-ant",
+                "https://gw.example";,
+                {},
+                "AsyncAnthropic",
+                "Anthropic",
+                {"api_key": "sk-ant", "base_url": "https://gw.example"},
+                id="anthropic",
+            ),
+            pytest.param(
+                "from-password",
+                None,
+                {"anthropic_client_kwargs": {"api_key": "from-extra", 
"max_retries": 5}},
+                "AsyncAnthropic",
+                "Anthropic",
+                {"api_key": "from-extra", "base_url": None, "max_retries": 5},
+                id="anthropic-client-kwargs",
+            ),
+            pytest.param(
+                None,
+                None,
+                {},
+                "AsyncAnthropic",
+                "Anthropic",
+                {"base_url": None},
+                id="anthropic-sdk-resolves-credentials",
+            ),
+            pytest.param(
+                None,
+                None,
+                {"platform": "bedrock", "aws_region": "us-east-1"},
+                "AsyncAnthropicBedrock",
+                "AnthropicBedrock",
+                {"aws_region": "us-east-1"},
+                id="bedrock",
+            ),
+            pytest.param(
+                None,
+                None,
+                {"platform": "vertex", "project_id": "p1", "region": 
"us-central1"},
+                "AsyncAnthropicVertex",
+                "AnthropicVertex",
+                {"project_id": "p1", "region": "us-central1"},
+                id="vertex",
+            ),
+            pytest.param(
+                None,
+                None,
+                {"platform": "AWS", "aws_region": "us-east-1"},
+                "AsyncAnthropicAWS",
+                "AnthropicAWS",
+                {"aws_region": "us-east-1"},
+                id="aws",
+            ),
+            pytest.param(
+                "azkey",
+                None,
+                {"platform": "foundry", "resource": "r1"},
+                "AsyncAnthropicFoundry",
+                "AnthropicFoundry",
+                {"api_key": "azkey", "resource": "r1"},
+                id="foundry",
+            ),
+        ],
+    )
+    async def test_builds_the_async_client_for_each_platform(
+        self,
+        mock_get_async_connection,
+        mock_get_connection,
+        password,
+        host,
+        extra,
+        async_name,
+        sync_name,
+        expected_kwargs,
+    ):
+        mock_get_async_connection.return_value = _conn(password=password, 
host=host, extra=extra)
+
+        with (
+            mock.patch(f"{HOOK_PATH}.{async_name}", autospec=True) as 
mock_async_client,
+            mock.patch(f"{HOOK_PATH}.{sync_name}", autospec=True) as 
mock_sync_client,
+        ):
+            client = await AnthropicHook().get_async_conn()
+
+        mock_async_client.assert_called_once_with(**expected_kwargs)
+        assert client is mock_async_client.return_value
+        mock_sync_client.assert_not_called()
+        mock_get_connection.assert_not_called()
+
+    async def test_looks_up_the_connection_once_and_asynchronously(
+        self, mock_get_async_connection, mock_get_connection
+    ):
+        mock_get_async_connection.return_value = _conn(extra={"model": 
"claude-from-conn"})
+        hook = AnthropicHook(conn_id="my_anthropic")
+
+        with mock.patch(f"{HOOK_PATH}.AsyncAnthropic", autospec=True):
+            await hook.get_async_conn()
+            await hook.get_async_conn()
+
+        # Through the hook, so a subclass's own connection lookup is honoured.
+        mock_get_async_connection.assert_awaited_once_with("my_anthropic", 
hook=hook)
+        # Later reads, such as the model default, use the same connection.
+        assert hook.default_model == "claude-from-conn"
+        mock_get_connection.assert_not_called()
+
+    async def test_async_client_accepts_the_workload_identity_credential(
+        self, mock_get_async_connection, mock_get_connection
+    ):
+        # Unmocked SDK classes: the async client takes the same synchronous 
WIF credential
+        # the sync client does (the SDK runs its token exchange in a worker 
thread). Building
+        # the client reads no token file and exchanges nothing, so no file or 
network is needed.
+        mock_get_async_connection.return_value = _conn(password=None, 
extra=WIF_EXTRA)
+
+        client = await AnthropicHook().get_async_conn()
+
+        try:
+            assert isinstance(client, AsyncAnthropic)
+            assert isinstance(client.credentials, WorkloadIdentityCredentials)
+        finally:
+            await client.close()
+
+
 class TestAnthropicHookFeatures:
     def _hook_with_client(self, extra=None):
         hook = AnthropicHook()

Reply via email to