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