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 023f859e5a60fabfc98594ecfbada2656d8123fe Author: Kaxil Naik <[email protected]> AuthorDate: Wed Sep 30 07:04:23 2026 +0100 Mask secrets in tool results before they reach the model (#73897) The SQL, hook, DataFusion, MCP, sandbox and managed-agent toolsets now pass what they return, and any exception they raise, through Airflow's secret masker. A connection password that shows up in a database error or a hook's return value previously went to the model, its provider and traces as is. Structured results are masked before they are serialized, since JSON escaping would hide a secret containing a quote, backslash or non-ASCII character from the masker. An exception keeps its type, but its message is masked and its cause chain is dropped; pydantic-ai's approval and deferral signals pass through as control flow. AgentOperator also masks the output of the toolsets passed in toolsets, in agent_params["toolsets"] and in a Toolset capability, including toolsets the Dag author wrote. Blocking hook calls now run in a worker thread, one at a time per process, instead of on the event loop. --- providers/common/ai/docs/agent_security.rst | 16 ++ providers/common/ai/docs/toolsets/hook.rst | 9 +- .../airflow/providers/common/ai/operators/agent.py | 13 +- .../airflow/providers/common/ai/policies/retry.py | 7 +- .../providers/common/ai/toolsets/datafusion.py | 25 +- .../airflow/providers/common/ai/toolsets/hook.py | 16 +- .../providers/common/ai/toolsets/managed_agent.py | 17 +- .../airflow/providers/common/ai/toolsets/mcp.py | 18 +- .../providers/common/ai/toolsets/sandbox.py | 11 +- .../airflow/providers/common/ai/toolsets/sql.py | 41 +-- .../providers/common/ai/utils/file_analysis.py | 9 +- .../airflow/providers/common/ai/utils/masking.py | 120 +++++++++ .../providers/common/ai/utils/query_results.py | 7 +- .../providers/common/ai/utils/tool_definition.py | 13 +- .../providers/common/ai/utils/toolset_base.py | 211 ++++++++++++++++ .../tests/unit/common/ai/decorators/test_agent.py | 3 +- .../tests/unit/common/ai/operators/test_agent.py | 99 +++++++- .../unit/common/ai/toolsets/test_datafusion.py | 16 +- .../ai/tests/unit/common/ai/toolsets/test_hook.py | 45 +++- .../tests/unit/common/ai/toolsets/test_sandbox.py | 18 +- .../ai/tests/unit/common/ai/toolsets/test_sql.py | 54 ++++ .../unit/common/ai/utils/test_file_analysis.py | 21 ++ .../ai/tests/unit/common/ai/utils/test_masking.py | 117 +++++++++ .../unit/common/ai/utils/test_query_results.py | 9 + .../unit/common/ai/utils/test_tool_definition.py | 21 +- .../unit/common/ai/utils/test_toolset_base.py | 281 +++++++++++++++++++++ .../in_container/run_provider_yaml_files_check.py | 2 + 27 files changed, 1127 insertions(+), 92 deletions(-) diff --git a/providers/common/ai/docs/agent_security.rst b/providers/common/ai/docs/agent_security.rst index 41b1e236245..ed1bd46b22a 100644 --- a/providers/common/ai/docs/agent_security.rst +++ b/providers/common/ai/docs/agent_security.rst @@ -76,6 +76,22 @@ No single layer is sufficient on its own. They work together. The LLM agent cannot see API keys or database passwords. - Does not prevent the agent from using the connection to access data the connection has access to. + * - **Secret masking of tool output** + - What a tool returns, and the error text handed to the model so it can correct a + call, pass through Airflow's secret masker first. A connection password that shows + up in a database error or a hook's return value reaches the model, the model + provider and any trace as ``***``. ``AgentOperator`` applies this to the toolsets + you pass in ``toolsets``, in ``agent_params["toolsets"]`` and in a ``Toolset`` + capability, including your own. The SQL, hook, DataFusion, MCP, sandbox and + managed-agent toolsets apply it wherever they run, including in a Pydantic AI agent + you build yourself. + - Masks only secrets Airflow has registered, such as connection passwords and + sensitive connection extras. A credential that exists only in the data itself is + not recognized. Not masked: prompts, model output, function tools passed as + ``agent_params["tools"]``, a ``Toolset`` capability built by a function, a + framework's own tools, and MCP servers a framework connects to itself. In a + sandbox, the model writes the commands, so it can print a secret in a form the + masker does not recognize; masking there guards against accidents only. * - **HookToolset: explicit allow-list** - Only methods listed in ``allowed_methods`` are exposed as tools. Auto-discovery is not supported. Methods are validated at Dag parse diff --git a/providers/common/ai/docs/toolsets/hook.rst b/providers/common/ai/docs/toolsets/hook.rst index 5f387b872ee..48b2fdd014e 100644 --- a/providers/common/ai/docs/toolsets/hook.rst +++ b/providers/common/ai/docs/toolsets/hook.rst @@ -108,15 +108,16 @@ reflection-based adapter, so the work is choosing the method list. states this outright. Choose methods whose worst case you accept, not methods you intend to constrain later. - Its calls act as barriers. The tools are registered with ``sequential=True`` - because hook methods perform synchronous I/O, so a slow call holds up every - other tool the model emitted in that step, not only this toolset's. This is - not specific to ``HookToolset``; see :ref:`toolset-call-barriers`. + and each hook method runs in a worker thread, one blocking hook call at a time + in the task process, so a slow call holds up every other tool the model emitted + in that step, not only this toolset's. This is not specific to ``HookToolset``; + see :ref:`toolset-call-barriers`. - It returns exactly one shape. Every result goes through ``serialize_for_llm`` and comes back as a JSON-encoded string; there is no structured error type and no ``ModelRetry`` wrapper, so a hook exception fails the agent run, and the task with it, instead of giving the model something it can correct. ``SQLToolset``, by contrast, hands the database's own error back as a retry. -- Its ``call_tool`` calls the method and serializes what comes back. The code +- It calls the method and serializes what comes back. The code contains no path that awaits a coroutine result, and none that checks for one, so an ``async def`` hook method is not a case this adapter is written to handle. Treat synchronous methods as the supported set. diff --git a/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py b/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py index 145304fb722..bbee364a2ef 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py +++ b/providers/common/ai/src/airflow/providers/common/ai/operators/agent.py @@ -56,6 +56,7 @@ from airflow.providers.common.ai.utils.logging import ( wrap_toolsets_for_logging, ) from airflow.providers.common.ai.utils.output_type import rehydrate_pydantic_output +from airflow.providers.common.ai.utils.toolset_base import with_masking from airflow.providers.common.ai.utils.toolsets import iter_toolsets from airflow.providers.common.ai.utils.usage import coerce_usage_limits from airflow.providers.common.ai.utils.usage_budget import ( @@ -654,13 +655,21 @@ class AgentOperator(CancellableAgentRunMixin, BaseOperator, HITLReviewMixin): storage = self._durable_storage counter = self._durable_counter if self.toolsets: - toolsets = self.toolsets + # Innermost, so the durable cache only ever stores masked results. + toolsets: list[AbstractToolset] = [with_masking(ts) for ts in self.toolsets] if self.durable and storage is not None and counter is not None: toolsets = self._build_durable_toolsets(toolsets, storage, counter) if self.enable_tool_logging: toolsets = wrap_toolsets_for_logging(toolsets, self.log) extra_kwargs["toolsets"] = toolsets - capabilities = list(extra_kwargs.get("capabilities") or []) + elif extra_kwargs.get("toolsets"): + extra_kwargs["toolsets"] = [with_masking(ts) for ts in extra_kwargs["toolsets"]] + capabilities = [ + replace(capability, toolset=with_masking(capability.toolset)) + if _is_concrete_toolset_capability(capability) + else capability + for capability in extra_kwargs.get("capabilities") or [] + ] if self.durable and storage is not None and counter is not None: # Tools supplied through a ``Toolset`` capability bypass the # ``toolsets=`` wrapping above, so their results would re-execute on diff --git a/providers/common/ai/src/airflow/providers/common/ai/policies/retry.py b/providers/common/ai/src/airflow/providers/common/ai/policies/retry.py index f3bd80c9fd8..ad21d52b569 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/policies/retry.py +++ b/providers/common/ai/src/airflow/providers/common/ai/policies/retry.py @@ -42,7 +42,7 @@ from collections.abc import Mapping from dataclasses import dataclass from datetime import timedelta from types import MappingProxyType -from typing import TYPE_CHECKING, Any, ClassVar, Literal, cast +from typing import TYPE_CHECKING, Any, ClassVar, Literal from pydantic import BaseModel @@ -56,7 +56,7 @@ from airflow.providers.common.ai.utils.decision import ( review_reason, threshold_for, ) -from airflow.providers.common.compat.sdk import redact +from airflow.providers.common.ai.utils.masking import mask_secrets try: from airflow.sdk.definitions.retry_policy import ( @@ -220,8 +220,7 @@ categories: those travel in the output schema with their descriptions, so a prom def redact_registered_secrets(message: str) -> str: """Mask values registered via ``mask_secret()``; the default ``redactor`` for the policies here.""" - # redact() is typed for arbitrary containers; a str in always yields a str out. - return cast("str", redact(message)) + return mask_secrets(message) _REDACTION_PARAMS_DOC = """ diff --git a/providers/common/ai/src/airflow/providers/common/ai/toolsets/datafusion.py b/providers/common/ai/src/airflow/providers/common/ai/toolsets/datafusion.py index 299d61fb9b7..a40cacad529 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/toolsets/datafusion.py +++ b/providers/common/ai/src/airflow/providers/common/ai/toolsets/datafusion.py @@ -18,7 +18,6 @@ from __future__ import annotations -import json import logging import re from typing import TYPE_CHECKING, Any @@ -34,14 +33,16 @@ except ImportError as e: from pydantic_ai.exceptions import ModelRetry from pydantic_ai.tools import ToolDefinition -from pydantic_ai.toolsets.abstract import AbstractToolset, ToolsetTool +from pydantic_ai.toolsets.abstract import ToolsetTool +from airflow.providers.common.ai.utils.masking import dumps_masked from airflow.providers.common.ai.utils.query_results import ( DEFAULT_MAX_RESULT_BYTES, QUERY_TOOL_DESCRIPTION as _QUERY_DESCRIPTION, build_query_result, ) from airflow.providers.common.ai.utils.tool_definition import build_args_validator +from airflow.providers.common.ai.utils.toolset_base import AirflowToolset if TYPE_CHECKING: from pydantic_ai._run_context import RunContext @@ -81,7 +82,7 @@ _RETRYABLE_QUERY_ERROR_PATTERNS = ( ) -class DataFusionToolset(AbstractToolset[Any]): +class DataFusionToolset(AirflowToolset): """ Curated toolset that gives an LLM agent SQL access to object-storage data via Apache DataFusion. @@ -174,7 +175,7 @@ class DataFusionToolset(AbstractToolset[Any]): ) return tools - async def call_tool( + async def _execute_tool( self, name: str, tool_args: dict[str, Any], @@ -182,21 +183,21 @@ class DataFusionToolset(AbstractToolset[Any]): tool: ToolsetTool[Any], ) -> Any: if name == "list_tables": - return self._list_tables() + return await self.run_blocking(self._list_tables) if name == "get_schema": - return self._get_schema(tool_args["table_name"]) + return await self.run_blocking(self._get_schema, tool_args["table_name"]) if name == "query": - return self._query(tool_args["sql"]) + return await self.run_blocking(self._query, tool_args["sql"]) raise ValueError(f"Unknown tool: {name!r}") def _list_tables(self) -> str: try: engine = self._get_engine() tables: list[str] = list(engine.session_context.catalog().schema().table_names()) - return json.dumps(tables) + return dumps_masked(tables) except Exception as ex: log.warning("list_tables failed: %s", ex) - return json.dumps({"error": str(ex)}) + return dumps_masked({"error": str(ex)}) def _get_schema(self, table_name: str) -> str: engine = self._get_engine() @@ -205,14 +206,14 @@ class DataFusionToolset(AbstractToolset[Any]): # When allow_writes is enabled, the agent may create temporary in-memory tables # that would not be captured there. if not engine.session_context.table_exist(table_name): - return json.dumps({"error": f"Table {table_name!r} is not available"}) + return dumps_masked({"error": f"Table {table_name!r} is not available"}) # Intentionally using session_context instead of engine.get_schema() — # the latter returns a pre-formatted string intended for other operators, # not a JSON-compatible format. # TODO: refactor engine.get_schema() to return JSON and update this accordingly table = engine.session_context.table(table_name) columns = [{"name": f.name, "type": str(f.type)} for f in table.schema()] - return json.dumps(columns) + return dumps_masked(columns) def _query(self, sql: str) -> str: try: @@ -246,7 +247,7 @@ class DataFusionToolset(AbstractToolset[Any]): raise ModelRetry( f"error: {ex!s}, Use get_schema and list_tables tools for more details." ) from ex - return json.dumps({"error": str(ex), "query": sql}) + return dumps_masked({"error": str(ex), "query": sql}) @staticmethod def _is_retryable_query_error(error: QueryExecutionException) -> bool: diff --git a/providers/common/ai/src/airflow/providers/common/ai/toolsets/hook.py b/providers/common/ai/src/airflow/providers/common/ai/toolsets/hook.py index 37ccf622782..382d11700d4 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/toolsets/hook.py +++ b/providers/common/ai/src/airflow/providers/common/ai/toolsets/hook.py @@ -25,7 +25,7 @@ import types from typing import TYPE_CHECKING, Any, Union, get_args, get_origin, get_type_hints from pydantic_ai.tools import ToolDefinition -from pydantic_ai.toolsets.abstract import AbstractToolset, ToolsetTool +from pydantic_ai.toolsets.abstract import ToolsetTool from airflow.providers.common.ai.tools._from_toolset import airflow_tools_from_toolset from airflow.providers.common.ai.utils.tool_definition import ( @@ -33,6 +33,7 @@ from airflow.providers.common.ai.utils.tool_definition import ( return_schema_kwargs, serialize_for_llm, ) +from airflow.providers.common.ai.utils.toolset_base import AirflowToolset if TYPE_CHECKING: from collections.abc import Callable, Sequence @@ -54,7 +55,7 @@ _TYPE_MAP: dict[type, dict[str, Any]] = { } -class HookToolset(AbstractToolset[Any]): +class HookToolset(AirflowToolset): """ Expose selected methods of an Airflow Hook as pydantic-ai tools. @@ -159,9 +160,10 @@ class HookToolset(AbstractToolset[Any]): if param_name in json_schema.get("properties", {}): json_schema["properties"][param_name]["description"] = param_desc - # sequential=True because hook methods perform synchronous I/O - # (network calls, DB queries) and should not run concurrently. - # return_schema is "string": call_tool serializes every result with + # sequential=True keeps pydantic-ai from running these calls concurrently + # within a turn; run_blocking's process-wide lock serializes them with the + # blocking calls of the other toolsets that use it. + # return_schema is "string": _execute_tool serializes every result with # serialize_for_llm, so the tool always returns a (JSON-encoded) # string regardless of the method's own return annotation. This lets # code mode render `-> str` instead of `-> Any`. @@ -180,7 +182,7 @@ class HookToolset(AbstractToolset[Any]): ) return tools - async def call_tool( + async def _execute_tool( self, name: str, tool_args: dict[str, Any], @@ -189,7 +191,7 @@ class HookToolset(AbstractToolset[Any]): ) -> Any: method_name = name.removeprefix(self._tool_name_prefix) if self._tool_name_prefix else name method: Callable[..., Any] = getattr(self._hook, method_name) - result = method(**tool_args) + result = await self.run_blocking(method, **tool_args) return serialize_for_llm(result) diff --git a/providers/common/ai/src/airflow/providers/common/ai/toolsets/managed_agent.py b/providers/common/ai/src/airflow/providers/common/ai/toolsets/managed_agent.py index aaddb2b0fa9..28627c9787e 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/toolsets/managed_agent.py +++ b/providers/common/ai/src/airflow/providers/common/ai/toolsets/managed_agent.py @@ -23,13 +23,14 @@ from typing import TYPE_CHECKING, Any from pydantic_ai.exceptions import ModelRetry from pydantic_ai.tools import ToolDefinition -from pydantic_ai.toolsets.abstract import AbstractToolset, ToolsetTool +from pydantic_ai.toolsets.abstract import ToolsetTool from airflow.providers.common.ai.utils.tool_definition import ( build_args_validator, return_schema_kwargs, serialize_for_llm, ) +from airflow.providers.common.ai.utils.toolset_base import AirflowToolset from airflow.providers.common.compat.sdk import Stats if TYPE_CHECKING: @@ -51,7 +52,7 @@ _PROMPT_SCHEMA: dict[str, Any] = { } -class BaseManagedAgentToolset(AbstractToolset[Any]): +class BaseManagedAgentToolset(AirflowToolset): """ Base class exposing a vendor-managed agent as a single pydantic-ai tool. @@ -207,11 +208,11 @@ class BaseManagedAgentToolset(AbstractToolset[Any]): name=self._tool_name, description=self._description, parameters_json_schema=_PROMPT_SCHEMA, - # HookToolset sets sequential=True because its tools call synchronous - # hook methods straight from the event loop. Here a blocking SDK goes - # through invoke_sync(), which the base class runs in a worker thread, - # and each call is an independent request to a remote service -- so - # two calls the model issues in one turn really can run at once. + # HookToolset sets sequential=True because its hook methods share one + # process-wide lock, so they run one at a time anyway. Here a blocking SDK + # goes through invoke_sync(), which the base class runs in a worker thread + # of its own, and each call is an independent request to a remote service, + # so two calls the model issues in one turn really can run at once. sequential=False, **return_schema_kwargs({"type": "string"}), ) @@ -230,7 +231,7 @@ class BaseManagedAgentToolset(AbstractToolset[Any]): ) } - async def call_tool( + async def _execute_tool( self, name: str, tool_args: dict[str, Any], diff --git a/providers/common/ai/src/airflow/providers/common/ai/toolsets/mcp.py b/providers/common/ai/src/airflow/providers/common/ai/toolsets/mcp.py index 5a1d7281477..c4ec7d991f2 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/toolsets/mcp.py +++ b/providers/common/ai/src/airflow/providers/common/ai/toolsets/mcp.py @@ -20,16 +20,18 @@ from __future__ import annotations from typing import TYPE_CHECKING, Any -from pydantic_ai.toolsets.abstract import AbstractToolset, ToolsetTool from typing_extensions import Self +from airflow.providers.common.ai.utils.toolset_base import AirflowToolset + if TYPE_CHECKING: from collections.abc import Callable, Sequence from pydantic_ai._run_context import RunContext + from pydantic_ai.toolsets.abstract import ToolsetTool -class MCPToolset(AbstractToolset[Any]): +class MCPToolset(AirflowToolset): """ Toolset that connects to an MCP server configured via an Airflow connection. @@ -109,8 +111,12 @@ class MCPToolset(AbstractToolset[Any]): self._server = hook.get_conn() return self._server + async def _server_once_resolved(self) -> Any: + # Resolving the connection talks to the supervisor, so it takes the blocking-call lock. + return self._server if self._server is not None else await self.run_blocking(self._get_server) + async def __aenter__(self) -> Self: - await self._get_server().__aenter__() + await (await self._server_once_resolved()).__aenter__() return self async def __aexit__(self, *args: Any) -> bool | None: @@ -119,13 +125,13 @@ class MCPToolset(AbstractToolset[Any]): return None async def get_tools(self, ctx: RunContext[Any]) -> dict[str, ToolsetTool[Any]]: - return await self._get_server().get_tools(ctx) + return await (await self._server_once_resolved()).get_tools(ctx) - async def call_tool( + async def _execute_tool( self, name: str, tool_args: dict[str, Any], ctx: RunContext[Any], tool: ToolsetTool[Any], ) -> Any: - return await self._get_server().call_tool(name, tool_args, ctx, tool) + return await (await self._server_once_resolved()).call_tool(name, tool_args, ctx, tool) diff --git a/providers/common/ai/src/airflow/providers/common/ai/toolsets/sandbox.py b/providers/common/ai/src/airflow/providers/common/ai/toolsets/sandbox.py index d0a92c20e17..7b2ec1468df 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/toolsets/sandbox.py +++ b/providers/common/ai/src/airflow/providers/common/ai/toolsets/sandbox.py @@ -44,11 +44,13 @@ from airflow.providers.common.ai.sandbox.output import ( render_file_window, truncate_output, ) +from airflow.providers.common.ai.utils.masking import mask_secrets from airflow.providers.common.ai.utils.tool_definition import ( build_args_validator, code_arg_kwargs, return_schema_kwargs, ) +from airflow.providers.common.ai.utils.toolset_base import AirflowToolset from airflow.providers.common.compat.sdk import get_current_context if TYPE_CHECKING: @@ -140,7 +142,7 @@ _DESCRIPTIONS = { } -class SandboxToolset(AbstractToolset[Any]): +class SandboxToolset(AirflowToolset): """ Give an agent shell and file access inside a disposable sandbox, off the Airflow worker. @@ -628,7 +630,7 @@ class SandboxToolset(AbstractToolset[Any]): ) return tools - async def call_tool( + async def _execute_tool( self, name: str, tool_args: dict[str, Any], @@ -725,8 +727,9 @@ class SandboxToolset(AbstractToolset[Any]): return output def _truncate(self, text: str, already_truncated: bool) -> str: + # Masked before it is cut, or a secret split at the cut would no longer match. return truncate_output( - text, + mask_secrets(text), max_lines=self._max_output_lines, max_bytes=self._max_output_bytes, already_truncated=already_truncated, @@ -740,7 +743,7 @@ class SandboxToolset(AbstractToolset[Any]): max_bytes=self._max_read_bytes, ) return render_file_window( - data, + mask_secrets(data), offset=tool_args.get("offset"), limit=tool_args.get("limit"), max_lines=self._max_output_lines, diff --git a/providers/common/ai/src/airflow/providers/common/ai/toolsets/sql.py b/providers/common/ai/src/airflow/providers/common/ai/toolsets/sql.py index 216e538705f..8814497989e 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/toolsets/sql.py +++ b/providers/common/ai/src/airflow/providers/common/ai/toolsets/sql.py @@ -18,7 +18,6 @@ from __future__ import annotations -import json from contextlib import suppress from typing import TYPE_CHECKING, Any @@ -39,15 +38,17 @@ except ImportError as e: from pydantic_ai.exceptions import ModelRetry from pydantic_ai.tools import ToolDefinition -from pydantic_ai.toolsets.abstract import AbstractToolset, ToolsetTool +from pydantic_ai.toolsets.abstract import ToolsetTool from airflow.providers.common.ai.tools._from_toolset import airflow_tools_from_toolset +from airflow.providers.common.ai.utils.masking import dumps_masked from airflow.providers.common.ai.utils.query_results import ( DEFAULT_MAX_RESULT_BYTES, QUERY_TOOL_DESCRIPTION as _QUERY_DESCRIPTION, build_query_result, ) from airflow.providers.common.ai.utils.tool_definition import build_args_validator, return_schema_kwargs +from airflow.providers.common.ai.utils.toolset_base import AirflowToolset from airflow.providers.common.compat.sdk import BaseHook if TYPE_CHECKING: @@ -164,7 +165,7 @@ class _CappedFetch: return rows -class SQLToolset(AbstractToolset[Any]): +class SQLToolset(AirflowToolset): """ Curated toolset that gives an LLM agent safe access to a SQL database. @@ -393,8 +394,9 @@ class SQLToolset(AbstractToolset[Any]): ("query", _QUERY_DESCRIPTION, _QUERY_SCHEMA), ("check_query", "Validate SQL syntax without executing it.", _CHECK_QUERY_SCHEMA), ): - # sequential=True because all tools use a shared DbApiHook with - # synchronous I/O — they must not run concurrently. + # sequential=True keeps pydantic-ai from running these calls concurrently + # within a turn; run_blocking's process-wide lock serializes them with the + # blocking calls of the other toolsets that use it. # return_schema is "string": every tool returns a JSON-encoded string # (json.dumps), so code mode renders `-> str` instead of `-> Any`. tool_def = ToolDefinition( @@ -412,7 +414,7 @@ class SQLToolset(AbstractToolset[Any]): ) return tools - async def call_tool( + async def _execute_tool( self, name: str, tool_args: dict[str, Any], @@ -422,13 +424,7 @@ class SQLToolset(AbstractToolset[Any]): if name not in ("list_tables", "get_schema", "query", "check_query"): raise ValueError(f"Unknown tool: {name!r}") try: - if name == "list_tables": - return self._list_tables() - if name == "get_schema": - return self._get_schema(tool_args["table_name"]) - if name == "query": - return self._query(tool_args["sql"]) - return self._check_query(tool_args["sql"]) + return await self.run_blocking(self._run_tool, name, tool_args) except Exception as e: # Hand the database's own error back to the agent as a retry so it can # read the message and fix its SQL within the run. pydantic-ai bounds @@ -441,6 +437,15 @@ class SQLToolset(AbstractToolset[Any]): "then fix the query and try again." ) from e + def _run_tool(self, name: str, tool_args: dict[str, Any]) -> str: + if name == "list_tables": + return self._list_tables() + if name == "get_schema": + return self._get_schema(tool_args["table_name"]) + if name == "query": + return self._query(tool_args["sql"]) + return self._check_query(tool_args["sql"]) + # ------------------------------------------------------------------ # Tool implementations # ------------------------------------------------------------------ @@ -478,15 +483,15 @@ class SQLToolset(AbstractToolset[Any]): for name in hook.inspector.get_table_names(schema=self._schema): add(self._schema, name, name) - return json.dumps(tables) + return dumps_masked(tables) def _get_schema(self, table_name: str) -> str: schema, table = self._split_table_identifier(table_name) if not self._is_ref_allowed("", schema, table): - return json.dumps({"error": f"Table {table_name!r} is not in the allowed tables list."}) + return dumps_masked({"error": f"Table {table_name!r} is not in the allowed tables list."}) hook = self._get_db_hook() columns = hook.get_table_schema(table, schema=schema) - return json.dumps(columns) + return dumps_masked(columns) def _dialect_for_validation(self) -> str | None: """Resolve the hook's sqlglot dialect so DESCRIBE/SHOW validate correctly.""" @@ -537,9 +542,9 @@ class SQLToolset(AbstractToolset[Any]): try: statements = _validate_sql(sql, dialect=dialect, allow_read_only_metadata=True) self._enforce_allowed_tables(statements) - return json.dumps({"valid": True}) + return dumps_masked({"valid": True}) except Exception as e: - return json.dumps({"valid": False, "error": str(e)}) + return dumps_masked({"valid": False, "error": str(e)}) def _enforce_allowed_tables(self, statements: list[Any]) -> None: """ diff --git a/providers/common/ai/src/airflow/providers/common/ai/utils/file_analysis.py b/providers/common/ai/src/airflow/providers/common/ai/utils/file_analysis.py index f53a23fa63a..9f0fbc4e241 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/utils/file_analysis.py +++ b/providers/common/ai/src/airflow/providers/common/ai/utils/file_analysis.py @@ -46,6 +46,7 @@ from airflow.providers.common.ai.exceptions import ( LLMFileAnalysisMultimodalRequiredError, LLMFileAnalysisUnsupportedFormatError, ) +from airflow.providers.common.ai.utils.masking import dumps_masked from airflow.providers.common.compat.sdk import AirflowOptionalProviderFeatureException, ObjectStoragePath if TYPE_CHECKING: @@ -445,7 +446,7 @@ def _render_json( estimated_rows = len(document) else: estimated_rows = None - pretty = json.dumps(document, indent=2, sort_keys=True, default=str) + pretty = dumps_masked(document, indent=2, sort_keys=True) return _RenderResult( text=_truncate_text(pretty), estimated_rows=estimated_rows, @@ -508,7 +509,7 @@ def _render_parquet(path: ObjectStoragePath, *, sample_rows: int, max_content_by group_rows = row_group.slice(0, remaining_rows).to_pylist() sampled_rows.extend(group_rows) remaining_rows -= len(group_rows) - payload = [f"Schema: {schema}", "Sample rows:", json.dumps(sampled_rows, indent=2, default=str)] + payload = [f"Schema: {schema}", "Sample rows:", dumps_masked(sampled_rows, indent=2)] return _RenderResult( text=_truncate_text("\n".join(payload)), estimated_rows=num_rows, @@ -547,9 +548,9 @@ def _render_avro(path: ObjectStoragePath, *, sample_rows: int, max_content_bytes else: fully_read = True payload = [ - f"Schema: {json.dumps(writer_schema, indent=2, default=str)}", + f"Schema: {dumps_masked(writer_schema, indent=2)}", "Sample rows:", - json.dumps(sampled_rows, indent=2, default=str), + dumps_masked(sampled_rows, indent=2), ] return _RenderResult( text=_truncate_text("\n".join(payload)), diff --git a/providers/common/ai/src/airflow/providers/common/ai/utils/masking.py b/providers/common/ai/src/airflow/providers/common/ai/utils/masking.py new file mode 100644 index 00000000000..9fe2de9ebdc --- /dev/null +++ b/providers/common/ai/src/airflow/providers/common/ai/utils/masking.py @@ -0,0 +1,120 @@ +# 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. +"""Apply Airflow's secret masker to what a tool hands back to a model.""" + +from __future__ import annotations + +import dataclasses +import json +from typing import Any, overload + +from pydantic import BaseModel +from pydantic_ai.messages import MULTI_MODAL_CONTENT_TYPES +from pydantic_core import to_jsonable_python + +from airflow.providers.common.compat.sdk import redact + +# Stands in for a container that contains itself, which would otherwise recurse forever. +_CYCLE = "<circular reference>" + + +@overload +def mask_secrets(value: str) -> str: ... + + +@overload +def mask_secrets(value: Any) -> Any: ... + + +def mask_secrets(value: Any) -> Any: + """ + Return ``value`` with every secret Airflow has registered replaced by ``***``. + + Strings are masked wherever they sit in nested dicts, lists, tuples, sets and dataclasses, + dict keys included, and the shape and types are kept. Bytes are masked as UTF-8 text. A + Pydantic model is turned into the JSON-compatible data the model would be shown. Images, + documents and other multimodal content, and any other object, pass through as they are. + Registered secrets are the ones Airflow knows about, such as connection passwords and + sensitive connection extras; a credential that only appears in the data itself is not + recognized. + + ``redact()`` does part of this, but stops descending at a fixed depth, and it hides every + string under a key that looks sensitive: a model reading a query result needs + ``{"auth_type": "oauth"}`` as it is. Two dict keys that both mask to ``***`` collapse into + one. + """ + return _mask(value, frozenset()) + + +def dumps_masked(value: Any, **kwargs: Any) -> str: + """ + Serialize ``value`` to JSON for a model, with registered secrets masked first. + + Masking a JSON string afterwards is not enough: JSON escapes quotes, backslashes, + control characters and, by default, non-ASCII characters, so a password containing + any of them no longer matches the registered value once it is inside the document. + Bytes become their UTF-8 text and dataclasses their fields; any other object JSON + cannot represent is rendered with ``str()``, after the values inside it are masked. + + :param kwargs: Passed to :func:`json.dumps`. + """ + return json.dumps(mask_secrets(value), default=_json_default, **kwargs) + + +def _json_default(value: object) -> Any: + if isinstance(value, bytes): + return value.decode("utf-8", "replace") + if dataclasses.is_dataclass(value) and not isinstance(value, type): + return {field.name: getattr(value, field.name) for field in dataclasses.fields(value)} + return mask_secrets(str(value)) + + +def _mask(value: Any, seen: frozenset[int]) -> Any: + if isinstance(value, str): + return redact(value) + if isinstance(value, bytes): + text = value.decode("utf-8", "surrogateescape") + masked = mask_secrets(text) + return value if masked == text else masked.encode("utf-8", "surrogateescape") + if isinstance(value, BaseModel): + return _mask(to_jsonable_python(value), seen) + if isinstance(value, (dict, list, tuple, set, frozenset)) or _is_masked_dataclass(value): + if id(value) in seen: + return _CYCLE + seen = seen | {id(value)} + if isinstance(value, dict): + return {_mask(key, seen): _mask(item, seen) for key, item in value.items()} + if isinstance(value, list): + return [_mask(item, seen) for item in value] + if isinstance(value, tuple): + return tuple(_mask(item, seen) for item in value) + if isinstance(value, (set, frozenset)): + return type(value)(_mask(item, seen) for item in value) + if _is_masked_dataclass(value): + # Rebuilt rather than turned into a dict, so a type the framework acts on, such as + # pydantic-ai's TextContent, still is one. + fields = [field for field in dataclasses.fields(value) if field.init] + return dataclasses.replace(value, **{f.name: _mask(getattr(value, f.name), seen) for f in fields}) + return value + + +def _is_masked_dataclass(value: object) -> bool: + return ( + dataclasses.is_dataclass(value) + and not isinstance(value, type) + and not isinstance(value, MULTI_MODAL_CONTENT_TYPES) + ) diff --git a/providers/common/ai/src/airflow/providers/common/ai/utils/query_results.py b/providers/common/ai/src/airflow/providers/common/ai/utils/query_results.py index 956c9de86e3..b8ecdb3e651 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/utils/query_results.py +++ b/providers/common/ai/src/airflow/providers/common/ai/utils/query_results.py @@ -31,10 +31,11 @@ cost is re-paid on every subsequent request. Two things here keep that bounded: from __future__ import annotations -import json from collections.abc import Sequence from typing import Any +from airflow.providers.common.ai.utils.masking import dumps_masked + # A policy default, not a limit imposed by any storage, protocol, or model layer: # roughly 16k tokens at 4 characters per token. Large enough that ordinary queries are # unaffected, small enough that no single tool result can dominate the context window. @@ -45,7 +46,7 @@ DEFAULT_MAX_RESULT_BYTES = 65_536 # the separators: escaping one CJK character to \uXXXX costs six bytes instead of three, # so an ASCII-escaped result is charged several times over against the budget and # truncated that much earlier than an equivalent English one. -_DUMP_KWARGS: dict[str, Any] = {"default": str, "separators": (",", ":"), "ensure_ascii": False} +_DUMP_KWARGS: dict[str, Any] = {"separators": (",", ":"), "ensure_ascii": False} #: Description for the ``query`` tool. States the columnar shape, since the model has #: to align each row's values to ``columns`` positionally, and the truncation contract, @@ -61,7 +62,7 @@ QUERY_TOOL_DESCRIPTION = ( def _dumps(payload: Any) -> str: - return json.dumps(payload, **_DUMP_KWARGS) + return dumps_masked(payload, **_DUMP_KWARGS) def _size(payload: Any) -> int: diff --git a/providers/common/ai/src/airflow/providers/common/ai/utils/tool_definition.py b/providers/common/ai/src/airflow/providers/common/ai/utils/tool_definition.py index 34425d04180..978581d0f6e 100644 --- a/providers/common/ai/src/airflow/providers/common/ai/utils/tool_definition.py +++ b/providers/common/ai/src/airflow/providers/common/ai/utils/tool_definition.py @@ -19,12 +19,13 @@ from __future__ import annotations import dataclasses -import json from typing import Any, Literal from pydantic_ai.tools import ToolDefinition from pydantic_core import SchemaValidator, core_schema +from airflow.providers.common.ai.utils.masking import dumps_masked, mask_secrets + # ``ToolDefinition.return_schema`` is newer than the provider's pydantic-ai # floor. Detect it once so callers can include the kwarg only when supported, # rather than raising ``TypeError`` on older installs. @@ -48,18 +49,20 @@ def return_schema_kwargs(schema: dict[str, Any]) -> dict[str, Any]: def serialize_for_llm(value: Any) -> str: """ - Convert a Python return value to a string suitable for an LLM. + Convert a Python return value to a string suitable for an LLM, with registered secrets masked. :param value: The tool's return value. """ if value is None: return "null" if isinstance(value, str): - return value + return mask_secrets(value) try: - return json.dumps(value, default=str) + return dumps_masked(value) except (TypeError, ValueError): - return str(value) + # Masked before str(), which escapes backslashes and quotes inside the strings it + # renders, so a secret containing either would no longer match once rendered. + return mask_secrets(str(mask_secrets(value))) _SUPPORTS_METADATA = any(f.name == "metadata" for f in dataclasses.fields(ToolDefinition)) 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 new file mode 100644 index 00000000000..1e5a2dd93b7 --- /dev/null +++ b/providers/common/ai/src/airflow/providers/common/ai/utils/toolset_base.py @@ -0,0 +1,211 @@ +# 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, 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.utils.masking import mask_secrets + +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 + +log = logging.getLogger(__name__) + +P = ParamSpec("P") +R = TypeVar("R") + +# 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 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) + + +def _call_locked(fn: Callable[P, R], /, *args: P.args, **kwargs: P.kwargs) -> R: + with _blocking_call_lock: + return fn(*args, **kwargs) + + +async def _masked(name: str, call: Awaitable[Any]) -> 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. + """ + error: Exception | None = None + try: + result = await call + except _CONTROL_FLOW as e: + log.debug("Tool %s returned no result", 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) + 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 = _stripped(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 _stripped(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 + + +class AirflowToolset(AbstractToolset[Any]): + """ + A toolset whose tool results are safe to hand to a model. + + Subclasses implement :meth:`_execute_tool`. :meth:`call_tool` runs it and passes what it + returns, and any exception it raises, through Airflow's secret masker, so a connection + password that ends up in a database error or a hook's return value is replaced with + ``***`` before the model, the model provider or a trace sees it. + """ + + async def call_tool( + self, + name: str, + tool_args: dict[str, Any], + ctx: RunContext[Any], + tool: ToolsetTool[Any], + ) -> Any: + return await _masked(name, self._execute_tool(name, tool_args, ctx, tool)) + + @abstractmethod + async def _execute_tool( + self, + name: str, + tool_args: dict[str, Any], + ctx: RunContext[Any], + tool: ToolsetTool[Any], + ) -> Any: + """Run tool ``name`` with validated ``tool_args``; :meth:`call_tool` masks what it returns.""" + + @staticmethod + async def run_blocking(fn: Callable[P, R], /, *args: P.args, **kwargs: P.kwargs) -> R: + """ + Run a blocking hook call in a worker thread, keeping the event loop free. + + Calls made through this method are serialized across the process. + """ + return await asyncio.to_thread(_call_locked, fn, *args, **kwargs) + + +@dataclass +class MaskingToolset(WrapperToolset[Any]): + """ + Apply the same masking as :class:`AirflowToolset` to any toolset. + + ``AgentOperator`` wraps every toolset it runs in one, so a toolset the Dag author wrote + gets masked output too. + """ + + async def call_tool( + self, + name: str, + tool_args: dict[str, Any], + ctx: RunContext[Any], + tool: ToolsetTool[Any], + ) -> Any: + return await _masked(name, self.wrapped.call_tool(name, tool_args, ctx, tool)) + + +def with_masking(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) diff --git a/providers/common/ai/tests/unit/common/ai/decorators/test_agent.py b/providers/common/ai/tests/unit/common/ai/decorators/test_agent.py index f9c706f62db..0ed8b5cd8d2 100644 --- a/providers/common/ai/tests/unit/common/ai/decorators/test_agent.py +++ b/providers/common/ai/tests/unit/common/ai/decorators/test_agent.py @@ -25,6 +25,7 @@ from pydantic_ai.toolsets.function import FunctionToolset from airflow.providers.common.ai.decorators.agent import _AgentDecoratedOperator from airflow.providers.common.ai.toolsets.logging import LoggingToolset +from airflow.providers.common.ai.utils.toolset_base import MaskingToolset try: from airflow.sdk.serde import SUPPORTS_OPERATOR_DESERIALIZATION_WALKER as _CORE_WALKER @@ -176,7 +177,7 @@ class TestAgentDecoratedOperator: passed_toolsets = create_call[1]["toolsets"] assert len(passed_toolsets) == 1 assert isinstance(passed_toolsets[0], LoggingToolset) - assert passed_toolsets[0].wrapped is toolset + assert passed_toolsets[0].wrapped == MaskingToolset(wrapped=toolset) @requires_typed_xcom @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook", autospec=True) diff --git a/providers/common/ai/tests/unit/common/ai/operators/test_agent.py b/providers/common/ai/tests/unit/common/ai/operators/test_agent.py index 79b3c6679e2..274fb3ce3d8 100644 --- a/providers/common/ai/tests/unit/common/ai/operators/test_agent.py +++ b/providers/common/ai/tests/unit/common/ai/operators/test_agent.py @@ -43,6 +43,7 @@ from pydantic_ai.messages import ( ) from pydantic_ai.models.function import AgentInfo, FunctionModel from pydantic_ai.models.wrapper import WrapperModel +from pydantic_ai.toolsets.abstract import AbstractToolset from pydantic_ai.toolsets.combined import CombinedToolset from pydantic_ai.toolsets.function import FunctionToolset from pydantic_ai.toolsets.wrapper import WrapperToolset @@ -68,6 +69,7 @@ from airflow.providers.common.ai.toolsets.logging import LoggingToolset from airflow.providers.common.ai.toolsets.mcp import MCPToolset from airflow.providers.common.ai.toolsets.sandbox import SandboxToolset from airflow.providers.common.ai.toolsets.sql import SQLToolset +from airflow.providers.common.ai.utils.toolset_base import MaskingToolset from airflow.providers.common.ai.utils.toolsets import find_toolset from airflow.providers.common.ai.utils.usage_budget import ( USAGE_BUDGET_KEY, @@ -748,12 +750,12 @@ class TestAgentOperatorExecute: @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook", autospec=True) def test_execute_passes_toolsets_in_agent_kwargs(self, mock_hook_cls, make_mock_run_result): - """Toolsets are passed through to the agent constructor.""" + """Toolsets reach the agent wrapped for masking, then for logging.""" mock_hook_cls.get_hook.return_value.create_agent.return_value = _make_mock_agent( "done", make_mock_run_result ) - mock_toolset = MagicMock() + mock_toolset = MagicMock(spec=AbstractToolset) op = AgentOperator( task_id="test", prompt="Do something", @@ -766,16 +768,17 @@ class TestAgentOperatorExecute: passed_toolsets = create_call[1]["toolsets"] assert len(passed_toolsets) == 1 assert isinstance(passed_toolsets[0], LoggingToolset) - assert passed_toolsets[0].wrapped is mock_toolset + assert isinstance(passed_toolsets[0].wrapped, MaskingToolset) + assert passed_toolsets[0].wrapped.wrapped is mock_toolset @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook", autospec=True) def test_enable_tool_logging_false_skips_wrapping(self, mock_hook_cls, make_mock_run_result): - """enable_tool_logging=False passes toolsets through unwrapped.""" + """enable_tool_logging=False skips the logging wrapper; masking still applies.""" mock_hook_cls.get_hook.return_value.create_agent.return_value = _make_mock_agent( "done", make_mock_run_result ) - mock_toolset = MagicMock() + mock_toolset = MagicMock(spec=AbstractToolset) op = AgentOperator( task_id="test", prompt="Do something", @@ -786,7 +789,7 @@ class TestAgentOperatorExecute: op.execute(context=MagicMock()) create_call = mock_hook_cls.get_hook.return_value.create_agent.call_args - assert create_call[1]["toolsets"] == [mock_toolset] + assert create_call[1]["toolsets"] == [MaskingToolset(wrapped=mock_toolset)] @patch("airflow.providers.common.ai.operators.agent.PydanticAIHook", autospec=True) def test_execute_passes_agent_params(self, mock_hook_cls, make_mock_run_result): @@ -3080,3 +3083,87 @@ class TestAgentOperatorDurableUsageBudgetEndToEnd: # second call's preamble raises; what matters is that neither ever exceeds the limit. assert scenario.get_last_saved_usage()["tool_calls"] <= 3 assert scenario.live_tool_calls <= 3 + + +def _echo_tool_result(messages, info: AgentInfo) -> ModelResponse: + """Call ``read_setting`` once, then answer with whatever it returned.""" + returns = [p for m in messages for p in m.parts if isinstance(p, ToolReturnPart)] + if returns: + return ModelResponse(parts=[TextPart(content=str(returns[-1].content))]) + return ModelResponse(parts=[ToolCallPart(tool_name="read_setting", args={}, tool_call_id="c1")]) + + [email protected]_redact +class TestAgentOperatorMasksToolOutput: + """What any tool hands the model is masked, however the toolset reaches the agent.""" + + @staticmethod + def _run(op: AgentOperator, storage=None) -> str: + if storage is not None: + op._durable_storage = storage + op._durable_counter = DurableStepCounter() + hook = MagicMock(spec=["create_agent"]) + hook.create_agent.side_effect = lambda **kw: Agent(FunctionModel(_echo_tool_result), **kw) + op.llm_hook = hook + return op._build_agent().run_sync("hi").output + + @staticmethod + def _dag_authors_toolset(secret: str) -> FunctionToolset: + def read_setting() -> str: + return f"api key: {secret}" + + return FunctionToolset(tools=[read_setting]) + + def test_a_toolset_passed_as_toolsets(self, registered_secret): + op = AgentOperator( + task_id="t", prompt="hi", llm_conn_id="c", toolsets=[self._dag_authors_toolset(registered_secret)] + ) + + assert self._run(op) == "api key: ***" + + def test_a_toolset_passed_through_agent_params(self, registered_secret): + op = AgentOperator( + task_id="t", + prompt="hi", + llm_conn_id="c", + agent_params={"toolsets": [self._dag_authors_toolset(registered_secret)]}, + ) + + assert self._run(op) == "api key: ***" + + def test_a_function_that_builds_a_toolset_per_run(self, registered_secret): + """pydantic-ai accepts such a function wherever it accepts a toolset.""" + op = AgentOperator( + task_id="t", + prompt="hi", + llm_conn_id="c", + agent_params={"toolsets": [lambda ctx: self._dag_authors_toolset(registered_secret)]}, + ) + + assert self._run(op) == "api key: ***" + + def test_a_toolset_capability_is_masked_before_the_durable_cache_stores_it(self, registered_secret): + storage = _InMemoryDurableStorage() + op = AgentOperator( + task_id="t", + prompt="hi", + llm_conn_id="c", + durable=True, + agent_params={"capabilities": [Toolset(self._dag_authors_toolset(registered_secret))]}, + ) + + assert self._run(op, storage) == "api key: ***" + cached = [value for value, _ in storage.tools.values()] + assert cached == ["api key: ***"] + + def test_a_toolset_that_masks_its_own_output_is_not_wrapped_again(self): + sql = SQLToolset("pg_default") + op = AgentOperator( + task_id="t", prompt="hi", llm_conn_id="c", toolsets=[sql], enable_tool_logging=False + ) + hook = MagicMock(spec=["create_agent"]) + op.llm_hook = hook + + op._build_agent() + + assert hook.create_agent.call_args.kwargs["toolsets"] == [sql] diff --git a/providers/common/ai/tests/unit/common/ai/toolsets/test_datafusion.py b/providers/common/ai/tests/unit/common/ai/toolsets/test_datafusion.py index 98a8d21f853..c97d6978acd 100644 --- a/providers/common/ai/tests/unit/common/ai/toolsets/test_datafusion.py +++ b/providers/common/ai/tests/unit/common/ai/toolsets/test_datafusion.py @@ -31,7 +31,6 @@ from airflow.providers.common.ai.toolsets.datafusion import ( _RETRYABLE_QUERY_ERROR_PATTERNS, DataFusionToolset, ) -from airflow.providers.common.ai.utils.sql_validation import SQLSafetyError from airflow.providers.common.sql.config import DataSourceConfig @@ -140,6 +139,19 @@ class TestDataFusionToolsetListTables: tables = json.loads(result) assert set(tables) == {"sales", "orders"} + @pytest.mark.enable_redact + def test_an_error_carrying_a_secret_that_json_escapes_is_masked(self, register_secret): + secret = register_secret('s3-se"cret-91c3') + ts = DataFusionToolset([_make_mock_datasource_config()]) + ts._engine = _make_mock_engine() + ts._engine.session_context.catalog.side_effect = RuntimeError(f"object store auth failed: {secret}") + + result = asyncio.run( + ts.call_tool("list_tables", {}, ctx=MagicMock(spec=RunContext), tool=MagicMock(spec=ToolsetTool)) + ) + + assert json.loads(result) == {"error": "object store auth failed: ***"} + class TestDataFusionToolsetGetSchema: def test_returns_column_info(self): @@ -251,7 +263,7 @@ class TestDataFusionToolsetQuery: tool=MagicMock(spec=ToolsetTool), ) ) - assert isinstance(exc_info.value.__cause__, SQLSafetyError) + assert "Only read-only SELECT-family queries are allowed" in exc_info.value.message def test_allows_create_table_when_writes_enabled(self): cfg = _make_mock_datasource_config() diff --git a/providers/common/ai/tests/unit/common/ai/toolsets/test_hook.py b/providers/common/ai/tests/unit/common/ai/toolsets/test_hook.py index 8ddc911b144..a54b0c6b051 100644 --- a/providers/common/ai/tests/unit/common/ai/toolsets/test_hook.py +++ b/providers/common/ai/tests/unit/common/ai/toolsets/test_hook.py @@ -17,9 +17,11 @@ from __future__ import annotations import asyncio +import threading from unittest.mock import MagicMock import pytest +from pydantic_ai._run_context import RunContext from pydantic_core import ValidationError from airflow.providers.common.ai.toolsets.hook import ( @@ -255,11 +257,52 @@ class TestHookToolsetCallTool: result = asyncio.run( ts.call_tool( - "storage_read_file", {"key": "test.txt"}, ctx=MagicMock(), tool=tools["storage_read_file"] + "storage_read_file", + {"key": "test.txt"}, + ctx=MagicMock(spec=RunContext), + tool=tools["storage_read_file"], ) ) assert result == "contents of test.txt" + @pytest.mark.enable_redact + def test_a_result_carrying_a_registered_secret_reaches_the_model_masked(self, registered_secret): + hook = _FakeHook() + ts = HookToolset(hook, allowed_methods=["read_file"]) + tools = asyncio.run(ts.get_tools(ctx=MagicMock(spec=RunContext))) + + result = asyncio.run( + ts.call_tool( + "read_file", + {"key": registered_secret}, + ctx=MagicMock(spec=RunContext), + tool=tools["read_file"], + ) + ) + + assert result == "contents of ***" + + def test_the_hook_method_runs_off_the_event_loop_thread(self): + calls: list[int] = [] + + class _ThreadRecordingHook: + def whoami(self) -> str: + """Report the calling thread.""" + calls.append(threading.get_ident()) + return "ok" + + ts = HookToolset(_ThreadRecordingHook(), allowed_methods=["whoami"]) + + async def call() -> int: + tools = await ts.get_tools(ctx=MagicMock(spec=RunContext)) + await ts.call_tool("whoami", {}, ctx=MagicMock(spec=RunContext), tool=tools["whoami"]) + return threading.get_ident() + + loop_thread = asyncio.run(call()) + + assert calls + assert calls[0] != loop_thread + class TestBuildJsonSchemaFromSignature: def test_basic_types(self): diff --git a/providers/common/ai/tests/unit/common/ai/toolsets/test_sandbox.py b/providers/common/ai/tests/unit/common/ai/toolsets/test_sandbox.py index 887ae7126ac..1852f8ea5d1 100644 --- a/providers/common/ai/tests/unit/common/ai/toolsets/test_sandbox.py +++ b/providers/common/ai/tests/unit/common/ai/toolsets/test_sandbox.py @@ -295,6 +295,19 @@ class TestNetworkNote: class TestRunCommand: + @pytest.mark.asyncio + @pytest.mark.enable_redact + async def test_output_is_masked_before_it_is_truncated(self, registered_secret): + """Cut first, a secret split at the cut would no longer match and would leak in part.""" + output = "x" * 40 + registered_secret + "y" * 40 + backend = _RecordingBackend(run_result=SandboxExecResult(exit_code=0, stdout=output, stderr="")) + ts = SandboxToolset(backend, max_output_bytes=60) + + async with ts: + result = await _call(ts, "run_command", {"command": "x"}) + + assert registered_secret[len(registered_secret) // 2 :] not in result + @pytest.mark.asyncio async def test_labels_streams_and_reports_a_nonzero_exit(self): backend = _RecordingBackend(run_result=SandboxExecResult(exit_code=3, stdout="hi\n", stderr="bad\n")) @@ -533,12 +546,9 @@ class TestErrorMapping: ts = SandboxToolset(backend) async with ts: - with pytest.raises( - SandboxTerminalError, match="Could not provision.*image pull timed out" - ) as caught: + with pytest.raises(SandboxTerminalError, match="Could not provision.*image pull timed out"): await _call(ts, "run_command", {"command": "x"}) - assert isinstance(caught.value.__cause__, SandboxError) assert backend.destroyed == [], "nothing was provisioned, so nothing is destroyed" @pytest.mark.asyncio diff --git a/providers/common/ai/tests/unit/common/ai/toolsets/test_sql.py b/providers/common/ai/tests/unit/common/ai/toolsets/test_sql.py index 57e0db38b0d..f62d2cb436e 100644 --- a/providers/common/ai/tests/unit/common/ai/toolsets/test_sql.py +++ b/providers/common/ai/tests/unit/common/ai/toolsets/test_sql.py @@ -452,6 +452,46 @@ class TestSQLToolsetQuery: assert "list_tables" in message assert "get_schema" in message + @pytest.mark.enable_redact + def test_a_database_error_carrying_the_connection_password_reaches_the_model_masked( + self, registered_secret + ): + ts = SQLToolset("pg_default") + ts._hook = _make_mock_db_hook() + ts._hook.run.side_effect = ConnectionError( + f'connection to "db:5432" failed: password "{registered_secret}" rejected' + ) + + with pytest.raises(ModelRetry) as exc_info: + asyncio.run( + ts.call_tool( + "query", + {"sql": "SELECT 1"}, + ctx=MagicMock(spec=RunContext), + tool=MagicMock(spec=ToolsetTool), + ) + ) + + assert registered_secret not in exc_info.value.message + assert 'password "***" rejected' in exc_info.value.message + + @pytest.mark.enable_redact + def test_a_row_carrying_the_connection_password_reaches_the_model_masked(self, registered_secret): + ts = SQLToolset("pg_default") + ts._hook = _make_mock_db_hook(records=[(1, f"dsn=postgres://app:{registered_secret}@db")]) + + result = asyncio.run( + ts.call_tool( + "query", + {"sql": "SELECT * FROM users"}, + ctx=MagicMock(spec=RunContext), + tool=MagicMock(spec=ToolsetTool), + ) + ) + + assert registered_secret not in result + assert "postgres://app:***@db" in result + class TestSQLToolsetCheckQuery: def test_valid_select(self): @@ -475,6 +515,20 @@ class TestSQLToolsetCheckQuery: assert data["valid"] is False assert "error" in data + @pytest.mark.enable_redact + def test_an_error_carrying_a_secret_that_json_escapes_is_masked(self, register_secret): + secret = register_secret('db-pa"ss-91c3') + ts = SQLToolset("pg_default") + ts._hook = _make_mock_db_hook() + + result = asyncio.run( + ts.call_tool("check_query", {"sql": f"SELECT '{secret}' FROM"}, ctx=MagicMock(), tool=MagicMock()) + ) + + data = json.loads(result) + assert data["valid"] is False + assert secret not in data["error"] + class TestSQLToolsetHookResolution: @patch("airflow.providers.common.ai.toolsets.sql.BaseHook", autospec=True) diff --git a/providers/common/ai/tests/unit/common/ai/utils/test_file_analysis.py b/providers/common/ai/tests/unit/common/ai/utils/test_file_analysis.py index b2a80186562..f2fbbaf0162 100644 --- a/providers/common/ai/tests/unit/common/ai/utils/test_file_analysis.py +++ b/providers/common/ai/tests/unit/common/ai/utils/test_file_analysis.py @@ -18,6 +18,7 @@ from __future__ import annotations import builtins import gzip +import json from pathlib import Path from unittest.mock import MagicMock, patch @@ -337,6 +338,26 @@ class TestBuildFileAnalysisRequest: assert '"a": 1' in request.user_content assert '"b": 2' in request.user_content + @pytest.mark.enable_redact + def test_json_holding_a_secret_that_json_escapes_is_masked(self, tmp_path, register_secret): + secret = register_secret("db-p\u00e4ss-91c3") + path = tmp_path / "config.json" + path.write_text(json.dumps({"password": secret}), encoding="utf-8") + + request = build_file_analysis_request( + file_path=str(path), + file_conn_id=None, + prompt="Analyze", + multi_modal=False, + max_files=1, + max_file_size_bytes=1024, + max_total_size_bytes=1024, + max_text_chars=500, + sample_rows=10, + ) + + assert '"password": "***"' in request.user_content + def test_text_context_truncation_is_marked(self, tmp_path): path = tmp_path / "huge.log" path.write_text("line\n" * 400, encoding="utf-8") diff --git a/providers/common/ai/tests/unit/common/ai/utils/test_masking.py b/providers/common/ai/tests/unit/common/ai/utils/test_masking.py new file mode 100644 index 00000000000..e6dffecaeec --- /dev/null +++ b/providers/common/ai/tests/unit/common/ai/utils/test_masking.py @@ -0,0 +1,117 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +import json +from dataclasses import dataclass + +import pytest +from pydantic import BaseModel +from pydantic_ai.messages import BinaryContent, TextContent + +from airflow.providers.common.ai.utils.masking import dumps_masked, mask_secrets + + [email protected]_redact +class TestMaskSecrets: + def test_masks_inside_every_container_and_keeps_the_shape(self, registered_secret): + value = { + "rows": [(1, registered_secret)], + "tags": {registered_secret}, + "note": f"key={registered_secret}", + } + + assert mask_secrets(value) == {"rows": [(1, "***")], "tags": {"***"}, "note": "key=***"} + + @pytest.mark.parametrize("value", [42, 3.5, True, None, b"raw bytes"]) + def test_leaves_non_text_values_alone(self, registered_secret, value): + assert mask_secrets(value) == value + + def test_keeps_values_under_keys_that_look_sensitive(self, registered_secret): + value = {"password_policy": "rotate every 90 days", "api_key_count": 3} + + assert mask_secrets(value) == value + + def test_masks_a_secret_used_as_a_key(self, registered_secret): + assert mask_secrets({registered_secret: 1}) == {"***": 1} + + def test_turns_a_pydantic_model_into_masked_data(self, registered_secret): + class Setting(BaseModel): + value: str + + assert mask_secrets([Setting(value=registered_secret)]) == [{"value": "***"}] + + def test_leaves_other_objects_such_as_images_alone(self, registered_secret): + image = BinaryContent(data=b"\x89PNG", media_type="image/png") + + assert mask_secrets({"screenshot": image})["screenshot"] is image + + def test_masks_a_dataclass_and_keeps_its_type(self, registered_secret): + @dataclass(frozen=True) + class Credentials: + user: str + password: str + + assert mask_secrets(Credentials("svc", registered_secret)) == Credentials("svc", "***") + assert mask_secrets(TextContent(content=f"key={registered_secret}")) == TextContent(content="key=***") + + def test_masks_bytes_as_text(self, registered_secret): + assert mask_secrets(f"key={registered_secret}".encode()) == b"key=***" + + def test_a_container_that_contains_itself_is_masked_without_recursing(self, registered_secret): + looped: list = [registered_secret] + looped.append(looped) + + assert mask_secrets(looped) == ["***", "<circular reference>"] + + +# Each changes representation inside a JSON string, so masking the dumped text would miss it. +SECRETS_JSON_ESCAPES = pytest.mark.parametrize( + "secret", + ['db-pa"ss-91c3', "db-pa\\ss-91c3", "db-p\u00e4ss-91c3"], + ids=["quote", "backslash", "non-ascii"], +) + + [email protected]_redact +class TestDumpsMasked: + @SECRETS_JSON_ESCAPES + def test_masks_a_secret_that_json_escapes(self, register_secret, secret): + register_secret(secret) + + dumped = dumps_masked({"password": secret, "rows": [[1, f"dsn={secret}"]]}) + + assert json.loads(dumped) == {"password": "***", "rows": [[1, "dsn=***"]]} + + def test_masks_an_object_json_can_only_render_as_text(self, registered_secret): + class Credentials: + def __str__(self) -> str: + return f"login with {registered_secret}" + + assert json.loads(dumps_masked({"auth": Credentials()})) == {"auth": "login with ***"} + + @SECRETS_JSON_ESCAPES + def test_masks_a_dataclass_or_bytes_holding_a_secret_that_escapes(self, register_secret, secret): + register_secret(secret) + + @dataclass + class Row: + dsn: str + + dumped = dumps_masked({"row": Row(f"pg://svc:{secret}@db"), "blob": f"key={secret}".encode()}) + + assert json.loads(dumped) == {"row": {"dsn": "pg://svc:***@db"}, "blob": "key=***"} diff --git a/providers/common/ai/tests/unit/common/ai/utils/test_query_results.py b/providers/common/ai/tests/unit/common/ai/utils/test_query_results.py index 0fb33b3b802..b5bebc921ac 100644 --- a/providers/common/ai/tests/unit/common/ai/utils/test_query_results.py +++ b/providers/common/ai/tests/unit/common/ai/utils/test_query_results.py @@ -160,3 +160,12 @@ class TestByteBudget: def test_empty_result_is_not_reported_as_truncated(self): data = _build(["id", "name"], []) assert data == {"columns": ["id", "name"], "rows": [], "row_count": 0} + + [email protected]_redact +def test_a_secret_that_json_escapes_is_masked_in_the_rows(register_secret): + secret = register_secret('db-pa"ss-91c3') + + result = _build(["user", "password"], [["admin", secret]]) + + assert result["rows"] == [["admin", "***"]] diff --git a/providers/common/ai/tests/unit/common/ai/utils/test_tool_definition.py b/providers/common/ai/tests/unit/common/ai/utils/test_tool_definition.py index 1664c453630..8c3707a516c 100644 --- a/providers/common/ai/tests/unit/common/ai/utils/test_tool_definition.py +++ b/providers/common/ai/tests/unit/common/ai/utils/test_tool_definition.py @@ -23,7 +23,11 @@ import pytest from pydantic_core import ValidationError from airflow.providers.common.ai.utils import tool_definition -from airflow.providers.common.ai.utils.tool_definition import build_args_validator, return_schema_kwargs +from airflow.providers.common.ai.utils.tool_definition import ( + build_args_validator, + return_schema_kwargs, + serialize_for_llm, +) def test_returns_kwarg_when_supported(): @@ -162,3 +166,18 @@ class TestBuildArgsValidator: validator = build_args_validator(schema) args = {"payload": {"any": 1, "deep": {"k": "v"}}} assert _validate(validator, args, use_json) == args + + [email protected]_redact +def test_serialize_for_llm_masks_a_secret_that_json_escapes(register_secret): + secret = register_secret('db-pa"ss-91c3') + + assert json.loads(serialize_for_llm({"password": secret})) == {"password": "***"} + + [email protected]_redact +def test_serialize_for_llm_masks_before_rendering_a_value_json_cannot_encode(register_secret): + """A tuple key makes json.dumps fail, and str() would escape the backslash in the secret.""" + secret = register_secret("db-pa\\ss-91c3") + + assert serialize_for_llm({("eu", "gold"): f"dsn={secret}"}) == "{('eu', 'gold'): 'dsn=***'}" diff --git a/providers/common/ai/tests/unit/common/ai/utils/test_toolset_base.py b/providers/common/ai/tests/unit/common/ai/utils/test_toolset_base.py new file mode 100644 index 00000000000..8154ce25016 --- /dev/null +++ b/providers/common/ai/tests/unit/common/ai/utils/test_toolset_base.py @@ -0,0 +1,281 @@ +# 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 +import sys +import threading +import time +import traceback +from typing import Any +from unittest.mock import MagicMock + +import pytest +from pydantic_ai import Agent +from pydantic_ai._run_context import RunContext +from pydantic_ai.exceptions import ApprovalRequired, ModelRetry, ToolFailed +from pydantic_ai.messages import ModelResponse, TextPart, ToolCallPart, ToolReturn, ToolReturnPart +from pydantic_ai.models.function import FunctionModel +from pydantic_ai.models.test import TestModel +from pydantic_ai.toolsets.abstract import ToolsetTool +from pydantic_ai.toolsets.function import FunctionToolset +from pydantic_ai.usage import RunUsage + +from airflow.providers.common.ai.utils.toolset_base import AirflowToolset, MaskingToolset, with_masking + + +class _ScriptedToolset(AirflowToolset): + """Returns or raises whatever the test hands it.""" + + def __init__(self, outcome: Any) -> None: + self._outcome = outcome + + @property + def id(self) -> str: + return "scripted" + + async def get_tools(self, ctx: RunContext[Any]) -> dict[str, ToolsetTool[Any]]: + return {} + + async def _execute_tool(self, name, tool_args, ctx, tool) -> Any: + if isinstance(self._outcome, BaseException): + raise self._outcome + return self._outcome + + +def _call(toolset: AirflowToolset) -> Any: + return asyncio.run( + toolset.call_tool("t", {}, ctx=MagicMock(spec=RunContext), tool=MagicMock(spec=ToolsetTool)) + ) + + [email protected]_redact +class TestMasking: + def test_masks_a_secret_in_a_text_result(self, registered_secret): + assert _call(_ScriptedToolset(f"password is {registered_secret}")) == "password is ***" + + def test_masks_a_secret_nested_deep_in_a_structured_result(self, registered_secret): + deep: Any = registered_secret + for _ in range(10): + deep = {"level": [deep]} + + masked = _call(_ScriptedToolset(deep)) + + for _ in range(10): + masked = masked["level"][0] + assert masked == "***" + + def test_masks_the_text_of_a_model_retry_and_drops_its_unmasked_cause(self, registered_secret): + """Tracing records a failed call's traceback, cause included.""" + cause = ConnectionError(f"login failed for {registered_secret}") + retry = ModelRetry(f"query failed: {cause}") + retry.__cause__ = cause + + with pytest.raises(ModelRetry) as caught: + _call(_ScriptedToolset(retry)) + + assert caught.value.message == "query failed: login failed for ***" + assert str(caught.value) == "query failed: login failed for ***" + assert registered_secret not in "".join(traceback.format_exception(caught.value)) + + def test_masks_the_text_of_a_tool_failure(self, registered_secret): + with pytest.raises(ToolFailed) as caught: + _call(_ScriptedToolset(ToolFailed(f"no such bucket for key {registered_secret}"))) + + assert caught.value.message == "no such bucket for key ***" + + def test_another_exception_keeps_its_type_with_its_message_masked(self, registered_secret): + """Retry policies match on the type; tracing records the message.""" + cause = OSError(f"socket closed by {registered_secret}") + error = PermissionError(f"denied: {registered_secret}") + error.__cause__ = cause + + with pytest.raises(PermissionError) as caught: + _call(_ScriptedToolset(error)) + + assert str(caught.value) == "denied: ***" + assert registered_secret not in "".join(traceback.format_exception(caught.value)) + + def test_masks_an_os_error_whose_message_comes_from_its_attributes(self, registered_secret): + """``OSError`` prints ``strerror`` and ``filename``, which masking ``args`` alone leaves as they were.""" + error = PermissionError(13, f"denied for {registered_secret}", f"/keys/{registered_secret}") + + with pytest.raises(PermissionError) as caught: + _call(_ScriptedToolset(error)) + + assert str(caught.value) == "[Errno 13] denied for ***: '/keys/***'" + + def test_masks_the_attributes_a_custom_message_is_built_from(self, registered_secret): + class LoginError(Exception): + def __init__(self, user: str, password: str) -> None: + super().__init__(user) + self.password = password + + def __str__(self) -> str: + return f"{self.args[0]} could not log in with {self.password}" + + with pytest.raises(LoginError) as caught: + _call(_ScriptedToolset(LoginError("admin", registered_secret))) + + assert str(caught.value) == "admin could not log in with ***" + + def test_an_exception_whose_message_cannot_be_masked_is_replaced(self, registered_secret): + class OpaqueError(Exception): + def __str__(self) -> str: + return f"handshake failed: {registered_secret}" + + with pytest.raises(RuntimeError) as caught: + _call(_ScriptedToolset(OpaqueError())) + + assert str(caught.value) == "OpaqueError: handshake failed: ***" + assert registered_secret not in "".join(traceback.format_exception(caught.value)) + + def test_masks_a_tool_return_and_keeps_it_one(self, registered_secret): + result = _call( + _ScriptedToolset( + ToolReturn(return_value=f"key={registered_secret}", content=f"for {registered_secret}") + ) + ) + + assert result == ToolReturn(return_value="key=***", content="for ***") + + def test_an_approval_request_keeps_its_type_and_is_not_logged_as_a_failure( + self, registered_secret, caplog + ): + with pytest.raises(ApprovalRequired) as caught: + _call(_ScriptedToolset(ApprovalRequired(metadata={"reason": registered_secret}))) + + assert caught.value.metadata == {"reason": "***"} + assert "failed" not in caplog.text + + @pytest.mark.skipif(sys.version_info < (3, 11), reason="ExceptionGroup is built in from Python 3.11") + def test_masks_the_exceptions_inside_an_exception_group(self, registered_secret): + cause = ConnectionError(f"socket closed by {registered_secret}") + inner = ValueError(f"login failed for {registered_secret}") + inner.__cause__ = cause + group = ExceptionGroup("tool calls failed", [inner]) # noqa: F821 + + with pytest.raises(ExceptionGroup) as caught: # noqa: F821 + _call(_ScriptedToolset(group)) + + assert [str(e) for e in caught.value.exceptions] == ["login failed for ***"] + assert registered_secret not in "".join(traceback.format_exception(caught.value)) + + def test_an_exception_that_cannot_be_masked_is_withheld(self): + class Unprintable(Exception): + def __str__(self) -> str: + raise RuntimeError("cannot render") + + with pytest.raises(RuntimeError, match="Unprintable: details withheld"): + _call(_ScriptedToolset(Unprintable())) + + +class TestRunBlocking: + def test_runs_off_the_event_loop_thread(self): + async def call() -> tuple[int, int]: + loop_thread = threading.get_ident() + return loop_thread, await AirflowToolset.run_blocking(threading.get_ident) + + loop_thread, call_thread = asyncio.run(call()) + + assert call_thread != loop_thread + + def test_blocking_calls_from_different_toolsets_never_overlap(self): + active = 0 + overlapped = False + + def blocking_call() -> None: + nonlocal active, overlapped + active += 1 + overlapped = overlapped or active > 1 + time.sleep(0.05) + active -= 1 + + async def call_concurrently() -> None: + first, second = _ScriptedToolset(None), _ScriptedToolset(None) + await asyncio.gather(*(ts.run_blocking(blocking_call) for ts in (first, second, first))) + + asyncio.run(call_concurrently()) + + assert not overlapped + + [email protected]_redact +class TestMaskingToolset: + def test_masks_a_toolset_that_does_not_mask_itself(self, registered_secret): + def read_setting() -> str: + return f"api key: {registered_secret}" + + masked = MaskingToolset(wrapped=FunctionToolset([read_setting])) + + async def call() -> Any: + ctx = RunContext(deps=None, model=TestModel(), usage=RunUsage()) + tools = await masked.get_tools(ctx) + return await masked.call_tool("read_setting", {}, ctx, tools["read_setting"]) + + assert asyncio.run(call()) == "api key: ***" + + +class _ApiToolset(FunctionToolset): + """A Dag author's toolset whose one tool returns a connection string.""" + + def __init__(self, secret: str) -> None: + def connection_string() -> str: + """Return the connection string.""" + return f"postgres://svc:{secret}@db" + + super().__init__([connection_string]) + + +def _run_agent(toolset) -> str: + def model(messages, info): + returns = [p for m in messages for p in m.parts if isinstance(p, ToolReturnPart)] + if returns: + return ModelResponse(parts=[TextPart(str(returns[0].content))]) + return ModelResponse(parts=[ToolCallPart("connection_string", {}, tool_call_id="c")]) + + return Agent(FunctionModel(model), toolsets=[toolset]).run_sync("go").output + + [email protected]_redact +class TestWithMasking: + def test_a_function_that_builds_a_toolset_per_run_is_masked_and_still_runs(self, registered_secret): + def per_run(ctx: RunContext[Any]) -> FunctionToolset: + return _ApiToolset(registered_secret) + + assert _run_agent(with_masking(per_run)) == "postgres://svc:***@db" + + def test_an_airflow_toolset_that_overrides_call_tool_is_wrapped(self, registered_secret): + class Overriding(_ScriptedToolset): + async def call_tool(self, name, tool_args, ctx, tool) -> Any: + return f"token {registered_secret}" + + masked = with_masking(Overriding("unused")) + + assert isinstance(masked, MaskingToolset) + assert _call(masked) == "token ***" + + def test_an_airflow_toolset_that_masks_itself_is_not_wrapped(self): + toolset = _ScriptedToolset("ok") + + assert with_masking(toolset) is toolset + + def test_a_failure_masked_by_two_layers_is_logged_once(self, caplog): + with pytest.raises(ValueError, match="boom"): + _call(MaskingToolset(wrapped=_ScriptedToolset(ValueError("boom")))) + + assert caplog.text.count("Tool t failed") == 1 diff --git a/scripts/in_container/run_provider_yaml_files_check.py b/scripts/in_container/run_provider_yaml_files_check.py index 2ab9dc02ee5..5ed46fe888c 100755 --- a/scripts/in_container/run_provider_yaml_files_check.py +++ b/scripts/in_container/run_provider_yaml_files_check.py @@ -87,6 +87,8 @@ INTERNAL_UNREGISTERED_TOOLSET_CLASSES = { # Wraps a toolset with per-step result caching for durable execution; applied # automatically by AgentOperator, not part of the public toolsets how-to guide. "airflow.providers.common.ai.durable.caching_toolset.CachingToolset", + # Masks what a Dag author's toolset returns; applied automatically by AgentOperator. + "airflow.providers.common.ai.utils.toolset_base.MaskingToolset", } if __name__ != "__main__":
