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 35507e99a6ef7593cd1c2ad704fd1551176a39c2 Author: Kaxil Naik <[email protected]> AuthorDate: Wed Sep 30 07:04:24 2026 +0100 Let HookToolset pin arguments the model must not choose (#73900) pinned_arguments fixes values such as the bucket a storage hook may read. A pinned argument is left out of the schema the model sees, passed to every allowed method, and refused if the model supplies it anyway. Every allowed method has to take each pinned argument by name: one that takes the same thing under another name, in a dict or through **kwargs would let the model choose it after all, so the toolset refuses to be built with it. --- providers/common/ai/docs/agent_security.rst | 8 +- providers/common/ai/docs/stability.rst | 8 +- providers/common/ai/docs/toolsets/hook.rst | 51 ++++++++- .../airflow/providers/common/ai/toolsets/hook.py | 52 ++++++++- .../ai/tests/unit/common/ai/toolsets/test_hook.py | 122 +++++++++++++++++++++ 5 files changed, 231 insertions(+), 10 deletions(-) diff --git a/providers/common/ai/docs/agent_security.rst b/providers/common/ai/docs/agent_security.rst index 5a169a869e5..08a6fbd96a5 100644 --- a/providers/common/ai/docs/agent_security.rst +++ b/providers/common/ai/docs/agent_security.rst @@ -104,7 +104,8 @@ No single layer is sufficient on its own. They work together. - Only methods listed in ``allowed_methods`` are exposed as tools. Auto-discovery is not supported. Methods are validated at Dag parse time. - - Does not restrict what arguments the agent passes to allowed methods. + - Restricts only the arguments named in ``pinned_arguments``, and only by parameter + name; the agent chooses every other argument of an allowed method. * - **SQLToolset: read-only by default** - ``allow_writes=False`` (default) validates every SQL query through ``validate_sql()``: SELECT-family and read-only metadata @@ -235,7 +236,8 @@ database user with the minimum privileges required. ``get_connection()``: these give broad access. - Prefer read-only methods (``list_*``, ``get_*``, ``describe_*``). - The agent controls arguments. If a method accepts a ``path`` parameter, - the agent can pass any path the hook has access to. + the agent can pass any path the hook has access to, unless the Dag author pins it with + ``pinned_arguments`` (see :doc:`toolsets/hook`). .. code-block:: python @@ -290,7 +292,7 @@ Before deploying an agent task to production: 2. **Database permissions**: Create a dedicated database user with minimum required grants. Don't reuse the admin connection. 3. **Tool allow-list**: Review ``allowed_methods`` / ``allowed_tables``. The - agent can call any exposed tool with any arguments. + agent can call any exposed tool with any arguments it does not pin. 4. **Read-only default**: Keep ``allow_writes=False`` unless the task specifically requires writes. 5. **Result limits**: Set ``max_rows`` and ``max_result_bytes`` appropriate to diff --git a/providers/common/ai/docs/stability.rst b/providers/common/ai/docs/stability.rst index ab798c360d7..78717560fcd 100644 --- a/providers/common/ai/docs/stability.rst +++ b/providers/common/ai/docs/stability.rst @@ -79,7 +79,8 @@ Pydantic AI toolsets, but the Pydantic AI class they inherit from can change. * - :class:`~airflow.providers.common.ai.toolsets.hook.HookToolset` - Exposes exactly the hook methods in ``allowed_methods``, each named after its method with ``tool_name_prefix`` in front, and raises an error when the toolset is created - if a listed method does not exist on the hook. + if a listed method does not exist on the hook. ``pinned_arguments`` is + experimental; see below. * - :class:`~airflow.providers.common.ai.toolsets.mcp.MCPToolset` and :class:`~airflow.providers.common.ai.hooks.mcp.MCPHook` - Exposes the tools of the MCP server configured by ``mcp_conn_id``, each named @@ -175,3 +176,8 @@ Everything this provider ships that is not in the table above is experimental. * - :class:`~airflow.providers.common.ai.toolsets.object_storage.ObjectStorageToolset` (:doc:`toolsets/object_storage`) - New; its tools and read limits may change after first use. + * - ``pinned_arguments`` on + :class:`~airflow.providers.common.ai.toolsets.hook.HookToolset` + (:doc:`toolsets/hook`) + - New; how a pinned argument is matched to each method's parameters may change after + first use. diff --git a/providers/common/ai/docs/toolsets/hook.rst b/providers/common/ai/docs/toolsets/hook.rst index 48b2fdd014e..9acf6ed2a7b 100644 --- a/providers/common/ai/docs/toolsets/hook.rst +++ b/providers/common/ai/docs/toolsets/hook.rst @@ -82,6 +82,47 @@ is not a connection ID yet. The same warning as for ``SQLToolset`` applies: buil the ID from values the Dag controls, not from ``params`` or ``dag_run.conf`` (see :ref:`sql-toolset-templated-connection`). +Fix arguments the model must not choose +--------------------------------------- + +.. note:: + + Experimental: ``pinned_arguments`` can change or be removed in a minor release of this + provider. + See :ref:`howto/stability`. + +Exposing a method lets the model pick every argument it takes. When some of them are +the Dag author's decision, such as which bucket a storage hook reads, pin them: + +.. code-block:: python + + from airflow.providers.amazon.aws.hooks.s3 import S3Hook + + from airflow.providers.common.ai.toolsets import HookToolset + + reports = HookToolset( + S3Hook(aws_conn_id="aws_default"), + allowed_methods=["list_keys", "read_key"], + pinned_arguments={"bucket_name": "acme-reports"}, + ) + +A pinned argument is left out of the schema the model sees and passed to every allowed +method. If the model supplies it anyway, the call is refused and the model is told the +argument is fixed. + +A pin binds one parameter name, so every allowed method has to take it by that name. +When one does not, the toolset raises ``ValueError`` when it is created: a method that +takes the same thing under another name, such as ``S3Hook.delete_objects``, which takes +``bucket``, or inside a dict or ``**kwargs``, such as ``S3Hook.generate_presigned_url``, +would let the model choose it after all. Expose such a method from a second +``HookToolset``, where what it can reach is visible in the Dag. A method that works out +the value again from another argument it is given, such as a full URL, is outside what +the pin controls. + +Pinned values are passed as written: they are not rendered as templates, and they are not +part of what ``AgentOperator(durable=True)`` fingerprints, so change one only between Dag +runs, not between the tries of one. + Parameters ---------- @@ -90,6 +131,8 @@ Parameters are validated with ``hasattr`` + ``callable`` at instantiation time. - ``tool_name_prefix``: Optional prefix prepended to each tool name (e.g. ``"s3_"`` produces ``"s3_list_keys"``). +- ``pinned_arguments``: Arguments fixed by the Dag author rather than chosen by the + model. See above. When to choose it ----------------- @@ -103,10 +146,10 @@ reflection-based adapter, so the work is choosing the method list. **What it cannot do** -- It allow-lists method *names*, not arguments. Once ``read_key`` is exposed, - the agent picks the key; the :ref:`defense-layer table <toolset-defense-layers>` - states this outright. Choose methods whose worst case you accept, not methods - you intend to constrain later. +- It allow-lists method *names*, and fixes only the arguments you pin. Once + ``read_key`` is exposed, the agent picks the key within the pinned bucket; the + :ref:`defense-layer table <toolset-defense-layers>` 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`` 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 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 84e2530ac6c..882bee88eb5 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 @@ -24,6 +24,7 @@ import re import types from typing import TYPE_CHECKING, Any, Union, get_args, get_origin, get_type_hints +from pydantic_ai.exceptions import ModelRetry from pydantic_ai.tools import ToolDefinition from pydantic_ai.toolsets.abstract import ToolsetTool @@ -35,7 +36,7 @@ from airflow.providers.common.ai.utils.tool_definition import ( from airflow.providers.common.ai.utils.toolset_base import AirflowToolset if TYPE_CHECKING: - from collections.abc import Callable, Sequence + from collections.abc import Callable, Iterable, Sequence from pydantic_ai._run_context import RunContext @@ -71,6 +72,14 @@ class HookToolset(AirflowToolset): auto-discovery is intentionally not supported for safety. :param tool_name_prefix: Optional prefix prepended to each tool name (e.g. ``"s3_"`` → ``"s3_list_keys"``). + :param pinned_arguments: Experimental. Arguments the Dag author fixes, such as the bucket a + storage hook may use: ``{"bucket_name": "reports"}``. Each is left out of the + arguments the model sees, refused if the model supplies it anyway, and passed to + every allowed method as it is written here, not rendered as a template. Every allowed + method must take each pinned argument as a named parameter: one that does not, such + as a method taking ``bucket`` or only ``**kwargs``, raises ``ValueError``, because + the model could still choose the value through it. Expose such a method from a + second ``HookToolset``. """ # Rendered, on a copy, by AgentOperator. Deliberately not ``template_fields``, which @@ -83,6 +92,7 @@ class HookToolset(AirflowToolset): *, allowed_methods: list[str], tool_name_prefix: str = "", + pinned_arguments: dict[str, Any] | None = None, ) -> None: if not allowed_methods: raise ValueError("allowed_methods must be a non-empty list.") @@ -96,6 +106,27 @@ class HookToolset(AirflowToolset): if not callable(getattr(hook, method_name)): raise ValueError(f"{hook_cls_name}.{method_name} is not callable.") + # Every allowed method has to name each pin as a parameter it can be passed by name. A + # method that takes the value under another name, inside a dict, or through **kwargs + # would let the model choose it after all, so it is refused rather than left unpinned. + pinned_arguments = pinned_arguments or {} + unpinned: dict[str, list[str]] = {} + for method_name in allowed_methods if pinned_arguments else (): + parameters = inspect.signature(getattr(hook, method_name)).parameters.values() + named = {p.name for p in parameters if p.kind in (p.POSITIONAL_OR_KEYWORD, p.KEYWORD_ONLY)} + if missing := sorted(set(pinned_arguments) - named): + unpinned[method_name] = missing + if unpinned: + details = "; ".join( + f"{method}() does not take {', '.join(args)}" for method, args in unpinned.items() + ) + raise ValueError( + f"Every allowed method of {hook_cls_name!r} has to take each pinned argument by name, or " + f"the model could still choose it through that method: {details}. Expose such a method " + "from a second HookToolset." + ) + self._pinned: dict[str, Any] = dict(pinned_arguments) + self._hook = hook self._allowed_methods = allowed_methods self._tool_name_prefix = tool_name_prefix @@ -142,6 +173,7 @@ class HookToolset(AirflowToolset): for param_name, param_desc in param_docs.items(): if param_name in json_schema.get("properties", {}): json_schema["properties"][param_name]["description"] = param_desc + _drop_properties(json_schema, self._pinned) # sequential=True keeps pydantic-ai from running these calls concurrently # within a turn; run_blocking's process-wide lock serializes them with the @@ -174,7 +206,14 @@ class HookToolset(AirflowToolset): ) -> 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 = await self.run_blocking(method, **tool_args) + if supplied := sorted(self._pinned.keys() & tool_args.keys()): + one = len(supplied) == 1 + raise ModelRetry( + f"{', '.join(supplied)} {'is' if one else 'are'} fixed for this tool: call it again " + f"without {'it' if one else 'them'}." + ) + # A copy per call, so a method that modifies an argument it is given cannot change the pin. + result = await self.run_blocking(method, **tool_args, **copy.deepcopy(self._pinned)) return serialize_for_llm(result) @@ -298,3 +337,12 @@ def _parse_param_docs(docstring: str) -> dict[str, str]: params[m.group(1)] = " ".join(m.group(2).split()) return params + + +def _drop_properties(schema: dict[str, Any], names: Iterable[str]) -> None: + for name in names: + schema["properties"].pop(name, None) + if name in schema.get("required", ()): + schema["required"].remove(name) + if "required" in schema and not schema["required"]: + del schema["required"] 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 a54b0c6b051..121e1a6d6d4 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,11 +17,15 @@ from __future__ import annotations import asyncio +import re import threading from unittest.mock import MagicMock import pytest +from pydantic_ai import Agent from pydantic_ai._run_context import RunContext +from pydantic_ai.messages import ModelResponse, RetryPromptPart, TextPart, ToolCallPart, ToolReturnPart +from pydantic_ai.models.function import FunctionModel from pydantic_core import ValidationError from airflow.providers.common.ai.toolsets.hook import ( @@ -437,3 +441,121 @@ class TestSerializeForLlm: obj = object() result = serialize_for_llm(obj) assert "object" in result + + +class TestHookToolsetPinnedArguments: + @staticmethod + def _tools(ts: HookToolset) -> dict: + return asyncio.run(ts.get_tools(ctx=MagicMock(spec=RunContext))) + + def test_a_pinned_argument_is_left_out_of_the_schema(self): + ts = HookToolset(_FakeHook(), allowed_methods=["list_keys"], pinned_arguments={"bucket": "reports"}) + + schema = self._tools(ts)["list_keys"].tool_def.parameters_json_schema + + assert "bucket" not in schema["properties"] + assert "required" not in schema + assert "prefix" in schema["properties"] + + def test_the_pinned_value_is_passed_on_every_call(self): + hook = _RecordingHook() + ts = HookToolset(hook, allowed_methods=["list_keys"], pinned_arguments={"bucket": "reports"}) + tools = self._tools(ts) + + asyncio.run( + ts.call_tool( + "list_keys", {"prefix": "2026/"}, ctx=MagicMock(spec=RunContext), tool=tools["list_keys"] + ) + ) + + assert hook.calls == [("reports", "2026/")] + + @pytest.mark.parametrize( + ("method", "missing"), + [ + pytest.param("read_file", "read_file() does not take bucket", id="another_name"), + pytest.param("request", "request() does not take bucket", id="kwargs_only"), + ], + ) + def test_every_allowed_method_has_to_take_the_pin_by_name(self, method, missing): + """A method that does not would let the model choose the value through it.""" + with pytest.raises(ValueError, match=re.escape(missing)): + HookToolset(_FakeHook(), allowed_methods=["list_keys", method], pinned_arguments={"bucket": "x"}) + + def test_a_pin_on_a_catch_all_parameter_is_rejected(self): + with pytest.raises(ValueError, match=r"request\(\) does not take kwargs"): + HookToolset(_FakeHook(), allowed_methods=["request"], pinned_arguments={"kwargs": {"b": "x"}}) + + def test_methods_that_all_take_the_pin_by_name_are_accepted(self): + ts = HookToolset( + _RecordingKwargsHook(), allowed_methods=["copy", "list_keys"], pinned_arguments={"bucket": "x"} + ) + + assert set(self._tools(ts)) == {"copy", "list_keys"} + + def test_the_model_cannot_override_it_in_a_real_run(self): + """A method taking **kwargs would accept the model's value, so the toolset refuses it.""" + hook = _RecordingKwargsHook() + ts = HookToolset(hook, allowed_methods=["list_keys"], pinned_arguments={"bucket": "reports"}) + attempts = iter([{"bucket": "payroll", "prefix": "x"}, {"prefix": "x"}]) + + def model(messages, info): + retried = [p for m in messages for p in m.parts if isinstance(p, RetryPromptPart)] + returned = [p for m in messages for p in m.parts if isinstance(p, ToolReturnPart)] + if returned: + return ModelResponse(parts=[TextPart(str(retried[0].content))]) + return ModelResponse(parts=[ToolCallPart("list_keys", next(attempts), tool_call_id="c")]) + + answer = Agent(FunctionModel(model), toolsets=[ts]).run_sync("list").output + + assert "bucket is fixed for this tool" in answer + assert hook.calls == [("reports", "x", {})] + + def test_a_method_that_modifies_its_argument_cannot_change_the_pin(self): + hook = _RecordingKwargsHook() + ts = HookToolset( + hook, allowed_methods=["copy"], pinned_arguments={"bucket": "reports", "tags": {"a": "1"}} + ) + tools = self._tools(ts) + ctx = MagicMock(spec=RunContext) + + for _ in range(2): + asyncio.run(ts.call_tool("copy", {"key": "k"}, ctx=ctx, tool=tools["copy"])) + + assert [call[2]["tags"] for call in hook.calls] == [{"a": "1"}, {"a": "1"}] + + +class _RecordingKwargsHook: + """Records its calls; ``list_keys`` takes **kwargs, so validation alone would let extra names in.""" + + def __init__(self) -> None: + self.calls: list[tuple[str, str | None, dict[str, object]]] = [] + + def list_keys(self, bucket: str, prefix: str | None = None, **kwargs: object) -> list[str]: + """List object keys in a bucket.""" + self.calls.append((bucket, prefix, kwargs)) + return [f"{bucket}/{prefix}"] + + def copy(self, bucket: str, key: str, tags: dict[str, str] | None = None) -> str: + """Copy an object, adding a tag as a side effect.""" + self.calls.append((bucket, key, {"tags": dict(tags or {})})) + if tags is not None: + tags["copied"] = "yes" + return key + + +class _RecordingHook: + """Records the arguments its method receives.""" + + def __init__(self) -> None: + self.calls: list[tuple[str, str | None]] = [] + + def list_keys(self, bucket: str, prefix: str | None = None) -> list[str]: + """ + List object keys in a bucket. + + :param bucket: Name of the bucket. + :param prefix: Key prefix to filter by. + """ + self.calls.append((bucket, prefix)) + return [f"{bucket}/{prefix}"]
