diff --git a/tests/pyathena/expected.py b/tests/pyathena/expected.py new file mode 100644 index 00000000..f6049cc8 --- /dev/null +++ b/tests/pyathena/expected.py @@ -0,0 +1,471 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT +"""Expected query results for the shared tables in ``tests.pyathena.tables``. + +How to read this module: + +- ``ExpectedResult`` is the entry point. It builds a query of a table's columns, + optionally followed by ``CastColumn`` items, and returns what a cursor should + give for it: ``rows()``, ``description()``, and ``dbapi_types()``. +- The ``*_VALUES`` dicts state how a cursor returns a value of each type. They + are keyed by base type, the Athena type without its parameters, such as + ``decimal`` for ``DECIMAL(10,1)``. Each rule takes the value from the table + definition and the column's Athena type. + +The rules state what the cursors return; they never call PyAthena's converters. +They cover the column types of the shared tables. A column of another type +raises ``KeyError`` or ``NotImplementedError`` until a rule is added here. +""" + +import json +import re +from collections.abc import Callable, Mapping +from dataclasses import dataclass +from datetime import timezone +from typing import Any + +from pyathena import BINARY, BOOLEAN, DATE, DATETIME, JSON, NUMBER, STRING, TIME +from tests.pyathena.tables import Column, Table + +# Athena types + + +def base_type(athena_type: str) -> str: + """Return an Athena type without its parameters. + + Args: + athena_type: An Athena type, such as ``DECIMAL(10,1)`` or ``ARRAY``. + + Returns: + The lowercase base type, such as ``decimal`` or ``array``. + """ + match = re.match(r"[a-z]+(?: [a-z]+)*", athena_type.strip().lower()) + assert match, athena_type + return match.group() + + +def type_parameters(athena_type: str) -> tuple[int, ...]: + """Return the numeric parameters of an Athena type. + + Args: + athena_type: An Athena type, such as ``DECIMAL(10,1)``. + + Returns: + The parameters of the outer type, such as ``(10, 1)``; empty for a type + without them, such as ``ARRAY``. + """ + match = re.match(r"\s*[a-zA-Z ]+\(([\d,\s]+)\)", athena_type) + return tuple(int(p) for p in match.group(1).split(",")) if match else () + + +def type_arguments(athena_type: str) -> list[str]: + """Return the types nested in an array, map, or struct type. + + Args: + athena_type: An Athena type, such as ``MAP``. + + Returns: + The element type of an array, the key and value types of a map, or the + field types of a struct; empty for other types. + """ + if base_type(athena_type) == "struct": + return [field_type for _, field_type in _struct_fields(athena_type)] + return _split_arguments(athena_type) + + +def _split_arguments(athena_type: str) -> list[str]: + """Split the text between a type's outer angle brackets at top-level commas. + + Args: + athena_type: An Athena type, such as ``STRUCT``. + + Returns: + The arguments, such as ``["a: int", "b: int"]``; empty for a type + without angle brackets. + """ + if "<" not in athena_type: + return [] + inner = athena_type[athena_type.index("<") + 1 : athena_type.rindex(">")] + arguments, depth, start = [], 0, 0 + for i, char in enumerate(inner): + if char in "<(": + depth += 1 + elif char in ">)": + depth -= 1 + elif char == "," and depth == 0: + arguments.append(inner[start:i].strip()) + start = i + 1 + arguments.append(inner[start:].strip()) + return arguments + + +def _struct_fields(athena_type: str) -> list[tuple[str, str]]: + """Return the fields of a struct type. + + Args: + athena_type: A struct type, such as ``STRUCT``. + + Returns: + ``(name, type)`` pairs. + """ + return [ + (name.strip(), field_type.strip()) + for name, field_type in (a.split(":", 1) for a in _split_arguments(athena_type)) + ] + + +# Cast columns + + +@dataclass(frozen=True) +class CastColumn: + """A result column that the query casts from a table column. + + Attributes: + name: The result column name. + source: The table column to cast. + athena_type: The type to cast to. + """ + + name: str + source: str + athena_type: str + + def sql(self) -> str: + """Return the select-list item. + + Returns: + The ``CAST`` expression with its alias. + """ + return f"CAST({self.source} AS {self.athena_type}) AS {self.name}" + + +TIMESTAMP_TZ = CastColumn("col_timestamp_tz", "col_timestamp", "timestamp with time zone") +TIME_OF_TIMESTAMP = CastColumn("col_time", "col_timestamp", "time") +ARRAY_JSON = CastColumn("col_array_json", "col_array", "json") +MAP_JSON = CastColumn("col_map_json", "col_map", "json") + +# Cursor descriptions + +# (type code, precision, scale) in the cursor description, by base type. The +# precision of varchar(n) and the precision and scale of decimal(p,s) come from +# the type parameters. +_DESCRIPTION: Mapping[str, tuple[str, int, int]] = { + "boolean": ("boolean", 0, 0), + "tinyint": ("tinyint", 3, 0), + "smallint": ("smallint", 5, 0), + "int": ("integer", 10, 0), + "bigint": ("bigint", 19, 0), + "float": ("float", 17, 0), + "double": ("double", 17, 0), + "string": ("varchar", 2147483647, 0), + "varchar": ("varchar", 0, 0), + "timestamp": ("timestamp", 3, 0), + "timestamp with time zone": ("timestamp with time zone", 3, 0), + "time": ("time", 3, 0), + "date": ("date", 0, 0), + "binary": ("varbinary", 1073741824, 0), + "array": ("array", 0, 0), + "map": ("map", 0, 0), + "struct": ("row", 0, 0), + "decimal": ("decimal", 0, 0), + "json": ("json", 0, 0), +} + +_DBAPI_TYPES = { + "boolean": BOOLEAN, + "tinyint": NUMBER, + "smallint": NUMBER, + "int": NUMBER, + "bigint": NUMBER, + "float": NUMBER, + "double": NUMBER, + "decimal": NUMBER, + "string": STRING, + "varchar": STRING, + "array": STRING, + "map": STRING, + "struct": STRING, + "timestamp": DATETIME, + "timestamp with time zone": DATETIME, + "time": TIME, + "date": DATE, + "binary": BINARY, + "json": JSON, +} + +# Value rules + +# A rule(value, athena_type) returns what a cursor gives for a value from the +# table definition. A cast column uses the rule of the type it casts to, with +# the value and type of its source column. +ValueRules = Mapping[str, Callable[[Any, str], Any]] + + +def _same(value: Any, athena_type: str) -> Any: + """Return the value unchanged. + + Args: + value: The value from the table definition. + athena_type: The value's Athena type. + + Returns: + The value. + """ + return value + + +# Scalar types that the cursors return as the Python value itself. +_SCALAR_VALUES: ValueRules = dict.fromkeys( + ( + "boolean", + "tinyint", + "smallint", + "int", + "bigint", + "float", + "double", + "string", + "varchar", + "timestamp", + "date", + "binary", + "decimal", + ), + _same, +) + + +def _athena_text(value: Any, athena_type: str) -> str: + """Return a value as Athena renders it in a CSV result. + + Args: + value: The value from the table definition. + athena_type: The value's Athena type. + + Returns: + The text, such as ``[1, 2]`` for an array, ``{1=2, 3=4}`` for a map, + and ``{a=1, b=2}`` for a struct. + + Raises: + NotImplementedError: For a nested scalar type other than integers and + strings, whose rendering has no rule here. + """ + name = base_type(athena_type) + if value is None: + return "null" + if name == "array": + (element_type,) = type_arguments(athena_type) + return f"[{', '.join(_athena_text(e, element_type) for e in value)}]" + if name == "map": + key_type, value_type = type_arguments(athena_type) + entries = (f"{_athena_text(k, key_type)}={_athena_text(x, value_type)}" for k, x in value) + return f"{{{', '.join(entries)}}}" + if name == "struct": + entries = (f"{k}={_athena_text(value[k], t)}" for k, t in _struct_fields(athena_type)) + return f"{{{', '.join(entries)}}}" + if name in ("tinyint", "smallint", "int", "bigint", "string", "varchar"): + return str(value) + raise NotImplementedError(f"No text rendering rule for {athena_type}.") + + +def _member_value(value: Any, athena_type: str) -> str | None: + """Return a value nested in an array, map, or struct as the cursors parse it. + + Args: + value: The nested value. + athena_type: The value's Athena type. + + Returns: + Its text rendering, or None for a null. + """ + return None if value is None else _athena_text(value, athena_type) + + +def _check_scalar_members(athena_type: str) -> None: + """Check that an array, map, or struct type nests only scalar types. + + Args: + athena_type: The array, map, or struct type. + + Raises: + NotImplementedError: For a nested array, map, or struct, which the + cursors parse differently. + """ + if any(base_type(t) in ("array", "map", "struct") for t in type_arguments(athena_type)): + raise NotImplementedError(f"No rule for the nesting in {athena_type}.") + + +def _python_array(value: Any, athena_type: str) -> list[Any]: + """Return an array as a cursor without type hints returns it. + + Args: + value: The value from the table definition. + athena_type: The array type. + + Returns: + The parsed JSON if Athena's rendering is valid JSON, such as ``[1, 2]``; + otherwise the elements' text, such as ``["a", "b"]`` for ``[a, b]``, + with None for a null element. + """ + _check_scalar_members(athena_type) + (element_type,) = type_arguments(athena_type) + try: + return json.loads(_athena_text(value, athena_type)) + except ValueError: + return [_member_value(e, element_type) for e in value] + + +def _python_map(value: Any, athena_type: str) -> dict[str, str | None]: + """Return a map as a cursor without type hints returns it. + + Args: + value: The value from the table definition. + athena_type: The map type. + + Returns: + The keys and values as their text renderings; a null value is None. + """ + _check_scalar_members(athena_type) + key_type, value_type = type_arguments(athena_type) + return {_athena_text(k, key_type): _member_value(x, value_type) for k, x in value} + + +def _python_struct(value: Any, athena_type: str) -> dict[str, str | None]: + """Return a struct as a cursor without type hints returns it. + + Args: + value: The value from the table definition. + athena_type: The struct type. + + Returns: + The field values as their text renderings, in declared order; a null + value is None. + """ + _check_scalar_members(athena_type) + return {k: _member_value(value[k], t) for k, t in _struct_fields(athena_type)} + + +def _parsed_json(value: Any, athena_type: str) -> Any: + """Return a value cast to JSON, after JSON parsing. + + Args: + value: The value from the table definition. + athena_type: The value's Athena type. + + Returns: + The 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)) + + +# Rows of Cursor, S3FSCursor, pyathena.pandas.util.as_pandas, and SQLAlchemy. +# Without type hints, the cursors parse Athena's text rendering of arrays, maps, +# and structs, so nested values that are not JSON are strings. +PYTHON_VALUES: ValueRules = { + **_SCALAR_VALUES, + "timestamp with time zone": lambda v, t: v.replace(tzinfo=timezone.utc), + "time": lambda v, t: v.time(), + "array": _python_array, + "map": _python_map, + "struct": _python_struct, + "json": _parsed_json, +} + +# The expected result of a query + + +@dataclass(frozen=True) +class ExpectedResult: + """A query of a table's columns, followed by cast columns, and its expected result. + + Attributes: + table: The table. + casts: The cast columns, after the table columns. + """ + + table: Table + casts: tuple[CastColumn, ...] = () + + @property + def columns(self) -> list[Column]: + """The selected table columns.""" + return list(self.table.columns) + + @property + def names(self) -> list[str]: + """The result column names.""" + return [c.name for c in self.columns] + [c.name for c in self.casts] + + @property + def sql(self) -> str: + """The query.""" + items = [c.name for c in self.columns] + [c.sql() for c in self.casts] + return f"SELECT {', '.join(items)} FROM {self.table.name}" + + def _sources(self) -> list[tuple[str, int, str]]: + """Return where each result column's value comes from. + + Returns: + ``(athena_type, source_index, source_type)`` per result column: the + column's type, and the index and type of the table column whose value + it holds. A table column is its own source. + """ + index = {c.name: i for i, c in enumerate(self.table.columns)} + by_name = {c.name: c for c in self.table.columns} + return [(c.athena_type, index[c.name], c.athena_type) for c in self.columns] + [ + (c.athena_type, index[c.source], by_name[c.source].athena_type) for c in self.casts + ] + + def _types(self) -> list[str]: + """Return the Athena type of each result column. + + Returns: + The types in result order. + """ + return [athena_type for athena_type, _, _ in self._sources()] + + def description(self) -> list[tuple[Any, ...]]: + """Return the expected cursor description. + + 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")) + return result + + def dbapi_types(self) -> list[Any]: + """Return the expected DB API type object of each result column. + + Returns: + The type objects in result order. + """ + return [_DBAPI_TYPES[base_type(t)] for t in self._types()] + + def rows(self, values: ValueRules) -> list[tuple[Any, ...]]: + """Return the expected rows. + + Args: + values: The value rules of the cursor under test, such as + ``PYTHON_VALUES``. + + Returns: + One tuple per row of the table. + """ + sources = self._sources() + return [ + tuple( + values[base_type(athena_type)](row[index], source_type) + for athena_type, index, source_type in sources + ) + for row in self.table.rows + ] diff --git a/tests/pyathena/pandas/test_util.py b/tests/pyathena/pandas/test_util.py index 1acbc07e..6382d637 100644 --- a/tests/pyathena/pandas/test_util.py +++ b/tests/pyathena/pandas/test_util.py @@ -1,7 +1,6 @@ import textwrap import uuid from datetime import date, datetime -from decimal import Decimal import numpy as np import pandas as pd @@ -16,6 +15,14 @@ to_sql, ) from tests import ENV +from tests.pyathena.expected import ( + ARRAY_JSON, + MAP_JSON, + PYTHON_VALUES, + TIME_OF_TIMESTAMP, + ExpectedResult, +) +from tests.pyathena.tables import ONE_ROW_COMPLEX def test_get_chunks(): @@ -52,77 +59,11 @@ def test_reset_index(): def test_as_pandas(cursor): - cursor.execute( - """ - SELECT - col_boolean - , col_tinyint - , col_smallint - , col_int - , col_bigint - , col_float - , col_double - , col_string - , col_timestamp - , CAST(col_timestamp AS time) AS col_time - , col_date - , col_binary - , col_array - , CAST(col_array AS json) AS col_array_json - , col_map - , CAST(col_map AS json) AS col_map_json - , col_struct - , col_decimal - FROM one_row_complex - """ - ) + expected = ExpectedResult(ONE_ROW_COMPLEX, casts=(TIME_OF_TIMESTAMP, ARRAY_JSON, MAP_JSON)) + cursor.execute(expected.sql) df = as_pandas(cursor) - rows = [ - ( - row["col_boolean"], - row["col_tinyint"], - row["col_smallint"], - row["col_int"], - row["col_bigint"], - row["col_float"], - row["col_double"], - row["col_string"], - row["col_timestamp"], - row["col_time"], - row["col_date"], - row["col_binary"], - row["col_array"], - row["col_array_json"], - row["col_map"], - row["col_map_json"], - row["col_struct"], - row["col_decimal"], - ) - for _, row in df.iterrows() - ] - expected = [ - ( - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - datetime(2017, 1, 1, 0, 0, 0), - datetime(2017, 1, 1, 0, 0, 0).time(), - date(2017, 1, 2), - b"123", - [1, 2], - [1, 2], - {"1": "2", "3": "4"}, - {"1": 2, "3": 4}, - {"a": "1", "b": "2"}, - Decimal("0.1"), - ) - ] - assert rows == expected + assert list(df.columns) == expected.names + assert [tuple(row) for _, row in df.iterrows()] == expected.rows(PYTHON_VALUES) def test_as_pandas_integer_na_values(cursor): diff --git a/tests/pyathena/s3fs/test_async_cursor.py b/tests/pyathena/s3fs/test_async_cursor.py index ffc96dfd..eaf7eca4 100644 --- a/tests/pyathena/s3fs/test_async_cursor.py +++ b/tests/pyathena/s3fs/test_async_cursor.py @@ -8,8 +8,6 @@ import contextlib import random import time -from datetime import datetime -from decimal import Decimal import pytest @@ -19,6 +17,14 @@ from pyathena.s3fs.result_set import AthenaS3FSResultSet from tests import ENV from tests.pyathena.conftest import connect +from tests.pyathena.expected import ( + ARRAY_JSON, + MAP_JSON, + PYTHON_VALUES, + TIME_OF_TIMESTAMP, + ExpectedResult, +) +from tests.pyathena.tables import ONE_ROW_COMPLEX class TestAsyncS3FSCursor: @@ -61,76 +67,11 @@ def test_invalid_arraysize(self, async_s3fs_cursor): async_s3fs_cursor.arraysize = -1 def test_complex(self, async_s3fs_cursor): - query_id, future = async_s3fs_cursor.execute( - """ - SELECT - col_boolean - ,col_tinyint - ,col_smallint - ,col_int - ,col_bigint - ,col_float - ,col_double - ,col_string - ,col_varchar - ,col_timestamp - ,CAST(col_timestamp AS time) AS col_time - ,col_date - ,col_binary - ,col_array - ,CAST(col_array AS json) AS col_array_json - ,col_map - ,CAST(col_map AS json) AS col_map_json - ,col_struct - ,col_decimal - FROM one_row_complex - """ - ) + expected = ExpectedResult(ONE_ROW_COMPLEX, casts=(TIME_OF_TIMESTAMP, ARRAY_JSON, MAP_JSON)) + query_id, future = async_s3fs_cursor.execute(expected.sql) result_set = future.result() - assert result_set.description == [ - ("col_boolean", "boolean", None, None, 0, 0, "UNKNOWN"), - ("col_tinyint", "tinyint", None, None, 3, 0, "UNKNOWN"), - ("col_smallint", "smallint", None, None, 5, 0, "UNKNOWN"), - ("col_int", "integer", None, None, 10, 0, "UNKNOWN"), - ("col_bigint", "bigint", None, None, 19, 0, "UNKNOWN"), - ("col_float", "float", None, None, 17, 0, "UNKNOWN"), - ("col_double", "double", None, None, 17, 0, "UNKNOWN"), - ("col_string", "varchar", None, None, 2147483647, 0, "UNKNOWN"), - ("col_varchar", "varchar", None, None, 10, 0, "UNKNOWN"), - ("col_timestamp", "timestamp", None, None, 3, 0, "UNKNOWN"), - ("col_time", "time", None, None, 3, 0, "UNKNOWN"), - ("col_date", "date", None, None, 0, 0, "UNKNOWN"), - ("col_binary", "varbinary", None, None, 1073741824, 0, "UNKNOWN"), - ("col_array", "array", None, None, 0, 0, "UNKNOWN"), - ("col_array_json", "json", None, None, 0, 0, "UNKNOWN"), - ("col_map", "map", None, None, 0, 0, "UNKNOWN"), - ("col_map_json", "json", None, None, 0, 0, "UNKNOWN"), - ("col_struct", "row", None, None, 0, 0, "UNKNOWN"), - ("col_decimal", "decimal", None, None, 10, 1, "UNKNOWN"), - ] - assert result_set.fetchall() == [ - ( - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - "varchar", - datetime(2017, 1, 1, 0, 0, 0), - datetime(2017, 1, 1, 0, 0, 0).time(), - datetime(2017, 1, 2).date(), - b"123", - [1, 2], - [1, 2], - {"1": "2", "3": "4"}, - {"1": 2, "3": 4}, - {"a": "1", "b": "2"}, - Decimal("0.1"), - ) - ] + assert result_set.description == expected.description() + assert result_set.fetchall() == expected.rows(PYTHON_VALUES) def test_cancel(self, async_s3fs_cursor): query_id, future = async_s3fs_cursor.execute( diff --git a/tests/pyathena/s3fs/test_cursor.py b/tests/pyathena/s3fs/test_cursor.py index 40edaa07..74929fbb 100644 --- a/tests/pyathena/s3fs/test_cursor.py +++ b/tests/pyathena/s3fs/test_cursor.py @@ -3,8 +3,6 @@ import string import time from concurrent.futures import ThreadPoolExecutor -from datetime import datetime -from decimal import Decimal import pytest @@ -14,6 +12,14 @@ from pyathena.s3fs.result_set import AthenaS3FSResultSet from tests import ENV from tests.pyathena.conftest import connect +from tests.pyathena.expected import ( + ARRAY_JSON, + MAP_JSON, + PYTHON_VALUES, + TIME_OF_TIMESTAMP, + ExpectedResult, +) +from tests.pyathena.tables import ONE_ROW_COMPLEX class TestS3FSCursor: @@ -55,75 +61,10 @@ def test_invalid_arraysize(self, s3fs_cursor): s3fs_cursor.arraysize = -1 def test_complex(self, s3fs_cursor): - s3fs_cursor.execute( - """ - SELECT - col_boolean - ,col_tinyint - ,col_smallint - ,col_int - ,col_bigint - ,col_float - ,col_double - ,col_string - ,col_varchar - ,col_timestamp - ,CAST(col_timestamp AS time) AS col_time - ,col_date - ,col_binary - ,col_array - ,CAST(col_array AS json) AS col_array_json - ,col_map - ,CAST(col_map AS json) AS col_map_json - ,col_struct - ,col_decimal - FROM one_row_complex - """ - ) - assert s3fs_cursor.description == [ - ("col_boolean", "boolean", None, None, 0, 0, "UNKNOWN"), - ("col_tinyint", "tinyint", None, None, 3, 0, "UNKNOWN"), - ("col_smallint", "smallint", None, None, 5, 0, "UNKNOWN"), - ("col_int", "integer", None, None, 10, 0, "UNKNOWN"), - ("col_bigint", "bigint", None, None, 19, 0, "UNKNOWN"), - ("col_float", "float", None, None, 17, 0, "UNKNOWN"), - ("col_double", "double", None, None, 17, 0, "UNKNOWN"), - ("col_string", "varchar", None, None, 2147483647, 0, "UNKNOWN"), - ("col_varchar", "varchar", None, None, 10, 0, "UNKNOWN"), - ("col_timestamp", "timestamp", None, None, 3, 0, "UNKNOWN"), - ("col_time", "time", None, None, 3, 0, "UNKNOWN"), - ("col_date", "date", None, None, 0, 0, "UNKNOWN"), - ("col_binary", "varbinary", None, None, 1073741824, 0, "UNKNOWN"), - ("col_array", "array", None, None, 0, 0, "UNKNOWN"), - ("col_array_json", "json", None, None, 0, 0, "UNKNOWN"), - ("col_map", "map", None, None, 0, 0, "UNKNOWN"), - ("col_map_json", "json", None, None, 0, 0, "UNKNOWN"), - ("col_struct", "row", None, None, 0, 0, "UNKNOWN"), - ("col_decimal", "decimal", None, None, 10, 1, "UNKNOWN"), - ] - assert s3fs_cursor.fetchall() == [ - ( - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - "varchar", - datetime(2017, 1, 1, 0, 0, 0), - datetime(2017, 1, 1, 0, 0, 0).time(), - datetime(2017, 1, 2).date(), - b"123", - [1, 2], - [1, 2], - {"1": "2", "3": "4"}, - {"1": 2, "3": 4}, - {"a": "1", "b": "2"}, - Decimal("0.1"), - ) - ] + expected = ExpectedResult(ONE_ROW_COMPLEX, casts=(TIME_OF_TIMESTAMP, ARRAY_JSON, MAP_JSON)) + s3fs_cursor.execute(expected.sql) + assert s3fs_cursor.description == expected.description() + assert s3fs_cursor.fetchall() == expected.rows(PYTHON_VALUES) def test_cancel(self, s3fs_cursor): def cancel(c): diff --git a/tests/pyathena/sqlalchemy/test_base.py b/tests/pyathena/sqlalchemy/test_base.py index e086e7d6..b838a9cf 100644 --- a/tests/pyathena/sqlalchemy/test_base.py +++ b/tests/pyathena/sqlalchemy/test_base.py @@ -34,7 +34,9 @@ ) from pyathena.util import RetryConfig from tests.pyathena.conftest import ENV -from tests.pyathena.util import decorated, throttle_metadata_api +from tests.pyathena.expected import PYTHON_VALUES, ExpectedResult +from tests.pyathena.tables import ONE_ROW_COMPLEX +from tests.pyathena.util import assert_sqlalchemy_type, decorated, throttle_metadata_api # Amazon S3 Tables tests need a pre-provisioned table-bucket catalog; the session # creates its own namespace in it. @@ -1394,53 +1396,12 @@ def test_filter_func(self, engine): def test_reflect_select(self, engine): engine, conn = engine one_row_complex = Table("one_row_complex", MetaData(schema=ENV.schema), autoload_with=conn) - assert len(one_row_complex.c) == 16 + assert [c.name for c in one_row_complex.c] == [c.name for c in ONE_ROW_COMPLEX.columns] assert isinstance(one_row_complex.c.col_string, Column) rows = conn.execute(one_row_complex.select()).fetchall() - assert len(rows) == 1 - assert list(rows[0]) == [ - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - "varchar", - datetime(2017, 1, 1, 0, 0, 0), - date(2017, 1, 2), - b"123", - [1, 2], - {"1": "2", "3": "4"}, # map type now converted to dict - {"a": "1", "b": "2"}, # row type now converted to dict - Decimal("0.1"), - ] - assert isinstance(one_row_complex.c.col_boolean.type, types.BOOLEAN) - assert isinstance(one_row_complex.c.col_tinyint.type, TINYINT) - assert isinstance(one_row_complex.c.col_smallint.type, types.SMALLINT) - assert isinstance(one_row_complex.c.col_int.type, types.INTEGER) - assert isinstance(one_row_complex.c.col_bigint.type, types.BIGINT) - assert isinstance(one_row_complex.c.col_float.type, types.FLOAT) - assert isinstance(one_row_complex.c.col_double.type, types.DOUBLE) - assert isinstance(one_row_complex.c.col_string.type, types.String) - assert isinstance(one_row_complex.c.col_varchar.type, types.VARCHAR) - assert one_row_complex.c.col_varchar.type.length == 10 - assert isinstance(one_row_complex.c.col_timestamp.type, types.TIMESTAMP) - assert isinstance(one_row_complex.c.col_date.type, types.DATE) - assert isinstance(one_row_complex.c.col_binary.type, types.BINARY) - assert isinstance(one_row_complex.c.col_array.type, AthenaArray) - assert isinstance(one_row_complex.c.col_array.type.item_type, types.INTEGER) - assert isinstance(one_row_complex.c.col_map.type, types.String) - # With struct support, col_struct should now be recognized as AthenaStruct - - assert isinstance(one_row_complex.c.col_struct.type, AthenaStruct) - assert isinstance( - one_row_complex.c.col_decimal.type, - types.DECIMAL, - ) - assert one_row_complex.c.col_decimal.type.precision == 10 - assert one_row_complex.c.col_decimal.type.scale == 1 + assert [tuple(row) for row in rows] == ExpectedResult(ONE_ROW_COMPLEX).rows(PYTHON_VALUES) + for column in ONE_ROW_COMPLEX.columns: + assert_sqlalchemy_type(one_row_complex.c[column.name].type, column) def test_select_offset_limit(self, engine): engine, conn = engine diff --git a/tests/pyathena/tables.py b/tests/pyathena/tables.py index bc8a54fc..55100442 100644 --- a/tests/pyathena/tables.py +++ b/tests/pyathena/tables.py @@ -156,6 +156,58 @@ def _to_text(value: Any) -> str: return str(value) +# One row with a value of each column type. Tests derive their expected results +# from this definition with tests.pyathena.expected. Keep strings nested in +# arrays, maps, and structs to words of letters, digits, and underscores +# separated by single spaces, other than null, because the cursors parse +# Athena's text rendering of those types. +ONE_ROW_COMPLEX = Table( + "one_row_complex", + ( + Column("col_boolean", "BOOLEAN", pa.bool_()), + Column("col_tinyint", "TINYINT", pa.int8()), + Column("col_smallint", "SMALLINT", pa.int16()), + Column("col_int", "INT", pa.int32()), + Column("col_bigint", "BIGINT", pa.int64()), + Column("col_float", "FLOAT", pa.float32()), + Column("col_double", "DOUBLE", pa.float64()), + Column("col_string", "STRING", pa.string()), + Column("col_varchar", "VARCHAR(10)", pa.string()), + Column("col_timestamp", "TIMESTAMP", pa.timestamp("ms")), + Column("col_date", "DATE", pa.date32()), + Column("col_binary", "BINARY", pa.binary()), + Column("col_array", "ARRAY", pa.list_(pa.int32())), + Column("col_map", "MAP", pa.map_(pa.int32(), pa.int32())), + Column( + "col_struct", + "STRUCT", + pa.struct([("a", pa.int32()), ("b", pa.int32())]), + ), + Column("col_decimal", "DECIMAL(10,1)", pa.decimal128(10, 1)), + ), + rows=( + ( + True, + 127, + 32767, + 2147483647, + 9223372036854775807, + 0.5, + 0.25, + "a string", + "varchar", + datetime(2017, 1, 1, 0, 0, 0), + date(2017, 1, 2), + b"123", + [1, 2], + [(1, 2), (3, 4)], + {"a": 1, "b": 2}, + Decimal("0.1"), + ), + ), +) + + TABLES = ( # A text table: the reflection tests assert its SerDe and delimiters. Table( @@ -173,51 +225,7 @@ def _to_text(value: Any) -> str: rows=tuple((i,) for i in range(10000)), storage="text", ), - Table( - "one_row_complex", - ( - Column("col_boolean", "BOOLEAN", pa.bool_()), - Column("col_tinyint", "TINYINT", pa.int8()), - Column("col_smallint", "SMALLINT", pa.int16()), - Column("col_int", "INT", pa.int32()), - Column("col_bigint", "BIGINT", pa.int64()), - Column("col_float", "FLOAT", pa.float32()), - Column("col_double", "DOUBLE", pa.float64()), - Column("col_string", "STRING", pa.string()), - Column("col_varchar", "VARCHAR(10)", pa.string()), - Column("col_timestamp", "TIMESTAMP", pa.timestamp("ms")), - Column("col_date", "DATE", pa.date32()), - Column("col_binary", "BINARY", pa.binary()), - Column("col_array", "ARRAY", pa.list_(pa.int32())), - Column("col_map", "MAP", pa.map_(pa.int32(), pa.int32())), - Column( - "col_struct", - "STRUCT", - pa.struct([("a", pa.int32()), ("b", pa.int32())]), - ), - Column("col_decimal", "DECIMAL(10,1)", pa.decimal128(10, 1)), - ), - rows=( - ( - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - "varchar", - datetime(2017, 1, 1, 0, 0, 0), - date(2017, 1, 2), - b"123", - [1, 2], - [(1, 2), (3, 4)], - {"a": 1, "b": 2}, - Decimal("0.1"), - ), - ), - ), + ONE_ROW_COMPLEX, Table( "partition_table", (Column("a", "STRING", pa.string()),), diff --git a/tests/pyathena/test_cursor.py b/tests/pyathena/test_cursor.py index aa1832b3..ad9b65ac 100644 --- a/tests/pyathena/test_cursor.py +++ b/tests/pyathena/test_cursor.py @@ -9,7 +9,7 @@ import uuid from concurrent import futures from concurrent.futures.thread import ThreadPoolExecutor -from datetime import date, datetime, timezone +from datetime import datetime, timezone from decimal import Decimal from random import randint from unittest.mock import MagicMock, patch @@ -19,13 +19,6 @@ from pyathena import ( BINARY, - BOOLEAN, - DATE, - DATETIME, - JSON, - NUMBER, - STRING, - TIME, Binary, ExecuteOptions, ) @@ -36,6 +29,15 @@ from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect +from tests.pyathena.expected import ( + ARRAY_JSON, + MAP_JSON, + PYTHON_VALUES, + TIME_OF_TIMESTAMP, + TIMESTAMP_TZ, + ExpectedResult, +) +from tests.pyathena.tables import ONE_ROW_COMPLEX from tests.pyathena.util import throttle_metadata_api, unreachable_glue _logger = logging.getLogger(__name__) @@ -600,105 +602,18 @@ def test_query_execution_initial(self, cursor): assert cursor.effective_engine_version is None def test_complex(self, cursor): - cursor.execute( - """ - SELECT - col_boolean - ,col_tinyint - ,col_smallint - ,col_int - ,col_bigint - ,col_float - ,col_double - ,col_string - ,col_varchar - ,col_timestamp - ,CAST(col_timestamp AS timestamp with time zone) AS col_timestamp_tz - ,CAST(col_timestamp AS time) AS col_time - ,col_date - ,col_binary - ,col_array - ,CAST(col_array AS json) AS col_array_json - ,col_map - ,CAST(col_map AS json) AS col_map_json - ,col_struct - ,col_decimal - FROM one_row_complex - """ + expected = ExpectedResult( + ONE_ROW_COMPLEX, casts=(TIMESTAMP_TZ, TIME_OF_TIMESTAMP, ARRAY_JSON, MAP_JSON) ) - assert cursor.description == [ - ("col_boolean", "boolean", None, None, 0, 0, "UNKNOWN"), - ("col_tinyint", "tinyint", None, None, 3, 0, "UNKNOWN"), - ("col_smallint", "smallint", None, None, 5, 0, "UNKNOWN"), - ("col_int", "integer", None, None, 10, 0, "UNKNOWN"), - ("col_bigint", "bigint", None, None, 19, 0, "UNKNOWN"), - ("col_float", "float", None, None, 17, 0, "UNKNOWN"), - ("col_double", "double", None, None, 17, 0, "UNKNOWN"), - ("col_string", "varchar", None, None, 2147483647, 0, "UNKNOWN"), - ("col_varchar", "varchar", None, None, 10, 0, "UNKNOWN"), - ("col_timestamp", "timestamp", None, None, 3, 0, "UNKNOWN"), - ("col_timestamp_tz", "timestamp with time zone", None, None, 3, 0, "UNKNOWN"), - ("col_time", "time", None, None, 3, 0, "UNKNOWN"), - ("col_date", "date", None, None, 0, 0, "UNKNOWN"), - ("col_binary", "varbinary", None, None, 1073741824, 0, "UNKNOWN"), - ("col_array", "array", None, None, 0, 0, "UNKNOWN"), - ("col_array_json", "json", None, None, 0, 0, "UNKNOWN"), - ("col_map", "map", None, None, 0, 0, "UNKNOWN"), - ("col_map_json", "json", None, None, 0, 0, "UNKNOWN"), - ("col_struct", "row", None, None, 0, 0, "UNKNOWN"), - ("col_decimal", "decimal", None, None, 10, 1, "UNKNOWN"), - ] + cursor.execute(expected.sql) + assert cursor.description == expected.description() rows = cursor.fetchall() - expected = [ - ( - True, - 127, - 32767, - 2147483647, - 9223372036854775807, - 0.5, - 0.25, - "a string", - "varchar", - datetime(2017, 1, 1, 0, 0, 0), - datetime(2017, 1, 1, 0, 0, 0, tzinfo=timezone.utc), - datetime(2017, 1, 1, 0, 0, 0).time(), - date(2017, 1, 2), - b"123", - [1, 2], - [1, 2], - {"1": "2", "3": "4"}, - {"1": 2, "3": 4}, - {"a": "1", "b": "2"}, - Decimal("0.1"), - ) - ] - assert rows == expected + expected_rows = expected.rows(PYTHON_VALUES) + assert rows == expected_rows # catch unicode/str - assert list(map(type, rows[0])) == list(map(type, expected[0])) + assert list(map(type, rows[0])) == list(map(type, expected_rows[0])) # compare dbapi type object - assert [d[1] for d in cursor.description] == [ - BOOLEAN, - NUMBER, - NUMBER, - NUMBER, - NUMBER, - NUMBER, - NUMBER, - STRING, - STRING, - DATETIME, - DATETIME, - TIME, - DATE, - BINARY, - STRING, - JSON, - STRING, - JSON, - STRING, - NUMBER, - ] + assert [d[1] for d in cursor.description] == expected.dbapi_types() def test_complex_with_type_hints(self, cursor): # 1. Basic complex columns from one_row_complex diff --git a/tests/pyathena/util.py b/tests/pyathena/util.py index f5ac673d..2ff90f08 100644 --- a/tests/pyathena/util.py +++ b/tests/pyathena/util.py @@ -15,6 +15,9 @@ from pyathena.glue import GlueMetadataClient from pyathena.model import AthenaCalculationExecutionStatus +from pyathena.sqlalchemy.types import TINYINT, AthenaArray, AthenaStruct +from tests.pyathena.expected import base_type, type_arguments, type_parameters +from tests.pyathena.tables import Column _queries = Environment( loader=FileSystemLoader(Path(__file__).parents[1].resolve() / "resources" / "queries") @@ -133,3 +136,46 @@ def wait_for_spark_session_state(client, session_id, state, timeout=120): return time.sleep(1) raise AssertionError(f"Session {session_id} did not become {state} in {timeout} seconds.") + + +# The SQLAlchemy type class of a reflected column, by base type. +_SQLALCHEMY_TYPES = { + "boolean": types.BOOLEAN, + "tinyint": TINYINT, + "smallint": types.SMALLINT, + "int": types.INTEGER, + "bigint": types.BIGINT, + "float": types.FLOAT, + "double": types.DOUBLE, + "string": types.String, + "varchar": types.VARCHAR, + "timestamp": types.TIMESTAMP, + "date": types.DATE, + "binary": types.BINARY, + "array": AthenaArray, + "map": types.String, + "struct": AthenaStruct, + "decimal": types.DECIMAL, +} + + +def assert_sqlalchemy_type(sqlalchemy_type, column: Column) -> None: + """Assert that a reflected SQLAlchemy column type matches a table column's definition. + + Args: + sqlalchemy_type: The reflected column type. + column: The column in ``tests.pyathena.tables``. + + Raises: + AssertionError: If the type class or its parameters differ. + """ + name = base_type(column.athena_type) + assert isinstance(sqlalchemy_type, _SQLALCHEMY_TYPES[name]), (column.name, sqlalchemy_type) + if name == "varchar": + assert sqlalchemy_type.length == type_parameters(column.athena_type)[0], column.name + elif name == "decimal": + precision, scale = type_parameters(column.athena_type) + assert (sqlalchemy_type.precision, sqlalchemy_type.scale) == (precision, scale) + elif name == "array": + (element,) = type_arguments(column.athena_type) + assert isinstance(sqlalchemy_type.item_type, _SQLALCHEMY_TYPES[base_type(element)])