# 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.
"""
Give Airflow tools to a `Strands Agents <https://strandsagents.com/>`__ agent.
.. note:: Experimental; see :mod:`airflow.providers.common.ai.tools`.
"""
from __future__ import annotations
import copy
import itertools
from typing import TYPE_CHECKING, Any
try:
# Needed at runtime: @hook reads the event type from the method's annotation.
from strands.hooks import AfterToolCallEvent, BeforeInvocationEvent # noqa: TC002
from strands.plugins import Plugin, hook
from strands.tools.tools import PythonAgentTool
except ImportError as e:
from airflow.providers.common.compat.sdk import AirflowOptionalProviderFeatureException
raise AirflowOptionalProviderFeatureException(e)
from airflow.providers.common.ai.tools import AirflowTool, ToolCallError, collect_tools
from airflow.providers.common.ai.tools._from_toolset import tool_call_scope
from airflow.providers.common.ai.utils.tool_metrics import calling_framework
if TYPE_CHECKING:
from collections.abc import Callable
from pydantic import JsonValue
from strands.types.tools import (
AgentTool,
ToolResult as StrandsToolResult,
ToolResultContent,
ToolSpec,
ToolUse,
)
from airflow.providers.common.ai.tools import ToolProvider
__all__ = ["AirflowTools"]
# Strands refuses two plugins with the same name on one agent, so each instance gets its own.
_instance_numbers = itertools.count(1)
def _to_strands_tool(tool: AirflowTool, current_run: Callable[[], object]) -> PythonAgentTool:
spec: ToolSpec = {
"name": tool.name,
"description": tool.description,
# Strands fills in missing property types and descriptions in place when it
# registers a tool, and the source schema can be a toolset's module-level
# constant that the pydantic-ai path also uses, so hand it a copy.
"inputSchema": {"json": copy.deepcopy(tool.parameters)},
}
async def call_airflow_tool(tool_use: ToolUse, **invocation_state: Any) -> StrandsToolResult:
# Strands gives every model turn of the event loop its own cycle ID.
turn = invocation_state.get("event_loop_cycle_id")
with calling_framework("strands"), tool_call_scope(run=current_run(), turn=turn):
result = await tool.call(tool_use["input"])
return {
"toolUseId": tool_use["toolUseId"],
"status": "error" if result.is_error else "success",
"content": [_content_block(result.content)],
}
return PythonAgentTool(tool.name, spec, call_airflow_tool)
def _content_block(content: JsonValue) -> ToolResultContent:
if isinstance(content, str):
return {"text": content}
return {"json": content}