Lee-W commented on code in PR #73532:
URL: https://github.com/apache/airflow/pull/73532#discussion_r4153052951


##########
providers/amazon/src/airflow/providers/amazon/aws/hooks/bedrock.py:
##########
@@ -128,21 +195,179 @@ def __init__(self, *args, **kwargs) -> None:
         super().__init__(*args, **kwargs)
 
 
-class BedrockAgentCoreHook(AwsBaseHook):
+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``) may be specified and
-    are passed down to the underlying AwsBaseHook.
+    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`
     """
 
     client_type = "bedrock-agentcore"
+    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()
+
+    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)
+
+    def agent_capabilities(self, agent: str) -> ManagedAgentCapabilities:
+        return ManagedAgentCapabilities(sessions=True, structured_output=True, 
trace=True)
+
+    def invoke_agent(self, agent: str, request: ManagedAgentRequest) -> 
ManagedAgentResponse:
+        self.resolve_agent(agent)
+        reserved = _RESERVED_OPTIONS.intersection(request.vendor_options)
+        if reserved:
+            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._client_for(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._where(agent)}: 
{exc}") from exc
+            raise  # throttling, conflicts, server errors: Airflow's task 
retry is the right layer
+        raw = {**response, "response": body}
+        return ManagedAgentResponse(
+            text=self._text(agent, body, text_key),
+            raw=raw,

Review Comment:
   ```suggestion
           return ManagedAgentResponse(
               text=self._text(agent, body, text_key),
               raw={**response, "response": body},
   ```



##########
providers/amazon/src/airflow/providers/amazon/aws/hooks/bedrock.py:
##########
@@ -128,21 +195,179 @@ def __init__(self, *args, **kwargs) -> None:
         super().__init__(*args, **kwargs)
 
 
-class BedrockAgentCoreHook(AwsBaseHook):
+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``) may be specified and
-    are passed down to the underlying AwsBaseHook.
+    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`
     """
 
     client_type = "bedrock-agentcore"
+    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()
+
+    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)
+
+    def agent_capabilities(self, agent: str) -> ManagedAgentCapabilities:

Review Comment:
   ```suggestion
       def get_agent_capabilities(self, agent: str) -> ManagedAgentCapabilities:
   ```



##########
providers/amazon/src/airflow/providers/amazon/aws/hooks/bedrock.py:
##########
@@ -128,21 +195,179 @@ def __init__(self, *args, **kwargs) -> None:
         super().__init__(*args, **kwargs)
 
 
-class BedrockAgentCoreHook(AwsBaseHook):
+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``) may be specified and
-    are passed down to the underlying AwsBaseHook.
+    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`
     """
 
     client_type = "bedrock-agentcore"
+    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()
+
+    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)
+
+    def agent_capabilities(self, agent: str) -> ManagedAgentCapabilities:
+        return ManagedAgentCapabilities(sessions=True, structured_output=True, 
trace=True)
+
+    def invoke_agent(self, agent: str, request: ManagedAgentRequest) -> 
ManagedAgentResponse:
+        self.resolve_agent(agent)
+        reserved = _RESERVED_OPTIONS.intersection(request.vendor_options)
+        if reserved:
+            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._client_for(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._where(agent)}: 
{exc}") from exc
+            raise  # throttling, conflicts, server errors: Airflow's task 
retry is the right layer
+        raw = {**response, "response": body}
+        return ManagedAgentResponse(
+            text=self._text(agent, body, text_key),
+            raw=raw,
+            structured=None if isinstance(body, str) else body,
+            session_id=response.get("runtimeSessionId"),
+            trace_ref=response.get("ResponseMetadata", {}).get("RequestId"),
+        )
+
+    def _client_for(self, timeout: float | None) -> Any:

Review Comment:
   ```suggestion
       def _get_client(self, timeout: float | None) -> Any:
   ```



##########
providers/amazon/src/airflow/providers/amazon/aws/hooks/bedrock.py:
##########
@@ -128,21 +195,179 @@ def __init__(self, *args, **kwargs) -> None:
         super().__init__(*args, **kwargs)
 
 
