pankajkoti commented on code in PR #74017:
URL: https://github.com/apache/airflow/pull/74017#discussion_r4158022512
##########
providers/common/ai/tests/unit/common/ai/utils/test_query_results.py:
##########
@@ -169,3 +169,99 @@ 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 "error" not in data
+ assert "all 5 columns" in data["hint"]
+
+ def test_too_many_columns_are_summarized_not_listed(self):
+ data = _schema(_cols(250), max_columns=100)
+ 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}
+ assert len(data["sample_columns"]) == 100
+ assert data["sample_columns"][0] == {"name": "col_0", "type":
"VARCHAR"}
+ assert "name_contains" in data["hint"]
+
+ 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"
Review Comment:
Added a parametrized test asserting the serialized summary stays within
`max_result_bytes`, and a test per toolset that `get_schema` forwards its own
`max_columns` and `max_result_bytes` into `build_schema_result` (the module
defaults would fail it).
##########
providers/common/ai/tests/unit/common/ai/toolsets/test_sql.py:
##########
@@ -187,16 +188,48 @@ def test_returns_column_info(self):
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
+
def test_blocks_table_not_in_allowed_list(self):
+ """The allow-list guard fires before any introspection or filtering."""
Review Comment:
Added `ts._hook.get_table_schema.assert_not_called()` to that test.
--
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]