vatsrahul1001 commented on code in PR #73899: URL: https://github.com/apache/airflow/pull/73899#discussion_r4136806492
########## providers/common/ai/tests/unit/common/ai/toolsets/test_object_storage.py: ########## @@ -0,0 +1,301 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +import asyncio +import gzip +import json +from typing import Any +from unittest.mock import patch + +import pytest +from pydantic_ai import RunContext +from pydantic_ai.exceptions import ToolFailed +from pydantic_ai.models.test import TestModel +from pydantic_ai.usage import RunUsage + +from airflow.providers.common.ai.toolsets.object_storage import ObjectStorageToolset + + [email protected] +def storage(tmp_path): Review Comment: Also noticed all the scoping tests use a local `file://` root — so the encoded `..` / backslash cases never get exercised on an actual object store, where only the string check applies. A `memory://` case would cover that. ########## providers/common/ai/src/airflow/providers/common/ai/toolsets/object_storage.py: ########## @@ -0,0 +1,367 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +"""Read-only toolset giving an agent the files under one object-storage path.""" + +from __future__ import annotations + +import os +from datetime import datetime, timezone +from pathlib import PurePosixPath +from typing import TYPE_CHECKING, Any, Literal + +from fsspec.implementations.local import LocalFileSystem +from pydantic_ai.exceptions import ToolFailed +from pydantic_ai.tools import ToolDefinition +from pydantic_ai.toolsets.abstract import ToolsetTool + +from airflow.providers.common.ai.exceptions import LLMFileAnalysisError, LLMFileAnalysisLimitExceededError +from airflow.providers.common.ai.sandbox.output import format_size, render_file_window +from airflow.providers.common.ai.utils.file_analysis import ( + detect_compression, + detect_file_format, + read_bytes, + sample_columnar_file, +) +from airflow.providers.common.ai.utils.masking import mask_secrets +from airflow.providers.common.ai.utils.tool_definition import ( + build_args_validator, + return_schema_kwargs, + serialize_for_llm, +) +from airflow.providers.common.ai.utils.toolset_base import AirflowToolset +from airflow.providers.common.compat.sdk import AirflowOptionalProviderFeatureException, ObjectStoragePath + +if TYPE_CHECKING: + from collections.abc import Sequence + + from pydantic_ai._run_context import RunContext + +LIST_FILES = "list_files" +GET_FILE_INFO = "get_file_info" +READ_FILE = "read_file" + +_PATH_DESCRIPTION = "Path relative to the storage root, using / between parts. Omit for the root itself." + +_SCHEMAS: dict[str, dict[str, Any]] = { + LIST_FILES: { + "type": "object", + "properties": { + "path": {"type": "string", "description": _PATH_DESCRIPTION}, + "offset": { + "type": ["integer", "null"], + "description": "Entry to start listing from (0-indexed), to page through a large directory.", + }, + }, + "required": [], + }, + GET_FILE_INFO: { + "type": "object", + "properties": {"path": {"type": "string", "description": _PATH_DESCRIPTION}}, + "required": ["path"], + }, + READ_FILE: { + "type": "object", + "properties": { + "path": {"type": "string", "description": _PATH_DESCRIPTION}, + "offset": { + "type": ["integer", "null"], + "description": "Line number to start reading from (1-indexed).", + }, + "limit": {"type": ["integer", "null"], "description": "Maximum number of lines to read."}, + }, + "required": ["path"], + }, +} + +_DESCRIPTIONS = { + LIST_FILES: ( + "List the files and directories in one directory of the storage, sorted by name. " + "Directories end with a slash; list one of them to go deeper. A large directory is listed " + "a page at a time, and the result tells you the offset to continue from." + ), + GET_FILE_INFO: "Get the size and last-modified time of one file or directory.", + READ_FILE: ( + "Read a text file. Long files are returned a window at a time and the result tells you the " + "offset to continue from. For a Parquet or Avro file, returns its schema and first rows." + ), +} + +_MAX_OUTPUT_LINES = 2000 +_SAMPLE_ROWS = 20 +_COLUMNAR_FORMATS: tuple[Literal["parquet", "avro"], ...] = ("parquet", "avro") +_MEDIA_FORMATS = frozenset({"jpeg", "jpg", "pdf", "png"}) + + +class ObjectStorageToolset(AirflowToolset): + """ + Give an agent read-only access to the files under one object-storage path. + + .. note:: + + Experimental: this can change or be removed in a minor release of this provider. + See :ref:`howto/stability`. + + Exposes three tools, ``list_files``, ``get_file_info`` and ``read_file``, rooted at + ``path``, which is any location Airflow's + :class:`~airflow.sdk.ObjectStoragePath` can open: ``s3://``, ``gs://``, ``abfs://``, + ``file://`` and the rest, with credentials from ``conn_id``. The model names files by + paths relative to that root. It cannot write, delete or move anything, and a path that + is absolute, carries a scheme, or climbs out of the root with ``..`` is refused. + + ``read_file`` returns a text file a window of lines at a time, like the sandbox's own + ``read_file``, and a Parquet or Avro file as its schema and first rows. Compressed text + (``.gz``, ``.bz2``, ``.xz``) is decompressed. Images, PDFs and other binary files are + refused, as is any file larger than ``max_read_bytes``. So is a file that cannot be read, + such as a corrupt one, or one the connection may not open: the model is told why, and the + run goes on. On a local root, a symlink that leads out of the root is refused too. + + :param path: Root the agent may read under. Templated when the toolset is passed to + ``AgentOperator`` / ``@task.agent``. + :param conn_id: Airflow connection for the storage, or ``None`` for the default + credentials of its protocol. Templated like ``path``. + :param max_files: Most entries one ``list_files`` result holds; the model pages through + a larger directory. The whole directory is still listed from storage on each call. + Default ``200``. + :param max_read_bytes: Largest file ``read_file`` will open, after decompression. + Default 10 MiB. + :param max_output_bytes: Most bytes one ``read_file`` result holds; the model reads on + from the offset it is given. Default 50 KiB. + :param tool_prefix: Prefix for the three tool names, e.g. ``"reports"`` gives + ``reports_read_file``. Set this when one agent has another toolset with the same tool + names, such as a second ``ObjectStorageToolset`` or a ``SandboxToolset``, whose + ``read_file`` would collide, since duplicate tool names are rejected. + """ + + # Rendered, on a copy, by AgentOperator. Deliberately not ``template_fields``, which + # Airflow's templater would render in place wherever the toolset is nested. + agent_template_fields: Sequence[str] = ("_path", "_conn_id") + + def __init__( + self, + path: str, + *, + conn_id: str | None = None, + max_files: int = 200, + max_read_bytes: int = 10 * 1024 * 1024, + max_output_bytes: int = 50 * 1024, + tool_prefix: str = "", + ) -> None: + for name, value in ( + ("max_files", max_files), + ("max_read_bytes", max_read_bytes), + ("max_output_bytes", max_output_bytes), + ): + if value < 1: + raise ValueError(f"{name} must be at least 1, got {value}.") + if tool_prefix and not tool_prefix.isidentifier(): + raise ValueError(f"tool_prefix must be a valid Python identifier, got {tool_prefix!r}.") + self._path = path + self._conn_id = conn_id + self._max_files = max_files + self._max_read_bytes = max_read_bytes + self._max_output_bytes = max_output_bytes + self._tool_prefix = tool_prefix + + @property + def id(self) -> str: + suffix = f"-{self._tool_prefix}" if self._tool_prefix else "" + return f"object-storage-{self._conn_id or 'default'}{suffix}" + + def _tool_name(self, base: str) -> str: + return f"{self._tool_prefix}_{base}" if self._tool_prefix else base + + async def get_tools(self, ctx: RunContext[Any]) -> dict[str, ToolsetTool[Any]]: + tools: dict[str, ToolsetTool[Any]] = {} + for base, schema in _SCHEMAS.items(): + name = self._tool_name(base) + tools[name] = ToolsetTool( + toolset=self, + tool_def=ToolDefinition( + name=name, + description=_DESCRIPTIONS[base], + parameters_json_schema=schema, + **return_schema_kwargs({"type": "string"}), + ), + max_retries=1, + args_validator=build_args_validator(schema), + ) + return tools + + async def _execute_tool( + self, + name: str, + tool_args: dict[str, Any], + ctx: RunContext[Any], + tool: ToolsetTool[Any], + ) -> str: + base = name.removeprefix(f"{self._tool_prefix}_") if self._tool_prefix else name + relative = tool_args.get("path") or "" + try: + if base == LIST_FILES: + return await self.run_blocking( + self._list_files, relative, offset=tool_args.get("offset") or 0 + ) + if base == GET_FILE_INFO: + return await self.run_blocking(self._get_file_info, relative) + if base == READ_FILE: + return await self.run_blocking( + self._read_file, relative, offset=tool_args.get("offset"), limit=tool_args.get("limit") + ) + except ToolFailed: + raise + except (OSError, EOFError, ValueError, AirflowOptionalProviderFeatureException) as e: Review Comment: I think corrupt parquet/avro slips through here. This catches `OSError`/`EOFError`/`ValueError`, but pyarrow throws `ArrowInvalid` (and fastavro its own errors), none of which are those — so a malformed columnar file fails the task instead of coming back as a refusal the model can read. Corrupt gzip is fine since gzip raises `OSError`, parquet just isn't covered. Can we widen the catch to include those? Worth a corrupt-parquet test too, only gzip is tested right now. -- This is an automated message from the Apache Git Service. To respond to the message, please log on to GitHub and use the URL above to go to the specific comment. To unsubscribe, e-mail: [email protected] For queries about this service, please contact Infrastructure at: [email protected]