-class BedrockAgentCoreHook(AwsBaseHook):
+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``) may be specified and
-    are passed down to the underlying AwsBaseHook.
+    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`
     """
 
     client_type = "bedrock-agentcore"
+    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()
+
+    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)
+
+    def agent_capabilities(self, agent: str) -> ManagedAgentCapabilities:
+        return ManagedAgentCapabilities(sessions=True, structured_output=True, 
trace=True)
+
+    def invoke_agent(self, agent: str, request: ManagedAgentRequest) -> 
ManagedAgentResponse:
+        self.resolve_agent(agent)
+        reserved = _RESERVED_OPTIONS.intersection(request.vendor_options)
+        if reserved:

Review Comment:
   ```suggestion
           if (reserved := 
_RESERVED_OPTIONS.intersection(request.vendor_options)):
   ```



##########
providers/amazon/src/airflow/providers/amazon/aws/hooks/bedrock.py:
##########
@@ -128,21 +195,179 @@ def __init__(self, *args, **kwargs) -> None:
         super().__init__(*args, **kwargs)
 
 
-class BedrockAgentCoreHook(AwsBaseHook):
+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``) may be specified and
-    are passed down to the underlying AwsBaseHook.
+    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`
     """
 
     client_type = "bedrock-agentcore"
+    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()
+
+    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)
+
+    def agent_capabilities(self, agent: str) -> ManagedAgentCapabilities:
+        return ManagedAgentCapabilities(sessions=True, structured_output=True, 
trace=True)
+
+    def invoke_agent(self, agent: str, request: ManagedAgentRequest) -> 
ManagedAgentResponse:
+        self.resolve_agent(agent)
+        reserved = _RESERVED_OPTIONS.intersection(request.vendor_options)
+        if reserved:
+            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._client_for(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._where(agent)}: 
{exc}") from exc
+            raise  # throttling, conflicts, server errors: Airflow's task 
retry is the right layer
+        raw = {**response, "response": body}
+        return ManagedAgentResponse(
+            text=self._text(agent, body, text_key),
+            raw=raw,
+            structured=None if isinstance(body, str) else body,
+            session_id=response.get("runtimeSessionId"),
+            trace_ref=response.get("ResponseMetadata", {}).get("RequestId"),
+        )
+
+    def _client_for(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 _where(self, agent: str) -> str:
+        return f"AgentCore agent {agent} via connection {self.aws_conn_id!r}"

Review Comment:
   not sure why this method is named where



##########
providers/amazon/tests/unit/amazon/aws/hooks/test_bedrock.py:
##########
@@ -69,3 +91,245 @@ def test_not_found(self, mock_conn):
 
         hook = BedrockHook()
         assert hook.get_guardrail_id_by_name("nonexistent") is None
+
+
+ARN = "arn:aws:bedrock-agentcore:us-east-1:123456789012:runtime/investigator"
+SESSION = "caller-session-0123456789abcdef0123456789abcdef"  # AgentCore needs 
33 to 256 characters
+
+
[email protected]
+def invoke_agent_runtime():
+    """The real hook and botocore client, stubbed at the API call."""
+    with mock.patch.dict("os.environ", {"AWS_EC2_METADATA_DISABLED": "true"}):
+        client = Session().create_client(
+            "bedrock-agentcore",
+            region_name="us-east-1",
+            aws_access_key_id="test",
+            aws_secret_access_key="test",
+        )
+        with mock.patch.object(BedrockAgentCoreHook, "get_client_type", 
autospec=True, return_value=client):
+            with mock.patch.object(client, "invoke_agent_runtime", 
autospec=True) as call:
+                yield call
+        client.close()
+
+
+def respond(call, body, **extra):
+    stream = io.BytesIO(json.dumps(body).encode())
+    call.return_value = {
+        "response": stream,
+        "contentType": "application/json",
+        "runtimeSessionId": "session-from-aws",
+        "ResponseMetadata": {"RequestId": "req-1"},
+        **extra,
+    }
+    return stream
+
+
+def hook(**kwargs) -> BedrockAgentCoreHook:

Review Comment:
   ```suggestion
   def create_bedrock_agent_core_hook(**kwargs) -> BedrockAgentCoreHook:
   ```



##########
providers/common/ai/src/airflow/providers/common/ai/toolsets/managed_agent.py:
##########
@@ -239,132 +217,120 @@ async def execute_tool(
         ctx: RunContext[Any],
         tool: ToolsetTool[Any],
     ) -> Any:
-        ref = self.agent_ref
-        log.info("Consulting managed agent %s on %s", ref.get("name"), 
ref.get("platform"))
         result = await self.invoke(tool_args["prompt"])
+        # Identity is resolved after the call, never before it: a toolset 
whose identity comes
+        # from a misconfigured connection must not fail a call that would have 
succeeded, and
+        # a failover group's identity joins every member's, standbys included.
+        ref = self._safe_agent_ref()
+        log.info(
+            "Consulted managed agent %s",
+            f"{ref.name} on {ref.platform}"
+            if ref is not None
+            else f"<unresolved identity> for tool {self._tool_name}",
+        )
+        # Emitted once per answer so managed-agent call volume is observable 
next to the
+        # ``managed_agent.failover`` counter. Tagged by platform to bound 
cardinality.
+        Stats.incr(
+            "managed_agent.served",
+            tags={"tool": self._tool_name, "platform": ref.platform if ref is 
not None else "unknown"},
+        )
         return serialize_for_llm(result)
 
+    def _safe_agent_ref(self) -> ManagedAgentRef | None:

Review Comment:
   ```suggestion
       def _resolve_agent_ref(self) -> ManagedAgentRef | None:
   ```



##########
providers/amazon/src/airflow/providers/amazon/aws/hooks/bedrock.py:
##########
@@ -128,21 +195,179 @@ def __init__(self, *args, **kwargs) -> None:
         super().__init__(*args, **kwargs)
 
 
-class BedrockAgentCoreHook(AwsBaseHook):
+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``) may be specified and
-    are passed down to the underlying AwsBaseHook.
+    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`
     """
 
     client_type = "bedrock-agentcore"
