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