This is an automated email from the ASF dual-hosted git repository.
pankajkoti 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 e4d397f021e Bound get_schema results for very wide tables in the SQL
toolsets (#74017)
e4d397f021e is described below
commit e4d397f021e1ed571a0464f06afcab4ab1c9d122
Author: Pankaj Koti <[email protected]>
AuthorDate: Fri Oct 2 00:41:54 2026 +0530
Bound get_schema results for very wide tables in the SQL toolsets (#74017)
## What
`get_schema` in `SQLToolset` and `DataFusionToolset` returned every column
of a
table with no bound. On a table with a few thousand columns that is a
six-figure-byte tool result, and because an agent calls `get_schema` before
it
can write a query, the result enters message history early and its cost is
re-paid on every later model request. On a several-thousand-column table an
agent could exhaust its context before issuing a single query, even after
#71317 bounded the `query` tool.
#71317 deliberately left `get_schema` alone, because the same fix does not
transfer: column names are the information the agent needs to write SQL, so
truncating to the first N columns leaves it unable to reference or discover
the
rest.
## How
`get_schema` now mirrors the `query` tool's bounded-output contract:
- It returns a JSON object `{"columns": [{"name", "type"}, ...],
"column_count":
N}` instead of a bare list.
- It accepts an optional `name_contains` substring filter, so the agent can
ask
for the columns relevant to its question on a wide table.
- Above `max_columns` (new constructor argument, default 100) or the
existing
`max_result_bytes` budget, the full list is replaced by a bounded summary:
`column_count`, a `type_histogram`, a `sample_columns` preview, and
`truncated` / `truncated_by` / `hint` that name the limit hit and point
back
at `name_contains`.
Both toolsets move together. The shared logic lives in `build_schema_result`
next to `build_query_result` in `query_results.py`. The provider changelog
records the tool-result shape change, and the toolset docs gain a "Bounded
schema results" section.
---
##### Was generative AI tooling used to co-author this PR?
- [X] Yes - Claude Code (Opus 4.8)
Generated-by: Claude Code (Opus 4.8) following [the
guidelines](https://github.com/apache/airflow/blob/main/contributing-docs/05_pull_requests.rst#gen-ai-assisted-contributions)
---
🤖 Generated with [Claude Code](https://claude.com/claude-code)
---
providers/common/ai/docs/changelog.rst | 12 ++
providers/common/ai/docs/toolsets/datafusion.rst | 11 +-
providers/common/ai/docs/toolsets/sql.rst | 41 ++++-
.../providers/common/ai/toolsets/datafusion.py | 35 +++-
.../airflow/providers/common/ai/toolsets/sql.py | 33 +++-
.../providers/common/ai/utils/query_results.py | 202 ++++++++++++++++++++-
.../unit/common/ai/toolsets/test_datafusion.py | 55 +++++-
.../ai/tests/unit/common/ai/toolsets/test_sql.py | 62 ++++++-
.../unit/common/ai/utils/test_query_results.py | 167 ++++++++++++++++-
9 files changed, 585 insertions(+), 33 deletions(-)
diff --git a/providers/common/ai/docs/changelog.rst
b/providers/common/ai/docs/changelog.rst
index 8e2826633fb..96b4ea26f35 100644
--- a/providers/common/ai/docs/changelog.rst
+++ b/providers/common/ai/docs/changelog.rst
@@ -39,6 +39,18 @@ Changelog
attempt is checked and counted on its own, unchanged. See :ref:`the
cross-attempt usage
budget <agent-usage-budget>`.
+.. note::
+ ``get_schema`` on ``SQLToolset`` and ``DataFusionToolset`` now returns a
JSON object
+ ``{"columns": [{"name", "type"}, ...], "column_count": N}`` instead of a
bare JSON array of
+ columns. The tool also accepts an optional ``name_contains`` substring
filter, and on a table
+ with more columns than ``max_columns`` (default 100), or one whose
serialized columns exceed
+ ``max_result_bytes``, it returns a bounded summary (``column_count``, a
``type_histogram`` and a
+ ``sample_columns`` preview, with ``truncated``, ``truncated_by`` and a
``hint``) in place of the
+ full list. Update any system prompt or direct ``call_tool("get_schema",
...)`` caller that read
+ the old top-level array: check ``truncated`` first, then read
``result["columns"]`` on a full
+ result or ``result["sample_columns"]`` on a summary. A summary carries no
``columns`` key, so
+ ``result["columns"]`` raises ``KeyError`` on any table wide enough to be
summarized.
+
0.10.0
......
diff --git a/providers/common/ai/docs/toolsets/datafusion.rst
b/providers/common/ai/docs/toolsets/datafusion.rst
index 55570ab32f5..dcf5831891a 100644
--- a/providers/common/ai/docs/toolsets/datafusion.rst
+++ b/providers/common/ai/docs/toolsets/datafusion.rst
@@ -37,7 +37,9 @@ querying files on object stores (S3, GCS, local filesystem,
Iceberg) via Apache
* - ``list_tables``
- Lists registered table names
* - ``get_schema``
- - Returns column names and types for a table (Arrow schema)
+ - Returns a table's columns (Arrow schema) as JSON, with a
``name_contains``
+ filter and a bounded summary on very wide tables (see
+ :ref:`bounded-schema-results`)
* - ``query``
- Executes a SQL query and returns bounded, columnar JSON (see
:ref:`bounded-query-results`)
@@ -84,8 +86,11 @@ Parameters
permitted. DataFusion on object stores is mostly read-only, but it does
support DDL for in-memory tables; this guard blocks those by default.
- ``max_rows``: Maximum rows returned from the ``query`` tool. Default ``50``.
-- ``max_result_bytes``: Budget for the serialized ``query`` result. Default 64
KiB.
- See :ref:`bounded-query-results`.
+- ``max_result_bytes``: Budget for the serialized ``query`` result, and the
byte backstop
+ that also triggers the ``get_schema`` summary. Default 64 KiB.
+ See :ref:`bounded-query-results` and :ref:`bounded-schema-results`.
+- ``max_columns``: Maximum columns ``get_schema`` returns in full. Default
``100``.
+ Above it the result becomes a bounded summary. See
:ref:`bounded-schema-results`.
- ``max_retries``: How many times the model may correct a failed call to these
tools. Default ``None``, the agent's ``retries``. See
:ref:`toolset-retry-budget`.
diff --git a/providers/common/ai/docs/toolsets/sql.rst
b/providers/common/ai/docs/toolsets/sql.rst
index 1c8481ef705..e59a0e1a88a 100644
--- a/providers/common/ai/docs/toolsets/sql.rst
+++ b/providers/common/ai/docs/toolsets/sql.rst
@@ -30,7 +30,8 @@ Curated toolset wrapping
* - ``list_tables``
- Lists available table names (filtered by ``allowed_tables`` if set)
* - ``get_schema``
- - Returns column names and types for a table
+ - Returns a table's columns as JSON, with a ``name_contains`` filter and a
+ bounded summary on very wide tables (see :ref:`bounded-schema-results`)
* - ``query``
- Executes a SQL query and returns bounded, columnar JSON (see
:ref:`bounded-query-results`)
@@ -203,8 +204,11 @@ Parameters
- ``max_rows``: Maximum rows returned from the ``query`` tool. Default ``50``.
Rows beyond it are not read out of a DBAPI cursor; what the driver has
already
transferred is its own call. See :ref:`bounded-query-results`.
-- ``max_result_bytes``: Budget for the serialized ``query`` result. Default 64
KiB.
- See :ref:`bounded-query-results`.
+- ``max_result_bytes``: Budget for the serialized ``query`` result, and the
byte backstop
+ that also triggers the ``get_schema`` summary. Default 64 KiB.
+ See :ref:`bounded-query-results` and :ref:`bounded-schema-results`.
+- ``max_columns``: Maximum columns ``get_schema`` returns in full. Default
``100``.
+ Above it the result becomes a bounded summary. See
:ref:`bounded-schema-results`.
- ``max_retries``: How many times the model may correct a failed call to these
tools. Default ``None``, the agent's ``retries``. See
:ref:`toolset-retry-budget`.
@@ -263,6 +267,37 @@ result several-fold, so results that fit before still fit.
Lower ``max_result_by
when an agent makes many queries in one run, since every result is re-paid on
every
later request.
+.. _bounded-schema-results:
+
+Bounded schema results
+----------------------
+
+``get_schema`` has the same context problem as ``query`` but cannot be solved
the same
+way. Column names are what the agent needs to write SQL, so truncating to the
first N
+columns would leave it unable to reference or discover the rest. The tool
filters and
+summarizes instead.
+
+**Filter with** ``name_contains``. Pass a case-insensitive substring to get
back only
+the columns whose name contains it, so on a very wide table the agent asks for
the
+columns relevant to its question rather than all of them.
+
+**A summary replaces a very wide list.** Above ``max_columns`` (default
``100``), or when
+the serialized columns exceed ``max_result_bytes``, the full list is replaced
by a
+summary that reports the shape and points at the filter:
+
+.. code-block:: json
+
+ {"column_count": 3200, "truncated": true, "truncated_by": "max_columns",
+ "hint": "...", "type_histogram": {"VARCHAR": 2000, "NUMBER": 1200},
+ "sample_columns": [{"name": "id", "type": "NUMBER"}]}
+
+``truncated_by`` is ``max_columns`` or ``max_result_bytes``.
``type_histogram`` counts
+columns per type (capped to the most common, the tail folded into one entry),
and
+``sample_columns`` previews the first columns -- both only while they fit the
budget, so
+a pathologically small budget still returns the ``column_count`` and ``hint``.
A filtered
+call echoes ``name_contains`` and adds ``total_columns`` so a subset is never
mistaken for
+the whole table.
+
When to choose it
-----------------
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 e338cf9bb83..2eda37a280e 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
@@ -37,9 +37,12 @@ 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
@@ -61,6 +64,13 @@ _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"],
}
@@ -112,7 +122,8 @@ class DataFusionToolset(AirflowToolset):
: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.
+ :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
@@ -121,6 +132,11 @@ class DataFusionToolset(AirflowToolset):
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.
@@ -133,6 +149,7 @@ class DataFusionToolset(AirflowToolset):
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)
@@ -142,6 +159,7 @@ class DataFusionToolset(AirflowToolset):
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
@@ -164,7 +182,7 @@ class DataFusionToolset(AirflowToolset):
for name, description, schema in (
("list_tables", "List available table names.",
_LIST_TABLES_SCHEMA),
- ("get_schema", "Get column names and types for a table.",
_GET_SCHEMA_SCHEMA),
+ ("get_schema", _GET_SCHEMA_DESCRIPTION, _GET_SCHEMA_SCHEMA),
("query", _QUERY_DESCRIPTION, _QUERY_SCHEMA),
):
tool_def = ToolDefinition(
@@ -192,7 +210,9 @@ class DataFusionToolset(AirflowToolset):
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"])
+ 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}")
@@ -206,7 +226,7 @@ class DataFusionToolset(AirflowToolset):
log.warning("list_tables failed: %s", ex)
return dumps_masked({"error": str(ex)})
- def _get_schema(self, table_name: str) -> str:
+ 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.
@@ -220,7 +240,12 @@ class DataFusionToolset(AirflowToolset):
# 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 dumps_masked(columns)
+ 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:
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 d3271315f18..a18149ea951 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
@@ -42,9 +42,12 @@ 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, return_schema_kwargs
from airflow.providers.common.ai.utils.toolset_base import AirflowToolset,
validate_max_retries
@@ -73,6 +76,13 @@ _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"],
}
@@ -241,7 +251,8 @@ class SQLToolset(AirflowToolset):
the time the first row is read, so only the per-row Python conversion
is
skipped. Treat this as a bound on what the agent is shown, not as a
guarantee
that ``SELECT * FROM huge_table`` is cheap.
- :param max_result_bytes: Budget for the serialized ``query`` result, in
bytes.
+ :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
@@ -250,6 +261,11 @@ class SQLToolset(AirflowToolset):
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.
@@ -269,6 +285,7 @@ class SQLToolset(AirflowToolset):
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)
@@ -294,6 +311,7 @@ class SQLToolset(AirflowToolset):
self._allow_writes = allow_writes
self._max_rows = max_rows
self._max_result_bytes = max_result_bytes
+ self._max_columns = max_columns
self._hook: DbApiHook | None = None
# Canonical ``(catalog, schema, table)`` view of allowed_tables for
membership
@@ -379,7 +397,7 @@ class SQLToolset(AirflowToolset):
for name, description, schema in (
("list_tables", "List available table names in the database.",
_LIST_TABLES_SCHEMA),
- ("get_schema", "Get column names and types for a table.",
_GET_SCHEMA_SCHEMA),
+ ("get_schema", _GET_SCHEMA_DESCRIPTION, _GET_SCHEMA_SCHEMA),
("query", _QUERY_DESCRIPTION, _QUERY_SCHEMA),
("check_query", "Validate SQL syntax without executing it.",
_CHECK_QUERY_SCHEMA),
):
@@ -431,7 +449,7 @@ class SQLToolset(AirflowToolset):
if name == "list_tables":
return self._list_tables()
if name == "get_schema":
- return self._get_schema(tool_args["table_name"])
+ return self._get_schema(tool_args["table_name"],
tool_args.get("name_contains"))
if name == "query":
return self._query(tool_args["sql"])
return self._check_query(tool_args["sql"])
@@ -475,13 +493,18 @@ class SQLToolset(AirflowToolset):
return dumps_masked(tables)
- def _get_schema(self, table_name: str) -> str:
+ def _get_schema(self, table_name: str, name_contains: str | None = None)
-> str:
schema, table = self._split_table_identifier(table_name)
if not self._is_ref_allowed("", schema, table):
return dumps_masked({"error": f"Table {table_name!r} is not in the
allowed tables list."})
hook = self._get_db_hook()
columns = hook.get_table_schema(table, schema=schema)
- return dumps_masked(columns)
+ return build_schema_result(
+ columns,
+ max_columns=self._max_columns,
+ max_result_bytes=self._max_result_bytes,
+ name_contains=name_contains,
+ )
def _dialect_for_validation(self) -> str | None:
"""Resolve the hook's sqlglot dialect so DESCRIBE/SHOW validate
correctly."""
diff --git
a/providers/common/ai/src/airflow/providers/common/ai/utils/query_results.py
b/providers/common/ai/src/airflow/providers/common/ai/utils/query_results.py
index b8ecdb3e651..09b03c17cf4 100644
--- a/providers/common/ai/src/airflow/providers/common/ai/utils/query_results.py
+++ b/providers/common/ai/src/airflow/providers/common/ai/utils/query_results.py
@@ -15,22 +15,28 @@
# specific language governing permissions and limitations
# under the License.
"""
-Bounded, columnar payloads for the ``query`` tool of the SQL toolsets.
+Bounded payloads for the ``query`` and ``get_schema`` tools of the SQL
toolsets.
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. Two things here keep that bounded:
+cost is re-paid on every subsequent request. Three things here keep that
bounded:
-* **Columnar shape.** ``{"columns": [...], "rows": [[...], ...]}`` names each
column
- once instead of repeating it in a dict per row. On a table with thousands of
+* **Columnar query results.** ``{"columns": [...], "rows": [[...], ...]}``
names each
+ column once instead of repeating it in a dict per row. On a table with
thousands of
columns the repeated names, not the values, are the bulk of the payload.
-* **A byte budget.** ``max_rows`` caps rows, which says nothing about size -- a
- single row of a 3000-column table dwarfs a thousand rows of a narrow one. The
- budget here is what actually bounds context, and when it bites the payload
says
- so, in terms the agent can act on (narrow the projection).
+* **A byte budget.** ``max_rows`` and ``max_columns`` cap how many rows or
columns come
+ back, which says nothing about size -- a single row of a 3000-column table
dwarfs a
+ thousand rows of a narrow one. The budget here is what actually bounds
context, and
+ when it bites the payload says so, in terms the agent can act on (narrow the
+ projection, or filter the columns).
+* **A column cap for get_schema.** Column names are what the agent needs to
write SQL,
+ so above ``max_columns`` the full list is replaced by a summary (count, a
type
+ histogram, a sample) that points at the ``name_contains`` filter rather than
+ truncating blindly.
"""
from __future__ import annotations
+from collections import Counter
from collections.abc import Sequence
from typing import Any
@@ -42,6 +48,22 @@ from airflow.providers.common.ai.utils.masking import
dumps_masked
# Deployments that keep many results in history should lower it.
DEFAULT_MAX_RESULT_BYTES = 65_536
+# Column cap for ``get_schema`` above which the full column list is replaced
by a summary. The
+# byte budget is the real backstop. This is a deterministic, count-based
trigger that is easy for
+# the agent to reason about and keeps an ordinary-width table returning in
full.
+DEFAULT_MAX_COLUMNS = 100
+
+# Cap on distinct keys in a ``get_schema`` type histogram. Parametrized types
(``VARCHAR(255)``,
+# ``Decimal128(38,10)``) would otherwise produce a distinct key per
parameterization and defeat the
+# bound, so everything past the most common ``_SCHEMA_TYPE_HISTOGRAM_TOP_K``
is folded into one entry.
+_SCHEMA_TYPE_HISTOGRAM_TOP_K = 20
+_SCHEMA_TYPE_HISTOGRAM_OTHER_KEY = "(other)"
+
+# Columns previewed in a ``get_schema`` summary. A small fixed sample --
deliberately not the
+# first ``max_columns`` -- keeps the summary much smaller than the list it
replaces and reads as
+# "here is what the names look like, now filter", never as a usable prefix of
the table.
+_SCHEMA_SAMPLE_SIZE = 20
+
# Tool results are machine-read, so no whitespace. ensure_ascii=False matters
as much as
# the separators: escaping one CJK character to \uXXXX costs six bytes instead
of three,
# so an ASCII-escaped result is charged several times over against the budget
and
@@ -60,6 +82,22 @@ QUERY_TOOL_DESCRIPTION = (
"SQL rather than paging through the result."
)
+#: Description for the ``get_schema`` tool. Column names are what the agent
needs to write SQL, so
+#: unlike query rows they cannot be silently dropped: the description states
the ``name_contains``
+#: filter and the summary shape so a truncated result is read as "narrow your
request", not as a
+#: small table.
+GET_SCHEMA_TOOL_DESCRIPTION = (
+ "Get a table's columns. Returns JSON of the form "
+ '{"columns": [{"name": ..., "type": ...}, ...], "column_count": N}. Pass
`name_contains` to '
+ "return only the columns whose name contains that substring
(case-insensitive) -- use it to "
+ "find the columns relevant to your question on a wide table. When a table
has more columns "
+ "than can be returned at once, the full list is replaced by a summary
(`column_count`, and "
+ "where the budget allows a `type_histogram` and a `sample_columns`
preview) with `truncated` "
+ "set, `truncated_by` naming the limit that was hit, and a `hint`. Call
get_schema again with "
+ "a `name_contains` substring to retrieve the specific columns you need
instead of the whole "
+ "table."
+)
+
def _dumps(payload: Any) -> str:
return dumps_masked(payload, **_DUMP_KWARGS)
@@ -155,3 +193,151 @@ def build_query_result(
f"than paging through the result."
)
return _dumps(output)
+
+
+def _build_type_histogram(columns: Sequence[dict[str, str]]) -> dict[str, int]:
+ """
+ Count columns per type, capped to the most common
``_SCHEMA_TYPE_HISTOGRAM_TOP_K`` types.
+
+ Ordered by ``(-count, type)`` so the output is content-stable
(deterministic for byte
+ accounting and tests, not merely input-ordered). The long tail is folded
into a single
+ aggregate entry rather than listed, so parametrized types cannot inflate
the key count.
+
+ :param columns: ``{"name", "type"}`` dicts.
+ """
+ counts = Counter(col["type"] for col in columns)
+ ordered = sorted(counts.items(), key=lambda item: (-item[1], item[0]))
+ if len(ordered) <= _SCHEMA_TYPE_HISTOGRAM_TOP_K:
+ return dict(ordered)
+ histogram = dict(ordered[:_SCHEMA_TYPE_HISTOGRAM_TOP_K])
+ histogram[_SCHEMA_TYPE_HISTOGRAM_OTHER_KEY] = sum(
+ count for _, count in ordered[_SCHEMA_TYPE_HISTOGRAM_TOP_K:]
+ )
+ return histogram
+
+
+def build_schema_result(
+ columns: Sequence[dict[str, str]],
+ *,
+ max_columns: int,
+ max_result_bytes: int,
+ name_contains: str | None = None,
+) -> str:
+ """
+ Render a table's columns as a bounded JSON tool result.
+
+ Column names are the information an agent needs to write SQL, so unlike
query rows they cannot
+ simply be dropped: above ``max_columns`` (or the byte budget) the full
list is replaced by a
+ summary -- count, a type histogram, and a sample of columns -- that names
``name_contains`` as
+ the way to retrieve specific columns. ``columns`` is assumed to hold
distinct names (a table's
+ introspected columns are unique by construction).
+
+ :param columns: ``{"name", "type"}`` dicts in table order.
+ :param max_columns: Column count above which a summary replaces the full
list.
+ :param max_result_bytes: Budget for the serialized result.
+ :param name_contains: Case-insensitive substring. When given (and
non-empty), only matching
+ columns are considered and the value is echoed back so a filtered
subset is never mistaken
+ for the whole table.
+ """
+ name_contains = name_contains or None
+ total_columns = len(columns)
+ if name_contains is not None:
+ needle = name_contains.casefold()
+ selected = [col for col in columns if needle in col["name"].casefold()]
+ else:
+ selected = list(columns)
+
+ if name_contains is not None and not selected:
+ plural = "" if total_columns == 1 else "s"
+ # Do not tell the agent to call without name_contains when the full
list would itself be
+ # summarized (total_columns > max_columns) -- that would bounce it
back to this filter.
+ if total_columns > max_columns:
+ advice = (
+ f"No columns match name_contains={name_contains!r}. Try a
different substring "
+ f"(the table has {total_columns} column{plural}, too many to
list in full)."
+ )
+ else:
+ advice = (
+ f"No columns match name_contains={name_contains!r}. Call
get_schema without "
+ f"name_contains to list all {total_columns} column{plural}."
+ )
+ return _dumps(
+ {
+ "columns": [],
+ "column_count": 0,
+ "name_contains": name_contains,
+ "total_columns": total_columns,
+ "hint": advice,
+ }
+ )
+
+ full: dict[str, Any] = {"columns": selected, "column_count": len(selected)}
+ if name_contains is not None:
+ full["name_contains"] = name_contains
+ full["total_columns"] = total_columns
+ if len(selected) <= max_columns and _size(full) <= max_result_bytes:
+ return _dumps(full)
+
+ return _summarize_schema(
+ selected,
+ total_columns=total_columns,
+ max_columns=max_columns,
+ max_result_bytes=max_result_bytes,
+ name_contains=name_contains,
+ )
+
+
+def _summarize_schema(
+ selected: list[dict[str, str]],
+ *,
+ total_columns: int,
+ max_columns: int,
+ max_result_bytes: int,
+ name_contains: str | None,
+) -> str:
+ """Build the bounded summary returned when the full column list does not
fit."""
+ # Count is checked before bytes: it is the cheaper, more explainable
bound. A bytes-only
+ # truncation then means "count fits but names/types are pathologically
long", a rarer signal.
+ n = len(selected)
+ plural = "" if n == 1 else "s"
+ truncated_by = "max_columns" if n > max_columns else "max_result_bytes"
+ subject = f"name_contains={name_contains!r} matched" if name_contains is
not None else "This table has"
+ if truncated_by == "max_columns":
+ reason = f"{subject} {n} column{plural}, more than the max_columns
limit of {max_columns}."
+ else:
+ reason = f"{subject} {n} column{plural} that did not fit
max_result_bytes ({max_result_bytes})."
+ move = (
+ "Use a more specific name_contains substring to narrow to the columns
you need."
+ if name_contains is not None
+ else "Call get_schema with a name_contains substring to return only
the columns you need."
+ )
+
+ output: dict[str, Any] = {
+ "column_count": n,
+ "truncated": True,
+ "truncated_by": truncated_by,
+ "hint": f"{reason} {move}",
+ }
+ if name_contains is not None:
+ output["name_contains"] = name_contains
+ output["total_columns"] = total_columns
+
+ # The core above is the guaranteed-useful payload. Add the histogram and a
small sample of
+ # columns only while they fit, reserving the core first. If the histogram
alone would not fit,
+ # drop it but still fill the sample from what remains, so a wide struct
type cannot strip the
+ # preview off every other column. A pathologically small budget still
returns the core.
+ output["type_histogram"] = _build_type_histogram(selected)
+ output["sample_columns"] = []
+ if _size(output) > max_result_bytes:
+ del output["type_histogram"]
+ budget = max_result_bytes - _size(output)
+ for col in selected[:_SCHEMA_SAMPLE_SIZE]:
+ cost = _size(col) + (1 if output["sample_columns"] else 0)
+ # Skip a column too large for the remaining budget and keep going,
rather than stopping:
+ # the sample is an unordered preview, so one oversized column early (a
wide struct type)
+ # must not strip the preview off every column after it.
+ if cost > budget:
+ continue
+ budget -= cost
+ output["sample_columns"].append(col)
+ return _dumps(output)
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 c97d6978acd..9e6372a96f9 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
@@ -112,6 +112,7 @@ class TestDataFusionToolsetArgsValidation:
("tool_name", "valid_args"),
[
("get_schema", {"table_name": "sales_data"}),
+ ("get_schema", {"table_name": "sales_data", "name_contains":
"am"}),
("query", {"sql": "SELECT 1"}),
],
)
@@ -169,12 +170,58 @@ class TestDataFusionToolsetGetSchema:
tool=MagicMock(spec=ToolsetTool),
)
)
- columns = json.loads(result)
- assert columns == [
- {"name": "id", "type": "Int64"},
+ data = json.loads(result)
+ assert data == {
+ "columns": [
+ {"name": "id", "type": "Int64"},
+ {"name": "amount", "type": "Float64"},
+ {"name": "name", "type": "Utf8"},
+ ],
+ "column_count": 3,
+ }
+
+ def test_name_contains_filters_the_columns(self):
+ """``name_contains`` threads from the tool call through to the bounded
result."""
+ cfg = _make_mock_datasource_config()
+ ts = DataFusionToolset([cfg])
+ ts._engine = _make_mock_engine(
+ schema_fields=[("id", "Int64"), ("amount", "Float64"),
("item_name", "Utf8")]
+ )
+
+ result = asyncio.run(
+ ts.call_tool(
+ "get_schema",
+ {"table_name": "sales_data", "name_contains": "am"},
+ ctx=MagicMock(spec=RunContext),
+ tool=MagicMock(spec=ToolsetTool),
+ )
+ )
+ data = json.loads(result)
+ assert data["columns"] == [
{"name": "amount", "type": "Float64"},
- {"name": "name", "type": "Utf8"},
+ {"name": "item_name", "type": "Utf8"},
]
+ assert data["name_contains"] == "am"
+ assert data["total_columns"] == 3
+
+
@patch("airflow.providers.common.ai.toolsets.datafusion.build_schema_result",
return_value="{}")
+ def test_get_schema_forwards_the_toolsets_bounds(self, mock_build):
+ """The toolset's own max_columns/max_result_bytes reach
build_schema_result, not defaults."""
+ ts = DataFusionToolset([_make_mock_datasource_config()],
max_columns=7, max_result_bytes=123)
+ ts._engine = _make_mock_engine()
+
+ asyncio.run(
+ ts.call_tool(
+ "get_schema",
+ {"table_name": "sales_data", "name_contains": "id"},
+ ctx=MagicMock(spec=RunContext),
+ tool=MagicMock(spec=ToolsetTool),
+ )
+ )
+ kwargs = mock_build.call_args.kwargs
+ assert kwargs["max_columns"] == 7
+ assert kwargs["max_result_bytes"] == 123
+ assert kwargs["name_contains"] == "id"
class TestDataFusionToolsetQuery:
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 f62d2cb436e..53a978cf7e8 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
@@ -123,6 +123,7 @@ class TestSQLToolsetGetTools:
("name", "valid_args"),
[
("get_schema", {"table_name": "users"}),
+ ("get_schema", {"table_name": "users", "name_contains": "cust"}),
("query", {"sql": "SELECT 1"}),
("check_query", {"sql": "SELECT 1"}),
],
@@ -187,20 +188,73 @@ class TestSQLToolsetGetSchema:
result = asyncio.run(
ts.call_tool("get_schema", {"table_name": "users"},
ctx=MagicMock(), tool=MagicMock())
)
- columns = json.loads(result)
- assert columns == [{"name": "id", "type": "INTEGER"}, {"name": "name",
"type": "VARCHAR"}]
+ data = json.loads(result)
+ assert data == {
+ "columns": [{"name": "id", "type": "INTEGER"}, {"name": "name",
"type": "VARCHAR"}],
+ "column_count": 2,
+ }
mock_hook.get_table_schema.assert_called_once_with("users",
schema=None)
+ def test_name_contains_filters_the_columns(self):
+ """``name_contains`` threads from the tool call through to the bounded
result."""
+ ts = SQLToolset("pg_default")
+ ts._hook = _make_mock_db_hook(
+ table_schema=[
+ {"name": "id", "type": "INTEGER"},
+ {"name": "customer_name", "type": "VARCHAR"},
+ ]
+ )
+
+ result = asyncio.run(
+ ts.call_tool(
+ "get_schema",
+ {"table_name": "users", "name_contains": "name"},
+ ctx=MagicMock(),
+ tool=MagicMock(),
+ )
+ )
+ data = json.loads(result)
+ assert data["columns"] == [{"name": "customer_name", "type":
"VARCHAR"}]
+ assert data["name_contains"] == "name"
+ assert data["total_columns"] == 2
+
+ @patch("airflow.providers.common.ai.toolsets.sql.build_schema_result",
return_value="{}")
+ def test_get_schema_forwards_the_toolsets_bounds(self, mock_build):
+ """The toolset's own max_columns/max_result_bytes reach
build_schema_result, not defaults."""
+ ts = SQLToolset("pg_default", max_columns=7, max_result_bytes=123)
+ ts._hook = _make_mock_db_hook()
+
+ asyncio.run(
+ ts.call_tool(
+ "get_schema",
+ {"table_name": "users", "name_contains": "id"},
+ ctx=MagicMock(),
+ tool=MagicMock(),
+ )
+ )
+ kwargs = mock_build.call_args.kwargs
+ assert kwargs["max_columns"] == 7
+ assert kwargs["max_result_bytes"] == 123
+ assert kwargs["name_contains"] == "id"
+
def test_blocks_table_not_in_allowed_list(self):
+ """The allow-list guard fires before any introspection or filtering."""
ts = SQLToolset("pg_default", allowed_tables=["orders"])
ts._hook = _make_mock_db_hook()
result = asyncio.run(
- ts.call_tool("get_schema", {"table_name": "secrets"},
ctx=MagicMock(), tool=MagicMock())
+ ts.call_tool(
+ "get_schema",
+ {"table_name": "secrets", "name_contains": "pw"},
+ ctx=MagicMock(),
+ tool=MagicMock(),
+ )
)
data = json.loads(result)
assert "error" in data
assert "secrets" in data["error"]
+ # The guard must short-circuit before touching the database.
+ ts._hook.get_table_schema.assert_not_called()
def test_introspection_error_raises_model_retry(self):
"""A failure while reading a table's schema is returned to the agent
as a retry."""
@@ -629,7 +683,7 @@ class TestSQLToolsetMultiSchema:
)
)
)
- assert result == [{"name": "id", "type": "INTEGER"}]
+ assert result == {"columns": [{"name": "id", "type": "INTEGER"}],
"column_count": 1}
ts._hook.get_table_schema.assert_called_once_with("DEPLOYMENT_IMAGE_DETAILS",
schema="MODEL_ASTRO")
def test_get_schema_blocks_table_outside_allowed_schema(self):
diff --git
a/providers/common/ai/tests/unit/common/ai/utils/test_query_results.py
b/providers/common/ai/tests/unit/common/ai/utils/test_query_results.py
index b5bebc921ac..eee5a06c746 100644
--- a/providers/common/ai/tests/unit/common/ai/utils/test_query_results.py
+++ b/providers/common/ai/tests/unit/common/ai/utils/test_query_results.py
@@ -22,7 +22,11 @@ from decimal import Decimal
import pytest
-from airflow.providers.common.ai.utils.query_results import build_query_result
+from airflow.providers.common.ai.utils.query_results import (
+ _SCHEMA_SAMPLE_SIZE,
+ build_query_result,
+ build_schema_result,
+)
def _build(columns, rows, *, max_rows=50, max_result_bytes=65_536, more=False,
total=None) -> dict:
@@ -169,3 +173,164 @@ def
test_a_secret_that_json_escapes_is_masked_in_the_rows(register_secret):
result = _build(["user", "password"], [["admin", secret]])
assert result["rows"] == [["admin", "***"]]
+
+
+def _schema(columns, *, max_columns=100, max_result_bytes=65_536,
name_contains=None) -> dict:
+ return json.loads(
+ build_schema_result(
+ columns,
+ max_columns=max_columns,
+ max_result_bytes=max_result_bytes,
+ name_contains=name_contains,
+ )
+ )
+
+
+def _cols(n: int, *, type_: str = "VARCHAR", prefix: str = "col") ->
list[dict[str, str]]:
+ return [{"name": f"{prefix}_{i}", "type": type_} for i in range(n)]
+
+
+class TestSchemaResult:
+ def test_small_table_returns_every_column(self):
+ cols = _cols(3)
+ assert _schema(cols) == {"columns": cols, "column_count": 3}
+
+ def test_name_contains_filters_case_insensitively(self):
+ cols = [
+ {"name": "CustomerId", "type": "INT"},
+ {"name": "amount", "type": "NUMERIC"},
+ {"name": "customer_name", "type": "VARCHAR"},
+ ]
+ data = _schema(cols, name_contains="customer")
+ assert data["columns"] == [cols[0], cols[2]]
+ assert data["column_count"] == 2
+ assert data["name_contains"] == "customer"
+ assert data["total_columns"] == 3
+ assert "truncated" not in data
+
+ def test_empty_name_contains_is_treated_as_no_filter(self):
+ cols = _cols(3)
+ assert _schema(cols, name_contains="") == {"columns": cols,
"column_count": 3}
+
+ def test_name_contains_with_no_matches_guides_without_erroring(self):
+ data = _schema(_cols(5), name_contains="zzz")
+ assert data["columns"] == []
+ assert data["column_count"] == 0
+ assert data["name_contains"] == "zzz"
+ assert data["total_columns"] == 5
+ assert "error" not in data
+ assert "all 5 columns" in data["hint"]
+
+ def test_no_match_on_a_wide_table_does_not_point_back_at_the_filter(self):
+ """When the full list would itself summarize, the hint must not say
'call without it'."""
+ data = _schema(_cols(300), max_columns=100, name_contains="zzz")
+ assert data["columns"] == []
+ assert data["total_columns"] == 300
+ assert "without name_contains" not in data["hint"]
+ assert "different substring" in data["hint"]
+
+ def test_too_many_columns_are_summarized_not_listed(self):
+ raw = build_schema_result(_cols(250), max_columns=100,
max_result_bytes=65_536)
+ data = json.loads(raw)
+ assert data["truncated"] is True
+ assert data["truncated_by"] == "max_columns"
+ assert data["column_count"] == 250
+ assert "columns" not in data
+ assert data["type_histogram"] == {"VARCHAR": 250}
+ # A small fixed preview, not the first max_columns, so the summary is
far smaller than the
+ # list it replaces rather than a near-identical prefix with the tail
dropped.
+ assert len(data["sample_columns"]) == _SCHEMA_SAMPLE_SIZE
+ assert data["sample_columns"][0] == {"name": "col_0", "type":
"VARCHAR"}
+ assert len(raw.encode("utf-8")) <
len(json.dumps(_cols(250)).encode("utf-8"))
+ assert "name_contains" in data["hint"]
+
+ @pytest.mark.parametrize("budget", [500, 2000, 65_536])
+ def test_summary_never_exceeds_the_byte_budget(self, budget):
+ """The contract the docstring promises: a summary that fits the budget
stays within it."""
+ cols = [{"name": f"customer_attribute_{i}", "type": "VARCHAR(255)"}
for i in range(500)]
+ raw = build_schema_result(cols, max_columns=100,
max_result_bytes=budget)
+ assert json.loads(raw)["truncated"] is True
+ assert len(raw.encode("utf-8")) <= budget
+
+ def test_histogram_is_dropped_but_the_sample_is_still_filled(self):
+ """A histogram too big for the budget must not strip the preview off
every column too."""
+ # Many distinct long type names make the histogram large. The budget
holds the core plus a
+ # few short-named columns but not the histogram.
+ cols = [{"name": f"c{i}", "type":
f"CUSTOM_STRUCT_TYPE_NUMBER_{i:04d}"} for i in range(200)]
+ data = _schema(cols, max_columns=100, max_result_bytes=600)
+ assert data["truncated"] is True
+ assert "type_histogram" not in data
+ assert len(data["sample_columns"]) >= 1
+
+ def
test_an_oversized_column_early_in_the_window_is_skipped_not_fatal(self):
+ """A wide column at the front of the preview window is skipped, not a
full stop."""
+ # One ~2 KB struct type at position 0, then narrow columns. The
preview must skip the wide
+ # one and keep filling from the rest rather than coming back empty.
+ cols = [{"name": "wide", "type": "X" * 2000}] + [{"name": f"c{i}",
"type": "INT"} for i in range(150)]
+ data = _schema(cols, max_columns=100, max_result_bytes=1500)
+ assert data["truncated"] is True
+ assert len(data["sample_columns"]) >= 1
+ assert all(col["name"] != "wide" for col in data["sample_columns"])
+
+ def test_the_max_columns_boundary_is_exact(self):
+ """Exactly max_columns returns the full list, and one more summarizes
by count."""
+ at_limit = _schema(_cols(100), max_columns=100)
+ assert at_limit["column_count"] == 100
+ assert "columns" in at_limit
+ assert "truncated" not in at_limit
+
+ assert _schema(_cols(101), max_columns=100)["truncated_by"] ==
"max_columns"
+
+ # Exactly max_columns but over the byte budget is a bytes truncation,
not a count one.
+ over_bytes = _schema(_cols(100), max_columns=100, max_result_bytes=500)
+ assert over_bytes["truncated_by"] == "max_result_bytes"
+
+ def test_long_names_blow_the_byte_budget_despite_few_columns(self):
+ cols = [{"name": "x" * 500, "type": "VARCHAR"} for _ in range(10)]
+ data = _schema(cols, max_columns=100, max_result_bytes=512)
+ assert data["truncated"] is True
+ assert data["truncated_by"] == "max_result_bytes"
+
+ def test_max_columns_takes_precedence_when_both_bounds_are_exceeded(self):
+ data = _schema(_cols(300), max_columns=100, max_result_bytes=256)
+ assert data["truncated_by"] == "max_columns"
+
+ def test_summary_hint_is_worded_from_the_limit_it_hit(self):
+ over_count = _schema(_cols(250), max_columns=100)
+ assert "250 columns, more than the max_columns limit of 100" in
over_count["hint"]
+
+ # One column whose name alone blows a tiny budget: the count fits, the
bytes do not, so the
+ # wording must say so and read "1 column" (not "1 columns").
+ over_bytes = _schema([{"name": "x" * 500, "type": "VARCHAR"}],
max_columns=100, max_result_bytes=200)
+ assert over_bytes["truncated_by"] == "max_result_bytes"
+ assert "1 column that did not fit max_result_bytes" in
over_bytes["hint"]
+ assert "1 columns" not in over_bytes["hint"]
+
+ def test_filtered_result_still_too_wide_is_summarized(self):
+ cols = [{"name": f"customer_{i}", "type": "VARCHAR"} for i in
range(200)]
+ data = _schema(cols, max_columns=100, name_contains="customer")
+ assert data["truncated"] is True
+ assert data["truncated_by"] == "max_columns"
+ assert data["name_contains"] == "customer"
+ assert data["total_columns"] == 200
+ assert "more specific name_contains" in data["hint"]
+
+ def test_type_histogram_is_capped_and_ordered_by_count(self):
+ cols: list[dict[str, str]] = []
+ for t in range(30):
+ cols.extend({"name": f"c_{t}_{i}", "type": f"T{t:02d}"} for i in
range(t + 1))
+ histogram = _schema(cols, max_columns=1)["type_histogram"]
+ keys = list(histogram.keys())
+ assert len(histogram) == 21 # 20 most common types plus the folded
remainder
+ assert keys[0] == "T29" # highest count first
+ assert keys[-1] == "(other)"
+ assert histogram["(other)"] == sum(range(1, 11)) # T00..T09 => 1 + 2
+ ... + 10
+
+ def test_summary_core_survives_a_pathologically_small_budget(self):
+ data = _schema(_cols(300), max_columns=100, max_result_bytes=1)
+ assert data["column_count"] == 300
+ assert data["truncated"] is True
+ assert data["truncated_by"] == "max_columns"
+ assert data["sample_columns"] == []
+ assert "type_histogram" not in data
+ assert data["hint"]