# Licensed to the Apache Software Foundation (ASF) under one
# or more contributor license agreements. See the NOTICE file
# distributed with this work for additional information
# regarding copyright ownership. The ASF licenses this file
# to you under the Apache License, Version 2.0 (the
# "License"); you may not use this file except in compliance
# with the License. You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing,
# software distributed under the License is distributed on an
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
# KIND, either express or implied. See the License for the
# specific language governing permissions and limitations
# under the License.
from __future__ import annotations
import json
import threading
from contextlib import closing
from typing import TYPE_CHECKING, Any
from botocore.config import Config
from botocore.exceptions import ClientError
from airflow.providers.amazon.aws.hooks.base_aws import AwsBaseHook
from airflow.providers.common.compat.sdk import AirflowOptionalProviderFeatureException
if TYPE_CHECKING:
from airflow.providers.common.ai.exceptions import ManagedAgentInvocationError
from airflow.providers.common.ai.managed_agents.base import (
BaseManagedAgentHook,
ManagedAgentCapabilities,
ManagedAgentRef,
ManagedAgentRequest,
ManagedAgentResponse,
)
else:
try:
from airflow.providers.common.ai.exceptions import ManagedAgentInvocationError
from airflow.providers.common.ai.managed_agents.base import (
BaseManagedAgentHook,
ManagedAgentCapabilities,
ManagedAgentRef,
ManagedAgentResponse,
)
except ImportError:
# The Common AI provider is optional. This module still imports without it, and every
# managed-agent entry point on BedrockAgentCoreHook then says what is missing.
def _needs_common_ai(*args: Any, **kwargs: Any) -> Any:
raise AirflowOptionalProviderFeatureException(
"Consulting an AgentCore runtime as a managed agent needs the 'common.ai' extra of the "
"amazon provider: pip install 'apache-airflow-providers-amazon[common.ai]'."
)
class BaseManagedAgentHook:
"""Stand-in for the Common AI contract base; ``agent()`` names the missing extra."""
agent = _needs_common_ai
ManagedAgentCapabilities = ManagedAgentRef = ManagedAgentResponse = _needs_common_ai
ManagedAgentInvocationError = _needs_common_ai
# Fields the managed-agent contract already covers. Letting ``vendor_options`` carry them would
# let a caller re-target the call, which is a change of authority, not an option.
_RESERVED_OPTIONS = frozenset(
{"agentRuntimeArn", "payload", "contentType", "accept", "runtimeSessionId", "accountId", "mcpSessionId"}
)
# Error codes InvokeAgentRuntime can return that no retry or rephrase will fix. The container's own
# errors come back inside a 200 body, so AgentCore has no error that means "rephrase the prompt".
_TERMINAL_ERROR_CODES = frozenset(
{
"ValidationException",
"ResourceNotFoundException",
"AccessDeniedException",
"ServiceQuotaExceededException",
}
)
# Do not retry an invocation whose effects are unknown; Airflow's task-level retry is the right layer.
_NO_RETRIES = Config(retries={"total_max_attempts": 1})
# The service model's bounds for runtimeSessionId. botocore rejects a shorter value client-side with a
# ParamValidationError, which is not a ClientError and would otherwise escape the error classes.
_SESSION_ID_LENGTH = range(33, 257)
_MAX_RESPONSE_BYTES = 1024 * 1024
_TEXT_KEYS = ("output", "result", "text", "response")
[docs]
class BedrockHook(AwsBaseHook):
"""
Interact with Amazon Bedrock.
Provide thin wrapper around :external+boto3:py:class:`boto3.client("bedrock") <Bedrock.Client>`.
Additional arguments (such as ``aws_conn_id``) may be specified and
are passed down to the underlying AwsBaseHook.
.. seealso::
- :class:`airflow.providers.amazon.aws.hooks.base_aws.AwsBaseHook`
"""
[docs]
client_type = "bedrock"
def __init__(self, *args, **kwargs) -> None:
kwargs["client_type"] = self.client_type
super().__init__(*args, **kwargs)
[docs]
def get_guardrail_id_by_name(self, guardrail_name: str) -> str | None:
"""Get the guardrail ID by name, or None if not found."""
paginator = self.conn.get_paginator("list_guardrails")
for page in paginator.paginate():
for g in page.get("guardrails", []):
if g.get("name") == guardrail_name:
return g["id"]
return None
[docs]
class BedrockRuntimeHook(AwsBaseHook):
"""
Interact with the Amazon Bedrock Runtime.
Provide thin wrapper around :external+boto3:py:class:`boto3.client("bedrock-runtime") <BedrockRuntime.Client>`.
Additional arguments (such as ``aws_conn_id``) may be specified and
are passed down to the underlying AwsBaseHook.
.. seealso::
- :class:`airflow.providers.amazon.aws.hooks.base_aws.AwsBaseHook`
"""
[docs]
client_type = "bedrock-runtime"
def __init__(self, *args, **kwargs) -> None:
kwargs["client_type"] = self.client_type
super().__init__(*args, **kwargs)
[docs]
class BedrockAgentHook(AwsBaseHook):
"""
Interact with the Amazon Agents for Bedrock API.
Provide thin wrapper around :external+boto3:py:class:`boto3.client("bedrock-agent") <AgentsforBedrock.Client>`.
Additional arguments (such as ``aws_conn_id``) may be specified and
are passed down to the underlying AwsBaseHook.
.. seealso::
- :class:`airflow.providers.amazon.aws.hooks.base_aws.AwsBaseHook`
"""
[docs]
client_type = "bedrock-agent"
def __init__(self, *args, **kwargs) -> None:
kwargs["client_type"] = self.client_type
super().__init__(*args, **kwargs)
[docs]
class BedrockAgentRuntimeHook(AwsBaseHook):
"""
Interact with the Amazon Agents for Bedrock API.
Provide thin wrapper around :external+boto3:py:class:`boto3.client("bedrock-agent-runtime") <AgentsforBedrockRuntime.Client>`.
Additional arguments (such as ``aws_conn_id``) may be specified and
are passed down to the underlying AwsBaseHook.
.. seealso::
- :class:`airflow.providers.amazon.aws.hooks.base_aws.AwsBaseHook`
"""
[docs]
client_type = "bedrock-agent-runtime"
def __init__(self, *args, **kwargs) -> None:
kwargs["client_type"] = self.client_type
super().__init__(*args, **kwargs)
[docs]
class BedrockAgentCoreControlHook(AwsBaseHook):
"""
Interact with the Amazon Bedrock AgentCore control plane API.
Provide thin wrapper around :external+boto3:py:class:`boto3.client("bedrock-agentcore-control") <BedrockAgentCoreControl.Client>`.
Additional arguments (such as ``aws_conn_id``) may be specified and
are passed down to the underlying AwsBaseHook.
.. seealso::
- :class:`airflow.providers.amazon.aws.hooks.base_aws.AwsBaseHook`
"""
[docs]
client_type = "bedrock-agentcore-control"
def __init__(self, *args, **kwargs) -> None:
kwargs["client_type"] = self.client_type
super().__init__(*args, **kwargs)
[docs]
class BedrockAgentCoreHook(AwsBaseHook, BaseManagedAgentHook):
"""
Interact with the Amazon Bedrock AgentCore runtime plane API.
Provide thin wrapper around :external+boto3:py:class:`boto3.client("bedrock-agentcore") <BedrockAgentCore.Client>`.
Additional arguments (such as ``aws_conn_id`` and ``config``) may be specified and
are passed down to the underlying AwsBaseHook; the connection's ``config_kwargs`` apply
as they do for every other AWS hook.
With the ``common.ai`` extra installed, the hook also implements the Common AI
managed-agent contract, so ``hook.agent(runtime_arn)`` can be handed to a
``ManagedAgentToolset``. The agent is the runtime ARN; the session is AgentCore's
``runtimeSessionId``, which the service requires to be 33 to 256 characters long. A request
carrying a ``prompt`` is sent as ``{"prompt": ...}`` and a request carrying ``messages`` as
``{"messages": [...]}``, both as ``application/json``. The container behind the runtime
defines its own response shape, so the answer text is taken from the first of ``output``,
``result``, ``text`` or ``response`` that holds a string, or from the ``text_key`` vendor
option when the container's contract is known; otherwise the whole JSON body is returned as
text. The decoded body is always available on ``ManagedAgentResponse.raw``.
A remote invocation may have unknown effects, so the contract methods disable botocore's
retries unless the connection or the caller configured them, and let failures propagate to
Airflow's task-level retry instead. ``ManagedAgentRequest.timeout`` is honored as the botocore
connect and read timeout of the call; a client is built per distinct timeout and reused.
AgentCore has no error that means "rephrase the prompt" (a container's own errors arrive
inside a successful body), so the hook raises terminal errors or lets transient ones
propagate, never :class:`~airflow.providers.common.ai.exceptions.ManagedAgentRejected`.
.. code-block:: python
from airflow.providers.amazon.aws.hooks.bedrock import BedrockAgentCoreHook
from airflow.providers.common.ai.toolsets import ManagedAgentToolset
claims = BedrockAgentCoreHook(aws_conn_id="aws_prod", region_name="us-east-1").agent(
"arn:aws:bedrock-agentcore:us-east-1:123456789012:runtime/claims"
)
toolset = ManagedAgentToolset(
claims, tool_name="ask_claims_agent", description="...", vendor_options={"text_key": "answer"}
)
Two ``vendor_options`` are read by the hook itself rather than forwarded to
``InvokeAgentRuntime``: ``text_key`` names the response field that holds the answer text when
the container's contract is known (a body without a string there is an error), and
``max_response_bytes`` bounds the body read into worker memory (default 1 MiB). Every other
option is passed to the API call as is.
.. seealso::
- :class:`airflow.providers.amazon.aws.hooks.base_aws.AwsBaseHook`
- :class:`airflow.providers.common.ai.managed_agents.base.BaseManagedAgentHook`
"""
[docs]
client_type = "bedrock-agentcore"
def __init__(self, *args, **kwargs) -> None:
kwargs["client_type"] = self.client_type
super().__init__(*args, **kwargs)
self._clients: dict[float | None, Any] = {}
self._clients_lock = threading.Lock()
[docs]
def resolve_agent(self, agent: str) -> ManagedAgentRef:
if not agent.startswith("arn:") or ":runtime/" not in agent:
raise ValueError(f"An AgentCore agent is a runtime ARN, got {agent!r}.")
return ManagedAgentRef(platform=self.agent_platform, name=agent)
[docs]
def get_agent_capabilities(self, agent: str) -> ManagedAgentCapabilities:
return ManagedAgentCapabilities(sessions=True, structured_output=True, trace=True)
[docs]
def invoke_agent(self, agent: str, request: ManagedAgentRequest) -> ManagedAgentResponse:
self.resolve_agent(agent)
if reserved := _RESERVED_OPTIONS.intersection(request.vendor_options):
raise ValueError(f"vendor_options cannot override contract fields: {sorted(reserved)}")
if request.session_id is not None and len(request.session_id) not in _SESSION_ID_LENGTH:
raise ValueError(
f"AgentCore requires a session id of {_SESSION_ID_LENGTH.start} to {_SESSION_ID_LENGTH[-1]} "
f"characters; got {len(request.session_id)}."
)
payload = (
{"prompt": request.prompt} if request.prompt is not None else {"messages": request.as_messages()}
)
kwargs: dict[str, Any] = dict(request.vendor_options)
text_key = kwargs.pop("text_key", None)
if text_key is not None and not isinstance(text_key, str):
raise ValueError(f"vendor_options['text_key'] must be a string, got {text_key!r}.")
max_response_bytes = kwargs.pop("max_response_bytes", _MAX_RESPONSE_BYTES)
if not isinstance(max_response_bytes, int) or max_response_bytes <= 0:
raise ValueError(
f"vendor_options['max_response_bytes'] must be a positive integer, got {max_response_bytes!r}."
)
if request.session_id is not None:
kwargs["runtimeSessionId"] = request.session_id
try:
response = self._get_client(request.timeout).invoke_agent_runtime(
agentRuntimeArn=agent,
payload=json.dumps(payload).encode(),
contentType="application/json",
accept="application/json",
**kwargs,
)
body = self._read_json_body(agent, response, max_response_bytes)
except ClientError as exc:
if exc.response.get("Error", {}).get("Code") in _TERMINAL_ERROR_CODES:
raise ManagedAgentInvocationError(f"{self._describe_call(agent)}: {exc}") from exc
raise # throttling, conflicts, server errors: Airflow's task retry is the right layer
return ManagedAgentResponse(
text=self._extract_text(agent, body, text_key),
raw={**response, "response": body},
structured=None if isinstance(body, str) else body,
session_id=response.get("runtimeSessionId"),
trace_ref=response.get("ResponseMetadata", {}).get("RequestId"),
)
def _get_client(self, timeout: float | None) -> Any:
"""
One boto3 client per distinct request timeout, built on first use and reused.
The timeout is a client setting, so it cannot ride on the hook's shared client, and a
toolset uses one timeout, so this is one client per hook in practice. Clients are
thread-safe, which the toolset relies on when a model issues two calls in one turn.
"""
with self._clients_lock:
client = self._clients.get(timeout)
if client is None:
client = self._clients[timeout] = self.get_client_type(config=self._call_config(timeout))
return client
def _call_config(self, timeout: float | None) -> Config:
"""Return the connection's or caller's botocore config, with retries off unless they set them."""
base = self.config
if base.retries is None:
base = base.merge(_NO_RETRIES)
if timeout is None:
return base
return base.merge(Config(connect_timeout=timeout, read_timeout=timeout))
def _describe_call(self, agent: str) -> str:
return f"AgentCore agent {agent} via connection {self.aws_conn_id!r}"
def _read_json_body(self, agent: str, response: dict[str, Any], max_response_bytes: int) -> Any:
with closing(response["response"]) as stream:
content_type = response.get("contentType", "").split(";", 1)[0].strip().lower()
if content_type != "application/json":
raise ManagedAgentInvocationError(
f"{self._describe_call(agent)} returned {content_type or 'no Content-Type'}; "
"this hook handles application/json only."
)
data = stream.read(max_response_bytes + 1)
if len(data) > max_response_bytes:
raise ManagedAgentInvocationError(
f"{self._describe_call(agent)} returned more than max_response_bytes={max_response_bytes}."
)
try:
return json.loads(data)
except ValueError as exc:
raise ManagedAgentInvocationError(
f"{self._describe_call(agent)} returned a body that is not JSON: {exc}"
) from exc
def _extract_text(self, agent: str, body: Any, text_key: str | None) -> str:
if text_key is not None:
value = body.get(text_key) if isinstance(body, dict) else None
if not isinstance(value, str):
raise ManagedAgentInvocationError(
f"{self._describe_call(agent)} returned no string at text_key={text_key!r}."
)
return value
if isinstance(body, str):
return body
if isinstance(body, dict):
for key in _TEXT_KEYS:
if isinstance(body.get(key), str):
return body[key]
return json.dumps(body, default=str)