From a9f4db9b66fe6062a51df404c5998075ebdf46b4 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Mon, 28 Sep 2026 19:13:52 +0900 Subject: [PATCH 01/14] Derive the one_row_complex expectations of the Python-object cursors test_complex for Cursor, S3FSCursor, and AsyncS3FSCursor, test_as_pandas, and the SQLAlchemy test_reflect_select spelled out the one_row_complex row, description, DB API types, and reflected types by hand. They now take them from the table definition through tests/pyathena/expected.py, which states per Athena type family what these cursors return: map and struct values as strings, JSON casts as parsed JSON, and so on. The rules never call PyAthena's converters. The queries list the table columns first and the CAST columns after them. test_as_pandas now also covers col_varchar, and test_reflect_select compares the reflected column names instead of their count. Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/expected.py | 334 +++++++++++++++++++++++ tests/pyathena/pandas/test_util.py | 77 +----- tests/pyathena/s3fs/test_async_cursor.py | 77 +----- tests/pyathena/s3fs/test_cursor.py | 77 +----- tests/pyathena/sqlalchemy/test_base.py | 51 +--- tests/pyathena/tables.py | 97 ++++--- tests/pyathena/test_cursor.py | 117 ++------ 7 files changed, 426 insertions(+), 404 deletions(-) create mode 100644 tests/pyathena/expected.py diff --git a/tests/pyathena/expected.py b/tests/pyathena/expected.py new file mode 100644 index 00000000..6df955a0 --- /dev/null +++ b/tests/pyathena/expected.py @@ -0,0 +1,334 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT +"""Expected query results for the shared tables in ``tests.pyathena.tables``. + +The expectations are derived from a table's column types and row values with +explicit rules per Athena type family, such as ``array`` or ``decimal``. The +rules state what a cursor returns; they never call PyAthena's converters. +""" + +import json +import re +from collections.abc import Callable, Mapping +from dataclasses import dataclass +from datetime import timezone +from typing import Any + +from sqlalchemy import types + +from pyathena import BINARY, BOOLEAN, DATE, DATETIME, JSON, NUMBER, STRING, TIME +from pyathena.sqlalchemy.types import TINYINT, AthenaArray, AthenaStruct +from tests.pyathena.tables import Column, Table + + +def family(athena_type: str) -> str: + """Return the type family, the type name without its parameters. + + Args: + athena_type: An Athena type, such as ``DECIMAL(10,1)`` or ``ARRAY``. + + Returns: + The lowercase family, such as ``decimal`` or ``array``. + """ + match = re.match(r"[a-z]+(?: [a-z]+)*", athena_type.strip().lower()) + assert match, athena_type + return match.group() + + +def _parameters(athena_type: str) -> tuple[int, ...]: + """Return the numeric parameters of a type. + + Args: + athena_type: An Athena type, such as ``DECIMAL(10,1)``. + + Returns: + The parameters, such as ``(10, 1)``; empty for a type without them. + """ + match = re.search(r"\(([\d,\s]+)\)", athena_type) + return tuple(int(p) for p in match.group(1).split(",")) if match else () + + +def _element_type(athena_type: str) -> str: + """Return the element type of an array type. + + Args: + athena_type: An array type, such as ``ARRAY``. + + Returns: + The element type, such as ``int``. + """ + return athena_type[athena_type.index("<") + 1 : athena_type.rindex(">")] + + +@dataclass(frozen=True) +class Cast: + """A column the query casts from a table column. + + Attributes: + name: The result column name. + source: The table column to cast. + athena_type: The type to cast to. + """ + + name: str + source: str + athena_type: str + + def sql(self) -> str: + """Return the select-list item. + + Returns: + The ``CAST`` expression with its alias. + """ + return f"CAST({self.source} AS {self.athena_type}) AS {self.name}" + + +TIMESTAMP_TZ = Cast("col_timestamp_tz", "col_timestamp", "timestamp with time zone") +TIME_OF_TIMESTAMP = Cast("col_time", "col_timestamp", "time") +ARRAY_JSON = Cast("col_array_json", "col_array", "json") +MAP_JSON = Cast("col_map_json", "col_map", "json") + +# (type code, precision, scale) in the cursor description, by type family. +# The precision of varchar(n) and the precision and scale of decimal(p,s) come +# from the type parameters. +_DESCRIPTION: Mapping[str, tuple[str, int, int]] = { + "boolean": ("boolean", 0, 0), + "tinyint": ("tinyint", 3, 0), + "smallint": ("smallint", 5, 0), + "int": ("integer", 10, 0), + "bigint": ("bigint", 19, 0), + "float": ("float", 17, 0), + "double": ("double", 17, 0), + "string": ("varchar", 2147483647, 0), + "varchar": ("varchar", 0, 0), + "timestamp": ("timestamp", 3, 0), + "timestamp with time zone": ("timestamp with time zone", 3, 0), + "time": ("time", 3, 0), + "date": ("date", 0, 0), + "binary": ("varbinary", 1073741824, 0), + "array": ("array", 0, 0), + "map": ("map", 0, 0), + "struct": ("row", 0, 0), + "decimal": ("decimal", 0, 0), + "json": ("json", 0, 0), +} + +_DBAPI_TYPES = { + "boolean": BOOLEAN, + "tinyint": NUMBER, + "smallint": NUMBER, + "int": NUMBER, + "bigint": NUMBER, + "float": NUMBER, + "double": NUMBER, + "decimal": NUMBER, + "string": STRING, + "varchar": STRING, + "array": STRING, + "map": STRING, + "struct": STRING, + "timestamp": DATETIME, + "timestamp with time zone": DATETIME, + "time": TIME, + "date": DATE, + "binary": BINARY, + "json": JSON, +} + +# The SQLAlchemy type class a reflected column has, by type family. +_SQLALCHEMY_TYPES = { + "boolean": types.BOOLEAN, + "tinyint": TINYINT, + "smallint": types.SMALLINT, + "int": types.INTEGER, + "bigint": types.BIGINT, + "float": types.FLOAT, + "double": types.DOUBLE, + "string": types.String, + "varchar": types.VARCHAR, + "timestamp": types.TIMESTAMP, + "date": types.DATE, + "binary": types.BINARY, + "array": AthenaArray, + "map": types.String, + "struct": AthenaStruct, + "decimal": types.DECIMAL, +} + + +def _json_compatible(value: Any, athena_type: str) -> Any: + """Return a value as the JSON structure Athena renders for ``CAST(... AS json)``. + + Args: + value: The value from the table definition. + athena_type: The value's Athena type. + + Returns: + The value, with a map's key-value pairs as a dict. + """ + if family(athena_type) == "map": + return dict(value) + return value + + +def _parsed_json(value: Any, athena_type: str) -> Any: + """Return a value after a round trip through JSON text. + + Args: + value: The value from the table definition. + athena_type: The value's Athena type. + + Returns: + The parsed JSON value; map keys become strings. + """ + return json.loads(json.dumps(_json_compatible(value, athena_type))) + + +def _same(value: Any, athena_type: str) -> Any: + return value + + +def _as_list(value: Any, athena_type: str) -> Any: + return list(value) + + +# A representation maps a type family to a rule(value, athena_type) that +# returns what the cursor gives for that value. Cast columns use the rule of +# the target family with the source column's value and type. +Representation = Mapping[str, Callable[[Any, str], Any]] + +_SCALARS: Representation = dict.fromkeys( + ( + "boolean", + "tinyint", + "smallint", + "int", + "bigint", + "float", + "double", + "string", + "varchar", + "timestamp", + "date", + "binary", + "decimal", + ), + _same, +) + +# Cursor, DictCursor, S3FSCursor, pyathena.pandas.util.as_pandas, and SQLAlchemy +# result rows. Without type hints, map and struct values are strings. +PYTHON: Representation = { + **_SCALARS, + "timestamp with time zone": lambda v, t: v.replace(tzinfo=timezone.utc), + "time": lambda v, t: v.time(), + "array": _as_list, + "map": lambda v, t: {str(k): str(x) for k, x in v}, + "struct": lambda v, t: {k: str(x) for k, x in v.items()}, + "json": _parsed_json, +} + + +@dataclass(frozen=True) +class Selection: + """A query of a table's columns, followed by cast columns. + + Attributes: + table: The table. + casts: The cast columns, after the table columns. + """ + + table: Table + casts: tuple[Cast, ...] = () + + @property + def names(self) -> list[str]: + """The result column names.""" + return [c.name for c in self.table.columns] + [c.name for c in self.casts] + + @property + def sql(self) -> str: + """The query.""" + items = [c.name for c in self.table.columns] + [c.sql() for c in self.casts] + return f"SELECT {', '.join(items)} FROM {self.table.name}" + + def _items(self) -> list[tuple[str, str, int, str]]: + """Return how each result column is derived from the table. + + Returns: + ``(name, athena_type, source_index, source_type)`` per result column, + where the source is the table column the value comes from. A table + column is its own source. + """ + columns = self.table.columns + index = {c.name: i for i, c in enumerate(columns)} + return [(c.name, c.athena_type, i, c.athena_type) for i, c in enumerate(columns)] + [ + (c.name, c.athena_type, index[c.source], columns[index[c.source]].athena_type) + for c in self.casts + ] + + def description(self) -> list[tuple[Any, ...]]: + """Return the expected cursor description. + + Returns: + One DB API description tuple per result column. + """ + result = [] + for name, athena_type, _, _ in self._items(): + code, precision, scale = _DESCRIPTION[family(athena_type)] + if parameters := _parameters(athena_type): + precision, scale = (*parameters, 0)[:2] + result.append((name, code, None, None, precision, scale, "UNKNOWN")) + return result + + def dbapi_types(self) -> list[Any]: + """Return the expected DB API type object of each result column. + + Returns: + The type objects in result order. + """ + return [_DBAPI_TYPES[family(athena_type)] for _, athena_type, _, _ in self._items()] + + def rows(self, representation: Representation) -> list[tuple[Any, ...]]: + """Return the expected rows. + + Args: + representation: The rules of the cursor under test, such as ``PYTHON``. + + Returns: + One tuple per row of the table. + """ + items = self._items() + return [ + tuple( + representation[family(athena_type)](row[index], source_type) + for _, athena_type, index, source_type in items + ) + for row in self.table.rows + ] + + +def assert_sqlalchemy_type(sqlalchemy_type: Any, column: Column) -> None: + """Assert that a reflected SQLAlchemy type matches a table column. + + Args: + sqlalchemy_type: The reflected column type. + column: The column in the table definition. + + Raises: + AssertionError: If the type class or its parameters differ. + """ + name = family(column.athena_type) + assert isinstance(sqlalchemy_type, _SQLALCHEMY_TYPES[name]), (column.name, sqlalchemy_type) + if name == "varchar": + assert sqlalchemy_type.length == _parameters(column.athena_type)[0], column.name + elif name == "decimal": + precision, scale = _parameters(column.athena_type) + assert (sqlalchemy_type.precision, sqlalchemy_type.scale) == (precision, scale) + elif name == "array": + element = _element_type(column.athena_type) + assert isinstance(sqlalchemy_type.item_type, _SQLALCHEMY_TYPES[family(element)]) diff --git a/tests/pyathena/pandas/test_util.py b/tests/pyathena/pandas/test_util.py index 1acbc07e..4eccdc2d 100644 --- a/tests/pyathena/pandas/test_util.py +++ b/tests/pyathena/pandas/test_util.py @@ -1,7 +1,6 @@ import textwrap import uuid from datetime import date, datetime -from decimal import Decimal import numpy as np import pandas as pd @@ -16,6 +15,8 @@ to_sql, ) from tests import ENV +from tests.pyathena.expected import ARRAY_JSON, MAP_JSON, PYTHON, TIME_OF_TIMESTAMP, Selection +from tests.pyathena.tables import ONE_ROW_COMPLEX def test_get_chunks(): @@ -52,77 +53,11 @@ def test_reset_index(): def test_as_pandas(cursor): - cursor.execute( - """ - SELECT - col_boolean - , col_tinyint - , col_smallint - , col_int - , col_bigint - , col_float - , col_double - , col_string - , col_timestamp - , CAST(col_timestamp AS time) AS col_time - , col_date - , col_binary - , col_array - , CAST(col_array AS json) AS col_array_json - , col_map - , CAST(col_map AS json) AS col_map_json - , col_struct - , col_decimal - FROM one_row_complex - """ - ) + selection = Selection(ONE_ROW_COMPLEX, (TIME_OF_TIMESTAMP, ARRAY_JSON, MAP_JSON)) + cursor.execute(selection.sql) df = as_pandas(cursor) - rows = [ - ( - row["col_boolean"], - row["col_tinyint"], - row["col_smallint"], - row["col_int"], - row["col_bigint"], - row["col_float"], - row["col_double"], - row["col_string"], - row["col_timestamp"], - row["col_time"], - row["col_date"], - row["col_binary"], - row["col_array"], - row["col_array_json"], - row["col_map"], - row["col_map_json"], - row["col_struct"], - row["col_decimal"], - ) - for _, row in df.iterrows() - ] - expected = [ - ( - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - datetime(2017, 1, 1, 0, 0, 0), - datetime(2017, 1, 1, 0, 0, 0).time(), - date(2017, 1, 2), - b"123", - [1, 2], - [1, 2], - {"1": "2", "3": "4"}, - {"1": 2, "3": 4}, - {"a": "1", "b": "2"}, - Decimal("0.1"), - ) - ] - assert rows == expected + assert list(df.columns) == selection.names + assert [tuple(row) for _, row in df.iterrows()] == selection.rows(PYTHON) def test_as_pandas_integer_na_values(cursor): diff --git a/tests/pyathena/s3fs/test_async_cursor.py b/tests/pyathena/s3fs/test_async_cursor.py index ffc96dfd..c3d7f372 100644 --- a/tests/pyathena/s3fs/test_async_cursor.py +++ b/tests/pyathena/s3fs/test_async_cursor.py @@ -8,8 +8,6 @@ import contextlib import random import time -from datetime import datetime -from decimal import Decimal import pytest @@ -19,6 +17,8 @@ from pyathena.s3fs.result_set import AthenaS3FSResultSet from tests import ENV from tests.pyathena.conftest import connect +from tests.pyathena.expected import ARRAY_JSON, MAP_JSON, PYTHON, TIME_OF_TIMESTAMP, Selection +from tests.pyathena.tables import ONE_ROW_COMPLEX class TestAsyncS3FSCursor: @@ -61,76 +61,11 @@ def test_invalid_arraysize(self, async_s3fs_cursor): async_s3fs_cursor.arraysize = -1 def test_complex(self, async_s3fs_cursor): - query_id, future = async_s3fs_cursor.execute( - """ - SELECT - col_boolean - ,col_tinyint - ,col_smallint - ,col_int - ,col_bigint - ,col_float - ,col_double - ,col_string - ,col_varchar - ,col_timestamp - ,CAST(col_timestamp AS time) AS col_time - ,col_date - ,col_binary - ,col_array - ,CAST(col_array AS json) AS col_array_json - ,col_map - ,CAST(col_map AS json) AS col_map_json - ,col_struct - ,col_decimal - FROM one_row_complex - """ - ) + selection = Selection(ONE_ROW_COMPLEX, (TIME_OF_TIMESTAMP, ARRAY_JSON, MAP_JSON)) + query_id, future = async_s3fs_cursor.execute(selection.sql) result_set = future.result() - assert result_set.description == [ - ("col_boolean", "boolean", None, None, 0, 0, "UNKNOWN"), - ("col_tinyint", "tinyint", None, None, 3, 0, "UNKNOWN"), - ("col_smallint", "smallint", None, None, 5, 0, "UNKNOWN"), - ("col_int", "integer", None, None, 10, 0, "UNKNOWN"), - ("col_bigint", "bigint", None, None, 19, 0, "UNKNOWN"), - ("col_float", "float", None, None, 17, 0, "UNKNOWN"), - ("col_double", "double", None, None, 17, 0, "UNKNOWN"), - ("col_string", "varchar", None, None, 2147483647, 0, "UNKNOWN"), - ("col_varchar", "varchar", None, None, 10, 0, "UNKNOWN"), - ("col_timestamp", "timestamp", None, None, 3, 0, "UNKNOWN"), - ("col_time", "time", None, None, 3, 0, "UNKNOWN"), - ("col_date", "date", None, None, 0, 0, "UNKNOWN"), - ("col_binary", "varbinary", None, None, 1073741824, 0, "UNKNOWN"), - ("col_array", "array", None, None, 0, 0, "UNKNOWN"), - ("col_array_json", "json", None, None, 0, 0, "UNKNOWN"), - ("col_map", "map", None, None, 0, 0, "UNKNOWN"), - ("col_map_json", "json", None, None, 0, 0, "UNKNOWN"), - ("col_struct", "row", None, None, 0, 0, "UNKNOWN"), - ("col_decimal", "decimal", None, None, 10, 1, "UNKNOWN"), - ] - assert result_set.fetchall() == [ - ( - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - "varchar", - datetime(2017, 1, 1, 0, 0, 0), - datetime(2017, 1, 1, 0, 0, 0).time(), - datetime(2017, 1, 2).date(), - b"123", - [1, 2], - [1, 2], - {"1": "2", "3": "4"}, - {"1": 2, "3": 4}, - {"a": "1", "b": "2"}, - Decimal("0.1"), - ) - ] + assert result_set.description == selection.description() + assert result_set.fetchall() == selection.rows(PYTHON) def test_cancel(self, async_s3fs_cursor): query_id, future = async_s3fs_cursor.execute( diff --git a/tests/pyathena/s3fs/test_cursor.py b/tests/pyathena/s3fs/test_cursor.py index 40edaa07..2dec95c3 100644 --- a/tests/pyathena/s3fs/test_cursor.py +++ b/tests/pyathena/s3fs/test_cursor.py @@ -3,8 +3,6 @@ import string import time from concurrent.futures import ThreadPoolExecutor -from datetime import datetime -from decimal import Decimal import pytest @@ -14,6 +12,8 @@ from pyathena.s3fs.result_set import AthenaS3FSResultSet from tests import ENV from tests.pyathena.conftest import connect +from tests.pyathena.expected import ARRAY_JSON, MAP_JSON, PYTHON, TIME_OF_TIMESTAMP, Selection +from tests.pyathena.tables import ONE_ROW_COMPLEX class TestS3FSCursor: @@ -55,75 +55,10 @@ def test_invalid_arraysize(self, s3fs_cursor): s3fs_cursor.arraysize = -1 def test_complex(self, s3fs_cursor): - s3fs_cursor.execute( - """ - SELECT - col_boolean - ,col_tinyint - ,col_smallint - ,col_int - ,col_bigint - ,col_float - ,col_double - ,col_string - ,col_varchar - ,col_timestamp - ,CAST(col_timestamp AS time) AS col_time - ,col_date - ,col_binary - ,col_array - ,CAST(col_array AS json) AS col_array_json - ,col_map - ,CAST(col_map AS json) AS col_map_json - ,col_struct - ,col_decimal - FROM one_row_complex - """ - ) - assert s3fs_cursor.description == [ - ("col_boolean", "boolean", None, None, 0, 0, "UNKNOWN"), - ("col_tinyint", "tinyint", None, None, 3, 0, "UNKNOWN"), - ("col_smallint", "smallint", None, None, 5, 0, "UNKNOWN"), - ("col_int", "integer", None, None, 10, 0, "UNKNOWN"), - ("col_bigint", "bigint", None, None, 19, 0, "UNKNOWN"), - ("col_float", "float", None, None, 17, 0, "UNKNOWN"), - ("col_double", "double", None, None, 17, 0, "UNKNOWN"), - ("col_string", "varchar", None, None, 2147483647, 0, "UNKNOWN"), - ("col_varchar", "varchar", None, None, 10, 0, "UNKNOWN"), - ("col_timestamp", "timestamp", None, None, 3, 0, "UNKNOWN"), - ("col_time", "time", None, None, 3, 0, "UNKNOWN"), - ("col_date", "date", None, None, 0, 0, "UNKNOWN"), - ("col_binary", "varbinary", None, None, 1073741824, 0, "UNKNOWN"), - ("col_array", "array", None, None, 0, 0, "UNKNOWN"), - ("col_array_json", "json", None, None, 0, 0, "UNKNOWN"), - ("col_map", "map", None, None, 0, 0, "UNKNOWN"), - ("col_map_json", "json", None, None, 0, 0, "UNKNOWN"), - ("col_struct", "row", None, None, 0, 0, "UNKNOWN"), - ("col_decimal", "decimal", None, None, 10, 1, "UNKNOWN"), - ] - assert s3fs_cursor.fetchall() == [ - ( - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - "varchar", - datetime(2017, 1, 1, 0, 0, 0), - datetime(2017, 1, 1, 0, 0, 0).time(), - datetime(2017, 1, 2).date(), - b"123", - [1, 2], - [1, 2], - {"1": "2", "3": "4"}, - {"1": 2, "3": 4}, - {"a": "1", "b": "2"}, - Decimal("0.1"), - ) - ] + selection = Selection(ONE_ROW_COMPLEX, (TIME_OF_TIMESTAMP, ARRAY_JSON, MAP_JSON)) + s3fs_cursor.execute(selection.sql) + assert s3fs_cursor.description == selection.description() + assert s3fs_cursor.fetchall() == selection.rows(PYTHON) def test_cancel(self, s3fs_cursor): def cancel(c): diff --git a/tests/pyathena/sqlalchemy/test_base.py b/tests/pyathena/sqlalchemy/test_base.py index e086e7d6..0ca2d81f 100644 --- a/tests/pyathena/sqlalchemy/test_base.py +++ b/tests/pyathena/sqlalchemy/test_base.py @@ -34,6 +34,8 @@ ) from pyathena.util import RetryConfig from tests.pyathena.conftest import ENV +from tests.pyathena.expected import PYTHON, Selection, assert_sqlalchemy_type +from tests.pyathena.tables import ONE_ROW_COMPLEX from tests.pyathena.util import decorated, throttle_metadata_api # Amazon S3 Tables tests need a pre-provisioned table-bucket catalog; the session @@ -1394,53 +1396,12 @@ def test_filter_func(self, engine): def test_reflect_select(self, engine): engine, conn = engine one_row_complex = Table("one_row_complex", MetaData(schema=ENV.schema), autoload_with=conn) - assert len(one_row_complex.c) == 16 + assert [c.name for c in one_row_complex.c] == [c.name for c in ONE_ROW_COMPLEX.columns] assert isinstance(one_row_complex.c.col_string, Column) rows = conn.execute(one_row_complex.select()).fetchall() - assert len(rows) == 1 - assert list(rows[0]) == [ - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - "varchar", - datetime(2017, 1, 1, 0, 0, 0), - date(2017, 1, 2), - b"123", - [1, 2], - {"1": "2", "3": "4"}, # map type now converted to dict - {"a": "1", "b": "2"}, # row type now converted to dict - Decimal("0.1"), - ] - assert isinstance(one_row_complex.c.col_boolean.type, types.BOOLEAN) - assert isinstance(one_row_complex.c.col_tinyint.type, TINYINT) - assert isinstance(one_row_complex.c.col_smallint.type, types.SMALLINT) - assert isinstance(one_row_complex.c.col_int.type, types.INTEGER) - assert isinstance(one_row_complex.c.col_bigint.type, types.BIGINT) - assert isinstance(one_row_complex.c.col_float.type, types.FLOAT) - assert isinstance(one_row_complex.c.col_double.type, types.DOUBLE) - assert isinstance(one_row_complex.c.col_string.type, types.String) - assert isinstance(one_row_complex.c.col_varchar.type, types.VARCHAR) - assert one_row_complex.c.col_varchar.type.length == 10 - assert isinstance(one_row_complex.c.col_timestamp.type, types.TIMESTAMP) - assert isinstance(one_row_complex.c.col_date.type, types.DATE) - assert isinstance(one_row_complex.c.col_binary.type, types.BINARY) - assert isinstance(one_row_complex.c.col_array.type, AthenaArray) - assert isinstance(one_row_complex.c.col_array.type.item_type, types.INTEGER) - assert isinstance(one_row_complex.c.col_map.type, types.String) - # With struct support, col_struct should now be recognized as AthenaStruct - - assert isinstance(one_row_complex.c.col_struct.type, AthenaStruct) - assert isinstance( - one_row_complex.c.col_decimal.type, - types.DECIMAL, - ) - assert one_row_complex.c.col_decimal.type.precision == 10 - assert one_row_complex.c.col_decimal.type.scale == 1 + assert [tuple(row) for row in rows] == Selection(ONE_ROW_COMPLEX).rows(PYTHON) + for column in ONE_ROW_COMPLEX.columns: + assert_sqlalchemy_type(one_row_complex.c[column.name].type, column) def test_select_offset_limit(self, engine): engine, conn = engine diff --git a/tests/pyathena/tables.py b/tests/pyathena/tables.py index bc8a54fc..6d728bfc 100644 --- a/tests/pyathena/tables.py +++ b/tests/pyathena/tables.py @@ -156,6 +156,57 @@ def _to_text(value: Any) -> str: return str(value) +# One row with a value of each column type. Tests derive their expected results +# from this definition with tests.pyathena.expected. Keep string values free of +# the separators in Athena's text rendering of arrays, maps, and structs +# (, = [ ] { }), because the cursors parse that rendering. +ONE_ROW_COMPLEX = Table( + "one_row_complex", + ( + Column("col_boolean", "BOOLEAN", pa.bool_()), + Column("col_tinyint", "TINYINT", pa.int8()), + Column("col_smallint", "SMALLINT", pa.int16()), + Column("col_int", "INT", pa.int32()), + Column("col_bigint", "BIGINT", pa.int64()), + Column("col_float", "FLOAT", pa.float32()), + Column("col_double", "DOUBLE", pa.float64()), + Column("col_string", "STRING", pa.string()), + Column("col_varchar", "VARCHAR(10)", pa.string()), + Column("col_timestamp", "TIMESTAMP", pa.timestamp("ms")), + Column("col_date", "DATE", pa.date32()), + Column("col_binary", "BINARY", pa.binary()), + Column("col_array", "ARRAY", pa.list_(pa.int32())), + Column("col_map", "MAP", pa.map_(pa.int32(), pa.int32())), + Column( + "col_struct", + "STRUCT", + pa.struct([("a", pa.int32()), ("b", pa.int32())]), + ), + Column("col_decimal", "DECIMAL(10,1)", pa.decimal128(10, 1)), + ), + rows=( + ( + True, + 127, + 32767, + 2147483647, + 9223372036854775807, + 0.5, + 0.25, + "a string", + "varchar", + datetime(2017, 1, 1, 0, 0, 0), + date(2017, 1, 2), + b"123", + [1, 2], + [(1, 2), (3, 4)], + {"a": 1, "b": 2}, + Decimal("0.1"), + ), + ), +) + + TABLES = ( # A text table: the reflection tests assert its SerDe and delimiters. Table( @@ -173,51 +224,7 @@ def _to_text(value: Any) -> str: rows=tuple((i,) for i in range(10000)), storage="text", ), - Table( - "one_row_complex", - ( - Column("col_boolean", "BOOLEAN", pa.bool_()), - Column("col_tinyint", "TINYINT", pa.int8()), - Column("col_smallint", "SMALLINT", pa.int16()), - Column("col_int", "INT", pa.int32()), - Column("col_bigint", "BIGINT", pa.int64()), - Column("col_float", "FLOAT", pa.float32()), - Column("col_double", "DOUBLE", pa.float64()), - Column("col_string", "STRING", pa.string()), - Column("col_varchar", "VARCHAR(10)", pa.string()), - Column("col_timestamp", "TIMESTAMP", pa.timestamp("ms")), - Column("col_date", "DATE", pa.date32()), - Column("col_binary", "BINARY", pa.binary()), - Column("col_array", "ARRAY", pa.list_(pa.int32())), - Column("col_map", "MAP", pa.map_(pa.int32(), pa.int32())), - Column( - "col_struct", - "STRUCT", - pa.struct([("a", pa.int32()), ("b", pa.int32())]), - ), - Column("col_decimal", "DECIMAL(10,1)", pa.decimal128(10, 1)), - ), - rows=( - ( - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - "varchar", - datetime(2017, 1, 1, 0, 0, 0), - date(2017, 1, 2), - b"123", - [1, 2], - [(1, 2), (3, 4)], - {"a": 1, "b": 2}, - Decimal("0.1"), - ), - ), - ), + ONE_ROW_COMPLEX, Table( "partition_table", (Column("a", "STRING", pa.string()),), diff --git a/tests/pyathena/test_cursor.py b/tests/pyathena/test_cursor.py index aa1832b3..915be028 100644 --- a/tests/pyathena/test_cursor.py +++ b/tests/pyathena/test_cursor.py @@ -9,7 +9,7 @@ import uuid from concurrent import futures from concurrent.futures.thread import ThreadPoolExecutor -from datetime import date, datetime, timezone +from datetime import datetime, timezone from decimal import Decimal from random import randint from unittest.mock import MagicMock, patch @@ -19,13 +19,6 @@ from pyathena import ( BINARY, - BOOLEAN, - DATE, - DATETIME, - JSON, - NUMBER, - STRING, - TIME, Binary, ExecuteOptions, ) @@ -36,6 +29,15 @@ from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect +from tests.pyathena.expected import ( + ARRAY_JSON, + MAP_JSON, + PYTHON, + TIME_OF_TIMESTAMP, + TIMESTAMP_TZ, + Selection, +) +from tests.pyathena.tables import ONE_ROW_COMPLEX from tests.pyathena.util import throttle_metadata_api, unreachable_glue _logger = logging.getLogger(__name__) @@ -600,105 +602,18 @@ def test_query_execution_initial(self, cursor): assert cursor.effective_engine_version is None def test_complex(self, cursor): - cursor.execute( - """ - SELECT - col_boolean - ,col_tinyint - ,col_smallint - ,col_int - ,col_bigint - ,col_float - ,col_double - ,col_string - ,col_varchar - ,col_timestamp - ,CAST(col_timestamp AS timestamp with time zone) AS col_timestamp_tz - ,CAST(col_timestamp AS time) AS col_time - ,col_date - ,col_binary - ,col_array - ,CAST(col_array AS json) AS col_array_json - ,col_map - ,CAST(col_map AS json) AS col_map_json - ,col_struct - ,col_decimal - FROM one_row_complex - """ + selection = Selection( + ONE_ROW_COMPLEX, (TIMESTAMP_TZ, TIME_OF_TIMESTAMP, ARRAY_JSON, MAP_JSON) ) - assert cursor.description == [ - ("col_boolean", "boolean", None, None, 0, 0, "UNKNOWN"), - ("col_tinyint", "tinyint", None, None, 3, 0, "UNKNOWN"), - ("col_smallint", "smallint", None, None, 5, 0, "UNKNOWN"), - ("col_int", "integer", None, None, 10, 0, "UNKNOWN"), - ("col_bigint", "bigint", None, None, 19, 0, "UNKNOWN"), - ("col_float", "float", None, None, 17, 0, "UNKNOWN"), - ("col_double", "double", None, None, 17, 0, "UNKNOWN"), - ("col_string", "varchar", None, None, 2147483647, 0, "UNKNOWN"), - ("col_varchar", "varchar", None, None, 10, 0, "UNKNOWN"), - ("col_timestamp", "timestamp", None, None, 3, 0, "UNKNOWN"), - ("col_timestamp_tz", "timestamp with time zone", None, None, 3, 0, "UNKNOWN"), - ("col_time", "time", None, None, 3, 0, "UNKNOWN"), - ("col_date", "date", None, None, 0, 0, "UNKNOWN"), - ("col_binary", "varbinary", None, None, 1073741824, 0, "UNKNOWN"), - ("col_array", "array", None, None, 0, 0, "UNKNOWN"), - ("col_array_json", "json", None, None, 0, 0, "UNKNOWN"), - ("col_map", "map", None, None, 0, 0, "UNKNOWN"), - ("col_map_json", "json", None, None, 0, 0, "UNKNOWN"), - ("col_struct", "row", None, None, 0, 0, "UNKNOWN"), - ("col_decimal", "decimal", None, None, 10, 1, "UNKNOWN"), - ] + cursor.execute(selection.sql) + assert cursor.description == selection.description() rows = cursor.fetchall() - expected = [ - ( - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - "varchar", - datetime(2017, 1, 1, 0, 0, 0), - datetime(2017, 1, 1, 0, 0, 0, tzinfo=timezone.utc), - datetime(2017, 1, 1, 0, 0, 0).time(), - date(2017, 1, 2), - b"123", - [1, 2], - [1, 2], - {"1": "2", "3": "4"}, - {"1": 2, "3": 4}, - {"a": "1", "b": "2"}, - Decimal("0.1"), - ) - ] + expected = selection.rows(PYTHON) assert rows == expected # catch unicode/str assert list(map(type, rows[0])) == list(map(type, expected[0])) # compare dbapi type object - assert [d[1] for d in cursor.description] == [ - BOOLEAN, - NUMBER, - NUMBER, - NUMBER, - NUMBER, - NUMBER, - NUMBER, - STRING, - STRING, - DATETIME, - DATETIME, - TIME, - DATE, - BINARY, - STRING, - JSON, - STRING, - JSON, - STRING, - NUMBER, - ] + assert [d[1] for d in cursor.description] == selection.dbapi_types() def test_complex_with_type_hints(self, cursor): # 1. Basic complex columns from one_row_complex From db7a7928443d6c8091e7dc4c755e8054f4e9132e Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Mon, 28 Sep 2026 19:14:53 +0900 Subject: [PATCH 02/14] Document the expectation rules and fold the JSON helper Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/expected.py | 33 +++++++++++++++++++-------------- 1 file changed, 19 insertions(+), 14 deletions(-) diff --git a/tests/pyathena/expected.py b/tests/pyathena/expected.py index 6df955a0..d9bf62f5 100644 --- a/tests/pyathena/expected.py +++ b/tests/pyathena/expected.py @@ -160,39 +160,44 @@ def sql(self) -> str: } -def _json_compatible(value: Any, athena_type: str) -> Any: - """Return a value as the JSON structure Athena renders for ``CAST(... AS json)``. +def _parsed_json(value: Any, athena_type: str) -> Any: + """Return a value as Athena renders it with ``CAST(... AS json)``, after JSON parsing. Args: value: The value from the table definition. athena_type: The value's Athena type. Returns: - The value, with a map's key-value pairs as a dict. + The parsed JSON value; a map becomes a dict with string keys. """ if family(athena_type) == "map": - return dict(value) - return value + value = dict(value) + return json.loads(json.dumps(value)) -def _parsed_json(value: Any, athena_type: str) -> Any: - """Return a value after a round trip through JSON text. +def _same(value: Any, athena_type: str) -> Any: + """Return the value unchanged. Args: value: The value from the table definition. athena_type: The value's Athena type. Returns: - The parsed JSON value; map keys become strings. + The value. """ - return json.loads(json.dumps(_json_compatible(value, athena_type))) - - -def _same(value: Any, athena_type: str) -> Any: return value def _as_list(value: Any, athena_type: str) -> Any: + """Return an array value as a list. + + Args: + value: The value from the table definition. + athena_type: The value's Athena type. + + Returns: + The elements as a list. + """ return list(value) @@ -220,8 +225,8 @@ def _as_list(value: Any, athena_type: str) -> Any: _same, ) -# Cursor, DictCursor, S3FSCursor, pyathena.pandas.util.as_pandas, and SQLAlchemy -# result rows. Without type hints, map and struct values are strings. +# Rows of Cursor, S3FSCursor, pyathena.pandas.util.as_pandas, and SQLAlchemy. +# Without type hints, map and struct values are strings. PYTHON: Representation = { **_SCALARS, "timestamp with time zone": lambda v, t: v.replace(tzinfo=timezone.utc), From 9b8cfac91a58e2812a0beba9184825b7af7457c0 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Mon, 28 Sep 2026 19:20:37 +0900 Subject: [PATCH 03/14] Derive nested values from Athena's text rendering Without type hints, the cursors parse Athena's text rendering of arrays, maps, and structs, so the expected nested values now come from that rendering: an array that is valid JSON is parsed, and other nested values are their text, such as true for a boolean. Type parameters are read from the outer type only, so ARRAY has none. Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/expected.py | 149 ++++++++++++++++++++++++++++++++----- 1 file changed, 129 insertions(+), 20 deletions(-) diff --git a/tests/pyathena/expected.py b/tests/pyathena/expected.py index d9bf62f5..b633c2b9 100644 --- a/tests/pyathena/expected.py +++ b/tests/pyathena/expected.py @@ -46,22 +46,52 @@ def _parameters(athena_type: str) -> tuple[int, ...]: athena_type: An Athena type, such as ``DECIMAL(10,1)``. Returns: - The parameters, such as ``(10, 1)``; empty for a type without them. + The parameters of the outer type, such as ``(10, 1)``; empty for a type + without them, such as ``ARRAY``. """ - match = re.search(r"\(([\d,\s]+)\)", athena_type) + match = re.match(r"\s*[a-zA-Z ]+\(([\d,\s]+)\)", athena_type) return tuple(int(p) for p in match.group(1).split(",")) if match else () -def _element_type(athena_type: str) -> str: - """Return the element type of an array type. +def _type_arguments(athena_type: str) -> list[str]: + """Return the type arguments of a complex type. Args: - athena_type: An array type, such as ``ARRAY``. + athena_type: An Athena type, such as ``MAP``. Returns: - The element type, such as ``int``. + The element type of an array, the key and value types of a map, or the + ``name: type`` fields of a struct; empty for other types. """ - return athena_type[athena_type.index("<") + 1 : athena_type.rindex(">")] + if "<" not in athena_type: + return [] + inner = athena_type[athena_type.index("<") + 1 : athena_type.rindex(">")] + arguments, depth, start = [], 0, 0 + for i, char in enumerate(inner): + if char in "<(": + depth += 1 + elif char in ">)": + depth -= 1 + elif char == "," and depth == 0: + arguments.append(inner[start:i].strip()) + start = i + 1 + arguments.append(inner[start:].strip()) + return arguments + + +def _struct_fields(athena_type: str) -> list[tuple[str, str]]: + """Return the fields of a struct type. + + Args: + athena_type: A struct type, such as ``STRUCT``. + + Returns: + ``(name, type)`` pairs. + """ + return [ + (name.strip(), field_type.strip()) + for name, field_type in (a.split(":", 1) for a in _type_arguments(athena_type)) + ] @dataclass(frozen=True) @@ -175,30 +205,108 @@ def _parsed_json(value: Any, athena_type: str) -> Any: return json.loads(json.dumps(value)) -def _same(value: Any, athena_type: str) -> Any: - """Return the value unchanged. +def _athena_text(value: Any, athena_type: str) -> str: + """Return a value as Athena renders it in a CSV result. Args: value: The value from the table definition. athena_type: The value's Athena type. Returns: - The value. + The text, such as ``[1, 2]`` for an array, ``{1=2, 3=4}`` for a map, + and ``{a=1, b=2}`` for a struct. """ - return value + name = family(athena_type) + arguments = _type_arguments(athena_type) + if value is None: + return "null" + if name == "array": + return f"[{', '.join(_athena_text(e, arguments[0]) for e in value)}]" + if name == "map": + key_type, value_type = arguments + entries = (f"{_athena_text(k, key_type)}={_athena_text(x, value_type)}" for k, x in value) + return f"{{{', '.join(entries)}}}" + if name == "struct": + field_types = dict(_struct_fields(athena_type)) + entries = (f"{k}={_athena_text(x, field_types[k])}" for k, x in value.items()) + return f"{{{', '.join(entries)}}}" + if name == "boolean": + return str(value).lower() + if name == "timestamp": + return value.isoformat(sep=" ", timespec="milliseconds") + return str(value) + + +def _element_text(value: Any, athena_type: str) -> str | None: + """Return a nested value as a cursor without type hints returns it. + + Args: + value: The nested value. + athena_type: The value's Athena type. + + Returns: + Athena's text rendering of the value, or None for a null. + """ + return None if value is None else _athena_text(value, athena_type) -def _as_list(value: Any, athena_type: str) -> Any: - """Return an array value as a list. +def _python_array(value: Any, athena_type: str) -> list[Any]: + """Return an array as a cursor without type hints returns it. + + Args: + value: The value from the table definition. + athena_type: The array type. + + Returns: + The parsed JSON if Athena's rendering is valid JSON, such as ``[1, 2]``; + otherwise the elements' text renderings. + """ + try: + return json.loads(_athena_text(value, athena_type)) + except ValueError: + (element_type,) = _type_arguments(athena_type) + return [_element_text(e, element_type) for e in value] + + +def _python_map(value: Any, athena_type: str) -> dict[str, Any]: + """Return a map as a cursor without type hints returns it. + + Args: + value: The value from the table definition. + athena_type: The map type. + + Returns: + The keys and values as their text renderings. + """ + key_type, value_type = _type_arguments(athena_type) + return {_athena_text(k, key_type): _element_text(x, value_type) for k, x in value} + + +def _python_struct(value: Any, athena_type: str) -> dict[str, Any]: + """Return a struct as a cursor without type hints returns it. + + Args: + value: The value from the table definition. + athena_type: The struct type. + + Returns: + The field values as their text renderings. + """ + field_types = dict(_struct_fields(athena_type)) + return {k: _element_text(x, field_types[k]) for k, x in value.items()} + + +def _same(value: Any, athena_type: str) -> Any: + """Return the value unchanged. Args: value: The value from the table definition. athena_type: The value's Athena type. Returns: - The elements as a list. + The value. """ - return list(value) + return value # A representation maps a type family to a rule(value, athena_type) that @@ -226,14 +334,15 @@ def _as_list(value: Any, athena_type: str) -> Any: ) # Rows of Cursor, S3FSCursor, pyathena.pandas.util.as_pandas, and SQLAlchemy. -# Without type hints, map and struct values are strings. +# Without type hints, the cursors parse Athena's text rendering of arrays, maps, +# and structs, so nested values that are not JSON are strings. PYTHON: Representation = { **_SCALARS, "timestamp with time zone": lambda v, t: v.replace(tzinfo=timezone.utc), "time": lambda v, t: v.time(), - "array": _as_list, - "map": lambda v, t: {str(k): str(x) for k, x in v}, - "struct": lambda v, t: {k: str(x) for k, x in v.items()}, + "array": _python_array, + "map": _python_map, + "struct": _python_struct, "json": _parsed_json, } @@ -335,5 +444,5 @@ def assert_sqlalchemy_type(sqlalchemy_type: Any, column: Column) -> None: precision, scale = _parameters(column.athena_type) assert (sqlalchemy_type.precision, sqlalchemy_type.scale) == (precision, scale) elif name == "array": - element = _element_type(column.athena_type) + (element,) = _type_arguments(column.athena_type) assert isinstance(sqlalchemy_type.item_type, _SQLALCHEMY_TYPES[family(element)]) From fdad71a78da5843c3cec6c84ae4550aa33fc2d66 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Mon, 28 Sep 2026 19:25:07 +0900 Subject: [PATCH 04/14] Parse map and struct elements of arrays in the expected values When an array's rendering is not JSON, the cursors parse its map and struct elements into dicts and leave nested arrays unparsed; the expected values now follow that. Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/expected.py | 23 +++++++++++++++++------ 1 file changed, 17 insertions(+), 6 deletions(-) diff --git a/tests/pyathena/expected.py b/tests/pyathena/expected.py index b633c2b9..7b519d00 100644 --- a/tests/pyathena/expected.py +++ b/tests/pyathena/expected.py @@ -250,7 +250,7 @@ def _element_text(value: Any, athena_type: str) -> str | None: return None if value is None else _athena_text(value, athena_type) -def _python_array(value: Any, athena_type: str) -> list[Any]: +def _python_array(value: Any, athena_type: str) -> Any: """Return an array as a cursor without type hints returns it. Args: @@ -258,14 +258,25 @@ def _python_array(value: Any, athena_type: str) -> list[Any]: athena_type: The array type. Returns: - The parsed JSON if Athena's rendering is valid JSON, such as ``[1, 2]``; - otherwise the elements' text renderings. + The parsed JSON if Athena's rendering is valid JSON, such as ``[1, 2]``. + Otherwise a list of the elements, where a map or struct element is a + dict and any other element is its text rendering; or the rendering + itself for nested arrays, which the cursors leave unparsed. """ + text = _athena_text(value, athena_type) try: - return json.loads(_athena_text(value, athena_type)) + return json.loads(text) except ValueError: - (element_type,) = _type_arguments(athena_type) - return [_element_text(e, element_type) for e in value] + pass + (element_type,) = _type_arguments(athena_type) + element_family = family(element_type) + if element_family == "array": + return text + if element_family == "map": + return [None if e is None else _python_map(e, element_type) for e in value] + if element_family == "struct": + return [None if e is None else _python_struct(e, element_type) for e in value] + return [_element_text(e, element_type) for e in value] def _python_map(value: Any, athena_type: str) -> dict[str, Any]: From cea2dd1df7acef919f9571d57ce9e236967c6b82 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Mon, 28 Sep 2026 19:29:40 +0900 Subject: [PATCH 05/14] Limit the nested-value rules to the shapes they model The rules for values nested in arrays, maps, and structs cover scalars and arrays of maps or structs of scalars. Deeper nesting, which the cursors parse with further heuristics, now raises NotImplementedError instead of producing an expectation the cursors would not meet. Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/expected.py | 51 ++++++++++++++++++++++++++------------ 1 file changed, 35 insertions(+), 16 deletions(-) diff --git a/tests/pyathena/expected.py b/tests/pyathena/expected.py index 7b519d00..c9280a45 100644 --- a/tests/pyathena/expected.py +++ b/tests/pyathena/expected.py @@ -237,22 +237,46 @@ def _athena_text(value: Any, athena_type: str) -> str: return str(value) +_COMPLEX_FAMILIES = ("array", "map", "struct") + + +def _scalar_type(athena_type: str) -> str: + """Return a nested type after checking that the rules support it. + + Args: + athena_type: The type of a value nested in an array, map, or struct. + + Returns: + The type. + + Raises: + NotImplementedError: If the type is an array, map, or struct. The cursors + parse deeper nesting with further rules; add a rule for it here + together with the column that needs it. + """ + if family(athena_type) in _COMPLEX_FAMILIES: + raise NotImplementedError(f"No expectation rule for nested {athena_type}.") + return athena_type + + def _element_text(value: Any, athena_type: str) -> str | None: - """Return a nested value as a cursor without type hints returns it. + """Return a nested scalar value as a cursor without type hints returns it. Args: value: The nested value. - athena_type: The value's Athena type. + athena_type: The value's Athena type, a scalar type. Returns: Athena's text rendering of the value, or None for a null. """ - return None if value is None else _athena_text(value, athena_type) + return None if value is None else _athena_text(value, _scalar_type(athena_type)) def _python_array(value: Any, athena_type: str) -> Any: """Return an array as a cursor without type hints returns it. + Supports arrays of scalars and arrays of maps or structs of scalars. + Args: value: The value from the table definition. athena_type: The array type. @@ -260,27 +284,22 @@ def _python_array(value: Any, athena_type: str) -> Any: Returns: The parsed JSON if Athena's rendering is valid JSON, such as ``[1, 2]``. Otherwise a list of the elements, where a map or struct element is a - dict and any other element is its text rendering; or the rendering - itself for nested arrays, which the cursors leave unparsed. + dict and any other element is its text rendering. """ - text = _athena_text(value, athena_type) - try: - return json.loads(text) - except ValueError: - pass (element_type,) = _type_arguments(athena_type) element_family = family(element_type) - if element_family == "array": - return text if element_family == "map": return [None if e is None else _python_map(e, element_type) for e in value] if element_family == "struct": return [None if e is None else _python_struct(e, element_type) for e in value] - return [_element_text(e, element_type) for e in value] + try: + return json.loads(_athena_text(value, athena_type)) + except ValueError: + return [_element_text(e, element_type) for e in value] def _python_map(value: Any, athena_type: str) -> dict[str, Any]: - """Return a map as a cursor without type hints returns it. + """Return a map of scalars as a cursor without type hints returns it. Args: value: The value from the table definition. @@ -290,11 +309,11 @@ def _python_map(value: Any, athena_type: str) -> dict[str, Any]: The keys and values as their text renderings. """ key_type, value_type = _type_arguments(athena_type) - return {_athena_text(k, key_type): _element_text(x, value_type) for k, x in value} + return {_athena_text(k, _scalar_type(key_type)): _element_text(x, value_type) for k, x in value} def _python_struct(value: Any, athena_type: str) -> dict[str, Any]: - """Return a struct as a cursor without type hints returns it. + """Return a struct of scalars as a cursor without type hints returns it. Args: value: The value from the table definition. From 562f7b54697510b7133a58ca84ba77fa20a33827 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Mon, 28 Sep 2026 19:33:17 +0900 Subject: [PATCH 06/14] Check the nesting the rules model by type, and keep nulls The nested-value rules check the whole type before looking at values, so an unsupported type raises even for a null or empty value, and a null array, map, or struct stays None. The one_row_complex note adds that string values must differ from null, which Athena renders the same way. Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/expected.py | 74 +++++++++++++++++++++++++------------- tests/pyathena/tables.py | 3 +- 2 files changed, 51 insertions(+), 26 deletions(-) diff --git a/tests/pyathena/expected.py b/tests/pyathena/expected.py index c9280a45..5c64d9c4 100644 --- a/tests/pyathena/expected.py +++ b/tests/pyathena/expected.py @@ -240,23 +240,39 @@ def _athena_text(value: Any, athena_type: str) -> str: _COMPLEX_FAMILIES = ("array", "map", "struct") -def _scalar_type(athena_type: str) -> str: - """Return a nested type after checking that the rules support it. +def _member_types(athena_type: str) -> list[str]: + """Return the types nested directly in a complex type. Args: - athena_type: The type of a value nested in an array, map, or struct. + athena_type: An array, map, or struct type. Returns: - The type. + The element type, the key and value types, or the field types. + """ + if family(athena_type) == "struct": + return [field_type for _, field_type in _struct_fields(athena_type)] + return _type_arguments(athena_type) + + +def _check_nesting(athena_type: str) -> None: + """Check that the rules model the nesting of a complex type. + + They model scalars nested in an array, map, or struct, and arrays of maps + or structs of scalars. + + Args: + athena_type: An array, map, or struct type. Raises: - NotImplementedError: If the type is an array, map, or struct. The cursors - parse deeper nesting with further rules; add a rule for it here - together with the column that needs it. + NotImplementedError: For deeper nesting, which the cursors parse with + further rules; add a rule for it here together with the column that + needs it. """ - if family(athena_type) in _COMPLEX_FAMILIES: - raise NotImplementedError(f"No expectation rule for nested {athena_type}.") - return athena_type + members = _member_types(athena_type) + if family(athena_type) == "array" and family(members[0]) in ("map", "struct"): + members = _member_types(members[0]) + if any(family(m) in _COMPLEX_FAMILIES for m in members): + raise NotImplementedError(f"No expectation rule for the nesting in {athena_type}.") def _element_text(value: Any, athena_type: str) -> str | None: @@ -269,59 +285,67 @@ def _element_text(value: Any, athena_type: str) -> str | None: Returns: Athena's text rendering of the value, or None for a null. """ - return None if value is None else _athena_text(value, _scalar_type(athena_type)) + return None if value is None else _athena_text(value, athena_type) def _python_array(value: Any, athena_type: str) -> Any: """Return an array as a cursor without type hints returns it. - Supports arrays of scalars and arrays of maps or structs of scalars. - Args: value: The value from the table definition. athena_type: The array type. Returns: - The parsed JSON if Athena's rendering is valid JSON, such as ``[1, 2]``. - Otherwise a list of the elements, where a map or struct element is a - dict and any other element is its text rendering. + None for a null. Otherwise the parsed JSON if Athena's rendering is + valid JSON, such as ``[1, 2]``, or else a list of the elements, where a + map or struct element is a dict and any other element is its text + rendering. """ + _check_nesting(athena_type) + if value is None: + return None (element_type,) = _type_arguments(athena_type) element_family = family(element_type) if element_family == "map": - return [None if e is None else _python_map(e, element_type) for e in value] + return [_python_map(e, element_type) for e in value] if element_family == "struct": - return [None if e is None else _python_struct(e, element_type) for e in value] + return [_python_struct(e, element_type) for e in value] try: return json.loads(_athena_text(value, athena_type)) except ValueError: return [_element_text(e, element_type) for e in value] -def _python_map(value: Any, athena_type: str) -> dict[str, Any]: - """Return a map of scalars as a cursor without type hints returns it. +def _python_map(value: Any, athena_type: str) -> dict[str, Any] | None: + """Return a map as a cursor without type hints returns it. Args: value: The value from the table definition. athena_type: The map type. Returns: - The keys and values as their text renderings. + None for a null; otherwise the keys and values as their text renderings. """ + _check_nesting(athena_type) + if value is None: + return None key_type, value_type = _type_arguments(athena_type) - return {_athena_text(k, _scalar_type(key_type)): _element_text(x, value_type) for k, x in value} + return {_athena_text(k, key_type): _element_text(x, value_type) for k, x in value} -def _python_struct(value: Any, athena_type: str) -> dict[str, Any]: - """Return a struct of scalars as a cursor without type hints returns it. +def _python_struct(value: Any, athena_type: str) -> dict[str, Any] | None: + """Return a struct as a cursor without type hints returns it. Args: value: The value from the table definition. athena_type: The struct type. Returns: - The field values as their text renderings. + None for a null; otherwise the field values as their text renderings. """ + _check_nesting(athena_type) + if value is None: + return None field_types = dict(_struct_fields(athena_type)) return {k: _element_text(x, field_types[k]) for k, x in value.items()} diff --git a/tests/pyathena/tables.py b/tests/pyathena/tables.py index 6d728bfc..eb4ac9c3 100644 --- a/tests/pyathena/tables.py +++ b/tests/pyathena/tables.py @@ -159,7 +159,8 @@ def _to_text(value: Any) -> str: # One row with a value of each column type. Tests derive their expected results # from this definition with tests.pyathena.expected. Keep string values free of # the separators in Athena's text rendering of arrays, maps, and structs -# (, = [ ] { }), because the cursors parse that rendering. +# (, = [ ] { }) and different from null, because the cursors parse that +# rendering. ONE_ROW_COMPLEX = Table( "one_row_complex", ( From 192839951079cc2f4e60bb9df0bb68418cd593e6 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Mon, 28 Sep 2026 19:36:54 +0900 Subject: [PATCH 07/14] Reject nested strings the cursors would not parse back Strings nested in arrays, maps, and structs, including map keys, may contain only letters, digits, underscores, and spaces and must not be null; the rules raise ValueError for others instead of expecting a value the cursors return differently. Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/expected.py | 15 +++++++++++++-- tests/pyathena/tables.py | 8 ++++---- 2 files changed, 17 insertions(+), 6 deletions(-) diff --git a/tests/pyathena/expected.py b/tests/pyathena/expected.py index 5c64d9c4..1ec16998 100644 --- a/tests/pyathena/expected.py +++ b/tests/pyathena/expected.py @@ -284,8 +284,19 @@ def _element_text(value: Any, athena_type: str) -> str | None: Returns: Athena's text rendering of the value, or None for a null. + + Raises: + ValueError: For a string that the cursors would not parse back as + itself: one with characters other than letters, digits, underscores, + and spaces, or the text null. """ - return None if value is None else _athena_text(value, athena_type) + if value is None: + return None + if family(athena_type) in ("string", "varchar", "char") and ( + not re.fullmatch(r"[\w ]*", value) or value.lower() == "null" + ): + raise ValueError(f"Unsupported nested string value: {value!r}") + return _athena_text(value, athena_type) def _python_array(value: Any, athena_type: str) -> Any: @@ -330,7 +341,7 @@ def _python_map(value: Any, athena_type: str) -> dict[str, Any] | None: if value is None: return None key_type, value_type = _type_arguments(athena_type) - return {_athena_text(k, key_type): _element_text(x, value_type) for k, x in value} + return {_element_text(k, key_type): _element_text(x, value_type) for k, x in value} def _python_struct(value: Any, athena_type: str) -> dict[str, Any] | None: diff --git a/tests/pyathena/tables.py b/tests/pyathena/tables.py index eb4ac9c3..32c30530 100644 --- a/tests/pyathena/tables.py +++ b/tests/pyathena/tables.py @@ -157,10 +157,10 @@ def _to_text(value: Any) -> str: # One row with a value of each column type. Tests derive their expected results -# from this definition with tests.pyathena.expected. Keep string values free of -# the separators in Athena's text rendering of arrays, maps, and structs -# (, = [ ] { }) and different from null, because the cursors parse that -# rendering. +# from this definition with tests.pyathena.expected. Strings nested in arrays, +# maps, and structs may contain only letters, digits, underscores, and spaces, +# and must not be null, because the cursors parse Athena's text rendering of +# those types; tests.pyathena.expected rejects other nested strings. ONE_ROW_COMPLEX = Table( "one_row_complex", ( From 4e205a3a684f85050187b7e0f6794909a749158b Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Mon, 28 Sep 2026 19:40:17 +0900 Subject: [PATCH 08/14] Restrict nested strings to space-separated words Nested strings must be words of letters, digits, and underscores separated by single spaces, which excludes empty strings and surrounding spaces that the cursors strip or skip. Array elements are checked before the JSON parse as well. Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/expected.py | 11 ++++++----- tests/pyathena/tables.py | 7 ++++--- 2 files changed, 10 insertions(+), 8 deletions(-) diff --git a/tests/pyathena/expected.py b/tests/pyathena/expected.py index 1ec16998..db183003 100644 --- a/tests/pyathena/expected.py +++ b/tests/pyathena/expected.py @@ -286,14 +286,14 @@ def _element_text(value: Any, athena_type: str) -> str | None: Athena's text rendering of the value, or None for a null. Raises: - ValueError: For a string that the cursors would not parse back as - itself: one with characters other than letters, digits, underscores, - and spaces, or the text null. + ValueError: For a string that the cursors might not parse back as + itself. Nested strings must be words of letters, digits, and + underscores separated by single spaces, and not the word null. """ if value is None: return None if family(athena_type) in ("string", "varchar", "char") and ( - not re.fullmatch(r"[\w ]*", value) or value.lower() == "null" + not re.fullmatch(r"\w+(?: \w+)*", value) or value.lower() == "null" ): raise ValueError(f"Unsupported nested string value: {value!r}") return _athena_text(value, athena_type) @@ -321,10 +321,11 @@ def _python_array(value: Any, athena_type: str) -> Any: return [_python_map(e, element_type) for e in value] if element_family == "struct": return [_python_struct(e, element_type) for e in value] + elements = [_element_text(e, element_type) for e in value] try: return json.loads(_athena_text(value, athena_type)) except ValueError: - return [_element_text(e, element_type) for e in value] + return elements def _python_map(value: Any, athena_type: str) -> dict[str, Any] | None: diff --git a/tests/pyathena/tables.py b/tests/pyathena/tables.py index 32c30530..a9453c48 100644 --- a/tests/pyathena/tables.py +++ b/tests/pyathena/tables.py @@ -158,9 +158,10 @@ def _to_text(value: Any) -> str: # One row with a value of each column type. Tests derive their expected results # from this definition with tests.pyathena.expected. Strings nested in arrays, -# maps, and structs may contain only letters, digits, underscores, and spaces, -# and must not be null, because the cursors parse Athena's text rendering of -# those types; tests.pyathena.expected rejects other nested strings. +# maps, and structs must be words of letters, digits, and underscores separated +# by single spaces, and not the word null, because the cursors parse Athena's +# text rendering of those types; tests.pyathena.expected rejects other nested +# strings. ONE_ROW_COMPLEX = Table( "one_row_complex", ( From e9efc09f1e40127a13a87afe0e814da6620a4316 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Mon, 28 Sep 2026 19:43:54 +0900 Subject: [PATCH 09/14] Render nested scalars only for the types with a rule Athena's text rendering of nested values is now stated per type: binary as hex bytes and decimal with its declared scale, alongside booleans, integers, strings, dates, and timestamps. Other types, such as float and double, raise NotImplementedError. Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/expected.py | 12 +++++++++++- 1 file changed, 11 insertions(+), 1 deletion(-) diff --git a/tests/pyathena/expected.py b/tests/pyathena/expected.py index db183003..ad387b2e 100644 --- a/tests/pyathena/expected.py +++ b/tests/pyathena/expected.py @@ -215,6 +215,10 @@ def _athena_text(value: Any, athena_type: str) -> str: Returns: The text, such as ``[1, 2]`` for an array, ``{1=2, 3=4}`` for a map, and ``{a=1, b=2}`` for a struct. + + Raises: + NotImplementedError: For a type without a rendering rule here, such as + float and double, whose rendering follows Java's formatting. """ name = family(athena_type) arguments = _type_arguments(athena_type) @@ -232,9 +236,15 @@ def _athena_text(value: Any, athena_type: str) -> str: return f"{{{', '.join(entries)}}}" if name == "boolean": return str(value).lower() + if name in ("tinyint", "smallint", "int", "bigint", "string", "varchar", "date"): + return str(value) if name == "timestamp": return value.isoformat(sep=" ", timespec="milliseconds") - return str(value) + if name == "binary": + return " ".join(f"{b:02x}" for b in value) + if name == "decimal": + return f"{value:.{_parameters(athena_type)[1]}f}" + raise NotImplementedError(f"No text rendering rule for {athena_type}.") _COMPLEX_FAMILIES = ("array", "map", "struct") From 756bdd14941a93a0a6037f2c6470698712da95b3 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Mon, 28 Sep 2026 19:47:00 +0900 Subject: [PATCH 10/14] Check nested scalar types by type and reject empty renderings The nesting check now also requires every nested scalar type to have a rendering rule, so an unsupported type raises even when its values are null or empty. A nested value rendered as empty text, such as empty binary, raises ValueError because the cursors skip empty array items. Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/expected.py | 40 +++++++++++++++++++++++++++----------- 1 file changed, 29 insertions(+), 11 deletions(-) diff --git a/tests/pyathena/expected.py b/tests/pyathena/expected.py index ad387b2e..8014b484 100644 --- a/tests/pyathena/expected.py +++ b/tests/pyathena/expected.py @@ -249,6 +249,21 @@ def _athena_text(value: Any, athena_type: str) -> str: _COMPLEX_FAMILIES = ("array", "map", "struct") +# The scalar types the rules render when nested in an array, map, or struct. +_NESTED_SCALAR_FAMILIES = ( + "boolean", + "tinyint", + "smallint", + "int", + "bigint", + "string", + "varchar", + "date", + "timestamp", + "binary", + "decimal", +) + def _member_types(athena_type: str) -> list[str]: """Return the types nested directly in a complex type. @@ -267,21 +282,20 @@ def _member_types(athena_type: str) -> list[str]: def _check_nesting(athena_type: str) -> None: """Check that the rules model the nesting of a complex type. - They model scalars nested in an array, map, or struct, and arrays of maps - or structs of scalars. + They model scalars of ``_NESTED_SCALAR_FAMILIES`` nested in an array, map, + or struct, and arrays of maps or structs of such scalars. Args: athena_type: An array, map, or struct type. Raises: - NotImplementedError: For deeper nesting, which the cursors parse with - further rules; add a rule for it here together with the column that - needs it. + NotImplementedError: For deeper nesting or another nested scalar type; + add a rule for it here together with the column that needs it. """ members = _member_types(athena_type) if family(athena_type) == "array" and family(members[0]) in ("map", "struct"): members = _member_types(members[0]) - if any(family(m) in _COMPLEX_FAMILIES for m in members): + if any(family(m) not in _NESTED_SCALAR_FAMILIES for m in members): raise NotImplementedError(f"No expectation rule for the nesting in {athena_type}.") @@ -296,17 +310,21 @@ def _element_text(value: Any, athena_type: str) -> str | None: Athena's text rendering of the value, or None for a null. Raises: - ValueError: For a string that the cursors might not parse back as - itself. Nested strings must be words of letters, digits, and - underscores separated by single spaces, and not the word null. + ValueError: For a value that the cursors might not parse back as + itself: a string other than words of letters, digits, and + underscores separated by single spaces, the word null, or a value + rendered as empty text, such as empty binary. """ if value is None: return None - if family(athena_type) in ("string", "varchar", "char") and ( + if family(athena_type) in ("string", "varchar") and ( not re.fullmatch(r"\w+(?: \w+)*", value) or value.lower() == "null" ): raise ValueError(f"Unsupported nested string value: {value!r}") - return _athena_text(value, athena_type) + text = _athena_text(value, athena_type) + if not text: + raise ValueError(f"Unsupported nested value rendered as empty text: {value!r}") + return text def _python_array(value: Any, athena_type: str) -> Any: From c0b2e3cb031bfbff0abdad431690ab958feeaa40 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Mon, 28 Sep 2026 19:52:03 +0900 Subject: [PATCH 11/14] Follow the declared struct fields in the expected values Struct values are rendered and converted field by field in declared order, with a missing field as null, as the generated Parquet data stores it; a field the type does not declare raises ValueError. Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/expected.py | 28 +++++++++++++++++++++++----- 1 file changed, 23 insertions(+), 5 deletions(-) diff --git a/tests/pyathena/expected.py b/tests/pyathena/expected.py index 8014b484..7dba5aa0 100644 --- a/tests/pyathena/expected.py +++ b/tests/pyathena/expected.py @@ -231,9 +231,8 @@ def _athena_text(value: Any, athena_type: str) -> str: entries = (f"{_athena_text(k, key_type)}={_athena_text(x, value_type)}" for k, x in value) return f"{{{', '.join(entries)}}}" if name == "struct": - field_types = dict(_struct_fields(athena_type)) - entries = (f"{k}={_athena_text(x, field_types[k])}" for k, x in value.items()) - return f"{{{', '.join(entries)}}}" + items = _struct_items(value, athena_type) + return f"{{{', '.join(f'{k}={_athena_text(x, t)}' for k, x, t in items)}}}" if name == "boolean": return str(value).lower() if name in ("tinyint", "smallint", "int", "bigint", "string", "varchar", "date"): @@ -265,6 +264,26 @@ def _athena_text(value: Any, athena_type: str) -> str: ) +def _struct_items(value: Mapping[str, Any], athena_type: str) -> list[tuple[str, Any, str]]: + """Return a struct value's fields in declared order. + + Args: + value: The struct value from the table definition; a missing field is + null, as in the generated Parquet data. + athena_type: The struct type. + + Returns: + ``(name, value, type)`` per declared field. + + Raises: + ValueError: If the value has a field that the type does not declare. + """ + fields = _struct_fields(athena_type) + if unknown := set(value) - {name for name, _ in fields}: + raise ValueError(f"Undeclared struct fields {sorted(unknown)} for {athena_type}.") + return [(name, value.get(name), field_type) for name, field_type in fields] + + def _member_types(athena_type: str) -> list[str]: """Return the types nested directly in a complex type. @@ -386,8 +405,7 @@ def _python_struct(value: Any, athena_type: str) -> dict[str, Any] | None: _check_nesting(athena_type) if value is None: return None - field_types = dict(_struct_fields(athena_type)) - return {k: _element_text(x, field_types[k]) for k, x in value.items()} + return {k: _element_text(x, t) for k, x, t in _struct_items(value, athena_type)} def _same(value: Any, athena_type: str) -> Any: From d1db7dc4fd7b1cee232b2ac3036095f802e68115 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Mon, 28 Sep 2026 23:47:01 +0900 Subject: [PATCH 12/14] Simplify and rename the expectation helpers tests/pyathena/expected.py now covers the column types of the shared tables instead of modeling every nesting the cursors can parse, and its names say what they are: ExpectedResult (was Selection), CastColumn (was Cast), base_type (was family), and PYTHON_VALUES (was PYTHON). A module docstring explains how to read it. assert_sqlalchemy_type is a test assertion rather than an expected value, so it moves to tests/pyathena/util.py. Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/expected.py | 428 +++++++++-------------- tests/pyathena/pandas/test_util.py | 16 +- tests/pyathena/s3fs/test_async_cursor.py | 16 +- tests/pyathena/s3fs/test_cursor.py | 16 +- tests/pyathena/sqlalchemy/test_base.py | 6 +- tests/pyathena/test_cursor.py | 20 +- tests/pyathena/util.py | 46 +++ 7 files changed, 254 insertions(+), 294 deletions(-) diff --git a/tests/pyathena/expected.py b/tests/pyathena/expected.py index 7dba5aa0..6334fa36 100644 --- a/tests/pyathena/expected.py +++ b/tests/pyathena/expected.py @@ -6,9 +6,19 @@ # SPDX-License-Identifier: MIT """Expected query results for the shared tables in ``tests.pyathena.tables``. -The expectations are derived from a table's column types and row values with -explicit rules per Athena type family, such as ``array`` or ``decimal``. The -rules state what a cursor returns; they never call PyAthena's converters. +How to read this module: + +- ``ExpectedResult`` is the entry point. It builds a query of a table's columns, + optionally followed by ``CastColumn`` items, and returns what a cursor should + give for it: ``rows()``, ``description()``, and ``dbapi_types()``. +- The ``*_VALUES`` dicts state how a cursor returns a value of each type. They + are keyed by base type, the Athena type without its parameters, such as + ``decimal`` for ``DECIMAL(10,1)``. Each rule takes the value from the table + definition and the column's Athena type. + +The rules state what the cursors return; they never call PyAthena's converters. +They cover the column types of the shared tables. A column of another type +raises ``KeyError`` or ``NotImplementedError`` until a rule is added here. """ import json @@ -18,29 +28,28 @@ from datetime import timezone from typing import Any -from sqlalchemy import types - from pyathena import BINARY, BOOLEAN, DATE, DATETIME, JSON, NUMBER, STRING, TIME -from pyathena.sqlalchemy.types import TINYINT, AthenaArray, AthenaStruct from tests.pyathena.tables import Column, Table +# Athena types + -def family(athena_type: str) -> str: - """Return the type family, the type name without its parameters. +def base_type(athena_type: str) -> str: + """Return an Athena type without its parameters. Args: athena_type: An Athena type, such as ``DECIMAL(10,1)`` or ``ARRAY``. Returns: - The lowercase family, such as ``decimal`` or ``array``. + The lowercase base type, such as ``decimal`` or ``array``. """ match = re.match(r"[a-z]+(?: [a-z]+)*", athena_type.strip().lower()) assert match, athena_type return match.group() -def _parameters(athena_type: str) -> tuple[int, ...]: - """Return the numeric parameters of a type. +def type_parameters(athena_type: str) -> tuple[int, ...]: + """Return the numeric parameters of an Athena type. Args: athena_type: An Athena type, such as ``DECIMAL(10,1)``. @@ -53,15 +62,30 @@ def _parameters(athena_type: str) -> tuple[int, ...]: return tuple(int(p) for p in match.group(1).split(",")) if match else () -def _type_arguments(athena_type: str) -> list[str]: - """Return the type arguments of a complex type. +def type_arguments(athena_type: str) -> list[str]: + """Return the types nested in an array, map, or struct type. Args: athena_type: An Athena type, such as ``MAP``. Returns: The element type of an array, the key and value types of a map, or the - ``name: type`` fields of a struct; empty for other types. + field types of a struct; empty for other types. + """ + if base_type(athena_type) == "struct": + return [field_type for _, field_type in _struct_fields(athena_type)] + return _split_arguments(athena_type) + + +def _split_arguments(athena_type: str) -> list[str]: + """Split the text between a type's outer angle brackets at top-level commas. + + Args: + athena_type: An Athena type, such as ``STRUCT``. + + Returns: + The arguments, such as ``["a: int", "b: int"]``; empty for a type + without angle brackets. """ if "<" not in athena_type: return [] @@ -90,13 +114,16 @@ def _struct_fields(athena_type: str) -> list[tuple[str, str]]: """ return [ (name.strip(), field_type.strip()) - for name, field_type in (a.split(":", 1) for a in _type_arguments(athena_type)) + for name, field_type in (a.split(":", 1) for a in _split_arguments(athena_type)) ] +# Cast columns + + @dataclass(frozen=True) -class Cast: - """A column the query casts from a table column. +class CastColumn: + """A result column that the query casts from a table column. Attributes: name: The result column name. @@ -117,14 +144,16 @@ def sql(self) -> str: return f"CAST({self.source} AS {self.athena_type}) AS {self.name}" -TIMESTAMP_TZ = Cast("col_timestamp_tz", "col_timestamp", "timestamp with time zone") -TIME_OF_TIMESTAMP = Cast("col_time", "col_timestamp", "time") -ARRAY_JSON = Cast("col_array_json", "col_array", "json") -MAP_JSON = Cast("col_map_json", "col_map", "json") +TIMESTAMP_TZ = CastColumn("col_timestamp_tz", "col_timestamp", "timestamp with time zone") +TIME_OF_TIMESTAMP = CastColumn("col_time", "col_timestamp", "time") +ARRAY_JSON = CastColumn("col_array_json", "col_array", "json") +MAP_JSON = CastColumn("col_map_json", "col_map", "json") + +# Cursor descriptions -# (type code, precision, scale) in the cursor description, by type family. -# The precision of varchar(n) and the precision and scale of decimal(p,s) come -# from the type parameters. +# (type code, precision, scale) in the cursor description, by base type. The +# precision of varchar(n) and the precision and scale of decimal(p,s) come from +# the type parameters. _DESCRIPTION: Mapping[str, tuple[str, int, int]] = { "boolean": ("boolean", 0, 0), "tinyint": ("tinyint", 3, 0), @@ -169,40 +198,46 @@ def sql(self) -> str: "json": JSON, } -# The SQLAlchemy type class a reflected column has, by type family. -_SQLALCHEMY_TYPES = { - "boolean": types.BOOLEAN, - "tinyint": TINYINT, - "smallint": types.SMALLINT, - "int": types.INTEGER, - "bigint": types.BIGINT, - "float": types.FLOAT, - "double": types.DOUBLE, - "string": types.String, - "varchar": types.VARCHAR, - "timestamp": types.TIMESTAMP, - "date": types.DATE, - "binary": types.BINARY, - "array": AthenaArray, - "map": types.String, - "struct": AthenaStruct, - "decimal": types.DECIMAL, -} +# Value rules +# A rule(value, athena_type) returns what a cursor gives for a value from the +# table definition. A cast column uses the rule of the type it casts to, with +# the value and type of its source column. +ValueRules = Mapping[str, Callable[[Any, str], Any]] -def _parsed_json(value: Any, athena_type: str) -> Any: - """Return a value as Athena renders it with ``CAST(... AS json)``, after JSON parsing. + +def _same(value: Any, athena_type: str) -> Any: + """Return the value unchanged. Args: value: The value from the table definition. athena_type: The value's Athena type. Returns: - The parsed JSON value; a map becomes a dict with string keys. + The value. """ - if family(athena_type) == "map": - value = dict(value) - return json.loads(json.dumps(value)) + return value + + +# Scalar types that the cursors return as the Python value itself. +_SCALAR_VALUES: ValueRules = dict.fromkeys( + ( + "boolean", + "tinyint", + "smallint", + "int", + "bigint", + "float", + "double", + "string", + "varchar", + "timestamp", + "date", + "binary", + "decimal", + ), + _same, +) def _athena_text(value: Any, athena_type: str) -> str: @@ -217,136 +252,41 @@ def _athena_text(value: Any, athena_type: str) -> str: and ``{a=1, b=2}`` for a struct. Raises: - NotImplementedError: For a type without a rendering rule here, such as - float and double, whose rendering follows Java's formatting. + NotImplementedError: For a nested scalar type other than integers and + strings, whose rendering has no rule here. """ - name = family(athena_type) - arguments = _type_arguments(athena_type) - if value is None: - return "null" + name = base_type(athena_type) if name == "array": - return f"[{', '.join(_athena_text(e, arguments[0]) for e in value)}]" + (element_type,) = type_arguments(athena_type) + return f"[{', '.join(_athena_text(e, element_type) for e in value)}]" if name == "map": - key_type, value_type = arguments + key_type, value_type = type_arguments(athena_type) entries = (f"{_athena_text(k, key_type)}={_athena_text(x, value_type)}" for k, x in value) return f"{{{', '.join(entries)}}}" if name == "struct": - items = _struct_items(value, athena_type) - return f"{{{', '.join(f'{k}={_athena_text(x, t)}' for k, x, t in items)}}}" - if name == "boolean": - return str(value).lower() - if name in ("tinyint", "smallint", "int", "bigint", "string", "varchar", "date"): + field_types = dict(_struct_fields(athena_type)) + entries = (f"{k}={_athena_text(x, field_types[k])}" for k, x in value.items()) + return f"{{{', '.join(entries)}}}" + if name in ("tinyint", "smallint", "int", "bigint", "string", "varchar"): return str(value) - if name == "timestamp": - return value.isoformat(sep=" ", timespec="milliseconds") - if name == "binary": - return " ".join(f"{b:02x}" for b in value) - if name == "decimal": - return f"{value:.{_parameters(athena_type)[1]}f}" raise NotImplementedError(f"No text rendering rule for {athena_type}.") -_COMPLEX_FAMILIES = ("array", "map", "struct") - -# The scalar types the rules render when nested in an array, map, or struct. -_NESTED_SCALAR_FAMILIES = ( - "boolean", - "tinyint", - "smallint", - "int", - "bigint", - "string", - "varchar", - "date", - "timestamp", - "binary", - "decimal", -) - - -def _struct_items(value: Mapping[str, Any], athena_type: str) -> list[tuple[str, Any, str]]: - """Return a struct value's fields in declared order. +def _check_scalar_members(athena_type: str) -> None: + """Check that an array, map, or struct type nests only scalar types. Args: - value: The struct value from the table definition; a missing field is - null, as in the generated Parquet data. - athena_type: The struct type. - - Returns: - ``(name, value, type)`` per declared field. + athena_type: The array, map, or struct type. Raises: - ValueError: If the value has a field that the type does not declare. + NotImplementedError: For a nested array, map, or struct, which the + cursors parse differently. """ - fields = _struct_fields(athena_type) - if unknown := set(value) - {name for name, _ in fields}: - raise ValueError(f"Undeclared struct fields {sorted(unknown)} for {athena_type}.") - return [(name, value.get(name), field_type) for name, field_type in fields] - - -def _member_types(athena_type: str) -> list[str]: - """Return the types nested directly in a complex type. - - Args: - athena_type: An array, map, or struct type. - - Returns: - The element type, the key and value types, or the field types. - """ - if family(athena_type) == "struct": - return [field_type for _, field_type in _struct_fields(athena_type)] - return _type_arguments(athena_type) + if any(base_type(t) in ("array", "map", "struct") for t in type_arguments(athena_type)): + raise NotImplementedError(f"No rule for the nesting in {athena_type}.") -def _check_nesting(athena_type: str) -> None: - """Check that the rules model the nesting of a complex type. - - They model scalars of ``_NESTED_SCALAR_FAMILIES`` nested in an array, map, - or struct, and arrays of maps or structs of such scalars. - - Args: - athena_type: An array, map, or struct type. - - Raises: - NotImplementedError: For deeper nesting or another nested scalar type; - add a rule for it here together with the column that needs it. - """ - members = _member_types(athena_type) - if family(athena_type) == "array" and family(members[0]) in ("map", "struct"): - members = _member_types(members[0]) - if any(family(m) not in _NESTED_SCALAR_FAMILIES for m in members): - raise NotImplementedError(f"No expectation rule for the nesting in {athena_type}.") - - -def _element_text(value: Any, athena_type: str) -> str | None: - """Return a nested scalar value as a cursor without type hints returns it. - - Args: - value: The nested value. - athena_type: The value's Athena type, a scalar type. - - Returns: - Athena's text rendering of the value, or None for a null. - - Raises: - ValueError: For a value that the cursors might not parse back as - itself: a string other than words of letters, digits, and - underscores separated by single spaces, the word null, or a value - rendered as empty text, such as empty binary. - """ - if value is None: - return None - if family(athena_type) in ("string", "varchar") and ( - not re.fullmatch(r"\w+(?: \w+)*", value) or value.lower() == "null" - ): - raise ValueError(f"Unsupported nested string value: {value!r}") - text = _athena_text(value, athena_type) - if not text: - raise ValueError(f"Unsupported nested value rendered as empty text: {value!r}") - return text - - -def _python_array(value: Any, athena_type: str) -> Any: +def _python_array(value: Any, athena_type: str) -> list[Any]: """Return an array as a cursor without type hints returns it. Args: @@ -354,28 +294,18 @@ def _python_array(value: Any, athena_type: str) -> Any: athena_type: The array type. Returns: - None for a null. Otherwise the parsed JSON if Athena's rendering is - valid JSON, such as ``[1, 2]``, or else a list of the elements, where a - map or struct element is a dict and any other element is its text - rendering. + The parsed JSON if Athena's rendering is valid JSON, such as ``[1, 2]``; + otherwise the elements' text, such as ``["a", "b"]`` for ``[a, b]``. """ - _check_nesting(athena_type) - if value is None: - return None - (element_type,) = _type_arguments(athena_type) - element_family = family(element_type) - if element_family == "map": - return [_python_map(e, element_type) for e in value] - if element_family == "struct": - return [_python_struct(e, element_type) for e in value] - elements = [_element_text(e, element_type) for e in value] + _check_scalar_members(athena_type) + (element_type,) = type_arguments(athena_type) try: return json.loads(_athena_text(value, athena_type)) except ValueError: - return elements + return [_athena_text(e, element_type) for e in value] -def _python_map(value: Any, athena_type: str) -> dict[str, Any] | None: +def _python_map(value: Any, athena_type: str) -> dict[str, str]: """Return a map as a cursor without type hints returns it. Args: @@ -383,16 +313,14 @@ def _python_map(value: Any, athena_type: str) -> dict[str, Any] | None: athena_type: The map type. Returns: - None for a null; otherwise the keys and values as their text renderings. + The keys and values as their text renderings. """ - _check_nesting(athena_type) - if value is None: - return None - key_type, value_type = _type_arguments(athena_type) - return {_element_text(k, key_type): _element_text(x, value_type) for k, x in value} + _check_scalar_members(athena_type) + key_type, value_type = type_arguments(athena_type) + return {_athena_text(k, key_type): _athena_text(x, value_type) for k, x in value} -def _python_struct(value: Any, athena_type: str) -> dict[str, Any] | None: +def _python_struct(value: Any, athena_type: str) -> dict[str, str]: """Return a struct as a cursor without type hints returns it. Args: @@ -400,56 +328,31 @@ def _python_struct(value: Any, athena_type: str) -> dict[str, Any] | None: athena_type: The struct type. Returns: - None for a null; otherwise the field values as their text renderings. + The field values as their text renderings. """ - _check_nesting(athena_type) - if value is None: - return None - return {k: _element_text(x, t) for k, x, t in _struct_items(value, athena_type)} + _check_scalar_members(athena_type) + field_types = dict(_struct_fields(athena_type)) + return {k: _athena_text(x, field_types[k]) for k, x in value.items()} -def _same(value: Any, athena_type: str) -> Any: - """Return the value unchanged. +def _parsed_json(value: Any, athena_type: str) -> Any: + """Return a value cast to JSON, after JSON parsing. Args: value: The value from the table definition. athena_type: The value's Athena type. Returns: - The value. + The parsed JSON value; a map becomes a dict with string keys. """ - return value + return json.loads(json.dumps(dict(value) if base_type(athena_type) == "map" else value)) -# A representation maps a type family to a rule(value, athena_type) that -# returns what the cursor gives for that value. Cast columns use the rule of -# the target family with the source column's value and type. -Representation = Mapping[str, Callable[[Any, str], Any]] - -_SCALARS: Representation = dict.fromkeys( - ( - "boolean", - "tinyint", - "smallint", - "int", - "bigint", - "float", - "double", - "string", - "varchar", - "timestamp", - "date", - "binary", - "decimal", - ), - _same, -) - # Rows of Cursor, S3FSCursor, pyathena.pandas.util.as_pandas, and SQLAlchemy. # Without type hints, the cursors parse Athena's text rendering of arrays, maps, # and structs, so nested values that are not JSON are strings. -PYTHON: Representation = { - **_SCALARS, +PYTHON_VALUES: ValueRules = { + **_SCALAR_VALUES, "timestamp with time zone": lambda v, t: v.replace(tzinfo=timezone.utc), "time": lambda v, t: v.time(), "array": _python_array, @@ -458,10 +361,12 @@ def _same(value: Any, athena_type: str) -> Any: "json": _parsed_json, } +# The expected result of a query + @dataclass(frozen=True) -class Selection: - """A query of a table's columns, followed by cast columns. +class ExpectedResult: + """A query of a table's columns, followed by cast columns, and its expected result. Attributes: table: The table. @@ -469,34 +374,46 @@ class Selection: """ table: Table - casts: tuple[Cast, ...] = () + casts: tuple[CastColumn, ...] = () + + @property + def columns(self) -> list[Column]: + """The selected table columns.""" + return list(self.table.columns) @property def names(self) -> list[str]: """The result column names.""" - return [c.name for c in self.table.columns] + [c.name for c in self.casts] + return [c.name for c in self.columns] + [c.name for c in self.casts] @property def sql(self) -> str: """The query.""" - items = [c.name for c in self.table.columns] + [c.sql() for c in self.casts] + items = [c.name for c in self.columns] + [c.sql() for c in self.casts] return f"SELECT {', '.join(items)} FROM {self.table.name}" - def _items(self) -> list[tuple[str, str, int, str]]: - """Return how each result column is derived from the table. + def _sources(self) -> list[tuple[str, int, str]]: + """Return where each result column's value comes from. Returns: - ``(name, athena_type, source_index, source_type)`` per result column, - where the source is the table column the value comes from. A table - column is its own source. + ``(athena_type, source_index, source_type)`` per result column: the + column's type, and the index and type of the table column whose value + it holds. A table column is its own source. """ - columns = self.table.columns - index = {c.name: i for i, c in enumerate(columns)} - return [(c.name, c.athena_type, i, c.athena_type) for i, c in enumerate(columns)] + [ - (c.name, c.athena_type, index[c.source], columns[index[c.source]].athena_type) - for c in self.casts + index = {c.name: i for i, c in enumerate(self.table.columns)} + by_name = {c.name: c for c in self.table.columns} + return [(c.athena_type, index[c.name], c.athena_type) for c in self.columns] + [ + (c.athena_type, index[c.source], by_name[c.source].athena_type) for c in self.casts ] + def _types(self) -> list[str]: + """Return the Athena type of each result column. + + Returns: + The types in result order. + """ + return [athena_type for athena_type, _, _ in self._sources()] + def description(self) -> list[tuple[Any, ...]]: """Return the expected cursor description. @@ -504,9 +421,9 @@ def description(self) -> list[tuple[Any, ...]]: One DB API description tuple per result column. """ result = [] - for name, athena_type, _, _ in self._items(): - code, precision, scale = _DESCRIPTION[family(athena_type)] - if parameters := _parameters(athena_type): + for name, athena_type in zip(self.names, self._types(), strict=True): + code, precision, scale = _DESCRIPTION[base_type(athena_type)] + if parameters := type_parameters(athena_type): precision, scale = (*parameters, 0)[:2] result.append((name, code, None, None, precision, scale, "UNKNOWN")) return result @@ -517,44 +434,23 @@ def dbapi_types(self) -> list[Any]: Returns: The type objects in result order. """ - return [_DBAPI_TYPES[family(athena_type)] for _, athena_type, _, _ in self._items()] + return [_DBAPI_TYPES[base_type(t)] for t in self._types()] - def rows(self, representation: Representation) -> list[tuple[Any, ...]]: + def rows(self, values: ValueRules) -> list[tuple[Any, ...]]: """Return the expected rows. Args: - representation: The rules of the cursor under test, such as ``PYTHON``. + values: The value rules of the cursor under test, such as + ``PYTHON_VALUES``. Returns: One tuple per row of the table. """ - items = self._items() + sources = self._sources() return [ tuple( - representation[family(athena_type)](row[index], source_type) - for _, athena_type, index, source_type in items + values[base_type(athena_type)](row[index], source_type) + for athena_type, index, source_type in sources ) for row in self.table.rows ] - - -def assert_sqlalchemy_type(sqlalchemy_type: Any, column: Column) -> None: - """Assert that a reflected SQLAlchemy type matches a table column. - - Args: - sqlalchemy_type: The reflected column type. - column: The column in the table definition. - - Raises: - AssertionError: If the type class or its parameters differ. - """ - name = family(column.athena_type) - assert isinstance(sqlalchemy_type, _SQLALCHEMY_TYPES[name]), (column.name, sqlalchemy_type) - if name == "varchar": - assert sqlalchemy_type.length == _parameters(column.athena_type)[0], column.name - elif name == "decimal": - precision, scale = _parameters(column.athena_type) - assert (sqlalchemy_type.precision, sqlalchemy_type.scale) == (precision, scale) - elif name == "array": - (element,) = _type_arguments(column.athena_type) - assert isinstance(sqlalchemy_type.item_type, _SQLALCHEMY_TYPES[family(element)]) diff --git a/tests/pyathena/pandas/test_util.py b/tests/pyathena/pandas/test_util.py index 4eccdc2d..6382d637 100644 --- a/tests/pyathena/pandas/test_util.py +++ b/tests/pyathena/pandas/test_util.py @@ -15,7 +15,13 @@ to_sql, ) from tests import ENV -from tests.pyathena.expected import ARRAY_JSON, MAP_JSON, PYTHON, TIME_OF_TIMESTAMP, Selection +from tests.pyathena.expected import ( + ARRAY_JSON, + MAP_JSON, + PYTHON_VALUES, + TIME_OF_TIMESTAMP, + ExpectedResult, +) from tests.pyathena.tables import ONE_ROW_COMPLEX @@ -53,11 +59,11 @@ def test_reset_index(): def test_as_pandas(cursor): - selection = Selection(ONE_ROW_COMPLEX, (TIME_OF_TIMESTAMP, ARRAY_JSON, MAP_JSON)) - cursor.execute(selection.sql) + expected = ExpectedResult(ONE_ROW_COMPLEX, casts=(TIME_OF_TIMESTAMP, ARRAY_JSON, MAP_JSON)) + cursor.execute(expected.sql) df = as_pandas(cursor) - assert list(df.columns) == selection.names - assert [tuple(row) for _, row in df.iterrows()] == selection.rows(PYTHON) + assert list(df.columns) == expected.names + assert [tuple(row) for _, row in df.iterrows()] == expected.rows(PYTHON_VALUES) def test_as_pandas_integer_na_values(cursor): diff --git a/tests/pyathena/s3fs/test_async_cursor.py b/tests/pyathena/s3fs/test_async_cursor.py index c3d7f372..eaf7eca4 100644 --- a/tests/pyathena/s3fs/test_async_cursor.py +++ b/tests/pyathena/s3fs/test_async_cursor.py @@ -17,7 +17,13 @@ from pyathena.s3fs.result_set import AthenaS3FSResultSet from tests import ENV from tests.pyathena.conftest import connect -from tests.pyathena.expected import ARRAY_JSON, MAP_JSON, PYTHON, TIME_OF_TIMESTAMP, Selection +from tests.pyathena.expected import ( + ARRAY_JSON, + MAP_JSON, + PYTHON_VALUES, + TIME_OF_TIMESTAMP, + ExpectedResult, +) from tests.pyathena.tables import ONE_ROW_COMPLEX @@ -61,11 +67,11 @@ def test_invalid_arraysize(self, async_s3fs_cursor): async_s3fs_cursor.arraysize = -1 def test_complex(self, async_s3fs_cursor): - selection = Selection(ONE_ROW_COMPLEX, (TIME_OF_TIMESTAMP, ARRAY_JSON, MAP_JSON)) - query_id, future = async_s3fs_cursor.execute(selection.sql) + expected = ExpectedResult(ONE_ROW_COMPLEX, casts=(TIME_OF_TIMESTAMP, ARRAY_JSON, MAP_JSON)) + query_id, future = async_s3fs_cursor.execute(expected.sql) result_set = future.result() - assert result_set.description == selection.description() - assert result_set.fetchall() == selection.rows(PYTHON) + assert result_set.description == expected.description() + assert result_set.fetchall() == expected.rows(PYTHON_VALUES) def test_cancel(self, async_s3fs_cursor): query_id, future = async_s3fs_cursor.execute( diff --git a/tests/pyathena/s3fs/test_cursor.py b/tests/pyathena/s3fs/test_cursor.py index 2dec95c3..74929fbb 100644 --- a/tests/pyathena/s3fs/test_cursor.py +++ b/tests/pyathena/s3fs/test_cursor.py @@ -12,7 +12,13 @@ from pyathena.s3fs.result_set import AthenaS3FSResultSet from tests import ENV from tests.pyathena.conftest import connect -from tests.pyathena.expected import ARRAY_JSON, MAP_JSON, PYTHON, TIME_OF_TIMESTAMP, Selection +from tests.pyathena.expected import ( + ARRAY_JSON, + MAP_JSON, + PYTHON_VALUES, + TIME_OF_TIMESTAMP, + ExpectedResult, +) from tests.pyathena.tables import ONE_ROW_COMPLEX @@ -55,10 +61,10 @@ def test_invalid_arraysize(self, s3fs_cursor): s3fs_cursor.arraysize = -1 def test_complex(self, s3fs_cursor): - selection = Selection(ONE_ROW_COMPLEX, (TIME_OF_TIMESTAMP, ARRAY_JSON, MAP_JSON)) - s3fs_cursor.execute(selection.sql) - assert s3fs_cursor.description == selection.description() - assert s3fs_cursor.fetchall() == selection.rows(PYTHON) + expected = ExpectedResult(ONE_ROW_COMPLEX, casts=(TIME_OF_TIMESTAMP, ARRAY_JSON, MAP_JSON)) + s3fs_cursor.execute(expected.sql) + assert s3fs_cursor.description == expected.description() + assert s3fs_cursor.fetchall() == expected.rows(PYTHON_VALUES) def test_cancel(self, s3fs_cursor): def cancel(c): diff --git a/tests/pyathena/sqlalchemy/test_base.py b/tests/pyathena/sqlalchemy/test_base.py index 0ca2d81f..b838a9cf 100644 --- a/tests/pyathena/sqlalchemy/test_base.py +++ b/tests/pyathena/sqlalchemy/test_base.py @@ -34,9 +34,9 @@ ) from pyathena.util import RetryConfig from tests.pyathena.conftest import ENV -from tests.pyathena.expected import PYTHON, Selection, assert_sqlalchemy_type +from tests.pyathena.expected import PYTHON_VALUES, ExpectedResult from tests.pyathena.tables import ONE_ROW_COMPLEX -from tests.pyathena.util import decorated, throttle_metadata_api +from tests.pyathena.util import assert_sqlalchemy_type, decorated, throttle_metadata_api # Amazon S3 Tables tests need a pre-provisioned table-bucket catalog; the session # creates its own namespace in it. @@ -1399,7 +1399,7 @@ def test_reflect_select(self, engine): assert [c.name for c in one_row_complex.c] == [c.name for c in ONE_ROW_COMPLEX.columns] assert isinstance(one_row_complex.c.col_string, Column) rows = conn.execute(one_row_complex.select()).fetchall() - assert [tuple(row) for row in rows] == Selection(ONE_ROW_COMPLEX).rows(PYTHON) + assert [tuple(row) for row in rows] == ExpectedResult(ONE_ROW_COMPLEX).rows(PYTHON_VALUES) for column in ONE_ROW_COMPLEX.columns: assert_sqlalchemy_type(one_row_complex.c[column.name].type, column) diff --git a/tests/pyathena/test_cursor.py b/tests/pyathena/test_cursor.py index 915be028..ad9b65ac 100644 --- a/tests/pyathena/test_cursor.py +++ b/tests/pyathena/test_cursor.py @@ -32,10 +32,10 @@ from tests.pyathena.expected import ( ARRAY_JSON, MAP_JSON, - PYTHON, + PYTHON_VALUES, TIME_OF_TIMESTAMP, TIMESTAMP_TZ, - Selection, + ExpectedResult, ) from tests.pyathena.tables import ONE_ROW_COMPLEX from tests.pyathena.util import throttle_metadata_api, unreachable_glue @@ -602,18 +602,18 @@ def test_query_execution_initial(self, cursor): assert cursor.effective_engine_version is None def test_complex(self, cursor): - selection = Selection( - ONE_ROW_COMPLEX, (TIMESTAMP_TZ, TIME_OF_TIMESTAMP, ARRAY_JSON, MAP_JSON) + expected = ExpectedResult( + ONE_ROW_COMPLEX, casts=(TIMESTAMP_TZ, TIME_OF_TIMESTAMP, ARRAY_JSON, MAP_JSON) ) - cursor.execute(selection.sql) - assert cursor.description == selection.description() + cursor.execute(expected.sql) + assert cursor.description == expected.description() rows = cursor.fetchall() - expected = selection.rows(PYTHON) - assert rows == expected + expected_rows = expected.rows(PYTHON_VALUES) + assert rows == expected_rows # catch unicode/str - assert list(map(type, rows[0])) == list(map(type, expected[0])) + assert list(map(type, rows[0])) == list(map(type, expected_rows[0])) # compare dbapi type object - assert [d[1] for d in cursor.description] == selection.dbapi_types() + assert [d[1] for d in cursor.description] == expected.dbapi_types() def test_complex_with_type_hints(self, cursor): # 1. Basic complex columns from one_row_complex diff --git a/tests/pyathena/util.py b/tests/pyathena/util.py index f5ac673d..2ff90f08 100644 --- a/tests/pyathena/util.py +++ b/tests/pyathena/util.py @@ -15,6 +15,9 @@ from pyathena.glue import GlueMetadataClient from pyathena.model import AthenaCalculationExecutionStatus +from pyathena.sqlalchemy.types import TINYINT, AthenaArray, AthenaStruct +from tests.pyathena.expected import base_type, type_arguments, type_parameters +from tests.pyathena.tables import Column _queries = Environment( loader=FileSystemLoader(Path(__file__).parents[1].resolve() / "resources" / "queries") @@ -133,3 +136,46 @@ def wait_for_spark_session_state(client, session_id, state, timeout=120): return time.sleep(1) raise AssertionError(f"Session {session_id} did not become {state} in {timeout} seconds.") + + +# The SQLAlchemy type class of a reflected column, by base type. +_SQLALCHEMY_TYPES = { + "boolean": types.BOOLEAN, + "tinyint": TINYINT, + "smallint": types.SMALLINT, + "int": types.INTEGER, + "bigint": types.BIGINT, + "float": types.FLOAT, + "double": types.DOUBLE, + "string": types.String, + "varchar": types.VARCHAR, + "timestamp": types.TIMESTAMP, + "date": types.DATE, + "binary": types.BINARY, + "array": AthenaArray, + "map": types.String, + "struct": AthenaStruct, + "decimal": types.DECIMAL, +} + + +def assert_sqlalchemy_type(sqlalchemy_type, column: Column) -> None: + """Assert that a reflected SQLAlchemy column type matches a table column's definition. + + Args: + sqlalchemy_type: The reflected column type. + column: The column in ``tests.pyathena.tables``. + + Raises: + AssertionError: If the type class or its parameters differ. + """ + name = base_type(column.athena_type) + assert isinstance(sqlalchemy_type, _SQLALCHEMY_TYPES[name]), (column.name, sqlalchemy_type) + if name == "varchar": + assert sqlalchemy_type.length == type_parameters(column.athena_type)[0], column.name + elif name == "decimal": + precision, scale = type_parameters(column.athena_type) + assert (sqlalchemy_type.precision, sqlalchemy_type.scale) == (precision, scale) + elif name == "array": + (element,) = type_arguments(column.athena_type) + assert isinstance(sqlalchemy_type.item_type, _SQLALCHEMY_TYPES[base_type(element)]) From 0b4d138eb82b9bf83981af7a22a9325e1cff1e24 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 00:01:32 +0900 Subject: [PATCH 13/14] Describe the nested-string constraint as one the data keeps Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/tables.py | 9 ++++----- 1 file changed, 4 insertions(+), 5 deletions(-) diff --git a/tests/pyathena/tables.py b/tests/pyathena/tables.py index a9453c48..55100442 100644 --- a/tests/pyathena/tables.py +++ b/tests/pyathena/tables.py @@ -157,11 +157,10 @@ def _to_text(value: Any) -> str: # One row with a value of each column type. Tests derive their expected results -# from this definition with tests.pyathena.expected. Strings nested in arrays, -# maps, and structs must be words of letters, digits, and underscores separated -# by single spaces, and not the word null, because the cursors parse Athena's -# text rendering of those types; tests.pyathena.expected rejects other nested -# strings. +# from this definition with tests.pyathena.expected. Keep strings nested in +# arrays, maps, and structs to words of letters, digits, and underscores +# separated by single spaces, other than null, because the cursors parse +# Athena's text rendering of those types. ONE_ROW_COMPLEX = Table( "one_row_complex", ( From 36988f4b22df9f425ae187f8c307c78a43fc0ce7 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 00:05:59 +0900 Subject: [PATCH 14/14] Render nested nulls and struct fields as Athena does A null nested in an array, map, or struct renders as null and comes back as None, and struct fields render in their declared order, not in the order of the definition's dict. Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/expected.py | 37 ++++++++++++++++++++++++++----------- 1 file changed, 26 insertions(+), 11 deletions(-) diff --git a/tests/pyathena/expected.py b/tests/pyathena/expected.py index 6334fa36..f6049cc8 100644 --- a/tests/pyathena/expected.py +++ b/tests/pyathena/expected.py @@ -256,6 +256,8 @@ def _athena_text(value: Any, athena_type: str) -> str: strings, whose rendering has no rule here. """ name = base_type(athena_type) + if value is None: + return "null" if name == "array": (element_type,) = type_arguments(athena_type) return f"[{', '.join(_athena_text(e, element_type) for e in value)}]" @@ -264,14 +266,26 @@ def _athena_text(value: Any, athena_type: str) -> str: entries = (f"{_athena_text(k, key_type)}={_athena_text(x, value_type)}" for k, x in value) return f"{{{', '.join(entries)}}}" if name == "struct": - field_types = dict(_struct_fields(athena_type)) - entries = (f"{k}={_athena_text(x, field_types[k])}" for k, x in value.items()) + entries = (f"{k}={_athena_text(value[k], t)}" for k, t in _struct_fields(athena_type)) return f"{{{', '.join(entries)}}}" if name in ("tinyint", "smallint", "int", "bigint", "string", "varchar"): return str(value) raise NotImplementedError(f"No text rendering rule for {athena_type}.") +def _member_value(value: Any, athena_type: str) -> str | None: + """Return a value nested in an array, map, or struct as the cursors parse it. + + Args: + value: The nested value. + athena_type: The value's Athena type. + + Returns: + Its text rendering, or None for a null. + """ + return None if value is None else _athena_text(value, athena_type) + + def _check_scalar_members(athena_type: str) -> None: """Check that an array, map, or struct type nests only scalar types. @@ -295,17 +309,18 @@ def _python_array(value: Any, athena_type: str) -> list[Any]: Returns: The parsed JSON if Athena's rendering is valid JSON, such as ``[1, 2]``; - otherwise the elements' text, such as ``["a", "b"]`` for ``[a, b]``. + otherwise the elements' text, such as ``["a", "b"]`` for ``[a, b]``, + with None for a null element. """ _check_scalar_members(athena_type) (element_type,) = type_arguments(athena_type) try: return json.loads(_athena_text(value, athena_type)) except ValueError: - return [_athena_text(e, element_type) for e in value] + return [_member_value(e, element_type) for e in value] -def _python_map(value: Any, athena_type: str) -> dict[str, str]: +def _python_map(value: Any, athena_type: str) -> dict[str, str | None]: """Return a map as a cursor without type hints returns it. Args: @@ -313,14 +328,14 @@ def _python_map(value: Any, athena_type: str) -> dict[str, str]: athena_type: The map type. Returns: - The keys and values as their text renderings. + The keys and values as their text renderings; a null value is None. """ _check_scalar_members(athena_type) key_type, value_type = type_arguments(athena_type) - return {_athena_text(k, key_type): _athena_text(x, value_type) for k, x in value} + return {_athena_text(k, key_type): _member_value(x, value_type) for k, x in value} -def _python_struct(value: Any, athena_type: str) -> dict[str, str]: +def _python_struct(value: Any, athena_type: str) -> dict[str, str | None]: """Return a struct as a cursor without type hints returns it. Args: @@ -328,11 +343,11 @@ def _python_struct(value: Any, athena_type: str) -> dict[str, str]: athena_type: The struct type. Returns: - The field values as their text renderings. + The field values as their text renderings, in declared order; a null + value is None. """ _check_scalar_members(athena_type) - field_types = dict(_struct_fields(athena_type)) - return {k: _athena_text(x, field_types[k]) for k, x in value.items()} + return {k: _member_value(value[k], t) for k, t in _struct_fields(athena_type)} def _parsed_json(value: Any, athena_type: str) -> Any: