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