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..6102a3ad 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,75 @@ 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 ``==``, 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``. + + Returns: + The row as a tuple. + """ + return tuple( + [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) + ) + + 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()