Source code for airflow.providers.common.ai.toolsets.object_storage
# 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 lzma
import os
import zlib
from datetime import UTC, datetime
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, validate_max_retries
from airflow.providers.common.compat.sdk import AirflowOptionalProviderFeatureException, ObjectStoragePath
if TYPE_CHECKING:
from collections.abc import Sequence
from pydantic_ai._run_context import RunContext
_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 row count, schema and first "
"rows; offset and limit do not apply to it."
),
}
_MAX_OUTPUT_LINES = 2000
_SAMPLE_ROWS = 20
_COLUMNAR_FORMATS: tuple[Literal["parquet", "avro"], ...] = ("parquet", "avro")
_MEDIA_FORMATS = frozenset({"jpeg", "jpg", "pdf", "png"})
[docs]
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 row count, 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.
:param max_retries: How many times the model may correct a call with invalid arguments
before the run fails. A failed read goes back to the model without using it.
``None`` (the default) uses the agent's tool retry budget, its ``retries``, as
pydantic-ai's own toolsets do.
"""
# Rendered, on a copy, by AgentOperator. Deliberately not ``template_fields``, which
# Airflow's templater would render in place wherever the toolset is nested.
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 = "",
max_retries: int | None = None,
) -> None:
self._max_retries = validate_max_retries(max_retries)
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
[docs]
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
[docs]
async def get_tools(self, ctx: RunContext[Any]) -> dict[str, ToolsetTool[Any]]:
max_retries = self._get_tool_max_retries(ctx)
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=max_retries,
args_validator=build_args_validator(schema),
)
return tools
[docs]
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,
zlib.error,
lzma.LZMAError,
AirflowOptionalProviderFeatureException,
) as e:
# Storage the connection may not read, a corrupt or mislabelled file, a codec this
# Python build lacks: final for this path, so the model is told, not the task failed.
raise ToolFailed(f"{relative or '/'!r} cannot be read: {type(e).__name__}: {e}") from None
raise ValueError(f"Unknown tool: {name!r}")
# ------------------------------------------------------------------
# Tool implementations. Each runs in a worker thread.
# ------------------------------------------------------------------
def _root(self) -> ObjectStoragePath:
# Built on use, not in __init__, so AgentOperator's rendering of _path and _conn_id applies.
return ObjectStoragePath(self._path, conn_id=self._conn_id)
def _resolve(self, relative: str) -> ObjectStoragePath:
"""Turn the model's relative path into a path under the root, or refuse it."""
if relative.startswith("/") or "://" in relative:
raise ToolFailed(f"{relative!r} is not a relative path. Name files relative to the storage root.")
parts = [part for part in PurePosixPath(relative).parts if part != "."]
if ".." in parts:
raise ToolFailed(f"{relative!r} leaves the storage root; '..' is not allowed.")
root = self._root()
target = root.joinpath(*parts) if parts else root
if not _inside_local_root(root, target.path):
raise ToolFailed(f"{relative!r} resolves outside the storage root.")
return target
def _list_files(self, relative: str, *, offset: int) -> str:
directory = self._resolve(relative)
if not directory.is_dir():
raise ToolFailed(f"{relative or '/'!r} is not a directory.")
root = self._root()
# One listing call with details, rather than a stat per entry, which on S3 or GCS is a
# request each. Some stores list a directory's own placeholder object; skip it.
own_path = directory.path.rstrip("/")
entries = sorted(
(
info
for info in directory.fs.ls(directory.path, detail=True)
if info["name"].rstrip("/") != own_path and _inside_local_root(root, info["name"])
),
key=_entry_name,
)
page = entries[offset : offset + self._max_files]
listed = [
{"name": f"{_entry_name(info)}/"}
if info["type"] == "directory"
else {"name": _entry_name(info), "size_bytes": info.get("size")}
for info in page
]
result: dict[str, Any] = {"path": relative or "/", "entries": listed}
if offset + len(page) < len(entries):
result["note"] = (
f"Showing entries {offset + 1} to {offset + len(page)} of {len(entries)}; list again "
f"with offset={offset + len(page)} for more."
)
return serialize_for_llm(result)
def _get_file_info(self, relative: str) -> str:
target = self._resolve(relative)
# One request: on S3 or GCS, exists(), is_dir() and stat() would each be a round trip.
try:
stat = target.stat()
except FileNotFoundError:
raise ToolFailed(f"{relative!r} does not exist.") from None
if stat.get("type") == "directory":
return serialize_for_llm({"path": relative, "type": "directory"})
info: dict[str, Any] = {"path": relative, "type": "file", "size_bytes": stat.st_size}
if modified := _as_iso(stat.st_mtime):
info["modified"] = modified
return serialize_for_llm(info)
def _read_file(self, relative: str, *, offset: int | None, limit: int | None) -> str:
target = self._resolve(relative)
if not target.is_file():
raise ToolFailed(f"{relative!r} is not a file.")
try:
file_format, compression = detect_file_format(target)
except LLMFileAnalysisError:
# An extension file analysis does not know, such as .sql or .yaml.gz: read it as text.
file_format, compression = "txt", detect_compression(target)
if file_format in _MEDIA_FORMATS:
raise ToolFailed(f"{relative!r} is a {file_format} file, which this tool cannot read as text.")
try:
for columnar in _COLUMNAR_FORMATS:
if file_format == columnar:
sample = sample_columnar_file(
target, file_format=columnar, sample_rows=_SAMPLE_ROWS, max_bytes=self._max_read_bytes
)
# First, so that cutting a long sample never drops it.
shown = min(_SAMPLE_ROWS, sample.total_rows)
header = (
f"Rows: {sample.total_rows}. The schema and the first {shown} rows follow; "
f"offset and limit do not apply to {columnar.capitalize()} files.\n"
)
return _cut(header + sample.text, self._max_output_bytes)
data = read_bytes(target, compression=compression, max_bytes=self._max_read_bytes)
except LLMFileAnalysisLimitExceededError:
raise ToolFailed(
f"{relative!r} is larger than the {format_size(self._max_read_bytes)} this tool reads."
) from None
if b"\x00" in data:
raise ToolFailed(f"{relative!r} is a binary file, which this tool cannot read as text.")
# Masked whole, before a window is cut: a secret spanning lines would otherwise leak a
# line at a time.
return render_file_window(
mask_secrets(data),
offset=offset,
limit=limit,
max_lines=_MAX_OUTPUT_LINES,
max_bytes=self._max_output_bytes,
long_line_hint="It cannot be read a line at a time.",
)
def _entry_name(info: dict[str, Any]) -> str:
return PurePosixPath(info["name"].rstrip("/")).name
def _inside_local_root(root: ObjectStoragePath, path: str) -> bool:
"""
Whether ``path`` stays inside ``root`` once symlinks are resolved, for a root on local disk.
Checked by filesystem, not scheme: a root written as a plain path has no scheme but is
still on the worker's disk, where a symlink under it could point anywhere. Object stores
have no symlinks, so any other root is not checked here.
"""
if not isinstance(root.fs, LocalFileSystem):
return True
real_root = os.path.realpath(root.path)
return os.path.commonpath([real_root, os.path.realpath(path)]) == real_root
def _as_iso(modified: Any) -> str | None:
"""Return a modification time as ISO 8601; stores report a timestamp, a datetime or nothing."""
if isinstance(modified, datetime):
return modified.isoformat()
if isinstance(modified, (int, float)) and modified > 0:
return datetime.fromtimestamp(modified, tz=UTC).isoformat()
return None
def _cut(text: str, max_bytes: int) -> str:
encoded = text.encode("utf-8")
if len(encoded) <= max_bytes:
return text
kept = encoded[:max_bytes].decode("utf-8", "ignore")
return f"{kept}\n[... cut at {format_size(max_bytes)}]"