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


The following commit(s) were added to refs/heads/main by this push:
     new 4ef67b4e0c6 Tell the model a pinned hook argument is fixed for every 
allowed method (#74380)
4ef67b4e0c6 is described below

commit 4ef67b4e0c686ce713721a2fb4988f3c0b9b19a1
Author: Kaxil Naik <[email protected]>
AuthorDate: Wed Oct 7 14:17:05 2026 +0100

    Tell the model a pinned hook argument is fixed for every allowed method 
(#74380)
    
    When the model supplied a pinned argument to a HookToolset method with
    named parameters, such as S3Hook.read_key, argument validation rejected it
    first with a generic "Extra inputs are not permitted" error, and the
    toolset's own "is fixed for this tool" message only ever reached methods
    that also take **kwargs. The toolset now validates against a schema that
    still accepts the pinned names, so every refusal tells the model the
    argument is fixed. The model is still shown the schema without them.
    
    * Refuse pinned hook arguments during validation
    
    The refusal ran in execute_tool, after an approval gate had already asked a
    person to approve the call with the pinned argument in it. It now runs as 
the
    tool's args_validator_func, so the call is refused before any gate or the 
hook
    sees it; execute_tool keeps the check for the framework bridges, which 
validate
    arguments without that hook. The validator schema now adds the pinned names
    back as untyped optional properties instead of re-deriving the required 
list,
    so a pinned argument sent with the wrong type is also told it is fixed.
---
 providers/common/ai/docs/toolsets/hook.rst         |  24 +----
 providers/common/ai/docs/toolsets/index.rst        |   2 +-
 .../airflow/providers/common/ai/toolsets/hook.py   |  28 ++++--
 .../providers/common/ai/utils/tool_definition.py   |   2 +-
 .../ai/tests/unit/common/ai/toolsets/test_hook.py  | 109 +++++++++++++--------
 5 files changed, 95 insertions(+), 70 deletions(-)

diff --git a/providers/common/ai/docs/toolsets/hook.rst 
b/providers/common/ai/docs/toolsets/hook.rst
index 89df4662326..f44b72f62cc 100644
--- a/providers/common/ai/docs/toolsets/hook.rst
+++ b/providers/common/ai/docs/toolsets/hook.rst
@@ -107,11 +107,9 @@ the Dag author's decision, such as which bucket a storage 
hook reads, pin them:
     )
 
 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. For a method 
with named
-parameters, such as ``S3Hook.read_key``, argument validation refuses it as an 
extra
-input (see :ref:`hook-toolset-restricted`). A method that names the parameter 
and also
-takes ``**kwargs`` would accept it, so the toolset refuses it itself and tells 
the
-model the argument is fixed.
+method. If the model supplies it anyway, the call is refused while its 
arguments are
+validated, before an approval gate or the hook sees it. When the rest of the 
call is
+valid, the model is told the argument is fixed (see 
:ref:`hook-toolset-restricted`).
 
 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
@@ -147,19 +145,7 @@ that bucket was refused before it reached S3, with this 
message to the model:
 
 .. code-block:: text
 
-    1 validation error:
-    ```json
-    [
-      {
-        "type": "extra_forbidden",
-        "loc": [
-          "bucket_name"
-        ],
-        "msg": "Extra inputs are not permitted",
-        "input": "acme-payroll"
-      }
-    ]
-    ```
+    bucket_name is fixed for this tool: call it again without it.
 
     Fix the errors and try again.
 
@@ -187,7 +173,7 @@ Parameters
 - ``pinned_arguments``: Arguments fixed by the Dag author rather than chosen 
by the
   model. See above.
 - ``max_retries``: How many times the model may correct a call with invalid 
arguments,
-  or one that changes a pinned argument. Default ``None``, the agent's 
``retries``. See
+  or one that supplies a pinned argument. Default ``None``, the agent's 
``retries``. See
   :ref:`toolset-retry-budget`.
 
 When to choose it
diff --git a/providers/common/ai/docs/toolsets/index.rst 
b/providers/common/ai/docs/toolsets/index.rst
index d4a87b5ce85..d9aacae2d5f 100644
--- a/providers/common/ai/docs/toolsets/index.rst
+++ b/providers/common/ai/docs/toolsets/index.rst
@@ -255,7 +255,7 @@ the call. What counts differs by toolset:
 
 - ``SQLToolset`` and ``DataFusionToolset`` turn every query error into 
``ModelRetry``,
   so a misspelled column and a dropped connection both count.
-- ``HookToolset`` counts invalid arguments and a call that tries to change a 
pinned
+- ``HookToolset`` counts invalid arguments and a call that supplies a pinned
   argument. An exception from the hook itself fails the run straight away.
 - ``ObjectStorageToolset`` counts invalid arguments only. A path that does not 
exist or
   cannot be read goes back to the model as a failed result without using the 
budget;
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 0dac4f5b080..66814946586 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
@@ -81,7 +81,7 @@ class HookToolset(AirflowToolset):
         the model could still choose the value through it. Expose such a 
method from a
         second ``HookToolset``.
     :param max_retries: How many times the model may correct a call with 
invalid arguments,
-        or one that changes a pinned argument, before the run fails. An 
exception from the
+        or one that supplies a pinned argument, before the run fails. An 
exception from the
         hook itself fails the run straight away. ``None`` (the default) uses 
the agent's
         tool retry budget, its ``retries``, as pydantic-ai's own toolsets do.
     """
@@ -181,6 +181,12 @@ class HookToolset(AirflowToolset):
                 if param_name in json_schema.get("properties", {}):
                     json_schema["properties"][param_name]["description"] = 
param_desc
             _drop_properties(json_schema, self._pinned)
+            # The validator accepts the pinned names, with any value, so a 
model that sends one
+            # anyway is told it is fixed rather than given a generic 
extra-input or type error.
+            args_schema = {
+                **json_schema,
+                "properties": json_schema["properties"] | {name: {} for name 
in 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
@@ -200,10 +206,20 @@ class HookToolset(AirflowToolset):
                 toolset=self,
                 tool_def=tool_def,
                 max_retries=max_retries,
-                args_validator=build_args_validator(json_schema),
+                args_validator=build_args_validator(args_schema),
+                # Refused during validation, so an approval gate never asks 
about such a call.
+                args_validator_func=self._refuse_pinned if self._pinned else 
None,
             )
         return tools
 
+    def _refuse_pinned(self, ctx: RunContext[Any], /, **tool_args: Any) -> 
None:
+        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'}."
+            )
+
     async def execute_tool(
         self,
         name: str,
@@ -214,12 +230,8 @@ 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)
-        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'}."
-            )
+        # The framework bridges validate arguments without 
args_validator_func, so check again here.
+        self._refuse_pinned(ctx, **tool_args)
         # 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)
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 978581d0f6e..29cfdc91b2c 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
@@ -158,5 +158,5 @@ def _object_fragment_to_core_schema(fragment: dict[str, 
Any]) -> core_schema.Cor
 
 
 def build_args_validator(parameters_json_schema: dict[str, Any]) -> 
SchemaValidator:
-    """Build an argument validator from the schema advertised to the model."""
+    """Build an argument validator from a tool's JSON parameters schema."""
     return 
SchemaValidator(_object_fragment_to_core_schema(parameters_json_schema))
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 121e1a6d6d4..d3154e27bf9 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
@@ -22,7 +22,7 @@ import threading
 from unittest.mock import MagicMock
 
 import pytest
-from pydantic_ai import Agent
+from pydantic_ai import Agent, DeferredToolRequests
 from pydantic_ai._run_context import RunContext
 from pydantic_ai.messages import ModelResponse, RetryPromptPart, TextPart, 
ToolCallPart, ToolReturnPart
 from pydantic_ai.models.function import FunctionModel
@@ -443,6 +443,42 @@ class TestSerializeForLlm:
         assert "object" in result
 
 
+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}"]
+
+
 class TestHookToolsetPinnedArguments:
     @staticmethod
     def _tools(ts: HookToolset) -> dict:
