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 a83bd8bbc8c56..78013918f8701 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 @@ 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 @@ 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 @@ async def get_tools(self, ctx: RunContext[Any]) -> dict[str, ToolsetTool[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 132c5dc9db205..63037e1a4f8f9 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 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 @@ 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 @@ async def get_tools(self, ctx: RunContext[Any]) -> dict[str, ToolsetTool[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 @@ async def call_tool( 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/sql.py b/providers/common/ai/src/airflow/providers/common/ai/toolsets/sql.py index 667e5696fef96..a0d42bc3bc727 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 @@ 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", @@ -257,7 +254,7 @@ async def get_tools(self, ctx: RunContext[Any]) -> dict[str, ToolsetTool[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 8cf984f43969e..bdd684f5f40de 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,70 @@ 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``. + """ + 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", "ignore"] = ( + "allow" if fragment.get("additionalProperties") is True else "ignore" + ) + 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 89959649cdd90..5bcf240afaaa7 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 @@ 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 @@ def test_tool_definitions_have_descriptions(self): 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 e1145e13455b9..a1fe32731ad2b 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 @@ 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 @@ 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 @@ def read_file(self, key: str) -> str: 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,31 @@ def test_param_docs_enriched_in_schema(self): 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, + } + assert list_keys_tool.args_validator.validate_python({"bucket": "my-bucket", "bogus": 1}) == { + "bucket": "my-bucket" + } + + 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 +215,20 @@ def fn(name: str, count: int, rate: float, active: bool): 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 +237,19 @@ def fn(items: list[str]): 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_sql.py b/providers/common/ai/tests/unit/common/ai/toolsets/test_sql.py index e56b6bc461a09..589aafa6ab223 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 @@ 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 @@ def test_tool_definitions_have_descriptions(self): 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 5059ffeab0f55..c29bc1b25ad11 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,130 @@ 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) + + +@pytest.mark.parametrize("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_dropped(self, use_json): + validator = build_args_validator(TOOL_SCHEMA) + assert _validate(validator, {"name": "a", "junk": 1}, use_json) == {"name": "a"} + + 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