# 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.
"""Behaviour shared by the toolsets this provider ships."""
from __future__ import annotations
import asyncio
import dataclasses
import logging
import threading
from abc import abstractmethod
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Literal, TypeVar
from pydantic_ai.exceptions import ApprovalRequired, CallDeferred, ModelRetry, ToolFailed
from pydantic_ai.messages import ToolReturn
from pydantic_ai.toolsets import DynamicToolset
from pydantic_ai.toolsets.abstract import AbstractToolset
from pydantic_ai.toolsets.wrapper import WrapperToolset
from typing_extensions import ParamSpec
from airflow.providers.common.ai.tools._from_toolset import airflow_tools_from_toolset
from airflow.providers.common.ai.utils.masking import mask_secrets
from airflow.providers.common.ai.utils.tool_metrics import record_tool_call
if TYPE_CHECKING:
from collections.abc import Awaitable, Callable
from pydantic_ai._run_context import RunContext
from pydantic_ai.toolsets import ToolsetFunc
from pydantic_ai.toolsets.abstract import ToolsetTool
from airflow.providers.common.ai.tools import AirflowTool
[docs]
log = logging.getLogger(__name__)
# One blocking call through AirflowToolset.run_blocking at a time in the process, across every
# toolset that uses it. Hooks are not thread-safe in general, and before Airflow 3.2 the channel
# to the supervisor that resolves connections and variables has no lock of its own. Agent
# frameworks run tool calls concurrently, so one lock per instance would not be enough.
_blocking_call_lock = threading.Lock()
# Set on an exception _masked has already stripped, so a second masking layer around the same
# toolset does not log it again.
_STRIPPED = "_airflow_secrets_masked"
# How pydantic-ai pauses a run until a person approves a call or the call runs elsewhere.
_PAUSED = (ApprovalRequired, CallDeferred)
def _call_locked(fn: Callable[P, R], /, *args: P.args, **kwargs: P.kwargs) -> R:
with _blocking_call_lock:
return fn(*args, **kwargs)
async def _mask_call(name: str, call: Awaitable[Any], *, count_as: str | None = None) -> Any:
"""
Await a tool call and mask everything it hands on: its result, or the exception it raised.
An exception usually keeps its type, so retry policies and pydantic-ai's own handling
still recognize it, but its message is masked and the chain of exceptions that caused it
is dropped: frameworks and tracing record a failed call's traceback, cause included. A
retry rule can therefore match the exception's type but not its cause. A failure is
logged first, with its cause, to the task log, which masks it on the way out.
"""
outcome: Literal["executed", "failed"] | None = "failed"
error: Exception | None = None
try:
result = await call
outcome = "executed"
except _PAUSED as e:
# The run pauses for approval or deferred execution; the call has not happened yet.
log.debug("Tool %s is waiting to run", name, exc_info=e)
outcome = None
error = _strip(e)
except (ModelRetry, ToolFailed) as e:
log.debug("Tool %s returned an error for the model", name, exc_info=e)
error = _strip(e)
except Exception as e:
if not getattr(e, _STRIPPED, False):
log.warning("Tool %s failed", name, exc_info=e)
error = _strip(e)
finally:
if count_as and outcome:
record_tool_call(count_as, outcome)
if error is not None:
# Raised outside the except blocks, so Python does not chain the original back on.
raise error
if isinstance(result, ToolReturn):
return dataclasses.replace(
result, return_value=mask_secrets(result.return_value), content=mask_secrets(result.content)
)
return mask_secrets(result)
def _strip(error: Exception) -> Exception:
"""
Mask what ``error`` would print and drop its cause chain.
An exception's message need not come from its ``args``: ``OSError`` formats its
``strerror`` and ``filename``, and a custom ``__str__`` can read any attribute. Those
are masked too, and so are the exceptions inside an exception group. If the message
still holds a registered secret after that, or masking it fails, a ``RuntimeError``
carrying only what could be masked is returned in its place.
"""
error.__cause__ = None
error.__context__ = None
try:
stripped = _mask_attributes(error)
message = str(stripped)
if (masked := mask_secrets(message)) != message:
stripped = RuntimeError(f"{type(error).__name__}: {masked}")
except Exception:
stripped = RuntimeError(f"{type(error).__name__}: details withheld, they could not be masked")
setattr(stripped, _STRIPPED, True)
return stripped
def _mask_attributes(error: Exception) -> Exception:
group = getattr(error, "exceptions", None)
if isinstance(group, tuple) and hasattr(error, "derive"):
# An exception group's own message and arguments are set when it is built.
return type(error)(mask_secrets(getattr(error, "message", "")), [_strip(e) for e in group])
error.args = mask_secrets(error.args)
for attribute, value in vars(error).items():
# Only text and containers: turning a model into a dict could break the __str__ that reads it.
if isinstance(value, (str, bytes, dict, list, tuple, set, frozenset)):
vars(error)[attribute] = mask_secrets(value)
if isinstance(error, OSError):
# Only those that are set: assigning None to an unset one changes how it prints.
for attribute in ("strerror", "filename", "filename2"):
if (value := getattr(error, attribute)) is not None:
setattr(error, attribute, mask_secrets(value))
return error
[docs]
def validate_max_retries(max_retries: int | None) -> int | None:
"""Return ``max_retries`` unchanged, or raise ``ValueError`` if it is negative."""
if max_retries is not None and max_retries < 0:
raise ValueError(f"max_retries must not be negative, got {max_retries}.")
return max_retries
@dataclass
[docs]
def ensure_masked(toolset: AbstractToolset[Any] | ToolsetFunc[Any]) -> AbstractToolset[Any]:
"""
Return ``toolset`` wrapped in :class:`MaskingToolset`, unless it already masks its own output.
A function that builds a toolset for each run, which pydantic-ai also accepts, is wrapped
too. So is an :class:`AirflowToolset` whose ``call_tool`` is overridden, since the override
can bypass the masking.
"""
if not isinstance(toolset, AbstractToolset):
toolset = DynamicToolset(toolset)
elif isinstance(toolset, AirflowToolset) and type(toolset).call_tool is AirflowToolset.call_tool:
return toolset
return MaskingToolset(wrapped=toolset)