+    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()
+
+    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)
+
+    def agent_capabilities(self, agent: str) -> ManagedAgentCapabilities:
+        return ManagedAgentCapabilities(sessions=True, structured_output=True, 
trace=True)
+
+    def invoke_agent(self, agent: str, request: ManagedAgentRequest) -> 
ManagedAgentResponse:
+        self.resolve_agent(agent)
+        reserved = _RESERVED_OPTIONS.intersection(request.vendor_options)
+        if reserved:
+            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._client_for(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._where(agent)}: 
{exc}") from exc
+            raise  # throttling, conflicts, server errors: Airflow's task 
retry is the right layer
+        raw = {**response, "response": body}
+        return ManagedAgentResponse(
+            text=self._text(agent, body, text_key),
+            raw=raw,
+            structured=None if isinstance(body, str) else body,
+            session_id=response.get("runtimeSessionId"),
+            trace_ref=response.get("ResponseMetadata", {}).get("RequestId"),
+        )
+
+    def _client_for(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 _where(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._where(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._where(agent)} returned more than 
max_response_bytes={max_response_bytes}."
+            )
+        try:
+            return json.loads(data)
+        except ValueError as exc:
+            raise ManagedAgentInvocationError(
+                f"{self._where(agent)} returned a body that is not JSON: {exc}"
+            ) from exc
+
+    def _text(self, agent: str, body: Any, text_key: str | None) -> str:

Review Comment:
   ```suggestion
       def _extract_text(self, agent: str, body: Any, text_key: str | None) -> 
str:
   ```



##########
providers/common/ai/src/airflow/providers/common/ai/managed_agents/failover.py:
##########
@@ -0,0 +1,145 @@
+# 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 logging
+from collections.abc import Sequence
+
+from airflow.providers.common.ai.exceptions import ManagedAgentRejected
+from airflow.providers.common.ai.managed_agents.base import (
+    ManagedAgentCapabilities,
+    ManagedAgentClient,
+    ManagedAgentRef,
+    ManagedAgentRequest,
+    ManagedAgentResponse,
+    describe,
+)
+from airflow.providers.common.compat.sdk import Stats
+
+log = logging.getLogger(__name__)
+
+
+class FailoverManagedAgentClient:
+    """
+    Active/passive failover across interchangeable managed agents.
+
+    Members are tried in order and the first answer wins. The group is itself a
+    
:class:`~airflow.providers.common.ai.managed_agents.base.ManagedAgentClient`, so
+    a toolset built over it presents a single tool and the calling model has 
no say in
+    which provider serves the request: the policy stays deterministic Python 
rather than a
+    prompt instruction a model may ignore. Groups nest.
+
+    Members must be *substitutable*: the same agent deployed twice, not two 
specialists over
+    different data. Two containerized agents built from one image qualify; an 
agent bound to
+    one platform's own objects does not, because there is nothing equivalent 
to fail over to.
+    The group cannot check this.
+
+    What it can check is conversation state. A failover starts a fresh 
conversation on the
+    standby, which is correct for a one-shot consultation and wrong for a 
multi-turn one, so
+    the group never reports ``capabilities.sessions`` and refuses a request 
that carries a
+    ``session_id``, whatever its members support.
+
+    :param members: Interchangeable clients, tried in order. At least two.
+    :param failover_on: Exception types that move to the next member. Defaults 
to
+        ``Exception`` because ``common.ai`` cannot enumerate the cloud SDKs' 
exception trees.
+        Narrow it when the members' exception types are known.
+        :class:`~airflow.providers.common.ai.exceptions.ManagedAgentRejected` 
never
+        triggers failover, whatever this is set to: the standby would reject 
the same prompt.
+    """
+
+    def __init__(
+        self,
+        members: Sequence[ManagedAgentClient],
+        *,
+        failover_on: tuple[type[BaseException], ...] = (Exception,),
+    ) -> None:
+        if len(members) < 2:
+            raise ValueError(
+                f"A failover group needs at least two members; got 
{len(members)}. "
+                "Use the member directly instead."
+            )
+        # Copied, not aliased: a caller holding the original list could 
otherwise empty it.
+        self._members = tuple(members)
+        self._failover_on = failover_on
+
+    @property
+    def members(self) -> tuple[ManagedAgentClient, ...]:
+        return self._members
+
+    @property
+    def ref(self) -> ManagedAgentRef:
+        return ManagedAgentRef(platform="failover", name=" -> 
".join(describe(m) for m in self._members))
+
+    @property
+    def capabilities(self) -> ManagedAgentCapabilities:
+        """The intersection of the members', except sessions, which a group 
never offers."""
+        caps = [m.capabilities for m in self._members]
+        return ManagedAgentCapabilities(
+            sessions=False,
+            structured_output=all(c.structured_output for c in caps),
+            usage=all(c.usage for c in caps),
+            trace=all(c.trace for c in caps),
+        )
+
+    def invoke(self, request: ManagedAgentRequest) -> ManagedAgentResponse:
+        if request.session_id is not None:
+            raise ValueError(
+                "A failover group cannot continue a conversation: the standby 
would start a fresh one. "
+                "Send session-bound requests to one member directly."
+            )
+        for position, member in enumerate(self._members[:-1]):
+            try:
+                response = member.invoke(request)
+            except ManagedAgentRejected:
+                raise
+            except self._failover_on:
+                standby = self._members[position + 1]
+                log.warning(
+                    "Managed agent %s (member %d) failed; failing over to %s 
(member %d)",
+                    describe(member),
+                    position,
+                    describe(standby),
+                    position + 1,
+                    exc_info=True,
+                )
+                # A failover is a success-shaped event: without a counter, a 
primary that has
+                # been down for a week looks identical to a healthy one. 
Tagged by platform
+                # rather than agent name to keep cardinality bounded.
+                Stats.incr(
+                    "managed_agent.failover",
+                    tags={"from_platform": _platform(member), "to_platform": 
_platform(standby)},
+                )
+                continue
+            return self._served(position, member, response)
+        # The last member gets no failover: whatever it raises is the group's 
answer.
+        last = len(self._members) - 1
+        return self._served(last, self._members[last], 
self._members[last].invoke(request))
+
+    @staticmethod
+    def _served(
+        position: int, member: ManagedAgentClient, response: 
ManagedAgentResponse
+    ) -> ManagedAgentResponse:
+        if position > 0:
+            log.info("Managed agent request served by standby %s (member %d)", 
describe(member), position)
+        return response
+
+
+def _platform(client: ManagedAgentClient) -> str:

Review Comment:
   ```suggestion
   def _get_platform(client: ManagedAgentClient) -> str:
   ```



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]

Reply via email to