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 77838ed374f Validate common AI tool-call arguments to enable 
pydantic-ai retries (#70096)
77838ed374f is described below

commit 77838ed374f1dca38004785056ca1dd24061214e
Author: Yossi Eliaz <[email protected]>
AuthorDate: Thu Jul 30 20:23:40 2026 +0300

    Validate common AI tool-call arguments to enable pydantic-ai retries 
(#70096)
    
    Malformed model-generated arguments currently bypass validation in the 
Common AI toolsets. Missing or mistyped arguments can therefore reach tool 
implementations and fail with `KeyError` or `TypeError` instead of producing 
the bounded retry that pydantic-ai supports.
    
    Invalid payloads are now rejected at validation; pydantic-ai then 
regenerates a new tool call — the original malformed payload is never 
re-executed.
    
    Concretely: if the model calls the `query` tool but omits the required 
`sql` argument, the call previously crashed the run with `KeyError: 'sql'`; it 
now raises a `ValidationError` that pydantic-ai turns into one bounded retry, 
and the model regenerates the call with `sql` supplied.
    
    This change derives each tool's argument validator from the same JSON 
Schema advertised to the model and applies it to `SQLToolset`, 
`DataFusionToolset`, and `HookToolset`. Hook introspection keeps optional and 
union parameters accurate, leaves unsupported annotations permissive, rejects 
unexpected arguments for fixed signatures (so a mistyped field becomes a 
bounded retry, matching native pydantic-ai), and preserves additional arguments 
for methods that accept `**kwargs`.
    
    The LangChain bridge mirrors the same two-stage semantics: 
`ValidationError` from arg validation is fed back as a retry; `ValidationError` 
raised inside `call_tool` propagates so a non-idempotent tool is not re-invoked.
---
 .../providers/common/ai/toolsets/datafusion.py     |   7 +-
 .../airflow/providers/common/ai/toolsets/hook.py   |  30 ++---
 .../common/ai/toolsets/langchain_bridge.py         |  42 +++++--
 .../airflow/providers/common/ai/toolsets/sql.py    |   7 +-
 .../providers/common/ai/utils/tool_definition.py   |  74 +++++++++++-
 .../unit/common/ai/toolsets/test_datafusion.py     |  19 +++
 .../ai/tests/unit/common/ai/toolsets/test_hook.py  |  59 ++++++++-
 .../common/ai/toolsets/test_langchain_bridge.py    |  46 +++++++
 .../ai/tests/unit/common/ai/toolsets/test_sql.py   |  18 +++
 .../unit/common/ai/utils/test_tool_definition.py   | 134 ++++++++++++++++++++-
 10 files changed, 392 insertions(+), 44 deletions(-)

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 a83bd8bbc8c..78013918f87 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
@@ -35,7 +35,8 @@ 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_core import SchemaValidator, core_schema
+
+from airflow.providers.common.ai.utils.tool_definition import 
build_args_validator
 
 if TYPE_CHECKING:
     from pydantic_ai._run_context import RunContext
@@ -44,8 +45,6 @@ if TYPE_CHECKING:
 
 log = logging.getLogger(__name__)
 
-_PASSTHROUGH_VALIDATOR = SchemaValidator(core_schema.any_schema())
-
 # JSON Schemas for the three DataFusion tools.
 _LIST_TABLES_SCHEMA: dict[str, Any] = {
     "type": "object",
@@ -146,7 +145,7 @@ class DataFusionToolset(AbstractToolset[Any]):
                 toolset=self,
                 tool_def=tool_def,
                 max_retries=1,
-                args_validator=_PASSTHROUGH_VALIDATOR,
+                args_validator=build_args_validator(schema),
             )
         return tools
 
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 132c5dc9db2..63037e1a4f8 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
@@ -26,9 +26,8 @@ from typing import TYPE_CHECKING, Any, Union, get_args, 
get_origin, get_type_hin
 
 from pydantic_ai.tools import ToolDefinition
 from pydantic_ai.toolsets.abstract import AbstractToolset, ToolsetTool
-from pydantic_core import SchemaValidator, core_schema
 
-from airflow.providers.common.ai.utils.tool_definition import 
return_schema_kwargs
+from airflow.providers.common.ai.utils.tool_definition import 
build_args_validator, return_schema_kwargs
 
 if TYPE_CHECKING:
     from collections.abc import Callable
@@ -37,9 +36,6 @@ if TYPE_CHECKING:
 
     from airflow.providers.common.compat.sdk import BaseHook
 
-# Single shared validator — accepts any JSON-decoded dict from the LLM.
-_PASSTHROUGH_VALIDATOR = SchemaValidator(core_schema.any_schema())
-
 # Maps Python types to JSON Schema fragments.
 _TYPE_MAP: dict[type, dict[str, Any]] = {
     str: {"type": "string"},
@@ -127,7 +123,7 @@ class HookToolset(AbstractToolset[Any]):
                 toolset=self,
                 tool_def=tool_def,
                 max_retries=1,
-                args_validator=_PASSTHROUGH_VALIDATOR,
+                args_validator=build_args_validator(json_schema),
             )
         return tools
 
@@ -152,17 +148,16 @@ class HookToolset(AbstractToolset[Any]):
 def _python_type_to_json_schema(annotation: Any) -> dict[str, Any]:
     """Convert a Python type annotation to a JSON Schema fragment."""
     if annotation is inspect.Parameter.empty or annotation is Any:
-        return {"type": "string"}
+        return {}
+
+    if annotation is type(None):
+        return {"type": "null"}
 
     origin = get_origin(annotation)
     args = get_args(annotation)
 
-    # Optional[X] is Union[X, None] — handle both types.UnionType (3.10+) and 
typing.Union
     if origin is types.UnionType or origin is Union:
-        non_none = [a for a in args if a is not type(None)]
-        if len(non_none) == 1:
-            return _python_type_to_json_schema(non_none[0])
-        return {"type": "string"}
+        return {"anyOf": [_python_type_to_json_schema(arg) for arg in args]}
 
     # list[X]
     if origin is list:
@@ -175,7 +170,7 @@ def _python_type_to_json_schema(annotation: Any) -> 
dict[str, Any]:
 
     # Always return a fresh copy — callers may mutate the dict (e.g. adding 
"description").
     schema = _TYPE_MAP.get(annotation)
-    return dict(schema) if schema else {"type": "string"}
+    return dict(schema) if schema else {}
 
 
 def _build_json_schema_from_signature(method: Callable[..., Any]) -> dict[str, 
Any]:
@@ -189,12 +184,15 @@ def _build_json_schema_from_signature(method: 
Callable[..., Any]) -> dict[str, A
 
     properties: dict[str, Any] = {}
     required: list[str] = []
+    allows_additional_properties = False
 
     for name, param in sig.parameters.items():
         if name in ("self", "cls"):
             continue
-        # Skip **kwargs and *args
-        if param.kind in (param.VAR_POSITIONAL, param.VAR_KEYWORD):
+        if param.kind is param.VAR_POSITIONAL:
+            continue
+        if param.kind is param.VAR_KEYWORD:
+            allows_additional_properties = True
             continue
 
         annotation = hints.get(name, param.annotation)
@@ -207,6 +205,8 @@ def _build_json_schema_from_signature(method: Callable[..., 
Any]) -> dict[str, A
     schema: dict[str, Any] = {"type": "object", "properties": properties}
     if required:
         schema["required"] = required
+    if allows_additional_properties:
+        schema["additionalProperties"] = True
     return schema
 
 
diff --git 
a/providers/common/ai/src/airflow/providers/common/ai/toolsets/langchain_bridge.py
 
b/providers/common/ai/src/airflow/providers/common/ai/toolsets/langchain_bridge.py
index 35876d82930..2bb4c2094e7 100644
--- 
a/providers/common/ai/src/airflow/providers/common/ai/toolsets/langchain_bridge.py
+++ 
b/providers/common/ai/src/airflow/providers/common/ai/toolsets/langchain_bridge.py
@@ -35,6 +35,7 @@ import asyncio
 import concurrent.futures
 from typing import TYPE_CHECKING, Any
 
+from pydantic import ValidationError
 from pydantic_ai import RunContext
 from pydantic_ai.exceptions import ModelRetry
 from pydantic_ai.models.test import TestModel
@@ -85,10 +86,19 @@ def airflow_toolset_to_langchain_tools(
     works regardless of how the agent handles tool errors. Raising instead 
would
     abort the run under ``create_agent``'s default tool-error handling.
 
+    Argument validation failures are handled the same way: each tool validates
+    its arguments with the toolset's ``args_validator`` before dispatch, and a
+    :exc:`pydantic.ValidationError` from that step is fed back to the model as
+    the tool output so it can correct the call. This mirrors pydantic-ai's
+    native two-stage behaviour: only arg-validation ``ValidationError`` is
+    retried; a ``ValidationError`` raised inside ``call_tool`` (for example 
from
+    a Hook method or MCP client) propagates rather than being fed back, so a
+    non-idempotent tool that already ran a side effect is not re-invoked.
+
     The retry message is bounded by the tool's ``max_retries``: a tool that 
keeps
-    raising ``ModelRetry`` (for example an unrecoverable connection error) 
stops
-    being fed back and propagates once the budget is exhausted, so the run 
fails
-    instead of looping forever. The count resets after a successful call.
+    raising ``ModelRetry`` (or keeps failing arg validation) stops being fed 
back
+    and propagates once the budget is exhausted, so the run fails instead of
+    looping forever. The count resets after a successful call.
 
     The toolset's ``get_tools`` is invoked eagerly here to enumerate the tools.
 
@@ -153,16 +163,16 @@ def _build_structured_tool(
         # the args unchanged; a typed one coerces them (e.g. "5" -> 5).
         return toolset_tool.args_validator.validate_python(kwargs)
 
-    # ModelRetry is a "feed this back to the model and retry" signal, so the 
bridge
-    # returns its message as the tool output instead of raising (see 
docstring).
-    # Bound it the way native pydantic-ai does, via the tool's max_retries: a 
tool
-    # that keeps raising ModelRetry (e.g. an unrecoverable connection error) 
must
-    # eventually propagate so the run fails rather than looping forever. The 
count
-    # resets on the first successful call.
+    # Mirror pydantic-ai's ToolManager two-stage handling: ValidationError from
+    # arg validation is a "feed this back and retry" signal; ModelRetry from
+    # call_tool is too; a ValidationError raised inside call_tool is not (it
+    # would otherwise re-invoke a non-idempotent tool that already ran). Bound
+    # retries via the tool's max_retries so a tool that keeps failing 
eventually
+    # propagates. The count resets on the first successful call.
     max_retries = toolset_tool.max_retries if toolset_tool.max_retries is not 
None else 1
     retries = {"count": 0}
 
-    def _handle_retry(error: ModelRetry) -> str:
+    def _handle_retry(error: ModelRetry | ValidationError) -> str:
         retries["count"] += 1
         if retries["count"] > max_retries:
             # Reset before propagating so a reused tool starts the next run 
with a
@@ -173,7 +183,11 @@ def _build_structured_tool(
 
     def _sync_call(**kwargs: Any) -> Any:
         try:
-            result = _run_coro_sync(toolset.call_tool(name, _validate(kwargs), 
ctx, toolset_tool))
+            validated = _validate(kwargs)
+        except ValidationError as e:
+            return _handle_retry(e)
+        try:
+            result = _run_coro_sync(toolset.call_tool(name, validated, ctx, 
toolset_tool))
         except ModelRetry as e:
             return _handle_retry(e)
         retries["count"] = 0
@@ -181,7 +195,11 @@ def _build_structured_tool(
 
     async def _async_call(**kwargs: Any) -> Any:
         try:
-            result = await toolset.call_tool(name, _validate(kwargs), ctx, 
toolset_tool)
+            validated = _validate(kwargs)
+        except ValidationError as e:
+            return _handle_retry(e)
+        try:
+            result = await toolset.call_tool(name, validated, ctx, 
toolset_tool)
         except ModelRetry as e:
             return _handle_retry(e)
         retries["count"] = 0
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 befc56d1c24..f7b0f32001f 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
@@ -39,16 +39,13 @@ 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_core import SchemaValidator, core_schema
 
-from airflow.providers.common.ai.utils.tool_definition import 
return_schema_kwargs
+from airflow.providers.common.ai.utils.tool_definition import 
build_args_validator, return_schema_kwargs
 from airflow.providers.common.compat.sdk import BaseHook
 
 if TYPE_CHECKING:
     from pydantic_ai._run_context import RunContext
 
-_PASSTHROUGH_VALIDATOR = SchemaValidator(core_schema.any_schema())
-
 # JSON Schemas for the four SQL tools.
 _LIST_TABLES_SCHEMA: dict[str, Any] = {
     "type": "object",
@@ -272,7 +269,7 @@ class SQLToolset(AbstractToolset[Any]):
                 toolset=self,
                 tool_def=tool_def,
                 max_retries=1,
-                args_validator=_PASSTHROUGH_VALIDATOR,
+                args_validator=build_args_validator(schema),
             )
         return tools
 
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 8cf984f4396..9f4fb04b87b 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,9 +19,10 @@
 from __future__ import annotations
 
 import dataclasses
-from typing import Any
+from typing import Any, Literal
 
 from pydantic_ai.tools import ToolDefinition
+from pydantic_core import SchemaValidator, core_schema
 
 # ``ToolDefinition.return_schema`` is newer than the provider's pydantic-ai
 # floor. Detect it once so callers can include the kwarg only when supported,
@@ -42,3 +43,74 @@ def return_schema_kwargs(schema: dict[str, Any]) -> 
dict[str, Any]:
     if _SUPPORTS_RETURN_SCHEMA:
         return {"return_schema": schema}
     return {}
+
+
+def _fragment_to_core_schema(fragment: dict[str, Any]) -> 
core_schema.CoreSchema:
+    any_of = fragment.get("anyOf")
+    if isinstance(any_of, list):
+        choices: list[core_schema.CoreSchema | tuple[core_schema.CoreSchema, 
str]] = [
+            _fragment_to_core_schema(choice) for choice in any_of if 
isinstance(choice, dict)
+        ]
+        return core_schema.union_schema(choices) if choices else 
core_schema.any_schema()
+
+    schema_type = fragment.get("type")
+    if isinstance(schema_type, list):
+        choices = [
+            _fragment_to_core_schema({**fragment, "type": item})
+            for item in schema_type
+            if isinstance(item, str)
+        ]
+        return core_schema.union_schema(choices) if choices else 
core_schema.any_schema()
+
+    match schema_type:
+        case "string":
+            return core_schema.str_schema()
+        case "integer":
+            return core_schema.int_schema()
+        case "number":
+            return core_schema.float_schema()
+        case "boolean":
+            return core_schema.bool_schema()
+        case "null":
+            return core_schema.none_schema()
+        case "array":
+            items = fragment.get("items")
+            return core_schema.list_schema(
+                _fragment_to_core_schema(items) if isinstance(items, dict) 
else None
+            )
+        case "object":
+            return _object_fragment_to_core_schema(fragment)
+        case _:
+            return core_schema.any_schema()
+
+
+def _object_fragment_to_core_schema(fragment: dict[str, Any]) -> 
core_schema.CoreSchema:
+    """
+    Convert a JSON Schema ``object`` fragment to a core schema.
+
+    A fragment with no ``properties`` key is an untyped object (e.g. from a
+    ``dict[K, V]`` annotation): accept any dict rather than stripping its
+    contents. When ``properties`` is present, build a typed-dict that validates
+    each declared field recursively — nested objects are handled the same way
+    arrays already recurse into ``items``.
+
+    Undeclared keys follow native pydantic-ai: ``forbid`` for fixed signatures
+    (so a mistyped field name becomes a bounded retry), ``allow`` only when the
+    schema sets ``additionalProperties: true`` (methods that accept 
``**kwargs``).
+    """
+    if "properties" not in fragment:
+        return core_schema.dict_schema()
+    required = set(fragment.get("required", []))
+    fields = {
+        name: core_schema.typed_dict_field(_fragment_to_core_schema(prop), 
required=name in required)
+        for name, prop in fragment["properties"].items()
+    }
+    extra_behavior: Literal["allow", "forbid"] = (
+        "allow" if fragment.get("additionalProperties") is True else "forbid"
+    )
+    return core_schema.typed_dict_schema(fields, extra_behavior=extra_behavior)
+
+
+def build_args_validator(parameters_json_schema: dict[str, Any]) -> 
SchemaValidator:
+    """Build an argument validator from the schema advertised to the model."""
+    return 
SchemaValidator(_object_fragment_to_core_schema(parameters_json_schema))
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 89959649cdd..5bcf240afaa 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
@@ -25,6 +25,7 @@ import pytest
 from pydantic_ai._run_context import RunContext
 from pydantic_ai.exceptions import ModelRetry
 from pydantic_ai.toolsets.abstract import ToolsetTool
+from pydantic_core import ValidationError
 
 from airflow.providers.common.ai.toolsets.datafusion import (
     _RETRYABLE_QUERY_ERROR_PATTERNS,
@@ -107,6 +108,24 @@ class TestDataFusionToolsetGetTools:
             assert tool.tool_def.description
 
 
+class TestDataFusionToolsetArgsValidation:
+    @pytest.mark.parametrize(
+        ("tool_name", "valid_args"),
+        [
+            ("get_schema", {"table_name": "sales_data"}),
+            ("query", {"sql": "SELECT 1"}),
+        ],
+    )
+    def test_enforces_required_args(self, tool_name, valid_args):
+        cfg = _make_mock_datasource_config()
+        ts = DataFusionToolset([cfg])
+        tools = asyncio.run(ts.get_tools(ctx=MagicMock(spec=RunContext)))
+        validator = tools[tool_name].args_validator
+        assert validator.validate_python(valid_args) == valid_args
+        with pytest.raises(ValidationError):
+            validator.validate_python({})
+
+
 class TestDataFusionToolsetListTables:
     def test_returns_registered_tables(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 e1145e13455..ae2d4c6f867 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
@@ -20,6 +20,7 @@ import asyncio
 from unittest.mock import MagicMock
 
 import pytest
+from pydantic_core import ValidationError
 
 from airflow.providers.common.ai.toolsets.hook import (
     HookToolset,
@@ -34,13 +35,13 @@ from airflow.providers.common.ai.utils.tool_definition 
import _SUPPORTS_RETURN_S
 class _FakeHook:
     """Fake hook for testing HookToolset introspection."""
 
-    def list_keys(self, bucket: str, prefix: str = "") -> list[str]:
+    def list_keys(self, bucket: str, prefix: str | None = None) -> list[str]:
         """List object keys in a bucket.
 
         :param bucket: Name of the S3 bucket.
         :param prefix: Key prefix to filter by.
         """
-        return [f"{prefix}file1.txt", f"{prefix}file2.txt"]
+        return [f"{prefix or ''}file1.txt", f"{prefix or ''}file2.txt"]
 
     def read_file(self, key: str) -> str:
         """Read a file from storage."""
@@ -49,6 +50,11 @@ class _FakeHook:
     def no_docstring(self, x: int) -> int:
         return x * 2
 
+    def request(
+        self, endpoint: str | None = None, data: dict[str, object] | str | 
None = None, **kwargs: object
+    ) -> dict[str, object]:
+        return {"endpoint": endpoint, "data": data, **kwargs}
+
 
 class TestHookToolsetInit:
     def test_requires_non_empty_allowed_methods(self):
@@ -143,6 +149,32 @@ class TestHookToolsetGetTools:
         assert "S3 bucket" in props["bucket"]["description"]
 
 
+class TestHookToolsetArgsValidator:
+    @pytest.fixture
+    def list_keys_tool(self):
+        ts = HookToolset(_FakeHook(), allowed_methods=["list_keys"])
+        return asyncio.run(ts.get_tools(ctx=MagicMock()))["list_keys"]
+
+    def test_enforces_method_signature(self, list_keys_tool):
+        with pytest.raises(ValidationError, match="bucket"):
+            list_keys_tool.args_validator.validate_python({"prefix": "data/"})
+
+        assert list_keys_tool.args_validator.validate_python({"bucket": 
"my-bucket", "prefix": None}) == {
+            "bucket": "my-bucket",
+            "prefix": None,
+        }
+
+    def test_rejects_undeclared_args(self, list_keys_tool):
+        with pytest.raises(ValidationError, match="bogus"):
+            list_keys_tool.args_validator.validate_python({"bucket": 
"my-bucket", "bogus": 1})
+
+    def test_preserves_kwargs_accepted_by_method(self):
+        ts = HookToolset(_FakeHook(), allowed_methods=["request"])
+        tool = asyncio.run(ts.get_tools(ctx=MagicMock()))["request"]
+        args = {"endpoint": None, "data": {"key": "value"}, "timeout": 10}
+        assert tool.args_validator.validate_python(args) == args
+
+
 class TestHookToolsetCallTool:
     def test_dispatches_to_hook_method(self):
         hook = _FakeHook()
@@ -184,12 +216,20 @@ class TestBuildJsonSchemaFromSignature:
         assert schema["properties"]["active"] == {"type": "boolean"}
         assert set(schema["required"]) == {"name", "count", "rate", "active"}
 
-    def test_optional_params_not_required(self):
-        def fn(name: str, prefix: str = ""):
+    def test_optional_params_accept_null(self):
+        def fn(name: str, prefix: str | None = None):
             pass
 
         schema = _build_json_schema_from_signature(fn)
         assert schema["required"] == ["name"]
+        assert schema["properties"]["prefix"] == {"anyOf": [{"type": 
"string"}, {"type": "null"}]}
+
+    def test_union_types(self):
+        def fn(data: dict[str, object] | str):
+            pass
+
+        schema = _build_json_schema_from_signature(fn)
+        assert schema["properties"]["data"] == {"anyOf": [{"type": "object"}, 
{"type": "string"}]}
 
     def test_list_type(self):
         def fn(items: list[str]):
@@ -198,12 +238,19 @@ class TestBuildJsonSchemaFromSignature:
         schema = _build_json_schema_from_signature(fn)
         assert schema["properties"]["items"] == {"type": "array", "items": 
{"type": "string"}}
 
-    def test_no_annotation_defaults_to_string(self):
+    def test_no_annotation_is_untyped(self):
         def fn(x):
             pass
 
         schema = _build_json_schema_from_signature(fn)
-        assert schema["properties"]["x"] == {"type": "string"}
+        assert schema["properties"]["x"] == {}
+
+    def test_kwargs_allow_additional_properties(self):
+        def fn(x: int, **kwargs):
+            pass
+
+        schema = _build_json_schema_from_signature(fn)
+        assert schema["additionalProperties"] is True
 
     def test_skips_self_and_cls(self):
         class Foo:
diff --git 
a/providers/common/ai/tests/unit/common/ai/toolsets/test_langchain_bridge.py 
b/providers/common/ai/tests/unit/common/ai/toolsets/test_langchain_bridge.py
index 89bd866772e..0955694c67b 100644
--- a/providers/common/ai/tests/unit/common/ai/toolsets/test_langchain_bridge.py
+++ b/providers/common/ai/tests/unit/common/ai/toolsets/test_langchain_bridge.py
@@ -24,6 +24,7 @@ import pytest
 
 pytest.importorskip("langchain_core")
 
+from pydantic import ValidationError
 from pydantic_ai.exceptions import ModelRetry
 from pydantic_ai.tools import ToolDefinition
 from pydantic_ai.toolsets.abstract import AbstractToolset, ToolsetTool
@@ -192,6 +193,51 @@ class TestAirflowToolsetToLangChainTools:
 
         assert asyncio.run(boom.ainvoke({})) == "fix your input and try again"
 
+    def test_validation_error_returned_as_tool_output_sync(self):
+        toolset = FakeToolset()
+        add_one = {t.name: t for t in 
airflow_toolset_to_langchain_tools(toolset)}["add_one"]
+
+        result = add_one.invoke({"n": "not a number"})
+
+        assert "validation error" in result
+        assert toolset.calls == []
+
+    def test_validation_error_returned_as_tool_output_async(self):
+        toolset = FakeToolset()
+        add_one = {t.name: t for t in 
airflow_toolset_to_langchain_tools(toolset)}["add_one"]
+
+        result = asyncio.run(add_one.ainvoke({"n": "not a number"}))
+
+        assert "validation error" in result
+        assert toolset.calls == []
+
+    def test_repeated_validation_error_propagates_when_budget_exhausted(self):
+        add_one = {t.name: t for t in 
airflow_toolset_to_langchain_tools(FakeToolset())}["add_one"]
+
+        assert "validation error" in add_one.invoke({"n": "bad"})
+        with pytest.raises(ValidationError):
+            add_one.invoke({"n": "bad"})
+
+    def test_tool_body_validation_error_propagates(self):
+        # A ValidationError raised inside call_tool must not be fed back as a
+        # retry — that would re-invoke a tool that may already have run a side
+        # effect. Only arg-validation ValidationError is retried.
+        class BodyValidationToolset(FakeToolset):
+            async def call_tool(self, name, tool_args, ctx, tool) -> Any:
+                if name == "boom":
+                    raise ValidationError.from_exception_data(
+                        "BodyValidation",
+                        [{"type": "missing", "loc": ("response",), "input": 
{}}],
+                    )
+                return await super().call_tool(name, tool_args, ctx, tool)
+
+        boom = {t.name: t for t in 
airflow_toolset_to_langchain_tools(BodyValidationToolset())}["boom"]
+
+        with pytest.raises(ValidationError, match="response"):
+            boom.invoke({})
+        with pytest.raises(ValidationError, match="response"):
+            asyncio.run(boom.ainvoke({}))
+
     def test_deps_are_exposed_on_the_run_context(self):
         sentinel = object()
         captured: dict[str, Any] = {}
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 f39d97b95b8..9ff50ce142c 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
@@ -22,6 +22,7 @@ from unittest.mock import MagicMock, PropertyMock, patch
 
 import pytest
 from pydantic_ai.exceptions import ModelRetry
+from pydantic_core import ValidationError
 
 from airflow.providers.common.ai.toolsets.sql import SQLToolset
 from airflow.providers.common.ai.utils.tool_definition import 
_SUPPORTS_RETURN_SCHEMA
@@ -65,6 +66,23 @@ class TestSQLToolsetGetTools:
         for tool in tools.values():
             assert tool.tool_def.description
 
+    @pytest.mark.parametrize(
+        ("name", "valid_args"),
+        [
+            ("get_schema", {"table_name": "users"}),
+            ("query", {"sql": "SELECT 1"}),
+            ("check_query", {"sql": "SELECT 1"}),
+        ],
+    )
+    def test_args_validator_enforces_required_keys(self, name, valid_args):
+        ts = SQLToolset("pg_default")
+        tools = asyncio.run(ts.get_tools(ctx=MagicMock()))
+        validator = tools[name].args_validator
+
+        assert validator.validate_python(valid_args) == valid_args
+        with pytest.raises(ValidationError):
+            validator.validate_python({})
+
     @pytest.mark.skipif(
         not _SUPPORTS_RETURN_SCHEMA, reason="pydantic-ai too old for 
ToolDefinition.return_schema"
     )
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 5059ffeab0f..1664c453630 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
@@ -16,10 +16,14 @@
 # under the License.
 from __future__ import annotations
 
+import json
 from unittest.mock import patch
 
+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 
return_schema_kwargs
+from airflow.providers.common.ai.utils.tool_definition import 
build_args_validator, return_schema_kwargs
 
 
 def test_returns_kwarg_when_supported():
@@ -30,3 +34,131 @@ def test_returns_kwarg_when_supported():
 def test_returns_empty_when_unsupported():
     with patch.object(tool_definition, "_SUPPORTS_RETURN_SCHEMA", False):
         assert return_schema_kwargs({"type": "string"}) == {}
+
+
+TOOL_SCHEMA = {
+    "type": "object",
+    "properties": {
+        "name": {"type": "string"},
+        "count": {"type": "integer"},
+        "ratio": {"type": "number"},
+        "enabled": {"type": "boolean"},
+        "tags": {"type": "array", "items": {"type": "string"}},
+        "options": {"type": "object"},
+        "nullable_name": {"type": ["string", "null"]},
+        "payload": {"anyOf": [{"type": "string"}, {"type": "object"}, {"type": 
"null"}]},
+        "anything": {},
+    },
+    "required": ["name"],
+}
+
+
+NESTED_SCHEMA = {
+    "type": "object",
+    "properties": {
+        "config": {
+            "type": "object",
+            "properties": {"port": {"type": "integer"}},
+            "required": ["port"],
+        },
+    },
+    "required": ["config"],
+}
+
+
+def _validate(validator, args: dict, use_json: bool):
+    if use_json:
+        return validator.validate_json(json.dumps(args))
+    return validator.validate_python(args)
+
+
[email protected]("use_json", [False, True], ids=["python", "json"])
+class TestBuildArgsValidator:
+    @pytest.mark.parametrize(
+        ("args", "expected"),
+        [
+            ({"name": "a"}, {"name": "a"}),
+            (
+                {
+                    "name": "a",
+                    "count": 2,
+                    "ratio": 0.5,
+                    "enabled": True,
+                    "tags": ["x"],
+                    "options": {"k": "v"},
+                    "nullable_name": None,
+                    "payload": {"key": "value"},
+                    "anything": [1, "b"],
+                },
+                {
+                    "name": "a",
+                    "count": 2,
+                    "ratio": 0.5,
+                    "enabled": True,
+                    "tags": ["x"],
+                    "options": {"k": "v"},
+                    "nullable_name": None,
+                    "payload": {"key": "value"},
+                    "anything": [1, "b"],
+                },
+            ),
+            ({"name": "a", "count": "5"}, {"name": "a", "count": 5}),
+        ],
+    )
+    def test_valid_args_accepted(self, args, expected, use_json):
+        validator = build_args_validator(TOOL_SCHEMA)
+        assert _validate(validator, args, use_json) == expected
+
+    @pytest.mark.parametrize(
+        "args",
+        [
+            {"count": 1},
+            {"name": "a", "count": "not-an-int"},
+            {"name": "a", "tags": "not-a-list"},
+            {"name": "a", "nullable_name": 1},
+            {"name": "a", "payload": []},
+        ],
+    )
+    def test_invalid_args_rejected(self, args, use_json):
+        validator = build_args_validator(TOOL_SCHEMA)
+        with pytest.raises(ValidationError):
+            _validate(validator, args, use_json)
+
+    def test_extra_keys_rejected(self, use_json):
+        validator = build_args_validator(TOOL_SCHEMA)
+        with pytest.raises(ValidationError):
+            _validate(validator, {"name": "a", "junk": 1}, use_json)
+
+    def test_additional_properties_preserved(self, use_json):
+        schema = {**TOOL_SCHEMA, "additionalProperties": True}
+        validator = build_args_validator(schema)
+        assert _validate(validator, {"name": "a", "extra": 1}, use_json) == {
+            "name": "a",
+            "extra": 1,
+        }
+
+    def test_empty_properties_accepts_empty_args(self, use_json):
+        validator = build_args_validator({"type": "object", "properties": {}, 
"required": []})
+        assert _validate(validator, {}, use_json) == {}
+
+    def test_nested_object_validated_recursively(self, use_json):
+        validator = build_args_validator(NESTED_SCHEMA)
+        assert _validate(validator, {"config": {"port": "5432"}}, use_json) == 
{"config": {"port": 5432}}
+
+    @pytest.mark.parametrize(
+        "args",
+        [
+            {"config": {"port": "not-an-int"}},
+            {"config": {}},
+        ],
+    )
+    def test_nested_object_invalid_rejected(self, args, use_json):
+        validator = build_args_validator(NESTED_SCHEMA)
+        with pytest.raises(ValidationError):
+            _validate(validator, args, use_json)
+
+    def test_untyped_nested_object_passthrough(self, use_json):
+        schema = {"type": "object", "properties": {"payload": {"type": 
"object"}}}
+        validator = build_args_validator(schema)
+        args = {"payload": {"any": 1, "deep": {"k": "v"}}}
+        assert _validate(validator, args, use_json) == args

Reply via email to