This is an automated email from the ASF dual-hosted git repository. kaxil pushed a commit to branch main in repository https://gitbox.apache.org/repos/asf/airflow.git
commit 40ed903b6d6ffa8d3074744d10bc6f8b0bf248bf Author: Kaxil Naik <[email protected]> AuthorDate: Wed Sep 30 07:04:25 2026 +0100 Count and trace tool calls made by native agent frameworks (#73901) Every call to a toolset this provider ships increments common_ai.tool_calls, tagged with the toolset class, the agent framework making the call (pydantic_ai, strands, adk, langchain or none) and the outcome (executed, failed or replayed). agent_framework_tracing() wraps the code that builds and runs a Strands or Google ADK agent. Inside it the framework leaves prompts, completions and tool inputs and outputs out of its spans unless [common.ai] capture_content is on, spans carry the task's identity like AgentOperator's, and with worker tracing off the Dag run's unsampled trace context no longer drops the framework's spans. --- providers/common/ai/docs/frameworks/adk.rst | 9 + providers/common/ai/docs/frameworks/index.rst | 5 +- providers/common/ai/docs/frameworks/strands.rst | 12 ++ providers/common/ai/docs/observability.rst | 90 +++++++- providers/common/ai/docs/stability.rst | 5 + .../providers/common/ai/durable/caching_toolset.py | 17 +- .../common/ai/example_dags/example_adk_agent.py | 5 +- .../ai/example_dags/example_strands_agent.py | 28 +-- .../airflow/providers/common/ai/tools/__init__.py | 5 +- .../src/airflow/providers/common/ai/tools/adk.py | 3 +- .../airflow/providers/common/ai/tools/strands.py | 4 +- .../airflow/providers/common/ai/tools/tracing.py | 192 +++++++++++++++++ .../common/ai/toolsets/langchain_bridge.py | 4 +- .../providers/common/ai/utils/tool_metrics.py | 63 ++++++ .../providers/common/ai/utils/toolset_base.py | 28 ++- .../ai/tests/unit/common/ai/tools/test_adk.py | 10 + .../ai/tests/unit/common/ai/tools/test_strands.py | 10 + .../ai/tests/unit/common/ai/tools/test_tracing.py | 230 +++++++++++++++++++++ .../common/ai/toolsets/test_langchain_bridge.py | 14 ++ .../unit/common/ai/utils/test_tool_metrics.py | 179 ++++++++++++++++ .../observability/metrics/metrics_template.yaml | 8 + 21 files changed, 890 insertions(+), 31 deletions(-) diff --git a/providers/common/ai/docs/frameworks/adk.rst b/providers/common/ai/docs/frameworks/adk.rst index 350e0feb420..a4583bbb344 100644 --- a/providers/common/ai/docs/frameworks/adk.rst +++ b/providers/common/ai/docs/frameworks/adk.rst @@ -77,6 +77,15 @@ passes its result through Airflow's secret masker. Outside ``AgentOperator``, a toolset's connection ID is used as written: it is not rendered as a template. +Tracing +------- + +The example runs the agent inside +:func:`~airflow.providers.common.ai.tools.tracing.agent_framework_tracing`, so ADK's +OpenTelemetry spans carry the task's identity and leave out prompts, completions and tool +inputs and outputs unless ``[common.ai] capture_content`` is on. See +:doc:`../observability`. + Differences from ``AgentOperator`` ---------------------------------- diff --git a/providers/common/ai/docs/frameworks/index.rst b/providers/common/ai/docs/frameworks/index.rst index b118172385e..9ed5632091b 100644 --- a/providers/common/ai/docs/frameworks/index.rst +++ b/providers/common/ai/docs/frameworks/index.rst @@ -75,8 +75,9 @@ is Airflow's and which part stays yours. - No agent - ``[llamaindex]`` extra -The Strands and ADK integrations, and the framework-neutral tool interface under them, are -experimental: they can change or be removed in a minor release of this provider. +The Strands and ADK integrations, the framework-neutral tool interface under them, and +the tracing helper are experimental: they can change or be removed in a minor release of +this provider. Tested versions --------------- diff --git a/providers/common/ai/docs/frameworks/strands.rst b/providers/common/ai/docs/frameworks/strands.rst index be26120f22a..41fd9184512 100644 --- a/providers/common/ai/docs/frameworks/strands.rst +++ b/providers/common/ai/docs/frameworks/strands.rst @@ -147,6 +147,18 @@ the masker. To give one of your own functions the same treatment, wrap it as an ) agent = Agent(model=model, plugins=[AirflowTools(warehouse, lookup)]) +Tracing +------- + +The example runs the agent inside +:func:`~airflow.providers.common.ai.tools.tracing.agent_framework_tracing`, so Strands' +OpenTelemetry spans carry the task's identity and leave out prompts, completions and tool +inputs and outputs unless ``[common.ai] capture_content`` is on. Create the ``Agent`` +inside the block: Strands reads the switch once per process, when it creates its one +tracer, so an ``Agent`` created earlier in the process, such as at module level, keeps +content capture on. See +:doc:`../observability`. + Differences from ``AgentOperator`` ---------------------------------- diff --git a/providers/common/ai/docs/observability.rst b/providers/common/ai/docs/observability.rst index 0484bc1dcb9..828b28d1dda 100644 --- a/providers/common/ai/docs/observability.rst +++ b/providers/common/ai/docs/observability.rst @@ -81,8 +81,9 @@ How it works tool approval (see :doc:`tool_approval`) continues as ``<task-instance id>-resumed``, which is the ``run_id`` the operator pushes; ``usage`` covers both parts. -* **Scope.** The ``airflow.*`` identity attributes and the ``run_id`` / ``usage`` - XComs come only from ``AgentOperator`` and ``@task.agent``. The other LLM +* **Scope.** The ``run_id`` / ``usage`` XComs come only from ``AgentOperator`` and + ``@task.agent``, and so do the ``airflow.*`` identity attributes, apart from a Strands or + ADK agent run inside ``agent_framework_tracing`` (see below). The other LLM operators still emit GenAI spans correlated to the task span by nesting, but without the identity attributes or the run join key. * **Content is off by default.** Only token counts, model id, latency, tool @@ -143,4 +144,89 @@ outputs (``gen_ai.input.messages`` / ``gen_ai.output.messages``), set: in a trusted environment. It has no effect unless ``otel_export_enabled`` is ``True``. +Agents built with other frameworks +---------------------------------- + +.. note:: + + Experimental: ``agent_framework_tracing`` can change or be removed in a minor release + of this provider. + See :ref:`howto/stability`. + +Strands Agents and Google ADK emit OpenTelemetry spans of their own. Run the agent inside +:func:`~airflow.providers.common.ai.tools.tracing.agent_framework_tracing` so those spans +follow the same rules as ``AgentOperator``'s: + +.. code-block:: python + + from airflow.providers.common.ai.tools.strands import AirflowTools + from airflow.providers.common.ai.tools.tracing import agent_framework_tracing + + with agent_framework_tracing(): + agent = Agent(model=model, plugins=[AirflowTools(warehouse)]) + answer = agent(question) + +Inside the block the framework leaves prompts, completions and tool arguments and results +out of its spans unless ``[common.ai] otel_export_enabled`` and ``capture_content`` are +both on, using the switch each +framework reads: the ``gen_ai_unredacted_attributes`` token of +``OTEL_SEMCONV_STABILITY_OPT_IN`` for Strands, and ``ADK_CAPTURE_MESSAGE_CONTENT_IN_SPANS`` +and ``OTEL_INSTRUMENTATION_GENAI_CAPTURE_MESSAGE_CONTENT`` for ADK. A value your deployment +already set wins. Strands reads its switch once per process, when it creates its one tracer, +so create the first ``Agent`` of the task inside the block, not at module level; ADK reads +its switches when a ``TelemetryConfig`` is built, so build that inside the block too. + +Every span started inside the block under the worker's tracer provider carries the same +``airflow.*`` identity attributes as ``AgentOperator``'s spans. When no tracer provider made +a span for the task, the Dag run's trace context is not made the parent of the framework's +spans: that context is marked as not sampled, and a parent-based sampler, the OpenTelemetry +default, would drop every span a tracer provider the framework installs starts. When core +tracing or auto-instrumentation made the task's span, the Dag run's sampling decision +holds. + +The identity attributes go on spans of the tracer provider that is installed when the block +starts, so set up the framework's own telemetry, such as ``StrandsTelemetry``, before +entering it. They follow the task through ``asyncio``; ADK's synchronous ``Runner.run`` +runs the agent on a thread of its own that they do not reach, so use ``run_async``. + +Counting tool calls +------------------- + +.. note:: + + Experimental: the ``common_ai.tool_calls`` metric and its tags can change or be + removed in a minor release of this provider. + See :ref:`howto/stability`. + +Every call to one of this provider's connection-backed toolsets increments the +``common_ai.tool_calls`` counter through Airflow's metrics, whether the call comes from ``AgentOperator``, a +Pydantic AI agent you build yourself, or another framework through its adapter. It +answers which toolsets are used and from where without reading task logs. The counter +carries three tags: + +.. list-table:: + :header-rows: 1 + :widths: 20 80 + + * - Tag + - Values + * - ``toolset`` + - The toolset class, such as ``SQLToolset`` or ``ObjectStorageToolset``. + * - ``framework`` + - ``pydantic_ai`` for ``AgentOperator`` and your own Pydantic AI agents; + ``strands``, ``adk`` or ``langchain`` for a call through that framework's adapter; + ``none`` for a call to a toolset's ``airflow_tools()`` made without an adapter. + * - ``outcome`` + - ``executed`` when the call returned, ``failed`` when it raised (including a + failure the model is asked to correct), and ``replayed`` when + ``AgentOperator(durable=True)`` served the result from its cache on a retry. A call + whose arguments fail validation never reaches the toolset and is not counted, and + neither is a call that pauses the run until a person approves it. + +Tags never include arguments, connection IDs, table names or paths. They reach backends +that support them: OpenTelemetry metrics (``[metrics] otel_on``), or StatsD with +``[metrics] statsd_datadog_enabled`` or ``statsd_influxdb_enabled``. The Agent Skills +toolset, toolsets you write yourself and hand-built ``AirflowTool`` objects are not +counted. + See :doc:`configurations-ref` for the full list of options. diff --git a/providers/common/ai/docs/stability.rst b/providers/common/ai/docs/stability.rst index 78717560fcd..e47ed969296 100644 --- a/providers/common/ai/docs/stability.rst +++ b/providers/common/ai/docs/stability.rst @@ -181,3 +181,8 @@ Everything this provider ships that is not in the table above is experimental. (:doc:`toolsets/hook`) - New; how a pinned argument is matched to each method's parameters may change after first use. + * - The ``common_ai.tool_calls`` metric and + :func:`~airflow.providers.common.ai.tools.tracing.agent_framework_tracing` + (:doc:`observability`) + - The tracing helper follows the agent frameworks' own telemetry, which is still + changing; the metric's tags may change as more frameworks get adapters. diff --git a/providers/common/ai/src/airflow/providers/common/ai/durable/caching_toolset.py b/providers/common/ai/src/airflow/providers/common/ai/durable/caching_toolset.py index 1e1a430d3e6..83a5598ae0f 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/durable/caching_toolset.py +++ b/providers/common/ai/src/airflow/providers/common/ai/durable/caching_toolset.py @@ -26,9 +26,11 @@ from pydantic_ai.toolsets.wrapper import WrapperToolset from airflow.providers.common.ai.durable.base import build_tool_step_key from airflow.providers.common.ai.durable.fingerprint import fingerprint_tool_call +from airflow.providers.common.ai.utils.tool_metrics import record_tool_call +from airflow.providers.common.ai.utils.toolset_base import AirflowToolset if TYPE_CHECKING: - from pydantic_ai.toolsets.abstract import ToolsetTool + from pydantic_ai.toolsets.abstract import AbstractToolset, ToolsetTool from airflow.providers.common.ai.durable.base import DurableStorageProtocol from airflow.providers.common.ai.durable.replay_usage import ReplayUsageLedger @@ -83,6 +85,12 @@ class CachingToolset(WrapperToolset[Any]): log.debug("Durable: replayed cached tool result", step=step, tool=name) if self.replay_usage is not None: self.replay_usage.record_tool_replay(step) + leaf = _innermost(self.wrapped) + if not isinstance(leaf, AirflowToolset): + # Inside a combined or dynamic toolset, the tool knows which one it came from. + leaf = _innermost(tool.toolset) + if isinstance(leaf, AirflowToolset): + record_tool_call(type(leaf).__name__, "replayed") return cached log.warning( "Durable: cached tool result does not match the current tool call; " @@ -113,3 +121,10 @@ class CachingToolset(WrapperToolset[Any]): tool=name, ) return result + + +def _innermost(toolset: AbstractToolset[Any]) -> AbstractToolset[Any]: + """Return the toolset under any wrappers, such as the masking wrapper AgentOperator adds.""" + while isinstance(toolset, WrapperToolset): + toolset = toolset.wrapped + return toolset diff --git a/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_adk_agent.py b/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_adk_agent.py index 68c94afc3c9..fc4bcf888d9 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_adk_agent.py +++ b/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_adk_agent.py @@ -63,6 +63,7 @@ def example_adk_agent(): from google.genai import types from airflow.providers.common.ai.tools.adk import AirflowTools + from airflow.providers.common.ai.tools.tracing import agent_framework_tracing from airflow.providers.common.ai.toolsets.sql import SQLToolset llm = BaseHook.get_connection(LLM_CONN_ID) @@ -91,7 +92,9 @@ def example_adk_agent(): answer = "".join(part.text or "" for part in event.content.parts) return answer - return asyncio.run(ask()) + # Spans carry the task's identity and no prompt text; see the tracing section of the guide. + with agent_framework_tracing(): + return asyncio.run(ask()) run_adk_agent() diff --git a/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_strands_agent.py b/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_strands_agent.py index 7482618002d..faebd874dba 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_strands_agent.py +++ b/providers/common/ai/src/airflow/providers/common/ai/example_dags/example_strands_agent.py @@ -39,7 +39,7 @@ from __future__ import annotations import os -from airflow.providers.common.compat.sdk import dag, task +from airflow.providers.common.compat.sdk import BaseHook, dag, task LLM_CONN_ID = os.environ.get("LLM_CONN_ID", "anthropic_default") LLM_MODEL = os.environ.get("LLM_MODEL", "claude-sonnet-5") @@ -59,8 +59,8 @@ def example_strands_agent(): from strands.models.anthropic import AnthropicModel from airflow.providers.common.ai.tools.strands import AirflowTools + from airflow.providers.common.ai.tools.tracing import agent_framework_tracing from airflow.providers.common.ai.toolsets.sql import SQLToolset - from airflow.providers.common.compat.sdk import BaseHook llm = BaseHook.get_connection(LLM_CONN_ID) model = AnthropicModel( @@ -68,17 +68,19 @@ def example_strands_agent(): model_id=LLM_MODEL, max_tokens=2048, ) - agent = Agent( - model=model, - plugins=[AirflowTools(SQLToolset(db_conn_id=DB_CONN_ID))], - system_prompt=( - "You are a SQL analyst. Use list_tables and get_schema to explore " - "the database, then run read-only queries to answer the question." - ), - # Strands streams the reply to stdout by default; the task returns it instead. - callback_handler=None, - ) - return str(agent(question)) + # Spans carry the task's identity and no prompt text; see the tracing section of the guide. + with agent_framework_tracing(): + agent = Agent( + model=model, + plugins=[AirflowTools(SQLToolset(db_conn_id=DB_CONN_ID))], + system_prompt=( + "You are a SQL analyst. Use list_tables and get_schema to explore " + "the database, then run read-only queries to answer the question." + ), + # Strands streams the reply to stdout by default; the task returns it instead. + callback_handler=None, + ) + return str(agent(question)) run_strands_agent() diff --git a/providers/common/ai/src/airflow/providers/common/ai/tools/__init__.py b/providers/common/ai/src/airflow/providers/common/ai/tools/__init__.py index 5cd50b72beb..b6317aeb691 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/tools/__init__.py +++ b/providers/common/ai/src/airflow/providers/common/ai/tools/__init__.py @@ -45,6 +45,7 @@ from dataclasses import dataclass from typing import TYPE_CHECKING, Any, Protocol from airflow.providers.common.ai.utils.masking import mask_secrets +from airflow.providers.common.ai.utils.tool_metrics import calling_framework, current_framework if TYPE_CHECKING: from pydantic import JsonValue @@ -121,7 +122,9 @@ class AirflowTool: start = time.monotonic() failure: str | None = None try: - result = await self.function(arguments) + # A call that reaches here without a framework adapter is counted as "none". + with calling_framework(current_framework() or "none"): + result = await self.function(arguments) except Exception as e: log.warning("Tool %s failed after %.2fs", self.name, time.monotonic() - start, exc_info=True) failure = ( diff --git a/providers/common/ai/src/airflow/providers/common/ai/tools/adk.py b/providers/common/ai/src/airflow/providers/common/ai/tools/adk.py index 613f8815b52..cef3731be98 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/tools/adk.py +++ b/providers/common/ai/src/airflow/providers/common/ai/tools/adk.py @@ -36,6 +36,7 @@ except ImportError as e: from airflow.providers.common.ai.tools import AirflowTool, 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 google.adk.agents.readonly_context import ReadonlyContext @@ -102,7 +103,7 @@ class _AirflowAdkTool(BaseTool): async def run_async(self, *, args: dict[str, Any], tool_context: ToolContext) -> dict[str, Any]: # ADK does not identify the model turn, so calls that run at the same time count once. - with tool_call_scope(run=tool_context.invocation_id): + with calling_framework("adk"), tool_call_scope(run=tool_context.invocation_id): result = await self._tool.call(args) return {"error": result.content} if result.is_error else {"result": result.content} diff --git a/providers/common/ai/src/airflow/providers/common/ai/tools/strands.py b/providers/common/ai/src/airflow/providers/common/ai/tools/strands.py index c0f0353deac..3df6777269e 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/tools/strands.py +++ b/providers/common/ai/src/airflow/providers/common/ai/tools/strands.py @@ -38,6 +38,7 @@ except ImportError as 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 @@ -136,7 +137,8 @@ def _to_strands_tool(tool: AirflowTool, current_run: Callable[[], object]) -> Py 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. - with tool_call_scope(run=current_run(), turn=invocation_state.get("event_loop_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"], diff --git a/providers/common/ai/src/airflow/providers/common/ai/tools/tracing.py b/providers/common/ai/src/airflow/providers/common/ai/tools/tracing.py new file mode 100644 index 00000000000..4b4a59524e0 --- /dev/null +++ b/providers/common/ai/src/airflow/providers/common/ai/tools/tracing.py @@ -0,0 +1,192 @@ +# 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. +""" +Make the OpenTelemetry spans of an agent framework other than Pydantic AI part of the task. + +.. note:: Experimental; see :mod:`airflow.providers.common.ai.tools`. +""" + +from __future__ import annotations + +import os +import threading +import weakref +from contextlib import contextmanager +from contextvars import ContextVar +from typing import TYPE_CHECKING, Any + +from opentelemetry import context as otel_context, trace +from opentelemetry.sdk.trace import SpanProcessor +from opentelemetry.trace.propagation.tracecontext import TraceContextTextMapPropagator + +from airflow.providers.common.ai.observability import ( + _capture_content, + _live_tracer_provider, + _otel_export_enabled, + build_run_identity_attributes, +) +from airflow.providers.common.compat.sdk import get_current_context + +if TYPE_CHECKING: + from collections.abc import Iterator + + from opentelemetry.context import Context + from opentelemetry.sdk.trace import Span + +__all__ = ["agent_framework_tracing"] + +# Spans started while this is set carry the task's identity. A span processor cannot be +# removed from a provider once added, so it is added once and does nothing outside the block. +_task_identity: ContextVar[dict[str, Any] | None] = ContextVar("common_ai_task_identity", default=None) +_providers_with_identity: weakref.WeakSet[Any] = weakref.WeakSet() + +# The switches each framework reads to leave prompts, completions, and tool arguments and +# results out of its spans. Strands redacts every sensitive attribute when its +# ``gen_ai_unredacted_attributes`` token names none; ADK and OpenTelemetry's GenAI +# instrumentations read the other two. +_OPT_IN = "OTEL_SEMCONV_STABILITY_OPT_IN" +_REDACT_ALL = "gen_ai_unredacted_attributes=" +_CONTENT_OFF = { + "OTEL_INSTRUMENTATION_GENAI_CAPTURE_MESSAGE_CONTENT": "false", + "ADK_CAPTURE_MESSAGE_CONTENT_IN_SPANS": "false", +} + + +@contextmanager +def agent_framework_tracing() -> Iterator[None]: + """ + Make the OpenTelemetry spans of a Strands or Google ADK agent part of the Airflow task. + + Build and run the agent inside the block:: + + with agent_framework_tracing(): + agent = Agent(model=model, plugins=[AirflowTools(warehouse)]) + answer = agent(question) + + Inside it: + + - Prompts, completions, and tool arguments and results are left out of the + framework's spans unless ``[common.ai] otel_export_enabled`` and ``capture_content`` + are both on, as for ``AgentOperator``. A switch the deployment already set in the + environment wins. Strands reads its switch once, when it creates its tracer, and + ADK when a ``TelemetryConfig`` is built, so create both inside the block. + - Every span started under the worker's OpenTelemetry tracer provider carries the + task's Dag ID, run ID, task ID, map index, try number and task instance ID, the + attributes ``AgentOperator``'s spans carry. + - When no tracer provider made a span for the task, the Dag run's trace context is + not made the parent of the framework's spans. That context is marked as not + sampled, and a parent-based sampler, the OpenTelemetry default, would otherwise + drop every span a tracer provider the framework installs starts. When core tracing + or auto-instrumentation made the task's span, the Dag run's sampling decision + holds. + + Where the spans go is up to the tracer provider: core tracing's exporter when + ``[traces] otel_on`` is set, or the provider the framework's own telemetry setup + installs. + """ + provider = _live_tracer_provider() + if provider is not None and provider not in _providers_with_identity: + provider.add_span_processor(_TaskIdentityProcessor()) + _providers_with_identity.add(provider) + + ti = get_current_context()["ti"] + identity = _task_identity.set(build_run_identity_attributes(ti)) + detach = None + try: + current = trace.get_current_span().get_span_context() + if not current.trace_flags.sampled and _is_propagated_parent(current, ti): + # Only the span is replaced, so baggage and the instrumentation-suppression key + # stay in place. + detach = otel_context.attach(trace.set_span_in_context(trace.INVALID_SPAN)) + with _content_off: + yield + finally: + if detach is not None: + otel_context.detach(detach) + _task_identity.reset(identity) + + +def _content_switches() -> dict[str, str]: + # The same rule as AgentOperator: content is captured only when both settings are on. + if _otel_export_enabled() and _capture_content(): + return {} + switches = {name: value for name, value in _CONTENT_OFF.items() if name not in os.environ} + opt_in = os.environ.get(_OPT_IN, "") + if _REDACT_ALL not in opt_in: + switches[_OPT_IN] = ",".join(filter(None, (opt_in, _REDACT_ALL))) + return switches + + +def _is_propagated_parent(current: trace.SpanContext, ti: Any) -> bool: + """ + Whether the current span is the Dag run's propagated context rather than a span of the task. + + It is when no tracer provider made a span for the task, and the core still made the + propagated context current, as Airflow 3.2's task span does with the no-op tracer. + Airflow 3.0 and 3.1 propagate no context to the task. + """ + carrier = getattr(ti, "context_carrier", None) + if not current.is_valid or not carrier: + return False + propagated = trace.get_current_span(TraceContextTextMapPropagator().extract(carrier)).get_span_context() + return (current.trace_id, current.span_id) == (propagated.trace_id, propagated.span_id) + + +class _ContentOff: + """ + Hold the content-off switches in the process environment while any block is open. + + The frameworks read the switches only from the environment. Blocks can overlap, when a + task runs agents in threads or concurrent coroutines, so the switches are set when the + first block opens and restored when the last one closes, not by whichever exits first. + Changing the environment is safe here because a task runs in a process of its own. + """ + + def __init__(self) -> None: + self._lock = threading.Lock() + self._open_blocks = 0 + self._previous: dict[str, str | None] = {} + + def __enter__(self) -> None: + with self._lock: + if self._open_blocks == 0: + switches = _content_switches() + self._previous = {name: os.environ.get(name) for name in switches} + os.environ.update(switches) + self._open_blocks += 1 + + def __exit__(self, *args: object) -> None: + with self._lock: + self._open_blocks -= 1 + if self._open_blocks: + return + for name, value in self._previous.items(): + if value is None: + os.environ.pop(name, None) + else: + os.environ[name] = value + + +_content_off = _ContentOff() + + +class _TaskIdentityProcessor(SpanProcessor): + """Stamp the task's identity on spans started inside ``agent_framework_tracing``.""" + + def on_start(self, span: Span, parent_context: Context | None = None) -> None: + if (identity := _task_identity.get()) is not None: + span.set_attributes(identity) diff --git a/providers/common/ai/src/airflow/providers/common/ai/toolsets/langchain_bridge.py b/providers/common/ai/src/airflow/providers/common/ai/toolsets/langchain_bridge.py index a396c1c5161..c6dad5d8102 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/toolsets/langchain_bridge.py +++ b/providers/common/ai/src/airflow/providers/common/ai/toolsets/langchain_bridge.py @@ -39,6 +39,7 @@ from pydantic import JsonValue # noqa: TC002 from airflow.providers.common.ai.tools._from_toolset import airflow_tools_from_toolset from airflow.providers.common.ai.utils.coroutines import run_coroutine_sync +from airflow.providers.common.ai.utils.tool_metrics import calling_framework if TYPE_CHECKING: from langchain_core.tools import StructuredTool, ToolException @@ -118,7 +119,8 @@ def _to_structured_tool( tool_exception_cls: type[ToolException], ) -> StructuredTool: async def call(**kwargs: Any) -> JsonValue: - result = await tool.call(kwargs) + with calling_framework("langchain"): + result = await tool.call(kwargs) if result.is_error: # With handle_tool_error, LangChain hands this text to the model as an error result. raise tool_exception_cls(str(result.content)) diff --git a/providers/common/ai/src/airflow/providers/common/ai/utils/tool_metrics.py b/providers/common/ai/src/airflow/providers/common/ai/utils/tool_metrics.py new file mode 100644 index 00000000000..e2caf09afc5 --- /dev/null +++ b/providers/common/ai/src/airflow/providers/common/ai/utils/tool_metrics.py @@ -0,0 +1,63 @@ +# 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. +"""Count calls to the toolsets this provider ships, by toolset, agent framework and outcome.""" + +from __future__ import annotations + +from contextlib import contextmanager +from contextvars import ContextVar +from typing import TYPE_CHECKING, Literal + +from airflow.providers.common.compat.sdk import Stats + +if TYPE_CHECKING: + from collections.abc import Iterator + +#: The agent framework a tool call came through, as the ``framework`` tag reports it. +Framework = Literal["pydantic_ai", "strands", "adk", "langchain", "none"] + +# Set by each framework adapter around the calls it makes. +_framework: ContextVar[Framework | None] = ContextVar("common_ai_tool_framework", default=None) + + +@contextmanager +def calling_framework(name: Framework) -> Iterator[None]: + """Attribute the tool calls made inside the block to agent framework ``name``.""" + token = _framework.set(name) + try: + yield + finally: + _framework.reset(token) + + +def current_framework() -> Framework | None: + """Return the framework the current tool call is attributed to, if an adapter set one.""" + return _framework.get() + + +def record_tool_call(toolset: str, outcome: Literal["executed", "failed", "replayed"]) -> None: + """ + Count one call to a toolset. + + Tags stay low-cardinality: the toolset class, the framework and the outcome, never + arguments, connection IDs, table names or paths. A call made outside any adapter is a + Pydantic AI agent's, such as ``AgentOperator``'s. + """ + Stats.incr( + "common_ai.tool_calls", + tags={"toolset": toolset, "framework": _framework.get() or "pydantic_ai", "outcome": outcome}, + ) diff --git a/providers/common/ai/src/airflow/providers/common/ai/utils/toolset_base.py b/providers/common/ai/src/airflow/providers/common/ai/utils/toolset_base.py index f7ffd069440..7f4fe9ebbf0 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/utils/toolset_base.py +++ b/providers/common/ai/src/airflow/providers/common/ai/utils/toolset_base.py @@ -24,7 +24,7 @@ import logging import threading from abc import abstractmethod from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, TypeVar +from typing import TYPE_CHECKING, Any, Literal, TypeVar from pydantic_ai.exceptions import ApprovalRequired, CallDeferred, ModelRetry, ToolFailed from pydantic_ai.messages import ToolReturn @@ -35,6 +35,7 @@ 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 @@ -60,9 +61,8 @@ _blocking_call_lock = threading.Lock() # toolset does not log it again. _STRIPPED = "_airflow_secrets_masked" -# How the model or the run acts on a call without a result, rather than failures: pydantic-ai -# asks the model to correct its call, or pauses the run for approval or deferred execution. -_CONTROL_FLOW = (ModelRetry, ToolFailed, ApprovalRequired, CallDeferred) +# 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: @@ -70,7 +70,7 @@ def _call_locked(fn: Callable[P, R], /, *args: P.args, **kwargs: P.kwargs) -> R: return fn(*args, **kwargs) -async def _masked(name: str, call: Awaitable[Any]) -> Any: +async def _masked(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. @@ -80,16 +80,26 @@ async def _masked(name: str, call: Awaitable[Any]) -> Any: 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 - except _CONTROL_FLOW as e: - log.debug("Tool %s returned no result", name, exc_info=e) + 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 @@ -158,7 +168,9 @@ class AirflowToolset(AbstractToolset[Any]): ctx: RunContext[Any], tool: ToolsetTool[Any], ) -> Any: - return await _masked(name, self._execute_tool(name, tool_args, ctx, tool)) + return await _masked( + name, self._execute_tool(name, tool_args, ctx, tool), count_as=type(self).__name__ + ) @abstractmethod async def _execute_tool( diff --git a/providers/common/ai/tests/unit/common/ai/tools/test_adk.py b/providers/common/ai/tests/unit/common/ai/tools/test_adk.py index 9a4519965ca..00cb64ba766 100644 --- a/providers/common/ai/tests/unit/common/ai/tools/test_adk.py +++ b/providers/common/ai/tests/unit/common/ai/tools/test_adk.py @@ -20,6 +20,7 @@ import asyncio import copy import json from typing import Any +from unittest.mock import MagicMock, patch import pytest @@ -145,3 +146,12 @@ class TestAgentRun: toolset = AirflowTools(_tool(ToolResult(f"key={registered_secret}"))) assert _run_agent(toolset, "lookup", {"key": "a"}) == {"result": "key=***"} + + def test_the_toolsets_calls_are_counted_as_adk(self): + ts = SQLToolset("pg_default") + ts._hook = _make_mock_db_hook() + + with patch("airflow.providers.common.ai.utils.tool_metrics.Stats", MagicMock(spec=["incr"])) as stats: + _run_agent(AirflowTools(ts), "list_tables", {}) + + assert stats.incr.call_args.kwargs["tags"]["framework"] == "adk" diff --git a/providers/common/ai/tests/unit/common/ai/tools/test_strands.py b/providers/common/ai/tests/unit/common/ai/tools/test_strands.py index 65fe787955b..bc08c8b442d 100644 --- a/providers/common/ai/tests/unit/common/ai/tools/test_strands.py +++ b/providers/common/ai/tests/unit/common/ai/tools/test_strands.py @@ -20,6 +20,7 @@ import asyncio import copy import json from typing import Any +from unittest.mock import MagicMock, patch import pytest @@ -195,3 +196,12 @@ class TestAgentRun: assert registered_secret not in answer assert "key=***" in answer + + def test_the_toolsets_calls_are_counted_as_strands(self): + ts = SQLToolset("pg_default") + ts._hook = _make_mock_db_hook() + + with patch("airflow.providers.common.ai.utils.tool_metrics.Stats", MagicMock(spec=["incr"])) as stats: + _run_agent(AirflowTools(ts), "list_tables", {}) + + assert stats.incr.call_args.kwargs["tags"]["framework"] == "strands" diff --git a/providers/common/ai/tests/unit/common/ai/tools/test_tracing.py b/providers/common/ai/tests/unit/common/ai/tools/test_tracing.py new file mode 100644 index 00000000000..df2fafbd3fd --- /dev/null +++ b/providers/common/ai/tests/unit/common/ai/tools/test_tracing.py @@ -0,0 +1,230 @@ +# 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 os +import threading +import uuid +from types import SimpleNamespace +from unittest.mock import patch + +import pytest +from opentelemetry import baggage, context as otel_context, trace +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter +from opentelemetry.trace import NonRecordingSpan, SpanContext, TraceFlags + +from airflow.providers.common.ai.tools.tracing import agent_framework_tracing + +from tests_common.test_utils.config import conf_vars + +MODULE = "airflow.providers.common.ai.tools.tracing" +TI = SimpleNamespace( + dag_id="reports", task_id="summarize", run_id="manual__1", try_number=2, map_index=-1, id=uuid.uuid4() +) +UNSAMPLED = SpanContext( + trace_id=0x4BF92F3577B34DA6A3CE929D0E0E4736, + span_id=0x00F067AA0BA902B7, + is_remote=True, + trace_flags=TraceFlags(TraceFlags.DEFAULT), +) +# The Dag run's trace context, propagated to the task and marked as not sampled. +TI_UNSAMPLED = SimpleNamespace( + **vars(TI), context_carrier={"traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-00"} +) +CONTENT_SWITCHES = ( + "OTEL_INSTRUMENTATION_GENAI_CAPTURE_MESSAGE_CONTENT", + "ADK_CAPTURE_MESSAGE_CONTENT_IN_SPANS", + "OTEL_SEMCONV_STABILITY_OPT_IN", +) + + [email protected] +def exporter(): + """A tracer provider standing in for the worker's, recording what it would export.""" + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + with ( + patch(f"{MODULE}._live_tracer_provider", autospec=True, return_value=provider), + patch(f"{MODULE}.get_current_context", autospec=True, return_value={"ti": TI}), + ): + yield exporter, provider + + [email protected] +def clean_environment(monkeypatch): + for name in CONTENT_SWITCHES: + monkeypatch.delenv(name, raising=False) + + [email protected]("clean_environment") +class TestContentSwitches: + def test_content_is_off_inside_the_block_and_the_environment_is_restored(self, exporter): + with agent_framework_tracing(): + inside = {name: os.environ.get(name) for name in CONTENT_SWITCHES} + + assert inside == { + "OTEL_INSTRUMENTATION_GENAI_CAPTURE_MESSAGE_CONTENT": "false", + "ADK_CAPTURE_MESSAGE_CONTENT_IN_SPANS": "false", + "OTEL_SEMCONV_STABILITY_OPT_IN": "gen_ai_unredacted_attributes=", + } + assert all(name not in os.environ for name in CONTENT_SWITCHES) + + def test_keeps_the_deployments_own_switches(self, exporter, monkeypatch): + monkeypatch.setenv("ADK_CAPTURE_MESSAGE_CONTENT_IN_SPANS", "true") + monkeypatch.setenv("OTEL_SEMCONV_STABILITY_OPT_IN", "gen_ai_latest_experimental") + + with agent_framework_tracing(): + assert os.environ["ADK_CAPTURE_MESSAGE_CONTENT_IN_SPANS"] == "true" + assert os.environ["OTEL_SEMCONV_STABILITY_OPT_IN"] == ( + "gen_ai_latest_experimental,gen_ai_unredacted_attributes=" + ) + + assert os.environ["OTEL_SEMCONV_STABILITY_OPT_IN"] == "gen_ai_latest_experimental" + + def test_overlapping_blocks_keep_content_off_until_the_last_one_exits(self, exporter): + """Two agents running in threads of one task open blocks that need not close in order.""" + second_open, first_closed = threading.Event(), threading.Event() + seen_after_first_closed: list[str | None] = [] + + def first(): + with agent_framework_tracing(): + second_open.wait(5) + first_closed.set() + + def second(): + with agent_framework_tracing(): + second_open.set() + first_closed.wait(5) + seen_after_first_closed.append(os.environ.get("ADK_CAPTURE_MESSAGE_CONTENT_IN_SPANS")) + + threads = [threading.Thread(target=first), threading.Thread(target=second)] + threads[0].start() + threads[1].start() + for thread in threads: + thread.join() + + assert seen_after_first_closed == ["false"] + assert all(name not in os.environ for name in CONTENT_SWITCHES) + + @conf_vars({("common.ai", "otel_export_enabled"): "True", ("common.ai", "capture_content"): "True"}) + def test_captures_content_when_the_deployment_asks_for_it(self, exporter): + with agent_framework_tracing(): + assert all(name not in os.environ for name in CONTENT_SWITCHES) + + @conf_vars({("common.ai", "otel_export_enabled"): "False", ("common.ai", "capture_content"): "True"}) + def test_capture_content_alone_leaves_content_out(self, exporter): + """capture_content has no effect unless otel_export_enabled is on, as for AgentOperator.""" + with agent_framework_tracing(): + assert os.environ.get("ADK_CAPTURE_MESSAGE_CONTENT_IN_SPANS") == "false" + + +class TestTaskIdentity: + def test_spans_started_inside_the_block_carry_the_task(self, exporter): + spans, provider = exporter + tracer = provider.get_tracer("agent_framework") + + with agent_framework_tracing(): + tracer.start_span("inside").end() + tracer.start_span("outside").end() + + by_name = {span.name: dict(span.attributes) for span in spans.get_finished_spans()} + assert by_name["inside"]["airflow.dag_id"] == "reports" + assert by_name["inside"]["airflow.task_instance.try_number"] == 2 + assert by_name["inside"]["airflow.task_instance.id"] == str(TI.id) + assert "airflow.dag_id" not in by_name["outside"] + + def test_the_processor_is_added_once_per_provider(self, exporter): + _, provider = exporter + + with agent_framework_tracing(): + pass + with agent_framework_tracing(): + pass + + processors = provider._active_span_processor._span_processors + assert sum(type(p).__name__ == "_TaskIdentityProcessor" for p in processors) == 1 + + +def _current_inside_block(current: SpanContext, ti: SimpleNamespace, ctx=None): + token = otel_context.attach(trace.set_span_in_context(NonRecordingSpan(current), ctx)) + try: + with ( + patch(f"{MODULE}.get_current_context", new=lambda: {"ti": ti}), + agent_framework_tracing(), + ): + inside = trace.get_current_span().get_span_context() + inside_baggage = baggage.get_baggage("k") + after = trace.get_current_span().get_span_context() + finally: + otel_context.detach(token) + return inside, inside_baggage, after + + +class TestUnsampledParent: + def test_the_propagated_context_is_not_the_parent_when_no_provider_made_a_task_span(self, exporter): + inside, _, after = _current_inside_block(UNSAMPLED, TI_UNSAMPLED) + + assert not inside.is_valid + assert after == UNSAMPLED + + def test_a_task_span_under_the_propagated_context_keeps_the_sampling_decision(self, exporter): + """Core tracing or auto-instrumentation made the task's span; the Dag run was sampled out.""" + task_span = SpanContext( + trace_id=UNSAMPLED.trace_id, + span_id=0x1111111111111111, + is_remote=False, + trace_flags=TraceFlags(TraceFlags.DEFAULT), + ) + + inside, _, _ = _current_inside_block(task_span, TI_UNSAMPLED) + + assert inside == task_span + + def test_without_a_propagated_context_nothing_is_detached(self, exporter): + """Airflow 3.0 and 3.1 propagate no trace context to the task.""" + inside, _, _ = _current_inside_block(UNSAMPLED, TI) + + assert inside == UNSAMPLED + + +class TestParentContext: + def test_a_sampled_propagated_context_is_kept(self, exporter): + sampled = SpanContext( + trace_id=UNSAMPLED.trace_id, + span_id=UNSAMPLED.span_id, + is_remote=True, + trace_flags=TraceFlags(TraceFlags.SAMPLED), + ) + ti = SimpleNamespace( + **vars(TI), + context_carrier={"traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"}, + ) + + inside, _, _ = _current_inside_block(sampled, ti) + + assert inside == sampled + + def test_detaching_the_propagated_context_keeps_baggage(self, exporter): + inside, inside_baggage, _ = _current_inside_block( + UNSAMPLED, TI_UNSAMPLED, baggage.set_baggage("k", "v") + ) + + assert not inside.is_valid + assert inside_baggage == "v" diff --git a/providers/common/ai/tests/unit/common/ai/toolsets/test_langchain_bridge.py b/providers/common/ai/tests/unit/common/ai/toolsets/test_langchain_bridge.py index 973886232d0..011ee910a6d 100644 --- a/providers/common/ai/tests/unit/common/ai/toolsets/test_langchain_bridge.py +++ b/providers/common/ai/tests/unit/common/ai/toolsets/test_langchain_bridge.py @@ -19,6 +19,7 @@ from __future__ import annotations import asyncio import sys from typing import Any, get_type_hints +from unittest.mock import MagicMock, patch import pytest @@ -32,6 +33,9 @@ from pydantic_core import SchemaValidator, core_schema from airflow.providers.common.ai.tools import ToolCallError from airflow.providers.common.ai.toolsets.langchain_bridge import airflow_toolset_to_langchain_tools +from airflow.providers.common.ai.toolsets.sql import SQLToolset + +from unit.common.ai.toolsets.test_sql import _make_mock_db_hook _PASSTHROUGH = SchemaValidator(core_schema.any_schema()) # Coerces the ``n`` field to int so we can assert the args_validator runs. @@ -301,3 +305,13 @@ class TestErrorStatusAndMasking: echo = {t.name: t for t in airflow_toolset_to_langchain_tools(FakeToolset())}["echo"] assert asyncio.run(echo.ainvoke({"text": registered_secret})) == "echo: ***" + + def test_calls_are_counted_as_langchain(self): + ts = SQLToolset("pg_default") + ts._hook = _make_mock_db_hook() + list_tables = {t.name: t for t in airflow_toolset_to_langchain_tools(ts)}["list_tables"] + + with patch("airflow.providers.common.ai.utils.tool_metrics.Stats", MagicMock(spec=["incr"])) as stats: + list_tables.invoke({}) + + assert stats.incr.call_args.kwargs["tags"]["framework"] == "langchain" diff --git a/providers/common/ai/tests/unit/common/ai/utils/test_tool_metrics.py b/providers/common/ai/tests/unit/common/ai/utils/test_tool_metrics.py new file mode 100644 index 00000000000..e1a5fd23f57 --- /dev/null +++ b/providers/common/ai/tests/unit/common/ai/utils/test_tool_metrics.py @@ -0,0 +1,179 @@ +# 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 asyncio +from unittest.mock import MagicMock, call, patch + +import pytest +from pydantic_ai import RunContext +from pydantic_ai.exceptions import ApprovalRequired, CallDeferred, ModelRetry +from pydantic_ai.models.test import TestModel +from pydantic_ai.toolsets.combined import CombinedToolset +from pydantic_ai.toolsets.function import FunctionToolset +from pydantic_ai.usage import RunUsage + +from airflow.providers.common.ai.durable.caching_toolset import CachingToolset +from airflow.providers.common.ai.durable.step_counter import DurableStepCounter +from airflow.providers.common.ai.toolsets.sql import SQLToolset +from airflow.providers.common.ai.utils.tool_metrics import ( + calling_framework, + record_tool_call, +) +from airflow.providers.common.ai.utils.toolset_base import MaskingToolset, with_masking + +from unit.common.ai.operators.test_agent import _InMemoryDurableStorage +from unit.common.ai.toolsets.test_sql import _make_mock_db_hook +from unit.common.ai.utils.test_toolset_base import _call as _call_scripted, _ScriptedToolset + + [email protected] +def stats(): + with patch("airflow.providers.common.ai.utils.tool_metrics.Stats", MagicMock(spec=["incr"])) as mock: + yield mock + + +def _tags(outcome: str, framework: str = "pydantic_ai", toolset: str = "SQLToolset") -> dict[str, str]: + return {"toolset": toolset, "framework": framework, "outcome": outcome} + + +def _sql_toolset(**hook_kwargs) -> SQLToolset: + ts = SQLToolset("pg_default") + ts._hook = _make_mock_db_hook(**hook_kwargs) + return ts + + +def _call(toolset, name: str, args: dict): + async def run(): + ctx = RunContext(deps=None, model=TestModel(), usage=RunUsage()) + tools = await toolset.get_tools(ctx) + return await toolset.call_tool(name, args, ctx, tools[name]) + + return asyncio.run(run()) + + +class TestRecordToolCall: + def test_a_call_outside_any_adapter_is_pydantic_ais(self, stats): + record_tool_call("SQLToolset", "executed") + + stats.incr.assert_called_once_with("common_ai.tool_calls", tags=_tags("executed")) + + def test_an_adapter_names_its_framework_for_the_calls_inside(self, stats): + with calling_framework("strands"): + record_tool_call("SQLToolset", "executed") + record_tool_call("SQLToolset", "executed") + + assert stats.incr.call_args_list == [ + call("common_ai.tool_calls", tags=_tags("executed", framework="strands")), + call("common_ai.tool_calls", tags=_tags("executed")), + ] + + +class TestToolsetsCountTheirCalls: + def test_a_call_that_returns_is_executed(self, stats): + _call(_sql_toolset(), "list_tables", {}) + + stats.incr.assert_called_once_with("common_ai.tool_calls", tags=_tags("executed")) + + def test_a_call_that_raises_is_failed(self, stats): + ts = _sql_toolset() + ts._hook.run.side_effect = ConnectionError("down") + + with pytest.raises(ModelRetry): + _call(ts, "query", {"sql": "SELECT 1"}) + + stats.incr.assert_called_once_with("common_ai.tool_calls", tags=_tags("failed")) + + def test_a_call_through_the_neutral_interface_without_an_adapter_is_none(self, stats): + tool = {t.name: t for t in _sql_toolset().airflow_tools()}["list_tables"] + + asyncio.run(tool.call({})) + + stats.incr.assert_called_once_with("common_ai.tool_calls", tags=_tags("executed", framework="none")) + + def test_a_toolset_the_dag_author_wrote_is_not_counted(self, stats): + """The metric measures this provider's toolsets; the masking wrapper does not count.""" + + def ping() -> str: + return "pong" + + _call(MaskingToolset(wrapped=FunctionToolset([ping])), "ping", {}) + + stats.incr.assert_not_called() + + def test_a_durable_replay_is_counted_as_replayed_not_executed(self, stats): + storage = _InMemoryDurableStorage() + for _ in range(2): + cached = CachingToolset( + wrapped=with_masking(_sql_toolset()), storage=storage, counter=DurableStepCounter() + ) + ctx = RunContext(deps=None, model=TestModel(), usage=RunUsage(), tool_call_id="c1") + + async def run(toolset=cached, ctx=ctx): + tools = await toolset.get_tools(ctx) + return await toolset.call_tool("list_tables", {}, ctx, tools["list_tables"]) + + asyncio.run(run()) + + assert [c.kwargs["tags"]["outcome"] for c in stats.incr.call_args_list] == ["executed", "replayed"] + + +class TestOutcomes: + @pytest.mark.parametrize( + ("raised", "outcomes"), + [ + pytest.param(ApprovalRequired(), [], id="paused_for_approval"), + pytest.param(CallDeferred(), [], id="deferred"), + pytest.param(ModelRetry("fix it"), ["failed"], id="model_retry"), + pytest.param(RuntimeError("boom"), ["failed"], id="error"), + ], + ) + def test_a_paused_call_is_not_counted_and_a_failed_one_is(self, stats, raised, outcomes): + with pytest.raises(type(raised)): + _call_scripted(_ScriptedToolset(raised)) + + assert [c.kwargs["tags"]["outcome"] for c in stats.incr.call_args_list] == outcomes + + def test_a_replay_through_a_wrapper_counts_the_toolset_underneath(self, stats): + storage = _InMemoryDurableStorage() + for _ in range(2): + wrapped = with_masking(_sql_toolset().prefixed("wh")) + cached = CachingToolset(wrapped=wrapped, storage=storage, counter=DurableStepCounter()) + ctx = RunContext(deps=None, model=TestModel(), usage=RunUsage(), tool_call_id="c1") + + async def run(toolset=cached, ctx=ctx): + tools = await toolset.get_tools(ctx) + return await toolset.call_tool("wh_list_tables", {}, ctx, tools["wh_list_tables"]) + + asyncio.run(run()) + + assert [c.kwargs["tags"]["outcome"] for c in stats.incr.call_args_list] == ["executed", "replayed"] + + def test_a_replay_inside_a_combined_toolset_counts_the_toolset_it_came_from(self, stats): + storage = _InMemoryDurableStorage() + for _ in range(2): + combined = CombinedToolset([_sql_toolset(), FunctionToolset([])]) + cached = CachingToolset(wrapped=combined, storage=storage, counter=DurableStepCounter()) + ctx = RunContext(deps=None, model=TestModel(), usage=RunUsage(), tool_call_id="c1") + + async def run(toolset=cached, ctx=ctx): + tools = await toolset.get_tools(ctx) + return await toolset.call_tool("list_tables", {}, ctx, tools["list_tables"]) + + asyncio.run(run()) + + assert [c.kwargs["tags"]["outcome"] for c in stats.incr.call_args_list] == ["executed", "replayed"] diff --git a/shared/observability/src/airflow_shared/observability/metrics/metrics_template.yaml b/shared/observability/src/airflow_shared/observability/metrics/metrics_template.yaml index fde9ce33205..6e1182bf281 100644 --- a/shared/observability/src/airflow_shared/observability/metrics/metrics_template.yaml +++ b/shared/observability/src/airflow_shared/observability/metrics/metrics_template.yaml @@ -416,6 +416,14 @@ metrics: legacy_name: "-" name_variables: ["tool", "platform", "role", "position"] + - name: "common_ai.tool_calls" + description: "Number of calls to a toolset of the common.ai provider. Metric with toolset + (the toolset class), framework (``pydantic_ai``, ``strands``, ``adk``, ``langchain`` or + ``none``) and outcome (``executed``, ``failed`` or ``replayed``) tagging." + type: "counter" + legacy_name: "-" + name_variables: ["toolset", "framework", "outcome"] + # ========== # Gauges # ==========