@@ -493,9 +529,16 @@ class TestHookToolsetPinnedArguments:
 
         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()
+    @pytest.mark.parametrize(
+        ("hook_cls", "expected_call"),
+        [
+            pytest.param(_RecordingHook, ("reports", "x"), 
id="named_parameters"),
+            pytest.param(_RecordingKwargsHook, ("reports", "x", {}), 
id="also_kwargs"),
+        ],
+    )
+    def test_the_model_cannot_override_it_in_a_real_run(self, hook_cls, 
expected_call):
+        """The model is told the argument is fixed, whether or not the method 
also takes **kwargs."""
+        hook = hook_cls()
         ts = HookToolset(hook, allowed_methods=["list_keys"], 
pinned_arguments={"bucket": "reports"})
         attempts = iter([{"bucket": "payroll", "prefix": "x"}, {"prefix": 
"x"}])
 
@@ -509,7 +552,27 @@ class TestHookToolsetPinnedArguments:
         answer = Agent(FunctionModel(model), 
toolsets=[ts]).run_sync("list").output
 
         assert "bucket is fixed for this tool" in answer
-        assert hook.calls == [("reports", "x", {})]
+        assert hook.calls == [expected_call]
+
+    def test_an_approval_gate_is_never_asked_about_a_pinned_argument(self):
+        """The refusal happens during validation, before an approval gate sees 
the call."""
+        ts = HookToolset(
+            _RecordingHook(), allowed_methods=["list_keys"], 
pinned_arguments={"bucket": "reports"}
+        )
+        attempts = iter([{"bucket": "payroll", "prefix": "x"}, {"prefix": 
"x"}])
+
+        def model(messages, info):
+            return ModelResponse(parts=[ToolCallPart("list_keys", 
next(attempts), tool_call_id="c")])
+
+        agent = Agent(
+            FunctionModel(model),
+            toolsets=[ts.approval_required()],
+            output_type=[str, DeferredToolRequests],
+        )
+        output = agent.run_sync("list").output
+
+        assert isinstance(output, DeferredToolRequests)
+        assert [call.args for call in output.approvals] == [{"prefix": "x"}]
 
     def test_a_method_that_modifies_its_argument_cannot_change_the_pin(self):
         hook = _RecordingKwargsHook()
@@ -523,39 +586,3 @@ class TestHookToolsetPinnedArguments:
             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