Source code for airflow.providers.amazon.aws.hooks.bedrock

# 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"
[docs] agent_platform = "aws.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)

Was this entry helpful?