From 68270a4308a2529a41d50f8b19a88b6a05fb284d Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Mon, 28 Sep 2026 23:49:46 +0900 Subject: [PATCH 1/2] Derive the one_row_complex expectations of the pandas, Arrow, and Polars cursors The complex-row tests of PandasCursor, ArrowCursor, and PolarsCursor spelled out the one_row_complex rows, descriptions, pandas dtypes, Arrow schemas, and Polars dtypes by hand. They now derive them from the table definition with the value rules and type tables in tests/pyathena/expected.py: Athena's text rendering for the CSV paths, native values for the UNLOAD paths, and explicit pandas, Arrow, and Polars type tables. The PolarsCursor tests keep leaving out the binary and complex columns, now by type, and also assert the column names. Table.data_file rejects a row whose values change in the columns' Arrow types, such as a timestamp more precise than the type or a struct without all of its fields, because the expectations are derived from the rows as written. Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/arrow/test_cursor.py | 463 +++------------------------ tests/pyathena/expected.py | 329 ++++++++++++++++++- tests/pyathena/pandas/test_cursor.py | 423 ++---------------------- tests/pyathena/polars/test_cursor.py | 196 ++---------- tests/pyathena/tables.py | 14 +- 5 files changed, 418 insertions(+), 1007 deletions(-) diff --git a/tests/pyathena/arrow/test_cursor.py b/tests/pyathena/arrow/test_cursor.py index af1abca6..512b4e36 100644 --- a/tests/pyathena/arrow/test_cursor.py +++ b/tests/pyathena/arrow/test_cursor.py @@ -10,12 +10,7 @@ import string import time from concurrent.futures import ThreadPoolExecutor -from datetime import datetime -from decimal import Decimal -import pandas as pd -import polars as pl -import pyarrow as pa import pytest from pyathena.arrow.cursor import ArrowCursor @@ -23,6 +18,18 @@ from pyathena.error import DatabaseError, ProgrammingError from tests import ENV from tests.pyathena.conftest import connect +from tests.pyathena.expected import ( + ARRAY_JSON, + ARROW_TABLE_VALUES, + ARROW_VALUES, + MAP_JSON, + TIME_OF_TIMESTAMP, + UNLOAD_POLARS_VALUES, + UNLOAD_VALUES, + ExpectedResult, + polars_type, +) +from tests.pyathena.tables import ONE_ROW_COMPLEX class TestArrowCursor: @@ -112,75 +119,10 @@ def test_invalid_arraysize(self, arrow_cursor): arrow_cursor.arraysize = -1 def test_complex(self, arrow_cursor): - arrow_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 arrow_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 arrow_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"), - ) - ] + expected = ExpectedResult(ONE_ROW_COMPLEX, casts=(TIME_OF_TIMESTAMP, ARRAY_JSON, MAP_JSON)) + arrow_cursor.execute(expected.sql) + assert arrow_cursor.description == expected.description() + assert arrow_cursor.fetchall() == expected.rows(ARROW_VALUES) @pytest.mark.parametrize( "arrow_cursor", @@ -194,82 +136,10 @@ def test_complex(self, arrow_cursor): def test_complex_unload(self, arrow_cursor): # NOT_SUPPORTED: Unsupported Hive type: time # NOT_SUPPORTED: Unsupported Hive type: json - arrow_cursor.execute( - """ - SELECT - col_boolean - ,col_tinyint - ,col_smallint - ,col_int - ,col_bigint - ,col_float - ,col_double - ,col_string - ,col_varchar - ,col_timestamp - ,col_date - ,col_binary - ,col_array - ,col_map - ,col_struct - ,col_decimal - FROM one_row_complex - """ - ) - assert arrow_cursor.description == [ - ("col_boolean", "boolean", None, None, 0, 0, "NULLABLE"), - ( - "col_tinyint", - "tinyint", - None, - None, - 3, - 0, - "NULLABLE", - ), - ( - "col_smallint", - "smallint", - None, - None, - 5, - 0, - "NULLABLE", - ), - ("col_int", "integer", None, None, 10, 0, "NULLABLE"), - ("col_bigint", "bigint", None, None, 19, 0, "NULLABLE"), - ("col_float", "float", None, None, 17, 0, "NULLABLE"), - ("col_double", "double", None, None, 17, 0, "NULLABLE"), - ("col_string", "varchar", None, None, 2147483647, 0, "NULLABLE"), - ("col_varchar", "varchar", None, None, 2147483647, 0, "NULLABLE"), - ("col_timestamp", "timestamp", None, None, 3, 0, "NULLABLE"), - ("col_date", "date", None, None, 0, 0, "NULLABLE"), - ("col_binary", "varbinary", None, None, 1073741824, 0, "NULLABLE"), - ("col_array", "array", None, None, 0, 0, "NULLABLE"), - ("col_map", "map", None, None, 0, 0, "NULLABLE"), - ("col_struct", "row", None, None, 0, 0, "NULLABLE"), - ("col_decimal", "decimal", None, None, 10, 1, "NULLABLE"), - ] - assert arrow_cursor.fetchall() == [ - ( - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - "varchar", - pd.Timestamp(2017, 1, 1, 0, 0, 0), - datetime(2017, 1, 2).date(), - b"123", - [1, 2], - [(1, 2), (3, 4)], - {"a": 1, "b": 2}, - Decimal("0.1"), - ) - ] + expected = ExpectedResult(ONE_ROW_COMPLEX) + arrow_cursor.execute(expected.sql) + assert arrow_cursor.description == expected.description(unload=True) + assert arrow_cursor.fetchall() == expected.rows(UNLOAD_VALUES) def test_fetch_no_data(self, arrow_cursor): pytest.raises(ProgrammingError, arrow_cursor.fetchone) @@ -301,79 +171,12 @@ def test_many_as_arrow(self, arrow_cursor): assert list(zip(*table.to_pydict().values(), strict=False)) == [(i,) for i in range(10000)] def test_complex_as_arrow(self, arrow_cursor): - table = arrow_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 - """ - ).as_arrow() - assert table.shape[0] == 1 - assert table.shape[1] == 19 - assert table.schema == pa.schema( - [ - pa.field("col_boolean", pa.bool_()), - pa.field("col_tinyint", pa.int8()), - pa.field("col_smallint", pa.int16()), - pa.field("col_int", pa.int32()), - pa.field("col_bigint", pa.int64()), - pa.field("col_float", pa.float32()), - pa.field("col_double", pa.float64()), - pa.field("col_string", pa.string()), - pa.field("col_varchar", pa.string()), - pa.field("col_timestamp", pa.timestamp("ms")), - pa.field("col_time", pa.string()), - pa.field("col_date", pa.timestamp("ms")), - pa.field("col_binary", pa.string()), - pa.field("col_array", pa.string()), - pa.field("col_array_json", pa.string()), - pa.field("col_map", pa.string()), - pa.field("col_map_json", pa.string()), - pa.field("col_struct", pa.string()), - pa.field("col_decimal", pa.string()), - ] + expected = ExpectedResult(ONE_ROW_COMPLEX, casts=(TIME_OF_TIMESTAMP, ARRAY_JSON, MAP_JSON)) + table = arrow_cursor.execute(expected.sql).as_arrow() + assert table.schema == expected.arrow_schema() + assert list(zip(*table.to_pydict().values(), strict=True)) == expected.rows( + ARROW_TABLE_VALUES ) - assert list(zip(*table.to_pydict().values(), strict=False)) == [ - ( - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - "varchar", - datetime(2017, 1, 1, 0, 0, 0), - "00:00:00.000", - datetime(2017, 1, 2, 0, 0, 0), - "31 32 33", - "[1, 2]", - "[1,2]", - "{1=2, 3=4}", - '{"1":2,"3":4}', - "{a=1, b=2}", - "0.1", - ) - ] @pytest.mark.parametrize( "arrow_cursor", @@ -387,73 +190,10 @@ def test_complex_as_arrow(self, arrow_cursor): def test_complex_unload_as_arrow(self, arrow_cursor): # NOT_SUPPORTED: Unsupported Hive type: time # NOT_SUPPORTED: Unsupported Hive type: json - table = arrow_cursor.execute( - """ - SELECT - col_boolean - ,col_tinyint - ,col_smallint - ,col_int - ,col_bigint - ,col_float - ,col_double - ,col_string - ,col_varchar - ,col_timestamp - ,col_date - ,col_binary - ,col_array - ,col_map - ,col_struct - ,col_decimal - FROM one_row_complex - """ - ).as_arrow() - assert table.shape[0] == 1 - assert table.shape[1] == 16 - assert table.schema == pa.schema( - [ - pa.field("col_boolean", pa.bool_()), - pa.field("col_tinyint", pa.int8()), - pa.field("col_smallint", pa.int16()), - pa.field("col_int", pa.int32()), - pa.field("col_bigint", pa.int64()), - pa.field("col_float", pa.float32()), - pa.field("col_double", pa.float64()), - pa.field("col_string", pa.string()), - pa.field("col_varchar", pa.string()), - pa.field("col_timestamp", pa.timestamp("ns")), - pa.field("col_date", pa.date32()), - pa.field("col_binary", pa.binary()), - pa.field("col_array", pa.list_(pa.field("array_element", pa.int32()))), - pa.field("col_map", pa.map_(pa.int32(), pa.field("entries", pa.int32()))), - pa.field( - "col_struct", - pa.struct([pa.field("a", pa.int32()), pa.field("b", pa.int32())]), - ), - pa.field("col_decimal", pa.decimal128(10, 1)), - ] - ) - assert list(zip(*table.to_pydict().values(), strict=False)) == [ - ( - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - "varchar", - pd.Timestamp(2017, 1, 1, 0, 0, 0), - datetime(2017, 1, 2).date(), - b"123", - [1, 2], - [(1, 2), (3, 4)], - {"a": 1, "b": 2}, - Decimal("0.1"), - ) - ] + expected = ExpectedResult(ONE_ROW_COMPLEX) + table = arrow_cursor.execute(expected.sql).as_arrow() + assert table.schema == expected.arrow_schema(unload=True) + assert list(zip(*table.to_pydict().values(), strict=True)) == expected.rows(UNLOAD_VALUES) @pytest.mark.parametrize( "arrow_cursor", @@ -478,79 +218,12 @@ def test_many_as_polars(self, arrow_cursor): assert df.to_dicts() == [{"a": i} for i in range(10000)] def test_complex_as_polars(self, arrow_cursor): - df = arrow_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 - """ - ).as_polars() - assert df.height == 1 - assert df.width == 19 - dtypes = tuple(df.dtypes) - assert dtypes == ( - pl.Boolean, - pl.Int8, - pl.Int16, - pl.Int32, - pl.Int64, - pl.Float32, - pl.Float64, - pl.String, - pl.String, - pl.Datetime("ms"), - pl.String, - pl.Datetime("ms"), - pl.String, - pl.String, - pl.String, - pl.String, - pl.String, - pl.String, - pl.String, - ) - rows = df.to_dicts() - assert rows == [ - { - "col_boolean": True, - "col_tinyint": 127, - "col_smallint": 32767, - "col_int": 2147483647, - "col_bigint": 9223372036854775807, - "col_float": 0.5, - "col_double": 0.25, - "col_string": "a string", - "col_varchar": "varchar", - "col_timestamp": datetime(2017, 1, 1, 0, 0, 0), - "col_time": "00:00:00.000", - "col_date": datetime(2017, 1, 2, 0, 0, 0), - "col_binary": "31 32 33", - "col_array": "[1, 2]", - "col_array_json": "[1,2]", - "col_map": "{1=2, 3=4}", - "col_map_json": '{"1":2,"3":4}', - "col_struct": "{a=1, b=2}", - "col_decimal": "0.1", - } + expected = ExpectedResult(ONE_ROW_COMPLEX, casts=(TIME_OF_TIMESTAMP, ARRAY_JSON, MAP_JSON)) + df = arrow_cursor.execute(expected.sql).as_polars() + assert list(zip(df.columns, df.dtypes, strict=True)) == [ + (f.name, polars_type(f.type)) for f in expected.arrow_schema() ] + assert df.to_dicts() == expected.dicts(ARROW_TABLE_VALUES) @pytest.mark.parametrize( "arrow_cursor", @@ -564,70 +237,12 @@ def test_complex_as_polars(self, arrow_cursor): def test_complex_unload_as_polars(self, arrow_cursor): # NOT_SUPPORTED: Unsupported Hive type: time # NOT_SUPPORTED: Unsupported Hive type: json - df = arrow_cursor.execute( - """ - SELECT - col_boolean - ,col_tinyint - ,col_smallint - ,col_int - ,col_bigint - ,col_float - ,col_double - ,col_string - ,col_varchar - ,col_timestamp - ,col_date - ,col_binary - ,col_array - ,col_map - ,col_struct - ,col_decimal - FROM one_row_complex - """ - ).as_polars() - assert df.height == 1 - assert df.width == 16 - dtypes = tuple(df.dtypes) - assert dtypes == ( - pl.Boolean, - pl.Int8, - pl.Int16, - pl.Int32, - pl.Int64, - pl.Float32, - pl.Float64, - pl.String, - pl.String, - pl.Datetime("ns"), - pl.Date, - pl.Binary, - pl.List(pl.Int32), - pl.List(pl.Struct([pl.Field("key", pl.Int32), pl.Field("value", pl.Int32)])), - pl.Struct([pl.Field("a", pl.Int32), pl.Field("b", pl.Int32)]), - pl.Decimal(precision=10, scale=1), - ) - rows = df.to_dicts() - assert rows == [ - { - "col_boolean": True, - "col_tinyint": 127, - "col_smallint": 32767, - "col_int": 2147483647, - "col_bigint": 9223372036854775807, - "col_float": 0.5, - "col_double": 0.25, - "col_string": "a string", - "col_varchar": "varchar", - "col_timestamp": datetime(2017, 1, 1, 0, 0, 0), - "col_date": datetime(2017, 1, 2).date(), - "col_binary": b"123", - "col_array": [1, 2], - "col_map": [{"key": 1, "value": 2}, {"key": 3, "value": 4}], - "col_struct": {"a": 1, "b": 2}, - "col_decimal": Decimal("0.1"), - } + expected = ExpectedResult(ONE_ROW_COMPLEX) + df = arrow_cursor.execute(expected.sql).as_polars() + assert list(zip(df.columns, df.dtypes, strict=True)) == [ + (f.name, polars_type(f.type)) for f in expected.arrow_schema(unload=True) ] + assert df.to_dicts() == expected.dicts(UNLOAD_POLARS_VALUES) def test_cancel(self, arrow_cursor): def cancel(c): diff --git a/tests/pyathena/expected.py b/tests/pyathena/expected.py index f6049cc8..1ddd07f5 100644 --- a/tests/pyathena/expected.py +++ b/tests/pyathena/expected.py @@ -10,11 +10,14 @@ - ``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()``. + give for it: ``rows()``, ``description()``, ``dbapi_types()``, and the column + types of pandas, Arrow, and Polars results. - 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 ``_*_TYPES`` dicts, ``arrow_type()``, and ``polars_type()`` give the + column types of the pandas, Arrow, and Polars results. 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 @@ -25,9 +28,14 @@ import re from collections.abc import Callable, Mapping from dataclasses import dataclass -from datetime import timezone +from datetime import datetime, time, timezone from typing import Any +import numpy as np +import pandas as pd +import polars as pl +import pyarrow as pa + from pyathena import BINARY, BOOLEAN, DATE, DATETIME, JSON, NUMBER, STRING, TIME from tests.pyathena.tables import Column, Table @@ -350,6 +358,21 @@ def _python_struct(value: Any, athena_type: str) -> dict[str, str | None]: return {k: _member_value(value[k], t) for k, t in _struct_fields(athena_type)} +def _json_text(value: Any, athena_type: str) -> str: + """Return a value cast to JSON, as the compact JSON text in a CSV result. + + Args: + value: The value from the table definition. + athena_type: The value's Athena type. + + Returns: + The JSON text; a map becomes an object with string keys. + """ + return json.dumps( + dict(value) if base_type(athena_type) == "map" else value, separators=(",", ":") + ) + + def _parsed_json(value: Any, athena_type: str) -> Any: """Return a value cast to JSON, after JSON parsing. @@ -360,10 +383,26 @@ def _parsed_json(value: Any, athena_type: str) -> Any: Returns: The parsed JSON value; a map becomes a dict with string keys. """ - return json.loads(json.dumps(dict(value) if base_type(athena_type) == "map" else value)) + return json.loads(_json_text(value, athena_type)) + + +def _native(value: Any, athena_type: str) -> Any: + """Return an array, map, or struct as read from Parquet. + + Args: + value: The value from the table definition. + athena_type: The array, map, or struct type. + + Returns: + An array as a list, a map as a list of key-value tuples, and a struct as + a dict. + """ + _check_scalar_members(athena_type) + return dict(value) if base_type(athena_type) == "struct" else list(value) -# Rows of Cursor, S3FSCursor, pyathena.pandas.util.as_pandas, and SQLAlchemy. +# Rows of Cursor, S3FSCursor, PolarsCursor, 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_VALUES: ValueRules = { @@ -376,6 +415,197 @@ def _parsed_json(value: Any, athena_type: str) -> Any: "json": _parsed_json, } +_TEXT_VALUES: ValueRules = dict.fromkeys(("array", "map", "struct"), _athena_text) +_NATIVE_VALUES: ValueRules = dict.fromkeys(("array", "map", "struct"), _native) + +# Rows and DataFrames of PandasCursor. +PANDAS_VALUES: ValueRules = { + **_SCALAR_VALUES, + **_TEXT_VALUES, + "timestamp": lambda v, t: pd.Timestamp(v), + "date": lambda v, t: pd.Timestamp(v), + "time": lambda v, t: v.time(), + "json": _parsed_json, +} + +# Rows of ArrowCursor. +ARROW_VALUES: ValueRules = { + **_SCALAR_VALUES, + **_TEXT_VALUES, + "time": lambda v, t: v.time(), + "json": _parsed_json, +} + +# Arrow tables and Polars DataFrames of ArrowCursor. Types without an Arrow +# equivalent in the CSV reader are strings, and dates are timestamps. +ARROW_TABLE_VALUES: ValueRules = { + **_SCALAR_VALUES, + **_TEXT_VALUES, + "time": lambda v, t: v.strftime("%H:%M:%S.%f")[:-3], + "date": lambda v, t: datetime.combine(v, time()), + "binary": lambda v, t: " ".join(f"{b:02x}" for b in v), + "json": _json_text, + "decimal": lambda v, t: str(v), +} + +# Rows, Arrow tables, and DataFrames of the pandas and Arrow cursors with +# unload=True, which read Parquet. Compare pandas rows through +# ExpectedResult.with_array_lists(). +UNLOAD_VALUES: ValueRules = { + **_SCALAR_VALUES, + **_NATIVE_VALUES, + "timestamp": lambda v, t: pd.Timestamp(v), +} + +# Polars DataFrames of ArrowCursor with unload=True, which hold a map as a list +# of key-value dicts. +UNLOAD_POLARS_VALUES: ValueRules = { + **_SCALAR_VALUES, + **_NATIVE_VALUES, + "map": lambda v, t: [{"key": k, "value": x} for k, x in _native(v, t)], +} + +# Column types that the PolarsCursor tests leave out. +POLARS_EXCLUDED_TYPES = frozenset(("binary", "array", "map", "struct")) + +# pandas, Arrow, and Polars column types + +# pandas 3 infers its "str" dtype for strings; pandas 2 uses object columns. +STRING_TYPE = pd.Series(["a"]).dtype.type + +# The dtype.type of a PandasCursor column, by base type. +_PANDAS_TYPES = { + "boolean": np.bool_, + "tinyint": np.int64, + "smallint": np.int64, + "int": np.int64, + "bigint": np.int64, + "float": np.float64, + "double": np.float64, + "string": STRING_TYPE, + "varchar": STRING_TYPE, + "timestamp": np.datetime64, + "date": np.datetime64, + "time": np.object_, + "binary": np.object_, + "array": STRING_TYPE, + "map": STRING_TYPE, + "struct": STRING_TYPE, + "json": np.object_, + "decimal": np.object_, +} +_PANDAS_UNLOAD_TYPES = { + **_PANDAS_TYPES, + "tinyint": np.int8, + "smallint": np.int16, + "int": np.int32, + "float": np.float32, + **dict.fromkeys(("date", "binary", "array", "map", "struct", "decimal"), np.object_), +} + +# The Arrow type of an ArrowCursor column, by base type. +_ARROW_TYPES = { + "boolean": pa.bool_(), + "tinyint": pa.int8(), + "smallint": pa.int16(), + "int": pa.int32(), + "bigint": pa.int64(), + "float": pa.float32(), + "double": pa.float64(), + "string": pa.string(), + "varchar": pa.string(), + "timestamp": pa.timestamp("ms"), + "date": pa.timestamp("ms"), + **dict.fromkeys(("time", "binary", "array", "map", "struct", "json", "decimal"), pa.string()), +} +_ARROW_UNLOAD_TYPES = { + **_ARROW_TYPES, + "timestamp": pa.timestamp("ns"), + "date": pa.date32(), + "binary": pa.binary(), +} + +# The Polars type of a PolarsCursor column, by base type. +_POLARS_TYPES = { + "boolean": pl.Boolean, + "tinyint": pl.Int8, + "smallint": pl.Int16, + "int": pl.Int32, + "bigint": pl.Int64, + "float": pl.Float32, + "double": pl.Float64, + "string": pl.String, + "varchar": pl.String, + "timestamp": pl.Datetime("us"), + "date": pl.Date, +} + + +def arrow_type(athena_type: str, unload: bool = False) -> pa.DataType: + """Return the Arrow type of an ArrowCursor column. + + Args: + athena_type: The column's Athena type. + unload: Whether the cursor reads Parquet written by UNLOAD. + + Returns: + The Arrow type. + """ + name = base_type(athena_type) + if not unload: + return _ARROW_TYPES[name] + if name == "decimal": + return pa.decimal128(*type_parameters(athena_type)) + if name == "array": + (element_type,) = type_arguments(athena_type) + return pa.list_(pa.field("array_element", arrow_type(element_type, unload))) + if name == "map": + key_type, value_type = (arrow_type(t, unload) for t in type_arguments(athena_type)) + return pa.map_(key_type, pa.field("entries", value_type)) + if name == "struct": + return pa.struct( + [pa.field(n, arrow_type(t, unload)) for n, t in _struct_fields(athena_type)] + ) + return _ARROW_UNLOAD_TYPES[name] + + +def polars_type(arrow: pa.DataType) -> pl.DataType: + """Return the Polars type of a DataFrame column converted from an Arrow type. + + Args: + arrow: The Arrow type. + + Returns: + The Polars type. + """ + if pa.types.is_timestamp(arrow): + return pl.Datetime(arrow.unit) + if pa.types.is_decimal(arrow): + return pl.Decimal(arrow.precision, arrow.scale) + if pa.types.is_map(arrow): + entry = [ + pl.Field("key", polars_type(arrow.key_type)), + pl.Field("value", polars_type(arrow.item_type)), + ] + return pl.List(pl.Struct(entry)) + if pa.types.is_list(arrow): + return pl.List(polars_type(arrow.value_type)) + if pa.types.is_struct(arrow): + return pl.Struct([pl.Field(f.name, polars_type(f.type)) for f in arrow]) + return { + pa.bool_(): pl.Boolean, + pa.int8(): pl.Int8, + pa.int16(): pl.Int16, + pa.int32(): pl.Int32, + pa.int64(): pl.Int64, + pa.float32(): pl.Float32, + pa.float64(): pl.Float64, + pa.string(): pl.String, + pa.date32(): pl.Date, + pa.binary(): pl.Binary, + }[arrow] + + # The expected result of a query @@ -386,15 +616,17 @@ class ExpectedResult: Attributes: table: The table. casts: The cast columns, after the table columns. + exclude_types: Base types of the table columns to leave out. """ table: Table casts: tuple[CastColumn, ...] = () + exclude_types: frozenset[str] = frozenset() @property def columns(self) -> list[Column]: """The selected table columns.""" - return list(self.table.columns) + return [c for c in self.table.columns if base_type(c.athena_type) not in self.exclude_types] @property def names(self) -> list[str]: @@ -429,18 +661,26 @@ def _types(self) -> list[str]: """ return [athena_type for athena_type, _, _ in self._sources()] - def description(self) -> list[tuple[Any, ...]]: + def description(self, unload: bool = False) -> list[tuple[Any, ...]]: """Return the expected cursor description. + Args: + unload: Whether the cursor reads Parquet written by UNLOAD, whose + columns are nullable and whose varchar columns lose their length. + Returns: One DB API description tuple per result column. """ result = [] 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")) + name_type = base_type(athena_type) + if unload and name_type == "varchar": + name_type = "string" + code, precision, scale = _DESCRIPTION[name_type] + if name_type in ("varchar", "decimal"): + precision, scale = (*type_parameters(athena_type), 0)[:2] + null_ok = "NULLABLE" if unload else "UNKNOWN" + result.append((name, code, None, None, precision, scale, null_ok)) return result def dbapi_types(self) -> list[Any]: @@ -469,3 +709,72 @@ def rows(self, values: ValueRules) -> list[tuple[Any, ...]]: ) for row in self.table.rows ] + + def with_array_lists(self, row: Any) -> tuple[Any, ...]: + """Return a result row with the NumPy arrays of array columns as lists. + + pandas returns array values read from Parquet as NumPy arrays, which do + not compare with ``==``. Other columns keep their values. + + Args: + row: A result row, or a pandas row from ``DataFrame.iterrows``. + + Returns: + The row as a tuple. + """ + return tuple( + list(v) if base_type(t) == "array" and isinstance(v, np.ndarray) else v + for t, v in zip(self._types(), row, strict=True) + ) + + def dicts(self, values: ValueRules) -> list[dict[str, Any]]: + """Return the expected rows as dicts, as ``polars.DataFrame.to_dicts`` returns them. + + Args: + values: The value rules of the cursor under test. + + Returns: + One dict per row of the table. + """ + return [dict(zip(self.names, row, strict=True)) for row in self.rows(values)] + + def pandas_types(self, unload: bool = False) -> list[type]: + """Return the expected ``dtype.type`` of each PandasCursor column. + + Args: + unload: Whether the cursor reads Parquet written by UNLOAD. + + Returns: + The types in result order. + """ + types_ = _PANDAS_UNLOAD_TYPES if unload else _PANDAS_TYPES + return [types_[base_type(t)] for t in self._types()] + + def arrow_schema(self, unload: bool = False) -> pa.Schema: + """Return the expected schema of an ArrowCursor table. + + Args: + unload: Whether the cursor reads Parquet written by UNLOAD. + + Returns: + The schema. + """ + return pa.schema( + [ + pa.field(n, arrow_type(t, unload)) + for n, t in zip(self.names, self._types(), strict=True) + ] + ) + + def polars_types(self) -> list[pl.DataType]: + """Return the expected type of each PolarsCursor column. + + Returns: + The types in result order. + """ + return [ + pl.Decimal(*type_parameters(t)) + if base_type(t) == "decimal" + else _POLARS_TYPES[base_type(t)] + for t in self._types() + ] diff --git a/tests/pyathena/pandas/test_cursor.py b/tests/pyathena/pandas/test_cursor.py index 0e19ff18..54038ca9 100644 --- a/tests/pyathena/pandas/test_cursor.py +++ b/tests/pyathena/pandas/test_cursor.py @@ -5,8 +5,6 @@ import string import time from concurrent.futures import ThreadPoolExecutor -from datetime import datetime -from decimal import Decimal from unittest.mock import PropertyMock, patch import numpy as np @@ -19,10 +17,18 @@ from pyathena.pandas.result_set import AthenaPandasResultSet, PandasDataFrameIterator from tests import ENV from tests.pyathena.conftest import connect +from tests.pyathena.expected import ( + ARRAY_JSON, + MAP_JSON, + PANDAS_VALUES, + TIME_OF_TIMESTAMP, + UNLOAD_VALUES, + ExpectedResult, +) +from tests.pyathena.tables import ONE_ROW_COMPLEX # pandas 3 infers its "str" dtype for strings, which represents NULL as NaN; pandas 2 uses # object columns with None. -STRING_TYPE = pd.Series(["a"]).dtype.type STRING_NULL = pd.Series(["a", None]).iloc[1] @@ -347,76 +353,10 @@ def test_invalid_arraysize(self, pandas_cursor): indirect=["pandas_cursor"], ) def test_complex(self, pandas_cursor, chunksize): - pandas_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 - """, - chunksize=chunksize, - ) - assert pandas_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 pandas_cursor.fetchall() == [ - ( - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - "varchar", - pd.Timestamp(2017, 1, 1, 0, 0, 0), - datetime(2017, 1, 1, 0, 0, 0).time(), - pd.Timestamp(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 = ExpectedResult(ONE_ROW_COMPLEX, casts=(TIME_OF_TIMESTAMP, ARRAY_JSON, MAP_JSON)) + pandas_cursor.execute(expected.sql, chunksize=chunksize) + assert pandas_cursor.description == expected.description() + assert pandas_cursor.fetchall() == expected.rows(PANDAS_VALUES) @pytest.mark.parametrize( ("pandas_cursor", "parquet_engine"), @@ -428,106 +368,12 @@ def test_complex(self, pandas_cursor, chunksize): ) def test_complex_unload_pyarrow(self, pandas_cursor, parquet_engine): # NOT_SUPPORTED: Unsupported Hive type: time, json - pandas_cursor.execute( - """ - SELECT - col_boolean - ,col_tinyint - ,col_smallint - ,col_int - ,col_bigint - ,col_float - ,col_double - ,col_string - ,col_varchar - ,col_timestamp - ,col_date - ,col_binary - ,col_array - ,col_map - ,col_struct - ,col_decimal - FROM one_row_complex - """, - engine=parquet_engine, - ) - assert pandas_cursor.description == [ - ("col_boolean", "boolean", None, None, 0, 0, "NULLABLE"), - ( - "col_tinyint", - "tinyint", - None, - None, - 3, - 0, - "NULLABLE", - ), - ( - "col_smallint", - "smallint", - None, - None, - 5, - 0, - "NULLABLE", - ), - ("col_int", "integer", None, None, 10, 0, "NULLABLE"), - ("col_bigint", "bigint", None, None, 19, 0, "NULLABLE"), - ("col_float", "float", None, None, 17, 0, "NULLABLE"), - ("col_double", "double", None, None, 17, 0, "NULLABLE"), - ("col_string", "varchar", None, None, 2147483647, 0, "NULLABLE"), - ("col_varchar", "varchar", None, None, 2147483647, 0, "NULLABLE"), - ("col_timestamp", "timestamp", None, None, 3, 0, "NULLABLE"), - ("col_date", "date", None, None, 0, 0, "NULLABLE"), - ("col_binary", "varbinary", None, None, 1073741824, 0, "NULLABLE"), - ("col_array", "array", None, None, 0, 0, "NULLABLE"), - ("col_map", "map", None, None, 0, 0, "NULLABLE"), - ("col_struct", "row", None, None, 0, 0, "NULLABLE"), - ("col_decimal", "decimal", None, None, 10, 1, "NULLABLE"), - ] - rows = [ - ( - row[0], - row[1], - row[2], - row[3], - row[4], - row[5], - row[6], - row[7], - row[8], - row[9], - row[10], - row[11], - list(row[12]), - row[13], - row[14], - row[15], - ) - for row in pandas_cursor.fetchall() - ] - assert rows == [ - ( - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - "varchar", - pd.Timestamp(2017, 1, 1, 0, 0, 0), - datetime(2017, 1, 2).date(), - b"123", - # ValueError: The truth value of an array with more than one element is ambiguous. - # Use a.any() or a.all() - list(np.array([1, 2], dtype=np.int32)), - [(1, 2), (3, 4)], - {"a": 1, "b": 2}, - Decimal("0.1"), - ) - ] + expected = ExpectedResult(ONE_ROW_COMPLEX) + pandas_cursor.execute(expected.sql, engine=parquet_engine) + assert pandas_cursor.description == expected.description(unload=True) + assert [ + expected.with_array_lists(row) for row in pandas_cursor.fetchall() + ] == expected.rows(UNLOAD_VALUES) def test_fetch_no_data(self, pandas_cursor): pytest.raises(ProgrammingError, pandas_cursor.fetchone) @@ -587,125 +433,13 @@ def test_many_as_pandas(self, pandas_cursor, parquet_engine, chunksize): indirect=["pandas_cursor"], ) def test_complex_as_pandas(self, pandas_cursor, chunksize): - df = pandas_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 - """, - chunksize=chunksize, - ).as_pandas() + expected = ExpectedResult(ONE_ROW_COMPLEX, casts=(TIME_OF_TIMESTAMP, ARRAY_JSON, MAP_JSON)) + df = pandas_cursor.execute(expected.sql, chunksize=chunksize).as_pandas() if chunksize: df = pd.concat((d for d in df), ignore_index=True) - assert df.shape[0] == 1 - assert df.shape[1] == 19 - dtypes = ( - df["col_boolean"].dtype.type, - df["col_tinyint"].dtype.type, - df["col_smallint"].dtype.type, - df["col_int"].dtype.type, - df["col_bigint"].dtype.type, - df["col_float"].dtype.type, - df["col_double"].dtype.type, - df["col_string"].dtype.type, - df["col_varchar"].dtype.type, - df["col_timestamp"].dtype.type, - df["col_time"].dtype.type, - df["col_date"].dtype.type, - df["col_binary"].dtype.type, - df["col_array"].dtype.type, - df["col_array_json"].dtype.type, - df["col_map"].dtype.type, - df["col_map_json"].dtype.type, - df["col_struct"].dtype.type, - df["col_decimal"].dtype.type, - ) - assert dtypes == ( - np.bool_, - np.int64, - np.int64, - np.int64, - np.int64, - np.float64, - np.float64, - STRING_TYPE, - STRING_TYPE, - np.datetime64, - np.object_, - np.datetime64, - np.object_, - STRING_TYPE, - np.object_, - STRING_TYPE, - np.object_, - STRING_TYPE, - np.object_, - ) - 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_varchar"], - 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() - ] - assert rows == [ - ( - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - "varchar", - pd.Timestamp(2017, 1, 1, 0, 0, 0), - datetime(2017, 1, 1, 0, 0, 0).time(), - pd.Timestamp(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 list(df.columns) == expected.names + assert [df[n].dtype.type for n in expected.names] == expected.pandas_types() + assert [tuple(row) for _, row in df.iterrows()] == expected.rows(PANDAS_VALUES) @pytest.mark.parametrize( ("pandas_cursor", "parquet_engine"), @@ -717,110 +451,13 @@ def test_complex_as_pandas(self, pandas_cursor, chunksize): ) def test_complex_unload_as_pandas_pyarrow(self, pandas_cursor, parquet_engine): # NOT_SUPPORTED: Unsupported Hive type: time, json - df = pandas_cursor.execute( - """ - SELECT - col_boolean - ,col_tinyint - ,col_smallint - ,col_int - ,col_bigint - ,col_float - ,col_double - ,col_string - ,col_varchar - ,col_timestamp - ,col_date - ,col_binary - ,col_array - ,col_map - ,col_struct - ,col_decimal - FROM one_row_complex - """, - engine=parquet_engine, - ).as_pandas() - assert df.shape[0] == 1 - assert df.shape[1] == 16 - dtypes = ( - df["col_boolean"].dtype.type, - df["col_tinyint"].dtype.type, - df["col_smallint"].dtype.type, - df["col_int"].dtype.type, - df["col_bigint"].dtype.type, - df["col_float"].dtype.type, - df["col_double"].dtype.type, - df["col_string"].dtype.type, - df["col_varchar"].dtype.type, - df["col_timestamp"].dtype.type, - df["col_date"].dtype.type, - df["col_binary"].dtype.type, - df["col_array"].dtype.type, - df["col_map"].dtype.type, - df["col_struct"].dtype.type, - df["col_decimal"].dtype.type, + expected = ExpectedResult(ONE_ROW_COMPLEX) + df = pandas_cursor.execute(expected.sql, engine=parquet_engine).as_pandas() + assert list(df.columns) == expected.names + assert [df[n].dtype.type for n in expected.names] == expected.pandas_types(unload=True) + assert [expected.with_array_lists(row) for _, row in df.iterrows()] == expected.rows( + UNLOAD_VALUES ) - assert dtypes == ( - np.bool_, - np.int8, - np.int16, - np.int32, - np.int64, - np.float32, - np.float64, - STRING_TYPE, - STRING_TYPE, - np.datetime64, - np.object_, - np.object_, - np.object_, - np.object_, - np.object_, - np.object_, - ) - 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_varchar"], - row["col_timestamp"], - row["col_date"], - row["col_binary"], - list(row["col_array"]), - row["col_map"], - row["col_struct"], - row["col_decimal"], - ) - for _, row in df.iterrows() - ] - assert rows == [ - ( - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - "varchar", - pd.Timestamp(2017, 1, 1, 0, 0, 0), - datetime(2017, 1, 2).date(), - b"123", - # ValueError: The truth value of an array with more than one element is ambiguous. - # Use a.any() or a.all() - list(np.array([1, 2], dtype=np.int32)), - [(1, 2), (3, 4)], - {"a": 1, "b": 2}, - Decimal("0.1"), - ) - ] def test_cancel(self, pandas_cursor): def cancel(c): diff --git a/tests/pyathena/polars/test_cursor.py b/tests/pyathena/polars/test_cursor.py index c9da6169..ac0f99de 100644 --- a/tests/pyathena/polars/test_cursor.py +++ b/tests/pyathena/polars/test_cursor.py @@ -10,8 +10,6 @@ import string import time from concurrent.futures import ThreadPoolExecutor -from datetime import datetime -from decimal import Decimal import polars as pl import pytest @@ -21,6 +19,8 @@ from pyathena.polars.result_set import AthenaPolarsResultSet from tests import ENV from tests.pyathena.conftest import connect +from tests.pyathena.expected import POLARS_EXCLUDED_TYPES, PYTHON_VALUES, ExpectedResult +from tests.pyathena.tables import ONE_ROW_COMPLEX class TestPolarsCursor: @@ -119,54 +119,10 @@ def test_many_as_polars(self, polars_cursor): assert df.width == 1 def test_complex(self, polars_cursor): - polars_cursor.execute( - """ - SELECT - col_boolean - ,col_tinyint - ,col_smallint - ,col_int - ,col_bigint - ,col_float - ,col_double - ,col_string - ,col_varchar - ,col_timestamp - ,col_date - ,col_decimal - FROM one_row_complex - """ - ) - assert polars_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_date", "date", None, None, 0, 0, "UNKNOWN"), - ("col_decimal", "decimal", None, None, 10, 1, "UNKNOWN"), - ] - assert polars_cursor.fetchall() == [ - ( - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - "varchar", - datetime(2017, 1, 1, 0, 0, 0), - datetime(2017, 1, 2).date(), - Decimal("0.1"), - ) - ] + expected = ExpectedResult(ONE_ROW_COMPLEX, exclude_types=POLARS_EXCLUDED_TYPES) + polars_cursor.execute(expected.sql) + assert polars_cursor.description == expected.description() + assert polars_cursor.fetchall() == expected.rows(PYTHON_VALUES) @pytest.mark.parametrize( "polars_cursor", @@ -178,104 +134,19 @@ def test_complex(self, polars_cursor): indirect=["polars_cursor"], ) def test_complex_unload(self, polars_cursor): - polars_cursor.execute( - """ - SELECT - col_boolean - ,col_tinyint - ,col_smallint - ,col_int - ,col_bigint - ,col_float - ,col_double - ,col_string - ,col_varchar - ,col_timestamp - ,col_date - ,col_decimal - FROM one_row_complex - """ - ) - assert polars_cursor.description == [ - ("col_boolean", "boolean", None, None, 0, 0, "NULLABLE"), - ("col_tinyint", "tinyint", None, None, 3, 0, "NULLABLE"), - ("col_smallint", "smallint", None, None, 5, 0, "NULLABLE"), - ("col_int", "integer", None, None, 10, 0, "NULLABLE"), - ("col_bigint", "bigint", None, None, 19, 0, "NULLABLE"), - ("col_float", "float", None, None, 17, 0, "NULLABLE"), - ("col_double", "double", None, None, 17, 0, "NULLABLE"), - ("col_string", "varchar", None, None, 2147483647, 0, "NULLABLE"), - ("col_varchar", "varchar", None, None, 2147483647, 0, "NULLABLE"), - ("col_timestamp", "timestamp", None, None, 3, 0, "NULLABLE"), - ("col_date", "date", None, None, 0, 0, "NULLABLE"), - ("col_decimal", "decimal", None, None, 10, 1, "NULLABLE"), - ] - assert polars_cursor.fetchall() == [ - ( - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - "varchar", - datetime(2017, 1, 1, 0, 0, 0), - datetime(2017, 1, 2).date(), - Decimal("0.1"), - ) - ] + expected = ExpectedResult(ONE_ROW_COMPLEX, exclude_types=POLARS_EXCLUDED_TYPES) + polars_cursor.execute(expected.sql) + assert polars_cursor.description == expected.description(unload=True) + assert polars_cursor.fetchall() == expected.rows(PYTHON_VALUES) def test_complex_as_polars(self, polars_cursor): - df = polars_cursor.execute( - """ - SELECT - col_boolean - ,col_tinyint - ,col_smallint - ,col_int - ,col_bigint - ,col_float - ,col_double - ,col_string - ,col_varchar - ,col_timestamp - ,col_date - ,col_decimal - FROM one_row_complex - """ - ).as_polars() + expected = ExpectedResult(ONE_ROW_COMPLEX, exclude_types=POLARS_EXCLUDED_TYPES) + df = polars_cursor.execute(expected.sql).as_polars() assert isinstance(df, pl.DataFrame) - assert (df.height, df.width) == (1, 12) - assert df.schema == { - "col_boolean": pl.Boolean, - "col_tinyint": pl.Int8, - "col_smallint": pl.Int16, - "col_int": pl.Int32, - "col_bigint": pl.Int64, - "col_float": pl.Float32, - "col_double": pl.Float64, - "col_string": pl.String, - "col_varchar": pl.String, - "col_timestamp": pl.Datetime("us"), - "col_date": pl.Date, - "col_decimal": pl.Decimal(precision=10, scale=1), - } - assert df.row(0) == ( - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - "varchar", - datetime(2017, 1, 1, 0, 0, 0), - datetime(2017, 1, 2).date(), - Decimal("0.1"), + assert list(df.schema.items()) == list( + zip(expected.names, expected.polars_types(), strict=True) ) + assert df.rows() == expected.rows(PYTHON_VALUES) @pytest.mark.parametrize( "polars_cursor", @@ -287,40 +158,11 @@ def test_complex_as_polars(self, polars_cursor): indirect=["polars_cursor"], ) def test_complex_unload_as_polars(self, polars_cursor): - df = polars_cursor.execute( - """ - SELECT - col_boolean - ,col_tinyint - ,col_smallint - ,col_int - ,col_bigint - ,col_float - ,col_double - ,col_string - ,col_varchar - ,col_timestamp - ,col_date - ,col_decimal - FROM one_row_complex - """ - ).as_polars() + expected = ExpectedResult(ONE_ROW_COMPLEX, exclude_types=POLARS_EXCLUDED_TYPES) + df = polars_cursor.execute(expected.sql).as_polars() assert isinstance(df, pl.DataFrame) - assert (df.height, df.width) == (1, 12) - assert df.row(0) == ( - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - "varchar", - datetime(2017, 1, 1, 0, 0, 0), - datetime(2017, 1, 2).date(), - Decimal("0.1"), - ) + assert df.columns == expected.names + assert df.rows() == expected.rows(PYTHON_VALUES) @pytest.mark.parametrize( "polars_cursor", diff --git a/tests/pyathena/tables.py b/tests/pyathena/tables.py index 55100442..0812fdeb 100644 --- a/tests/pyathena/tables.py +++ b/tests/pyathena/tables.py @@ -101,6 +101,12 @@ def data_file(self) -> tuple[str, bytes] | None: Returns: The file name and content, or None for a table without rows. + + Raises: + ValueError: If a Parquet table's row does not round-trip through the + columns' Arrow types unchanged, for example a timestamp with more + precision than the type or a struct without all of its fields. + The tests derive their expectations from the rows as written. """ if not self.rows: return None @@ -108,9 +114,11 @@ def data_file(self) -> tuple[str, bytes] | None: lines = ["\t".join(_to_text(v) for v in row) for row in self.rows] return "data.tsv", "".join(f"{line}\n" for line in lines).encode() schema = pa.schema([(c.name, c.arrow_type) for c in self.columns]) - table = pa.Table.from_pylist( - [dict(zip(schema.names, row, strict=True)) for row in self.rows], schema=schema - ) + rows = [dict(zip(schema.names, row, strict=True)) for row in self.rows] + table = pa.Table.from_pylist(rows, schema=schema) + for row, stored in zip(rows, table.to_pylist(), strict=True): + if changed := [n for n in schema.names if row[n] != stored[n]]: + raise ValueError(f"{self.name}: values of {changed} change in their Arrow type.") buffer = io.BytesIO() pq.write_table(table, buffer) return "data.parquet", buffer.getvalue() From 0afa83165d2b8798a3763dec977c5162641b50d2 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 00:13:22 +0900 Subject: [PATCH 2/2] Compare nulls in pandas UNLOAD integer arrays as None pandas reads a null in an integer array from Parquet as NaN, so with_array_lists turns NaN in array columns into None, as the derived expectation has it. Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/expected.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/tests/pyathena/expected.py b/tests/pyathena/expected.py index 1ddd07f5..6102a3ad 100644 --- a/tests/pyathena/expected.py +++ b/tests/pyathena/expected.py @@ -714,7 +714,8 @@ def with_array_lists(self, row: Any) -> tuple[Any, ...]: """Return a result row with the NumPy arrays of array columns as lists. pandas returns array values read from Parquet as NumPy arrays, which do - not compare with ``==``. Other columns keep their values. + not compare with ``==``, and a null in an integer array as NaN. The + arrays become lists with None for NaN. Other columns keep their values. Args: row: A result row, or a pandas row from ``DataFrame.iterrows``. @@ -723,7 +724,9 @@ def with_array_lists(self, row: Any) -> tuple[Any, ...]: The row as a tuple. """ return tuple( - list(v) if base_type(t) == "array" and isinstance(v, np.ndarray) else v + [None if isinstance(x, float) and np.isnan(x) else x for x in v] + if base_type(t) == "array" and isinstance(v, np.ndarray) + else v for t, v in zip(self._types(), row, strict=True) )