Source code for airflow.providers.common.ai.toolsets.datafusion
# 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.
"""Curated SQL toolset wrapping DataFusionEngine for agentic object-store workflows."""
from __future__ import annotations
import logging
import re
from typing import TYPE_CHECKING, Any
try:
from airflow.providers.common.ai.utils.sql_validation import SQLSafetyError, validate_sql as _validate_sql
from airflow.providers.common.sql.datafusion.engine import DataFusionEngine
from airflow.providers.common.sql.datafusion.exceptions import QueryExecutionException
except ImportError as e:
from airflow.providers.common.compat.sdk import AirflowOptionalProviderFeatureException
raise AirflowOptionalProviderFeatureException(e)
from pydantic_ai.exceptions import ModelRetry
from pydantic_ai.tools import ToolDefinition
from pydantic_ai.toolsets.abstract import ToolsetTool
from airflow.providers.common.ai.utils.masking import dumps_masked
from airflow.providers.common.ai.utils.query_results import (
DEFAULT_MAX_COLUMNS,
DEFAULT_MAX_RESULT_BYTES,
GET_SCHEMA_TOOL_DESCRIPTION as _GET_SCHEMA_DESCRIPTION,
QUERY_TOOL_DESCRIPTION as _QUERY_DESCRIPTION,
build_query_result,
build_schema_result,
)
from airflow.providers.common.ai.utils.tool_definition import build_args_validator
from airflow.providers.common.ai.utils.toolset_base import AirflowToolset, validate_max_retries
if TYPE_CHECKING:
from pydantic_ai._run_context import RunContext
from airflow.providers.common.sql.config import DataSourceConfig
# JSON Schemas for the three DataFusion tools.
_LIST_TABLES_SCHEMA: dict[str, Any] = {
"type": "object",
"properties": {},
}
_GET_SCHEMA_SCHEMA: dict[str, Any] = {
"type": "object",
"properties": {
"table_name": {"type": "string", "description": "Name of the table to inspect."},
"name_contains": {
"type": "string",
"description": (
"Return only columns whose name contains this substring (case-insensitive). "
"Use it to find the relevant columns on a very wide table."
),
},
},
"required": ["table_name"],
}
_QUERY_SCHEMA: dict[str, Any] = {
"type": "object",
"properties": {
"sql": {"type": "string", "description": "SQL query to execute."},
},
"required": ["sql"],
}
# DataFusion python bindings don't expose any native exception types, it uses rust exceptions.
# So we have to rely on error message parsing with regex.
_RETRYABLE_IDENTIFIER = r"""(?:['"][^'"]+['"]|\w+)"""
_RETRYABLE_QUERY_ERROR_PATTERNS = (
re.compile(rf"""column\s+{_RETRYABLE_IDENTIFIER}\s+not\s+found""", re.IGNORECASE),
re.compile(rf"""table\s+{_RETRYABLE_IDENTIFIER}\s+not\s+found""", re.IGNORECASE),
)
[docs]
class DataFusionToolset(AirflowToolset):
"""
Curated toolset that gives an LLM agent SQL access to object-storage data via Apache DataFusion.
.. note::
Experimental: this can change or be removed in a minor release of this provider.
See :ref:`howto/stability`.
Provides three tools — ``list_tables``, ``get_schema``, and ``query`` —
backed by
:class:`~airflow.providers.common.sql.datafusion.engine.DataFusionEngine`.
Each :class:`~airflow.providers.common.sql.config.DataSourceConfig` entry
registers a table backed by Parquet, CSV, Avro, or Iceberg data on S3 or
local storage. Multiple configs can be registered so that SQL queries can
join across tables.
Requires the ``datafusion`` extra of ``apache-airflow-providers-common-sql``.
:param datasource_configs: One or more DataFusion data-source configurations.
:param allow_writes: Allow data-modifying SQL (CREATE TABLE, CREATE VIEW,
INSERT INTO, etc.). Default ``False`` — only SELECT-family statements
are permitted. ``EXPLAIN`` reaches the engine only with ``allow_writes=True``,
and fails there: the ``max_rows`` limit wraps the plan, and DataFusion requires
``EXPLAIN`` to be the root of the plan. The agent gets an error result, not the
plan.
:param max_rows: Maximum number of rows returned from the ``query`` tool.
Default ``50``. The query is limited to ``max_rows + 1`` rows, so a large
result is never fully materialized; the extra row only signals truncation.
:param max_result_bytes: Budget for the serialized ``query`` result, in bytes, and the
byte backstop that also triggers the ``get_schema`` summary (see ``max_columns``).
Default 64 KiB. ``max_rows`` bounds rows, which says nothing about size: one
row of a 3000-column table is larger than a thousand rows of a narrow one, and
a tool result stays in the model's message history for the rest of the run, so
its cost is re-paid on every subsequent request. Rows are returned as a
contiguous prefix, stopping at the first that does not fit the remaining budget
rather than skipping it and packing later ones, so one wide row early in the
result ends it. The result reports which limit it hit so the agent can narrow
its projection rather than page through the table.
:param max_columns: Maximum number of columns ``get_schema`` returns in full. Default
``100``. Above it -- or when the serialized columns exceed ``max_result_bytes`` --
the full list is replaced by a bounded summary (column count, a type histogram, a
sample of columns) that points the agent at the ``name_contains`` filter, so a
several-thousand-column table cannot exhaust the context before a query is written.
:param max_retries: How many times the model may correct a failed call to one of these
tools before the run fails. ``None`` (the default) uses the agent's tool retry
budget, its ``retries``, as pydantic-ai's own toolsets do.
"""
def __init__(
self,
datasource_configs: list[DataSourceConfig],
*,
allow_writes: bool = False,
max_rows: int = 50,
max_result_bytes: int = DEFAULT_MAX_RESULT_BYTES,
max_columns: int = DEFAULT_MAX_COLUMNS,
max_retries: int | None = None,
) -> None:
self._max_retries = validate_max_retries(max_retries)
if not datasource_configs:
raise ValueError("datasource_configs must contain at least one DataSourceConfig")
self._datasource_configs = datasource_configs
self._allow_writes = allow_writes
self._max_rows = max_rows
self._max_result_bytes = max_result_bytes
self._max_columns = max_columns
self._engine: DataFusionEngine | None = None
@property
[docs]
def id(self) -> str:
suffix = "_".join(config.table_name.replace("-", "_") for config in self._datasource_configs)
return f"sql_datafusion_{suffix}"
def _get_engine(self) -> DataFusionEngine:
"""Lazily create and configure a DataFusionEngine from *datasource_configs*."""
if self._engine is None:
engine = DataFusionEngine()
for config in self._datasource_configs:
engine.register_datasource(config)
self._engine = engine
return self._engine
[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 name, description, schema in (
("list_tables", "List available table names.", _LIST_TABLES_SCHEMA),
("get_schema", _GET_SCHEMA_DESCRIPTION, _GET_SCHEMA_SCHEMA),
("query", _QUERY_DESCRIPTION, _QUERY_SCHEMA),
):
tool_def = ToolDefinition(
name=name,
description=description,
parameters_json_schema=schema,
sequential=True,
)
tools[name] = ToolsetTool(
toolset=self,
tool_def=tool_def,
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],
) -> Any:
if name == "list_tables":
return await self.run_blocking(self._list_tables)
if name == "get_schema":
return await self.run_blocking(
self._get_schema, tool_args["table_name"], tool_args.get("name_contains")
)
if name == "query":
return await self.run_blocking(self._query, tool_args["sql"])
raise ValueError(f"Unknown tool: {name!r}")
def _list_tables(self) -> str:
try:
engine = self._get_engine()
tables: list[str] = list(engine.session_context.catalog().schema().table_names())
return dumps_masked(tables)
except Exception as ex:
log.warning("list_tables failed: %s", ex)
return dumps_masked({"error": str(ex)})
def _get_schema(self, table_name: str, name_contains: str | None = None) -> str:
engine = self._get_engine()
# session_context lookup is required here instead of engine.registered_tables,
# because registered_tables only tracks tables registered via datasource config.
# When allow_writes is enabled, the agent may create temporary in-memory tables
# that would not be captured there.
if not engine.session_context.table_exist(table_name):
return dumps_masked({"error": f"Table {table_name!r} is not available"})
# Intentionally using session_context instead of engine.get_schema() —
# the latter returns a pre-formatted string intended for other operators,
# not a JSON-compatible format.
# TODO: refactor engine.get_schema() to return JSON and update this accordingly
table = engine.session_context.table(table_name)
columns = [{"name": f.name, "type": str(f.type)} for f in table.schema()]
return build_schema_result(
columns,
max_columns=self._max_columns,
max_result_bytes=self._max_result_bytes,
name_contains=name_contains,
)
def _query(self, sql: str) -> str:
try:
if not self._allow_writes:
_validate_sql(sql)
engine = self._get_engine()
try:
pydict = engine.session_context.sql(sql).limit(self._max_rows + 1).to_pydict()
except Exception as e:
raise QueryExecutionException(f"Error while executing query: {e}") from e
col_names = list(pydict.keys())
num_rows = len(next(iter(pydict.values()), []))
rows = [[pydict[col][i] for col in col_names] for i in range(min(num_rows, self._max_rows))]
return build_query_result(
col_names,
rows,
max_rows=self._max_rows,
max_result_bytes=self._max_result_bytes,
more_rows_available=num_rows > self._max_rows,
)
except SQLSafetyError as ex:
log.warning("query failed SQL safety validation: %s", ex)
raise ModelRetry(
f"error: {ex!s}. Only read-only SELECT-family queries are allowed unless "
"allow_writes is enabled; check the SQL syntax and statement type, then try again."
) from ex
except QueryExecutionException as ex:
if self._is_retryable_query_error(ex):
raise ModelRetry(
f"error: {ex!s}, Use get_schema and list_tables tools for more details."
) from ex
return dumps_masked({"error": str(ex), "query": sql})
@staticmethod
def _is_retryable_query_error(error: QueryExecutionException) -> bool:
message = str(error)
return any(pattern.search(message) for pattern in _RETRYABLE_QUERY_ERROR_PATTERNS)