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}"]

Reply via email to