diff --git a/docs/arrow.md b/docs/arrow.md index 9cc66e73..5306f090 100644 --- a/docs/arrow.md +++ b/docs/arrow.md @@ -148,7 +148,7 @@ class CustomArrowTypeConverter(Converter): "char": pa.string(), "varchar": pa.string(), "string": pa.string(), - "timestamp": pa.timestamp("ms"), + "timestamp": pa.timestamp("us"), "date": pa.timestamp("ms"), "time": pa.string(), "varbinary": pa.string(), diff --git a/docs/usage.md b/docs/usage.md index 492e7c11..79d4b042 100644 --- a/docs/usage.md +++ b/docs/usage.md @@ -42,6 +42,8 @@ cursor = connect(work_group="YOUR_MANAGED_WORK_GROUP", With managed query result storage, query results are retrieved via the `GetQueryResults` API (1000 rows per request) instead of reading S3 files directly. This may be slower for large result sets. For large datasets, consider using customer-managed storage or the `UNLOAD` statement. +The pandas, Arrow, and Polars cursors read these rows as they read a CSV result file, with the +same types and converters, but not in chunks. ``` ## Cursor iteration @@ -709,8 +711,7 @@ column. You can mix both styles in the same dictionary. JSON-formatted output, which is parsed reliably. - **Arrow, Pandas, and Polars cursors** — These cursors accept `result_set_type_hints` but their converters do not currently use the hints because they rely on their own - type systems. The parameter is passed through for forward compatibility and for - result sets that fall back to the default conversion path. + type systems. The parameter is passed through for forward compatibility. ### Breaking change in 3.30.0 diff --git a/pyathena/arrow/converter.py b/pyathena/arrow/converter.py index 7da7bcef..13d63caf 100644 --- a/pyathena/arrow/converter.py +++ b/pyathena/arrow/converter.py @@ -5,9 +5,10 @@ import logging from collections.abc import Callable from copy import deepcopy -from typing import Any +from typing import TYPE_CHECKING, Any from pyathena.converter import ( + _TIMESTAMP_TEXT_LENGTHS, Converter, _to_binary, _to_date, @@ -20,6 +21,9 @@ ) from pyathena.util import override +if TYPE_CHECKING: + from pyarrow import ChunkedArray, TimestampType + _logger = logging.getLogger(__name__) @@ -34,6 +38,28 @@ } +def _to_timestamp(column: ChunkedArray, type_: TimestampType) -> ChunkedArray: + """Convert timestamp text to a timestamp type, truncating finer fractions. + + Athena writes up to 12 fractional digits, which pyarrow does not parse into a + timestamp type whose unit holds fewer. + + Args: + column: The timestamp text, with NULL as null or as an empty string. + type_: The timestamp type. + + Returns: + The timestamps. + """ + import pyarrow as pa + import pyarrow.compute as pc + + length = _TIMESTAMP_TEXT_LENGTHS[type_.unit] + if (pc.max(pc.utf8_length(column)).as_py() or 0) > length: + column = pc.utf8_slice_codeunits(column, 0, length) + return pc.if_else(pc.equal(column, ""), pa.scalar(None, pa.string()), column).cast(type_) + + class DefaultArrowTypeConverter(Converter): """Optimized type converter for Apache Arrow Table results. @@ -86,7 +112,7 @@ def _dtypes(self) -> dict[str, type[Any]]: "char": pa.string(), "varchar": pa.string(), "string": pa.string(), - "timestamp": pa.timestamp("ms"), + "timestamp": pa.timestamp("us"), "date": pa.timestamp("ms"), "time": pa.string(), "time with time zone": pa.string(), diff --git a/pyathena/arrow/result_set.py b/pyathena/arrow/result_set.py index d34a7cb5..2b1b9cbb 100644 --- a/pyathena/arrow/result_set.py +++ b/pyathena/arrow/result_set.py @@ -11,9 +11,9 @@ ) from pyathena import OperationalError +from pyathena.arrow.converter import _to_timestamp from pyathena.arrow.util import to_column_info -from pyathena.converter import _TEXT_VALUE_TYPES, Converter, _text_value_converter, _to_default -from pyathena.error import ProgrammingError +from pyathena.converter import Converter, _to_default from pyathena.model import AthenaQueryExecution from pyathena.result_set import AthenaResultSet from pyathena.util import RetryConfig, override, parse_output_location @@ -141,15 +141,13 @@ def __init__( if self.state == AthenaQueryExecution.STATE_SUCCEEDED and self.output_location: self._table = self._as_arrow() elif self.state == AthenaQueryExecution.STATE_SUCCEEDED: - self._table = self._as_arrow_from_api() + # Without a result file, as with managed query result storage, the rows from + # GetQueryResults are read as a CSV result file. + self._table = self._read_csv() else: import pyarrow as pa self._table = pa.Table.from_pydict({}) - # The fetch methods convert the values read from a result file. GetQueryResults - # values are already converted, except json and time with time zone values, - # which stay text. - self._convert_rows = bool(self.output_location) self._batches = iter(self._table.to_batches(arraysize)) def _create_s3_file_system(self): @@ -258,12 +256,7 @@ def _fetch(self) -> None: # converters property keep one column per name. columns = [column.to_pylist() for column in rows.columns] description = self.description if self.description else [] - converters = [ - self._converter.get(d[1]) - if self._convert_rows or d[1] in _TEXT_VALUE_TYPES - else _to_default - for d in description - ] + converters = [self._converter.get(d[1]) for d in description] if any(convert is not _to_default for convert in converters): processed_rows = [ tuple(convert(v) for convert, v in zip(converters, row, strict=False)) @@ -287,12 +280,18 @@ def fetchone( return self._rows.popleft() def _read_csv(self) -> Table: + """Read the CSV result file, or the GetQueryResults rows as one without it. + + Returns: + The Arrow Table of the results. + + Raises: + OperationalError: If reading the results fails. + """ import pyarrow as pa from pyarrow import csv - if not self.output_location: - raise ProgrammingError("OutputLocation is none or empty.") - if not self.output_location.endswith((".csv", ".txt")): + if self.output_location and not self.output_location.endswith((".csv", ".txt")): return pa.Table.from_pydict({}) if self.substatement_type and self.substatement_type.upper() in ( "UPDATE", @@ -301,7 +300,14 @@ def _read_csv(self) -> Table: "VACUUM_TABLE", ): return pa.Table.from_pydict({}) - length = self._get_content_length() + if self.output_location: + data = None + length = self._get_content_length() + location = "/".join(parse_output_location(self.output_location)) + else: + data = self._fetch_all_rows_as_csv() + length = len(data) + location = "the GetQueryResults rows" description = self.description if self.description else [] names = [d[0] for d in description] # pyarrow types every column with a name by its column_types entry, so columns @@ -318,8 +324,16 @@ def _read_csv(self) -> Table: else: column_names = names column_types = self.column_types + # Timestamp columns are read as text: pyarrow does not parse fractions finer + # than the unit, and on some platforms its strptime fallbacks accept such a + # value without its fraction. + timestamp_types = { + i: dtype + for i, (name, d) in enumerate(zip(column_names, description, strict=True)) + if d[1] == "timestamp" and isinstance(dtype := column_types.get(name), pa.TimestampType) + } binary_columns = {i for i, d in enumerate(description) if d[1] == "varbinary"} - if length and self.output_location.endswith(".txt"): + if length and self.output_location and self.output_location.endswith(".txt"): read_opts = csv.ReadOptions( skip_rows=0, column_names=column_names, @@ -332,7 +346,7 @@ def _read_csv(self) -> Table: double_quote=False, escape_char=False, ) - elif length and self.output_location.endswith(".csv"): + elif length: read_opts = csv.ReadOptions(skip_rows=0, block_size=self._block_size, use_threads=True) if has_duplicate_names: read_opts.column_names = column_names @@ -353,19 +367,27 @@ def _read_csv(self) -> Table: else: return pa.Table.from_pydict({}) - bucket, key = parse_output_location(self.output_location) try: table = csv.read_csv( - self._fs.open_input_stream(f"{bucket}/{key}"), + self._fs.open_input_stream(location) if data is None else pa.BufferReader(data), read_options=read_opts, parse_options=parse_opts, convert_options=csv.ConvertOptions( strings_can_be_null=bool(binary_columns), quoted_strings_can_be_null=False, timestamp_parsers=self.timestamp_parsers, - column_types=column_types, + column_types={ + **column_types, + **{column_names[i]: pa.string() for i in timestamp_types}, + }, ), ) + for index, type_ in timestamp_types.items(): + table = table.set_column( + index, + pa.field(table.schema.field(index).name, type_), + _to_timestamp(table.column(index), type_), + ) if has_duplicate_names: table = table.rename_columns(names) if binary_columns: @@ -377,7 +399,7 @@ def _read_csv(self) -> Table: table = table.set_column(index, field, table.column(index).fill_null("")) return table except Exception as e: - _logger.exception(f"Failed to read {bucket}/{key}.") + _logger.exception(f"Failed to read {location}.") raise OperationalError(*e.args) from e def _read_parquet(self) -> Table: @@ -406,27 +428,6 @@ def _as_arrow(self) -> Table: table = self._read_csv() return table - def _as_arrow_from_api(self, converter: Converter | None = None) -> Table: - """Build an Arrow Table from GetQueryResults API. - - Used as a fallback when ``output_location`` is not available - (e.g. managed query result storage). - - Args: - converter: Type converter for result values. Defaults to - ``DefaultTypeConverter`` with json and time with time zone values kept as - text, as in the CSV result file. Arrow has no type for JSON values or for - times with a time zone. - """ - import pyarrow as pa - - rows = self._fetch_all_rows(converter or _text_value_converter()) - if not rows: - return pa.Table.from_pydict({}) - description = self.description if self.description else [] - columns = [list(column) for column in zip(*rows, strict=True)] - return pa.table(columns, names=[d[0] for d in description]) - def as_arrow(self) -> Table: """Return the query results as an Apache Arrow Table. diff --git a/pyathena/converter.py b/pyathena/converter.py index 72d9d83d..644a50c0 100644 --- a/pyathena/converter.py +++ b/pyathena/converter.py @@ -68,6 +68,11 @@ def _to_datetime(varchar_value: str | None) -> datetime | None: return _parse_datetime(varchar_value) +# The length of Athena timestamp text, such as "2020-01-02 03:04:05.123456", that holds +# the fraction a time unit can represent. Athena writes up to 12 fractional digits. +_TIMESTAMP_TEXT_LENGTHS: dict[str, int] = {"s": 19, "ms": 23, "us": 26, "ns": 29} + + _UTC_OFFSET_PATTERN: re.Pattern[str] = re.compile(r"([+-])(\d{2}):(\d{2})") @@ -803,33 +808,3 @@ def _parse_type_hint(self, type_hint: str) -> TypeNode: if normalized not in self._parsed_hints: self._parsed_hints[normalized] = self._parser.parse(normalized) return self._parsed_hints[normalized] - - -# The types whose values the Arrow and Polars GetQueryResults fallbacks keep as text, -# as in a CSV result file, and convert when the rows are fetched. -_TEXT_VALUE_TYPES: tuple[str, ...] = ("json", "time with time zone", "timestamp with time zone") - - -def _text_value_converter() -> DefaultTypeConverter: - """Return a ``DefaultTypeConverter`` that keeps ``_TEXT_VALUE_TYPES`` values as text. - - Values nested in typed complex values keep only the time zone types as text, - because Arrow and Polars time and timestamp types hold one time zone per column; - nested JSON values are decoded as before. - - Returns: - The converter. - """ - converter = DefaultTypeConverter() - for type_ in _TEXT_VALUE_TYPES: - converter.set(type_, _to_default) - converter._typed_converter = TypedValueConverter( - converters={ - **_DEFAULT_CONVERTERS, - "time with time zone": _to_default, - "timestamp with time zone": _to_default, - }, - default_converter=_to_default, - struct_parser=_to_struct, - ) - return converter diff --git a/pyathena/pandas/result_set.py b/pyathena/pandas/result_set.py index 2ae91078..8e68adc6 100644 --- a/pyathena/pandas/result_set.py +++ b/pyathena/pandas/result_set.py @@ -8,7 +8,7 @@ from collections.abc import Callable, Iterable, Iterator from contextlib import ExitStack from functools import partial -from io import BufferedReader, IOBase, StringIO, TextIOWrapper +from io import BufferedReader, BytesIO, IOBase, StringIO, TextIOWrapper from multiprocessing import cpu_count from typing import ( TYPE_CHECKING, @@ -421,7 +421,6 @@ class AthenaPandasResultSet(AthenaResultSet): AUTO_CHUNK_SIZE_LARGE: int = 100_000 AUTO_CHUNK_SIZE_MEDIUM: int = 50_000 - _INTEGER_TYPES: ClassVar[tuple[str, ...]] = ("tinyint", "smallint", "integer", "bigint") _PARSE_DATES: ClassVar[list[str]] = [ "date", "time", @@ -546,16 +545,19 @@ def __init__( # The whole result when it was not read in chunks. self._df: DataFrame | None = None - if self.state == AthenaQueryExecution.STATE_SUCCEEDED and self.output_location: - result = self._as_pandas() - trunc_date = _no_trunc_date if self.is_unload else self._finish_csv_frame + if self.state == AthenaQueryExecution.STATE_SUCCEEDED: + if self.output_location: + result = self._as_pandas() + trunc_date = _no_trunc_date if self.is_unload else self._finish_csv_frame + else: + # Without a result file, as with managed query result storage, the rows + # from GetQueryResults are read as a CSV result file, but not in chunks. + result = self._read_csv() + trunc_date = self._finish_csv_frame if isinstance(result, pd.DataFrame): self._df = trunc_date(result) else: self._df_iter = PandasDataFrameIterator(result, trunc_date, self._csv_stream) - elif self.state == AthenaQueryExecution.STATE_SUCCEEDED: - # GetQueryResults values are already converted and need no time truncation. - self._df = self._as_pandas_from_api() else: self._df = pd.DataFrame() if self._df is not None: @@ -795,11 +797,19 @@ def fetchone( return tuple(row[1].values()) def _read_csv(self) -> TextFileReader | DataFrame: + """Read the CSV result file, or the GetQueryResults rows as one without it. + + The GetQueryResults rows are read as a whole, not in chunks. + + Returns: + The DataFrame of the results, or a reader of their chunks. + + Raises: + OperationalError: If reading the results fails. + """ import pandas as pd - if not self.output_location: - raise ProgrammingError("OutputLocation is none or empty.") - if not self.output_location.endswith((".csv", ".txt")): + if self.output_location and not self.output_location.endswith((".csv", ".txt")): return pd.DataFrame() if self.substatement_type and self.substatement_type.upper() in ( "UPDATE", @@ -808,15 +818,22 @@ def _read_csv(self) -> TextFileReader | DataFrame: "VACUUM_TABLE", ): return pd.DataFrame() - length = self._get_content_length() + if self.output_location: + data = None + length = self._get_content_length() + location = self.output_location + else: + data = self._fetch_all_rows_as_csv() + length = len(data) + location = "the GetQueryResults rows" if length == 0: return pd.DataFrame() # Chunksize determination with user preference priority - effective_chunksize = self._chunksize + effective_chunksize = self._chunksize if data is None else None # Only auto-optimize if user hasn't specified chunksize AND auto_optimize is enabled - if effective_chunksize is None and self._auto_optimize_chunksize: + if effective_chunksize is None and self._auto_optimize_chunksize and data is None: effective_chunksize = self._auto_determine_chunksize(length) if effective_chunksize: _logger.debug( @@ -832,13 +849,22 @@ def _read_csv(self) -> TextFileReader | DataFrame: try: with ExitStack() as stack: - source: str | IOBase = self.output_location + source: str | IOBase = location binary_columns = self._configure_binary_csv_read(read_csv_kwargs, labels) if labels is not None: # After _configure_binary_csv_read(), which checks for the header row. self._read_csv_header_as_labels(read_csv_kwargs, csv_engine) self._csv_converters = read_csv_kwargs.get("converters") or {} - if binary_columns: + if data is not None: + # The rows are in memory, with nothing to open with storage options. + read_csv_kwargs.pop("storage_options", None) + if binary_columns: + source = self._csv_stream = stack.enter_context( + self._open_binary_csv_stream(binary_columns, None, data) + ) + else: + source = BytesIO(data) + elif binary_columns: # Given storage_options, even None, open the file through fsspec # as pandas does. storage_options = None @@ -861,7 +887,7 @@ def _read_csv(self) -> TextFileReader | DataFrame: stack.pop_all() # Log performance information for large files - if length > self.LARGE_FILE_THRESHOLD_BYTES: + if data is None and length > self.LARGE_FILE_THRESHOLD_BYTES: mode = "chunked" if effective_chunksize else "full" chunksize = f" with chunksize={effective_chunksize}" if effective_chunksize else "" _logger.info( @@ -872,7 +898,7 @@ def _read_csv(self) -> TextFileReader | DataFrame: return result except Exception as e: - _logger.exception(f"Failed to read {self.output_location}.") + _logger.exception(f"Failed to read {location}.") raise OperationalError(*e.args) from e def _reads_csv_with_pyarrow(self) -> bool: @@ -984,8 +1010,7 @@ def _is_standard_csv_parsing(self, read_csv_kwargs: dict[str, Any]) -> bool: True if pandas reads the header and quoted fields of the file as written. """ return not ( - not self.output_location - or not self.output_location.endswith(".csv") + (self.output_location and not self.output_location.endswith(".csv")) or read_csv_kwargs.get("header") != 0 or read_csv_kwargs.get("skiprows") is not None or read_csv_kwargs.get("dialect") is not None @@ -1150,12 +1175,28 @@ def _configure_binary_csv_read( return binary_columns def _open_binary_csv_stream( - self, binary_columns: set[int], storage_options: dict[str, Any] | None + self, + binary_columns: set[int], + storage_options: dict[str, Any] | None, + data: bytes | None = None, ) -> TextIOWrapper: - """Open a stream that preserves binary NULL fields and original CSV newlines.""" + """Open a stream that preserves binary NULL fields and original CSV newlines. + + Args: + binary_columns: Zero-based indexes of the binary columns. + storage_options: The options to open the result file with through fsspec, + or None to open it with PyAthena's filesystem. + data: The CSV text to read in place of the result file. + + Returns: + The text stream for ``pandas.read_csv()``. + """ text_options: dict[str, Any] = {"mode": "rt", "encoding": "utf-8", "newline": ""} with ExitStack() as stack: - if storage_options is None: + source: IOBase + if data is not None: + source = StringIO(data.decode("utf-8"), newline="") + elif storage_options is None: source = stack.enter_context(self._fs.open(self.output_location, **text_options)) else: source = stack.enter_context( @@ -1227,37 +1268,6 @@ def _as_pandas(self) -> TextFileReader | DataFrame: df = self._read_csv() return df - def _as_pandas_from_api(self, converter: Converter | None = None) -> DataFrame: - """Build a DataFrame from GetQueryResults API. - - Used as a fallback when ``output_location`` is not available - (e.g. managed query result storage). - - Args: - converter: Type converter for result values. Defaults to - ``DefaultTypeConverter`` if not specified. - """ - import pandas as pd - - rows = self._fetch_all_rows(converter) - if not rows: - return pd.DataFrame() - description = self.description if self.description else [] - # Positional, so that columns with the same name keep their own values. - columns = [list(column) for column in zip(*rows, strict=True)] - # Integer columns get the dtype that the CSV result file reads them with, - # and json columns with NULL stay objects as there, so that NULL does not - # make their values floats. - data: dict[Any, Any] = {} - for name, values, d in zip(self._get_column_names(), columns, description, strict=True): - dtype = None - if d[1] in self._INTEGER_TYPES: - dtype = self._converter.get_dtype(d[1], d[4], d[5]) - elif d[1] == "json" and None in values: - dtype = object - data[name] = values if dtype is None else pd.array(values, dtype=dtype) - return pd.DataFrame(data) - def as_pandas(self) -> PandasDataFrameIterator | DataFrame: """Return the query results as a DataFrame or an iterator of DataFrame chunks. diff --git a/pyathena/polars/converter.py b/pyathena/polars/converter.py index 4d46c702..a40e9c7d 100644 --- a/pyathena/polars/converter.py +++ b/pyathena/polars/converter.py @@ -12,9 +12,10 @@ import logging from collections.abc import Callable from copy import deepcopy -from typing import Any +from typing import TYPE_CHECKING, Any from pyathena.converter import ( + _TIMESTAMP_TEXT_LENGTHS, Converter, _to_binary, _to_date, @@ -26,6 +27,9 @@ ) from pyathena.util import override +if TYPE_CHECKING: + import polars as pl + _logger = logging.getLogger(__name__) @@ -39,6 +43,38 @@ } +def _to_datetimes(df: pl.DataFrame, dtypes: dict[str, Any]) -> pl.DataFrame: + """Convert timestamp text columns to Datetime dtypes, truncating finer fractions. + + Athena writes up to 12 fractional digits, which Polars does not parse into a + Datetime whose time unit holds fewer. + + Args: + df: The DataFrame with the timestamp text columns, with NULL as null or as + an empty string. + dtypes: The Datetime dtypes keyed by column name. + + Returns: + The DataFrame with the columns converted. + """ + import polars as pl + + exprs = [] + for name, dtype in dtypes.items(): + if name not in df.columns: + # Not read, such as a column that the read options leave out. + continue + time_unit = (dtype() if isinstance(dtype, type) else dtype).time_unit + text = pl.col(name) + exprs.append( + pl.when(text != "") + .then(text.str.slice(0, _TIMESTAMP_TEXT_LENGTHS[time_unit])) + .str.to_datetime("%Y-%m-%d %H:%M:%S%.f", time_unit=time_unit) + .cast(dtype) + ) + return df.with_columns(exprs) if exprs else df + + class DefaultPolarsTypeConverter(Converter): """Optimized type converter for Polars DataFrame results. diff --git a/pyathena/polars/result_set.py b/pyathena/polars/result_set.py index 85ea4742..1ee5bc0b 100644 --- a/pyathena/polars/result_set.py +++ b/pyathena/polars/result_set.py @@ -22,9 +22,10 @@ ) from pyathena import OperationalError -from pyathena.converter import Converter, _text_value_converter +from pyathena.converter import Converter from pyathena.error import ProgrammingError from pyathena.model import AthenaQueryExecution +from pyathena.polars.converter import _to_datetimes from pyathena.polars.util import to_column_info from pyathena.result_set import AthenaResultSet from pyathena.util import RetryConfig, override @@ -278,13 +279,10 @@ def __init__( else: self._df_iter = self._create_dataframe_iterator() elif self.state == AthenaQueryExecution.STATE_SUCCEEDED: - self._df = self._as_polars_from_api() - # GetQueryResults values are already converted, except json and time with - # time zone values kept as text. - column_names = self._get_frame_column_names() - self._df_converters = self._text_value_converters( - self._get_converters(column_names), column_names - ) + # Without a result file, as with managed query result storage, the rows from + # GetQueryResults are read as a CSV result file, but not in chunks. + self._df = self._read_csv() + self._df_converters = self._get_converters(self._get_frame_column_names()) else: self._df = pl.DataFrame() if self._df is not None: @@ -432,6 +430,30 @@ def _get_converters( for name, d in zip(column_names, description, strict=True) } + def _get_timestamp_dtypes(self, has_header: bool) -> dict[str, Any]: + """Get the Datetime dtypes of the timestamp columns, which are read as text. + + Args: + has_header: Whether the CSV data has a header. Without one, or with + ``schema_overrides`` or ``with_column_names`` given to ``execute()``, + Polars does not name the columns by the header, and they are not read + as text. + + Returns: + The Datetime dtypes keyed by the header of a CSV file. + """ + import polars as pl + + if not has_header or self._kwargs.keys() & {"schema_overrides", "with_column_names"}: + return {} + dtypes = self._csv_dtypes + return { + name: dtype + for name, d in zip(self._get_column_names(), self.description or [], strict=True) + if d[1] == "timestamp" + and ((dtype := dtypes.get(name)) is pl.Datetime or isinstance(dtype, pl.Datetime)) + } + def _get_column_names(self) -> list[str]: """Get the names of the result columns in a DataFrame. @@ -462,7 +484,7 @@ def _get_frame_column_names(self) -> list[str]: """ names = self._get_column_names() new_columns = self._kwargs.get("new_columns") - if not new_columns or not self.output_location or self.is_unload: + if not new_columns or (self.output_location and self.is_unload): return names return [*new_columns[: len(names)], *names[len(new_columns) :]] @@ -507,15 +529,12 @@ def fetchone( def _is_csv_readable(self) -> bool: """Check if CSV output is available and can be read. + Without an output location, the GetQueryResults rows are read as CSV. + Returns: True if CSV data is available to read, False otherwise. - - Raises: - ProgrammingError: If output location is not set. """ - if not self.output_location: - raise ProgrammingError("OutputLocation is none or empty.") - if not self.output_location.endswith((".csv", ".txt")): + if self.output_location and not self.output_location.endswith((".csv", ".txt")): return False if self.substatement_type and self.substatement_type.upper() in ( "UPDATE", @@ -524,8 +543,7 @@ def _is_csv_readable(self) -> bool: "VACUUM_TABLE", ): return False - length = self._get_content_length() - return length != 0 + return not self.output_location or self._get_content_length() != 0 def _prepare_parquet_location(self) -> bool: """Prepare unload location for Parquet reading. @@ -541,25 +559,21 @@ def _prepare_parquet_location(self) -> bool: return True def _read_csv(self) -> pl.DataFrame: - """Read query results from CSV file in S3. + """Read query results from CSV file in S3, or the GetQueryResults rows as one without it. Returns: Polars DataFrame containing the CSV data. Raises: - ProgrammingError: If output location is not set. - OperationalError: If reading the CSV file fails. + OperationalError: If reading the CSV data fails. """ import polars as pl if not self._is_csv_readable(): return pl.DataFrame() - if self.output_location is None: - raise ProgrammingError("output_location is not available.") - separator, has_header, new_columns = self._get_csv_params() - read_kwargs = self._read_kwargs( + kwargs = self._read_kwargs( lambda: self._csv_storage_options, separator=separator, has_header=has_header, @@ -567,15 +581,37 @@ def _read_csv(self) -> pl.DataFrame: ) if "schema_overrides" not in self._kwargs: # Renamed after reading, so that Polars matches the types to the header. - read_kwargs.pop("new_columns", None) + kwargs.pop("new_columns", None) + source: str | bytes + if self.output_location: + source = location = self.output_location + else: + source = self._fetch_all_rows_as_csv() + if not source: + return pl.DataFrame() + location = "the GetQueryResults rows" + del kwargs["storage_options"] try: - df = pl.read_csv(self.output_location, **read_kwargs) + try: + df = pl.read_csv(source, **kwargs) + except pl.exceptions.ComputeError: + timestamp_dtypes = self._get_timestamp_dtypes(has_header) + if not timestamp_dtypes: + raise + # Athena writes up to 12 fractional digits, which Polars does not parse + # into a Datetime whose time unit holds fewer, so the data is read again + # with the timestamp columns as text. + kwargs["schema_overrides"] = { + **self._csv_dtypes, + **dict.fromkeys(timestamp_dtypes, pl.String), + } + df = _to_datetimes(pl.read_csv(source, **kwargs), timestamp_dtypes) if new_columns: df.columns = [*new_columns, *df.columns[len(new_columns) :]] return df except Exception as e: - _logger.exception(f"Failed to read {self.output_location}.") + _logger.exception(f"Failed to read {location}.") raise OperationalError(*e.args) from e def _read_parquet(self) -> pl.DataFrame: @@ -641,25 +677,6 @@ def _as_polars(self) -> pl.DataFrame: df = self._read_csv() return df - def _as_polars_from_api(self, converter: Converter | None = None) -> pl.DataFrame: - """Build a Polars DataFrame from GetQueryResults API. - - Used as a fallback when ``output_location`` is not available - (e.g. managed query result storage). - - Args: - converter: Type converter for result values. Defaults to - ``DefaultTypeConverter`` with json and time with time zone values kept as - text, as in the CSV result file. A Polars ``Time`` has no time zone. - """ - import polars as pl - - rows = self._fetch_all_rows(converter or _text_value_converter()) - if not rows: - return pl.DataFrame() - columns = [list(column) for column in zip(*rows, strict=True)] - return pl.DataFrame(dict(zip(self._get_column_names(), columns, strict=True))) - def as_polars(self) -> pl.DataFrame: """Return query results as a Polars DataFrame. diff --git a/pyathena/result_set.py b/pyathena/result_set.py index c0911778..fd97ac38 100644 --- a/pyathena/result_set.py +++ b/pyathena/result_set.py @@ -12,13 +12,13 @@ ) from pyathena.common import BaseCursor, CursorIterator -from pyathena.converter import _TEXT_VALUE_TYPES, Converter, DefaultTypeConverter +from pyathena.converter import Converter, DefaultTypeConverter from pyathena.error import DataError, OperationalError, ProgrammingError from pyathena.model import AthenaQueryExecution from pyathena.util import RetryConfig, override, parse_output_location, retry_api_call if TYPE_CHECKING: - from collections.abc import Callable + from collections.abc import Iterator from pyathena.connection import Connection @@ -675,29 +675,34 @@ def _is_first_row_column_labels(self, rows: list[dict[str, Any]]) -> bool: return False return True - def _text_value_converters( - self, - converters: dict[str, Callable[[str | None], Any | None]], - column_names: list[str] | None = None, - ) -> dict[str, Callable[[str | None], Any | None]]: - """Select the converters of the columns that the fallbacks keep as text. + def _iter_all_row_pages(self) -> Iterator[tuple[list[dict[str, Any]], int]]: + """Fetch all rows via GetQueryResults API from the beginning. - Args: - converters: The converters keyed by column name. - column_names: The names that ``converters`` uses for the columns, in - column order. Defaults to the names in the description. + Paginates through all results using MaxResults=1000. This is for subclass result + sets that need to fall back to the API when S3 output is not available (e.g., + managed query result storage). - Returns: - The converters of the columns whose Athena type is in ``_TEXT_VALUE_TYPES``. + Yields: + The rows of each page and the offset of their first data row, which is 1 when + the first page starts with the column labels. """ - description = self.description if self.description else [] - if column_names is None: - column_names = [d[0] for d in description] - return { - name: converters[name] - for name, d in zip(column_names, description, strict=True) - if d[1] in _TEXT_VALUE_TYPES - } + _logger.warning( + "output_location is not available (e.g. managed query result storage). " + "Falling back to GetQueryResults API. " + "This may be slow for large result sets." + ) + + next_token: str | None = None + first_page = True + while True: + response = self._get_query_results(self.DEFAULT_FETCH_SIZE, next_token) + rows, next_token = self._parse_result_rows(response) + # Only the first page starts with the column labels. + offset = 1 if first_page and rows and self._is_first_row_column_labels(rows) else 0 + first_page = False + yield rows, offset + if not next_token: + break def _fetch_all_rows( self, @@ -705,53 +710,63 @@ def _fetch_all_rows( ) -> list[tuple[Any | None, ...]]: """Fetch all rows via GetQueryResults API with type conversion. - Paginates through all results from the beginning using MaxResults=1000. - Defaults to ``DefaultTypeConverter`` for string-to-Python type conversion, - because subclass converters (e.g. Pandas/Arrow) are designed for S3 file - reading and may not handle API result strings. - - This method is intended for use by subclass result sets that need to - fall back to the API when S3 output is not available (e.g., managed - query result storage). - Args: converter: Type converter for result values. Defaults to ``DefaultTypeConverter`` if not specified. Returns: List of converted row tuples. + + Raises: + ProgrammingError: If the metadata is not available. """ - if self._metadata is None: + metadata = self._metadata + if metadata is None: raise ProgrammingError("Metadata is not available.") - - _logger.warning( - "output_location is not available (e.g. managed query result storage). " - "Falling back to GetQueryResults API. " - "This may be slow for large result sets." - ) - converter = converter or DefaultTypeConverter() all_rows: list[tuple[Any | None, ...]] = [] - next_token: str | None = None - - while True: - first_page = next_token is None - response = self._get_query_results(self.DEFAULT_FETCH_SIZE, next_token) - rows, next_token = self._parse_result_rows(response) - - # Only the first page can start with the column labels. - offset = 1 if first_page and rows and self._is_first_row_column_labels(rows) else 0 + for rows, offset in self._iter_all_row_pages(): all_rows.extend( cast( list[tuple[Any | None, ...]], - self._get_rows(offset, self._metadata, rows, converter), + self._get_rows(offset, metadata, rows, converter), ) ) + return all_rows - if not next_token: - break + def _fetch_all_rows_as_csv(self) -> bytes: + """Fetch all rows via GetQueryResults API as the text of a CSV result file. - return all_rows + GetQueryResults returns the same text for each value as the CSV result file + has, and this writes it in the same format: a header of the column labels, + each value quoted with its quotes doubled, NULL as an empty field, and a line + feed after each row. The cursors that read CSV result files can then read it + as one. + + Returns: + The CSV text encoded in UTF-8, or empty bytes when the result has no columns. + + Raises: + ProgrammingError: If the metadata is not available. + """ + if self._metadata is None: + raise ProgrammingError("Metadata is not available.") + description = self.description if self.description else [] + if not description: + return b"" + lines = [",".join('"' + d[0].replace('"', '""') + '"' for d in description)] + for rows, offset in self._iter_all_row_pages(): + lines.extend( + ",".join( + "" + if (value := data.get("VarCharValue")) is None + else '"' + value.replace('"', '""') + '"' + for data in row.get("Data", []) + ) + for row in rows[offset:] + ) + lines.append("") + return "\n".join(lines).encode("utf-8") def _get_content_length(self) -> int: if not self.output_location: diff --git a/tests/pyathena/arrow/test_converter.py b/tests/pyathena/arrow/test_converter.py index ae9c8cd9..5875e59c 100644 --- a/tests/pyathena/arrow/test_converter.py +++ b/tests/pyathena/arrow/test_converter.py @@ -5,7 +5,12 @@ # # SPDX-License-Identifier: MIT -from pyathena.arrow.converter import DefaultArrowUnloadTypeConverter +from datetime import datetime + +import pyarrow as pa +import pytest + +from pyathena.arrow.converter import DefaultArrowUnloadTypeConverter, _to_timestamp class TestDefaultArrowUnloadTypeConverter: @@ -13,3 +18,43 @@ def test_convert_delegates_to_default(self): """convert() dispatches through the default converter instead of returning None.""" converter = DefaultArrowUnloadTypeConverter() assert converter.convert("varchar", "hello") == "hello" + + +@pytest.mark.parametrize( + ("unit", "microseconds"), + [ + ("s", [0, 0, 0, 0, 0]), + ("ms", [0, 123000, 123000, 123000, 123000]), + ("us", [0, 123000, 123456, 123456, 123456]), + ], +) +def test_to_timestamp(unit, microseconds): + """Timestamp text with up to 12 fractional digits is truncated to the unit. + + NULL can be null or an empty string, depending on the CSV read options. + """ + column = pa.chunked_array( + [ + pa.array( + [ + "2020-01-02 03:04:05", + "2020-01-02 03:04:05.123", + "2020-01-02 03:04:05.123456", + "2020-01-02 03:04:05.123456789", + "2020-01-02 03:04:05.123456789012", + "0001-01-01 00:00:00.000", + "", + None, + ], + pa.string(), + ) + ] + ) + values = _to_timestamp(column, pa.timestamp(unit)) + assert values.type == pa.timestamp(unit) + assert values.to_pylist() == [ + *(datetime(2020, 1, 2, 3, 4, 5, us) for us in microseconds), + datetime(1, 1, 1), + None, + None, + ] diff --git a/tests/pyathena/arrow/test_cursor.py b/tests/pyathena/arrow/test_cursor.py index 54bc9253..54352fa9 100644 --- a/tests/pyathena/arrow/test_cursor.py +++ b/tests/pyathena/arrow/test_cursor.py @@ -27,7 +27,11 @@ from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect -from tests.pyathena.util import CONVERTED_VALUES_QUERY, CONVERTED_VALUES_ROW +from tests.pyathena.util import ( + CONVERTED_VALUES_QUERY, + CONVERTED_VALUES_ROW, + RESULT_FILE_VALUES_QUERY, +) class TestArrowCursor: @@ -373,7 +377,7 @@ def test_complex_as_arrow(self, arrow_cursor): 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_timestamp", pa.timestamp("us")), pa.field("col_time", pa.string()), pa.field("col_date", pa.timestamp("ms")), pa.field("col_binary", pa.string()), @@ -550,7 +554,7 @@ def test_complex_as_polars(self, arrow_cursor): pl.Float64, pl.String, pl.String, - pl.Datetime("ms"), + pl.Datetime("us"), pl.String, pl.Datetime("ms"), pl.String, @@ -1055,6 +1059,31 @@ def test_fetch_all_rows(self, arrow_cursor): "2024-02-29 23:59:58.123 +05:30" ] + @pytest.mark.skipif(not ENV.managed_work_group, reason="AWS_ATHENA_MANAGED_WORKGROUP not set") + def test_managed_results_match_result_file(self): + """Managed results have the types and values of the CSV result file. + + The cursor's converter applies to them, including a custom one. + """ + results = [] + for kwargs in ({}, {"work_group": ENV.managed_work_group, "s3_staging_dir": ""}): + converter = DefaultArrowTypeConverter() + converter.set("varchar", lambda value: value.upper() if value else value) + with ( + contextlib.closing(connect(**kwargs)) as conn, + conn.cursor(ArrowCursor, converter=converter) as cursor, + ): + cursor.execute(RESULT_FILE_VALUES_QUERY) + results.append((cursor.as_arrow(), cursor.fetchall())) + (table, rows), (managed_table, managed_rows) = results + assert managed_table.schema == table.schema + assert managed_table.equals(table) + assert managed_rows == rows + assert rows[0][1] == 'A,"B"\nC' + assert table.schema.field("col_timestamp_6").type == pa.timestamp("us") + # Fractions finer than microseconds are truncated. + assert (rows[0][5], rows[0][13]) == (datetime(2020, 1, 2, 3, 4, 5, 123456),) * 2 + @pytest.mark.parametrize( "arrow_cursor", [ diff --git a/tests/pyathena/arrow/test_result_set.py b/tests/pyathena/arrow/test_result_set.py index c79aa1e2..90cf33c9 100644 --- a/tests/pyathena/arrow/test_result_set.py +++ b/tests/pyathena/arrow/test_result_set.py @@ -4,8 +4,11 @@ # See LICENSE or https://opensource.org/licenses/MIT. # # SPDX-License-Identifier: MIT +from datetime import datetime from unittest.mock import MagicMock, patch +import pyarrow as pa + from pyathena.arrow.converter import DefaultArrowTypeConverter from pyathena.arrow.result_set import AthenaArrowResultSet from pyathena.model import AthenaQueryExecution @@ -27,3 +30,27 @@ def test_fetch_after_close(self): assert result_set.fetchone() is None assert result_set.fetchmany() == [] assert result_set.fetchall() == [] + + def test_read_csv_timestamps_with_the_same_name(self): + """Timestamp columns are converted by position, also with names that repeat. + + No AWS calls; the GetQueryResults rows are mocked. + """ + result_set = AthenaArrowResultSet.__new__(AthenaArrowResultSet) # bypass __init__ + result_set._query_execution = None + result_set._converter = DefaultArrowTypeConverter() + result_set._block_size = AthenaArrowResultSet.DEFAULT_BLOCK_SIZE + result_set._metadata = tuple( + {"Name": n, "Type": t, "Precision": 3, "Scale": 0, "Nullable": "UNKNOWN"} + for n, t in (("x", "varchar"), ("x", "timestamp"), ("y", "timestamp")) + ) + data = b'"x","x","y"\n"a","2020-01-02 03:04:05.123456789","2020-01-02 03:04:05.1"\n"b",,\n' + with patch.object(AthenaArrowResultSet, "_fetch_all_rows_as_csv", return_value=data): + table = result_set._read_csv() + assert table.column_names == ["x", "x", "y"] + assert table.schema.types == [pa.string(), pa.timestamp("us"), pa.timestamp("us")] + rows = list(zip(*(column.to_pylist() for column in table.columns), strict=True)) + assert rows == [ + ("a", datetime(2020, 1, 2, 3, 4, 5, 123456), datetime(2020, 1, 2, 3, 4, 5, 100000)), + ("b", None, None), + ] diff --git a/tests/pyathena/pandas/test_cursor.py b/tests/pyathena/pandas/test_cursor.py index 44bbf259..311c0bed 100644 --- a/tests/pyathena/pandas/test_cursor.py +++ b/tests/pyathena/pandas/test_cursor.py @@ -27,7 +27,12 @@ from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect -from tests.pyathena.util import CONVERTED_VALUES_QUERY, CONVERTED_VALUES_ROW, cached_file_systems +from tests.pyathena.util import ( + CONVERTED_VALUES_QUERY, + CONVERTED_VALUES_ROW, + RESULT_FILE_VALUES_QUERY, + cached_file_systems, +) def _pandas_converter_without_bigint_dtype(): @@ -1750,6 +1755,27 @@ def test_fetch_all_rows(self, pandas_cursor): pandas_cursor.execute(CONVERTED_VALUES_QUERY) assert pandas_cursor.fetchall() == [CONVERTED_VALUES_ROW] + @pytest.mark.skipif(not ENV.managed_work_group, reason="AWS_ATHENA_MANAGED_WORKGROUP not set") + def test_managed_results_match_result_file(self): + """Managed results have the dtypes and values of the CSV result file. + + The cursor's converter applies to them, including a custom one. + """ + results = [] + for kwargs in ({}, {"work_group": ENV.managed_work_group, "s3_staging_dir": ""}): + converter = DefaultPandasTypeConverter() + converter.set("varchar", lambda value: value.upper() if value else value) + with ( + contextlib.closing(connect(**kwargs)) as conn, + conn.cursor(PandasCursor, converter=converter) as cursor, + ): + cursor.execute(RESULT_FILE_VALUES_QUERY) + results.append((cursor.as_pandas(), cursor.fetchone())) + (df, row), (managed_df, managed_row) = results + pd.testing.assert_frame_equal(managed_df, df) + assert managed_row == row + assert row[1] == 'A,"B"\nC' + @pytest.mark.parametrize( "pandas_cursor", [ diff --git a/tests/pyathena/polars/test_converter.py b/tests/pyathena/polars/test_converter.py index 07b1b7de..f025e001 100644 --- a/tests/pyathena/polars/test_converter.py +++ b/tests/pyathena/polars/test_converter.py @@ -5,7 +5,12 @@ # # SPDX-License-Identifier: MIT -from pyathena.polars.converter import DefaultPolarsUnloadTypeConverter +from datetime import datetime + +import polars as pl +import pytest + +from pyathena.polars.converter import DefaultPolarsUnloadTypeConverter, _to_datetimes class TestDefaultPolarsUnloadTypeConverter: @@ -13,3 +18,41 @@ def test_convert_delegates_to_default(self): """convert() dispatches through the default converter instead of returning None.""" converter = DefaultPolarsUnloadTypeConverter() assert converter.convert("varchar", "hello") == "hello" + + +@pytest.mark.parametrize( + ("dtype", "microseconds"), + [ + (pl.Datetime, [0, 123000, 123456, 123456, 123456]), + (pl.Datetime("ms"), [0, 123000, 123000, 123000, 123000]), + (pl.Datetime("us"), [0, 123000, 123456, 123456, 123456]), + ], +) +def test_to_datetimes(dtype, microseconds): + """Timestamp text with up to 12 fractional digits is truncated to the time unit. + + NULL can be null or an empty string, depending on the read options. A column that + the DataFrame does not have, such as one not selected, is skipped. + """ + df = pl.DataFrame( + { + "t": [ + "2020-01-02 03:04:05", + "2020-01-02 03:04:05.123", + "2020-01-02 03:04:05.123456", + "2020-01-02 03:04:05.123456789", + "2020-01-02 03:04:05.123456789012", + "0001-01-01 00:00:00.000", + "", + None, + ] + } + ) + result = _to_datetimes(df, {"t": dtype, "missing": dtype}) + assert result.schema["t"] == dtype + assert result["t"].to_list() == [ + *(datetime(2020, 1, 2, 3, 4, 5, us) for us in microseconds), + datetime(1, 1, 1), + None, + None, + ] diff --git a/tests/pyathena/polars/test_cursor.py b/tests/pyathena/polars/test_cursor.py index d4c642eb..d739bc22 100644 --- a/tests/pyathena/polars/test_cursor.py +++ b/tests/pyathena/polars/test_cursor.py @@ -19,12 +19,18 @@ from pyathena.error import DatabaseError, ProgrammingError from pyathena.model import AthenaQueryExecution +from pyathena.polars.converter import DefaultPolarsTypeConverter from pyathena.polars.cursor import PolarsCursor from pyathena.polars.result_set import AthenaPolarsResultSet from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect -from tests.pyathena.util import CONVERTED_VALUES_QUERY, CONVERTED_VALUES_ROW, cached_file_systems +from tests.pyathena.util import ( + CONVERTED_VALUES_QUERY, + CONVERTED_VALUES_ROW, + RESULT_FILE_VALUES_QUERY, + cached_file_systems, +) class TestPolarsCursor: @@ -787,6 +793,30 @@ def test_fetch_all_rows(self, polars_cursor): "2024-02-29 23:59:58.123 +05:30" ] + @pytest.mark.skipif(not ENV.managed_work_group, reason="AWS_ATHENA_MANAGED_WORKGROUP not set") + def test_managed_results_match_result_file(self): + """Managed results have the types and values of the CSV result file. + + The cursor's converter applies to them, including a custom one. + """ + results = [] + for kwargs in ({}, {"work_group": ENV.managed_work_group, "s3_staging_dir": ""}): + converter = DefaultPolarsTypeConverter() + converter.set("varchar", lambda value: value.upper() if value else value) + with ( + contextlib.closing(connect(**kwargs)) as conn, + conn.cursor(PolarsCursor, converter=converter) as cursor, + ): + cursor.execute(RESULT_FILE_VALUES_QUERY) + results.append((cursor.as_polars(), cursor.fetchall())) + (df, rows), (managed_df, managed_rows) = results + assert managed_df.schema == df.schema + assert managed_df.equals(df) + assert managed_rows == rows + assert rows[0][1] == 'A,"B"\nC' + # Fractions finer than microseconds are truncated. + assert (rows[0][5], rows[0][13]) == (datetime(2020, 1, 2, 3, 4, 5, 123456),) * 2 + @pytest.mark.parametrize( "polars_cursor", [ diff --git a/tests/pyathena/polars/test_result_set.py b/tests/pyathena/polars/test_result_set.py index 231f8daa..a39fbcb1 100644 --- a/tests/pyathena/polars/test_result_set.py +++ b/tests/pyathena/polars/test_result_set.py @@ -5,12 +5,14 @@ # # SPDX-License-Identifier: MIT +from datetime import datetime from unittest.mock import PropertyMock, patch import polars as pl import pytest from pyathena.error import OperationalError +from pyathena.polars.converter import DefaultPolarsTypeConverter from pyathena.polars.result_set import AthenaPolarsResultSet, PolarsDataFrameIterator _ROWS_BEFORE_FAILURE = 300_000 @@ -25,6 +27,7 @@ def _chunked_result_set() -> AthenaPolarsResultSet: result_set = AthenaPolarsResultSet.__new__(AthenaPolarsResultSet) # bypass __init__ result_set._chunksize = 10_000 result_set._kwargs = {} + result_set._metadata = None return result_set @@ -206,6 +209,51 @@ def test_storage_options_replace_defaults(self, reader, function): list(result) assert read.call_args.kwargs["storage_options"] == {"anon": True} + @pytest.mark.parametrize( + ("kwargs", "expected"), + [ + ({}, {"t": [datetime(2020, 1, 2, 3, 4, 5, 123456), None]}), + ({"columns": ["v"]}, {}), + ({"new_columns": ["t2", "v2"]}, {"t2": [datetime(2020, 1, 2, 3, 4, 5, 123456), None]}), + ({"with_column_names": lambda names: [n.upper() for n in names]}, OperationalError), + ], + ) + def test_read_csv_truncates_timestamps(self, kwargs, expected): + """Timestamps that fail to parse are read again as text and truncated. + + With ``with_column_names`` given to execute(), Polars renames the columns as it + reads them, so they are not, and the read fails as before. No AWS calls; the + GetQueryResults rows are mocked. + """ + result_set = AthenaPolarsResultSet.__new__(AthenaPolarsResultSet) # bypass __init__ + result_set._query_execution = None + result_set._converter = DefaultPolarsTypeConverter() + result_set._kwargs = kwargs + result_set._metadata = tuple( + {"Name": n, "Type": t, "Precision": 3, "Scale": 0, "Nullable": "UNKNOWN"} + for n, t in (("t", "timestamp"), ("v", "varchar")) + ) + data = b'"t","v"\n"2020-01-02 03:04:05.123456789012","x"\n,"y"\n' + with ( + patch.object(AthenaPolarsResultSet, "_fetch_all_rows_as_csv", return_value=data), + patch.object( + AthenaPolarsResultSet, + "_csv_storage_options", + new_callable=PropertyMock, + return_value={}, + ), + ): + if expected is OperationalError: + with pytest.raises(OperationalError): + result_set._read_csv() + return + df = result_set._read_csv() + for name, values in expected.items(): + assert df.schema[name] == pl.Datetime("us") + assert df[name].to_list() == values + if not expected: + assert df.columns == ["v"] + class TestPolarsDataFrameIterator: @pytest.mark.parametrize( diff --git a/tests/pyathena/test_converter.py b/tests/pyathena/test_converter.py index 9293047f..fb32ba76 100644 --- a/tests/pyathena/test_converter.py +++ b/tests/pyathena/test_converter.py @@ -5,7 +5,6 @@ from pyathena.converter import ( DefaultTypeConverter, - _text_value_converter, _to_array, _to_datetime, _to_datetime_with_tz, @@ -675,21 +674,6 @@ def test_typed_time_with_tz_elements(): ) == {"a": time(12, 34, 56, tzinfo=jst)} -def test_text_value_converter(): - """The fallback converter keeps text values and nested time zones as text.""" - converter = _text_value_converter() - assert converter.convert("json", '{"a": 1}') == '{"a": 1}' - assert converter.convert("time with time zone", "12:34:56+09:00") == "12:34:56+09:00" - assert converter.convert( - "array", "[12:34:56.789+09:00]", type_hint="array(time with time zone)" - ) == ["12:34:56.789+09:00"] - assert converter.convert("array", '[{"a": 1}]', type_hint="array(json)") == [{"a": 1}] - # Other converters keep the default mappings. - assert DefaultTypeConverter().convert( - "array", "[12:34:56.789+09:00]", type_hint="array(time with time zone)" - ) == [time(12, 34, 56, 789000, tzinfo=timezone(timedelta(hours=9)))] - - @pytest.mark.parametrize( ("input_value", "expected"), [ diff --git a/tests/pyathena/test_result_set.py b/tests/pyathena/test_result_set.py index 4eb76777..27d7c3b1 100644 --- a/tests/pyathena/test_result_set.py +++ b/tests/pyathena/test_result_set.py @@ -23,6 +23,10 @@ def _page(values, next_token=None): return response +def _row(*values): + return {"Data": [{} if v is None else {"VarCharValue": v} for v in values]} + + class TestAthenaResultSet: def test_fetch_all_rows_skips_column_labels_only_on_first_page(self): """A later page can start with a data row equal to the column labels. @@ -44,6 +48,26 @@ def test_fetch_all_rows_skips_column_labels_only_on_first_page(self): with patch.object(result_set, "_get_query_results", side_effect=pages): assert result_set._fetch_all_rows() == [("1",), ("a",), ("2",)] + def test_fetch_all_rows_as_csv(self): + """The GetQueryResults rows are written as Athena writes a CSV result file. + + Only the first page starts with the column labels, so a later row equal to + them is data. No AWS calls; GetQueryResults is mocked. + """ + result_set = AthenaResultSet.__new__(AthenaResultSet) # bypass __init__ + result_set._query_execution = None + result_set._metadata = tuple( + {"Name": name, "Type": "varchar", "Precision": 0, "Scale": 0, "Nullable": "UNKNOWN"} + for name in ("v", 'w"') + ) + pages = [ + {"ResultSet": {"Rows": [_row("v", 'w"'), _row('a,"b"\nc', None)]}, "NextToken": "t"}, + {"ResultSet": {"Rows": [_row("v", 'w"'), _row("", "1")]}}, + ] + with patch.object(AthenaResultSet, "_get_query_results", side_effect=pages): + data = result_set._fetch_all_rows_as_csv() + assert data == b'"v","w"""\n"a,""b""\nc",\n"v","w"""\n"","1"\n' + class TestWithResultSet: def test_is_mixin(self): diff --git a/tests/pyathena/util.py b/tests/pyathena/util.py index cdd77799..cb194964 100644 --- a/tests/pyathena/util.py +++ b/tests/pyathena/util.py @@ -144,6 +144,21 @@ def unreachable_glue(connection): None, ) +# Values that the cursors reading CSV result files must read the same way from the +# GetQueryResults rows of managed query result storage. +RESULT_FILE_VALUES_QUERY = """ +SELECT * FROM (VALUES + (1, 'a,"b"' || chr(10) || 'c', '', DATE '2020-01-02', TIMESTAMP '2020-01-02 03:04:05.123', + CAST('2020-01-02 03:04:05.123456' AS TIMESTAMP(6)), CAST('12.30' AS DECIMAL(10, 2)), X'0102', + ARRAY[1, 2], MAP(ARRAY['k'], ARRAY[1]), CAST(ROW(1, 'x') AS ROW(a INTEGER, b VARCHAR)), + json_parse('{"a": [1, 2]}'), BIGINT '9007199254740993', + CAST('2020-01-02 03:04:05.123456789012' AS TIMESTAMP(12))), + (2, NULL, NULL, NULL, NULL, NULL, NULL, NULL, NULL, NULL, NULL, NULL, NULL, NULL) +) AS t(col_int, col_varchar, col_empty, col_date, col_timestamp, col_timestamp_6, col_decimal, + col_varbinary, col_array, col_map, col_row, col_json, col_bigint, col_timestamp_12) +ORDER BY col_int +""" + # A Spark job whose executor tasks sleep, so that StopCalculationExecution can cancel it. CANCELABLE_SPARK_JOB = """