diff --git a/pyathena/aio/arrow/cursor.py b/pyathena/aio/arrow/cursor.py index b8c915515..b66778bc9 100644 --- a/pyathena/aio/arrow/cursor.py +++ b/pyathena/aio/arrow/cursor.py @@ -144,6 +144,8 @@ async def execute( :class:`~pyathena.options.ExecuteOptions` instance. Individual keyword arguments take precedence over ``options`` fields. **kwargs: Additional execution parameters. + ``block_size`` sets the read block size for this query, and + ``connect_timeout`` and ``request_timeout`` override the cursor's values. Returns: Self reference for method chaining. @@ -182,8 +184,8 @@ async def execute( retry_config=self._retry_config, unload=self._unload, unload_location=unload_location, - connect_timeout=self._connect_timeout, - request_timeout=self._request_timeout, + connect_timeout=kwargs.pop("connect_timeout", self._connect_timeout), + request_timeout=kwargs.pop("request_timeout", self._request_timeout), result_set_type_hints=options.result_set_type_hints, **kwargs, ) diff --git a/pyathena/aio/pandas/cursor.py b/pyathena/aio/pandas/cursor.py index 43350c1f2..bca2824b2 100644 --- a/pyathena/aio/pandas/cursor.py +++ b/pyathena/aio/pandas/cursor.py @@ -167,6 +167,11 @@ async def execute( :class:`~pyathena.options.ExecuteOptions` instance. Individual keyword arguments take precedence over ``options`` fields. **kwargs: Additional pandas read_csv/read_parquet parameters. + ``engine``, ``chunksize``, ``block_size``, ``cache_type``, ``max_workers``, + and ``auto_optimize_chunksize`` override the cursor's values for this query. + ``storage_options`` and, for UNLOAD results, ``filesystem`` replace + PyAthena's S3 filesystem (see + :class:`~pyathena.pandas.result_set.AthenaPandasResultSet`). Returns: Self reference for method chaining. @@ -213,7 +218,9 @@ async def execute( block_size=kwargs.pop("block_size", self._block_size), cache_type=kwargs.pop("cache_type", self._cache_type), max_workers=kwargs.pop("max_workers", self._max_workers), - auto_optimize_chunksize=self._auto_optimize_chunksize, + auto_optimize_chunksize=kwargs.pop( + "auto_optimize_chunksize", self._auto_optimize_chunksize + ), result_set_type_hints=options.result_set_type_hints, **kwargs, ) diff --git a/pyathena/aio/polars/cursor.py b/pyathena/aio/polars/cursor.py index 8f584da6f..12cd3a64e 100644 --- a/pyathena/aio/polars/cursor.py +++ b/pyathena/aio/polars/cursor.py @@ -151,6 +151,11 @@ async def execute( :class:`~pyathena.options.ExecuteOptions` instance. Individual keyword arguments take precedence over ``options`` fields. **kwargs: Additional execution parameters passed to Polars read functions. + ``block_size``, ``cache_type``, ``max_workers``, and ``chunksize`` + override the cursor's values for this query. + Read function arguments replace the ones the result set chooses, such as + ``separator``, ``has_header``, ``schema_overrides``, and ``storage_options`` + (see :class:`~pyathena.polars.result_set.AthenaPolarsResultSet`). Returns: Self reference for method chaining. @@ -189,10 +194,10 @@ async def execute( retry_config=self._retry_config, unload=self._unload, unload_location=unload_location, - block_size=self._block_size, - cache_type=self._cache_type, - max_workers=self._max_workers, - chunksize=self._chunksize, + block_size=kwargs.pop("block_size", self._block_size), + cache_type=kwargs.pop("cache_type", self._cache_type), + max_workers=kwargs.pop("max_workers", self._max_workers), + chunksize=kwargs.pop("chunksize", self._chunksize), result_set_type_hints=options.result_set_type_hints, **kwargs, ) diff --git a/pyathena/aio/s3fs/cursor.py b/pyathena/aio/s3fs/cursor.py index 172348761..93ef181d6 100644 --- a/pyathena/aio/s3fs/cursor.py +++ b/pyathena/aio/s3fs/cursor.py @@ -147,6 +147,8 @@ async def execute( :class:`~pyathena.options.ExecuteOptions` instance. Individual keyword arguments take precedence over ``options`` fields. **kwargs: Additional execution parameters. + ``block_size`` sets the read block size for this query, and + ``csv_reader`` overrides the cursor's value. Returns: Self reference for method chaining. @@ -182,7 +184,7 @@ async def execute( query_execution=query_execution, arraysize=self.arraysize, retry_config=self._retry_config, - csv_reader=self._csv_reader, + csv_reader=kwargs.pop("csv_reader", self._csv_reader), filesystem_class=AioS3FileSystem, result_set_type_hints=options.result_set_type_hints, **kwargs, diff --git a/pyathena/arrow/async_cursor.py b/pyathena/arrow/async_cursor.py index 873a19090..340147bc7 100644 --- a/pyathena/arrow/async_cursor.py +++ b/pyathena/arrow/async_cursor.py @@ -171,8 +171,8 @@ def _collect_result_set( retry_config=self._retry_config, unload=self._unload, unload_location=unload_location, - connect_timeout=self._connect_timeout, - request_timeout=self._request_timeout, + connect_timeout=kwargs.pop("connect_timeout", self._connect_timeout), + request_timeout=kwargs.pop("request_timeout", self._request_timeout), result_set_type_hints=result_set_type_hints, **kwargs, ) @@ -212,6 +212,8 @@ def execute( :class:`~pyathena.options.ExecuteOptions` instance. Individual keyword arguments take precedence over ``options`` fields. **kwargs: Additional execution parameters. + ``block_size`` sets the read block size for this query, and + ``connect_timeout`` and ``request_timeout`` override the cursor's values. Returns: Tuple of (query_id, future) where future resolves to AthenaArrowResultSet. diff --git a/pyathena/arrow/cursor.py b/pyathena/arrow/cursor.py index e98890604..081350343 100644 --- a/pyathena/arrow/cursor.py +++ b/pyathena/arrow/cursor.py @@ -170,6 +170,8 @@ def execute( :class:`~pyathena.options.ExecuteOptions` instance. Individual keyword arguments take precedence over ``options`` fields. **kwargs: Additional execution parameters. + ``block_size`` sets the read block size for this query, and + ``connect_timeout`` and ``request_timeout`` override the cursor's values. Returns: Self reference for method chaining. @@ -210,8 +212,8 @@ def execute( retry_config=self._retry_config, unload=self._unload, unload_location=unload_location, - connect_timeout=self._connect_timeout, - request_timeout=self._request_timeout, + connect_timeout=kwargs.pop("connect_timeout", self._connect_timeout), + request_timeout=kwargs.pop("request_timeout", self._request_timeout), result_set_type_hints=options.result_set_type_hints, **kwargs, ) diff --git a/pyathena/pandas/async_cursor.py b/pyathena/pandas/async_cursor.py index 2328dc18a..a23cc9cfd 100644 --- a/pyathena/pandas/async_cursor.py +++ b/pyathena/pandas/async_cursor.py @@ -79,6 +79,9 @@ def __init__( chunksize: int | None = None, result_reuse_enable: bool = False, result_reuse_minutes: int = CursorIterator.DEFAULT_RESULT_REUSE_MINUTES, + block_size: int | None = None, + cache_type: str | None = None, + auto_optimize_chunksize: bool = False, **kwargs, ) -> None: """Initialize an AsyncPandasCursor. @@ -99,8 +102,13 @@ def __init__( unload: Whether to wrap queries in ``UNLOAD`` and read the Parquet output. engine: Parsing engine (``auto``, ``c``, ``python``, or ``pyarrow``). chunksize: Number of rows per DataFrame chunk when reading CSV results. + If set, it takes precedence over ``auto_optimize_chunksize``. result_reuse_enable: Whether to enable Athena query result reuse. result_reuse_minutes: Maximum age of a reused query result in minutes. + block_size: Default block size of the S3 filesystem that reads the results. + cache_type: Default cache type of the S3 filesystem that reads the results. + auto_optimize_chunksize: Whether to choose a chunk size from the size of the + CSV result file when ``chunksize`` is None. **kwargs: Other cursor arguments, such as ``connection`` and ``converter``, passed to ``AsyncCursor.__init__``. """ @@ -122,6 +130,9 @@ def __init__( self._unload = unload self._engine = engine self._chunksize = chunksize + self._block_size = block_size + self._cache_type = cache_type + self._auto_optimize_chunksize = auto_optimize_chunksize @staticmethod @override @@ -170,6 +181,11 @@ def _collect_result_set( unload_location=unload_location, engine=kwargs.pop("engine", self._engine), chunksize=kwargs.pop("chunksize", self._chunksize), + block_size=kwargs.pop("block_size", self._block_size), + cache_type=kwargs.pop("cache_type", self._cache_type), + auto_optimize_chunksize=kwargs.pop( + "auto_optimize_chunksize", self._auto_optimize_chunksize + ), result_set_type_hints=result_set_type_hints, **kwargs, ) @@ -215,6 +231,12 @@ def execute( :class:`~pyathena.options.ExecuteOptions` instance. Individual keyword arguments take precedence over ``options`` fields. **kwargs: Additional pandas read_csv/read_parquet parameters. + ``engine``, ``chunksize``, ``block_size``, ``cache_type``, and + ``auto_optimize_chunksize`` override the cursor's values for this query. + ``storage_options`` and, for UNLOAD results, ``filesystem`` replace + PyAthena's S3 filesystem (see + :class:`~pyathena.pandas.result_set.AthenaPandasResultSet`). + ``max_workers`` sets the number of S3 read workers for this query. Returns: Tuple of (query_id, future) where future resolves to AthenaPandasResultSet. diff --git a/pyathena/pandas/cursor.py b/pyathena/pandas/cursor.py index 97bf773f9..afcd12cec 100644 --- a/pyathena/pandas/cursor.py +++ b/pyathena/pandas/cursor.py @@ -193,6 +193,11 @@ def execute( :class:`~pyathena.options.ExecuteOptions` instance. Individual keyword arguments take precedence over ``options`` fields. **kwargs: Additional pandas read_csv/read_parquet parameters. + ``engine``, ``chunksize``, ``block_size``, ``cache_type``, ``max_workers``, + and ``auto_optimize_chunksize`` override the cursor's values for this query. + ``storage_options`` and, for UNLOAD results, ``filesystem`` replace + PyAthena's S3 filesystem (see + :class:`~pyathena.pandas.result_set.AthenaPandasResultSet`). Returns: Self reference for method chaining. @@ -242,7 +247,9 @@ def execute( block_size=kwargs.pop("block_size", self._block_size), cache_type=kwargs.pop("cache_type", self._cache_type), max_workers=kwargs.pop("max_workers", self._max_workers), - auto_optimize_chunksize=self._auto_optimize_chunksize, + auto_optimize_chunksize=kwargs.pop( + "auto_optimize_chunksize", self._auto_optimize_chunksize + ), result_set_type_hints=options.result_set_type_hints, **kwargs, ) diff --git a/pyathena/pandas/result_set.py b/pyathena/pandas/result_set.py index 6acd96cbd..638642723 100644 --- a/pyathena/pandas/result_set.py +++ b/pyathena/pandas/result_set.py @@ -310,6 +310,10 @@ def __init__( result_set_type_hints: Athena type signatures for complex-type columns, keyed by column name (case-insensitive) or zero-based column index. **kwargs: Additional arguments passed to pandas.read_csv/read_parquet. + A given ``storage_options``, even None, replaces PyAthena's S3 filesystem + for reading the result files, and so does ``filesystem`` for UNLOAD results. + The UNLOAD manifest is still read with the connection's S3 client, and the + schema with PyAthena's filesystem. """ super().__init__( connection=connection, @@ -778,23 +782,22 @@ def _read_parquet(self, engine) -> DataFrame: self._unload_location = "/".join(self._data_manifest[0].split("/")[:-1]) + "/" if engine == "pyarrow": - # pyarrow takes the path without the scheme with an fsspec filesystem. - bucket, key = parse_output_location(self._unload_location) - unload_location = f"{bucket}/{key}" - kwargs = { - "use_threads": True, - } + kwargs: dict[str, Any] = {"use_threads": True, **self._kwargs} + # Given storage_options, even None, pandas opens the files itself, + # as for CSV results. + if "filesystem" not in kwargs and "storage_options" not in kwargs: + kwargs["filesystem"] = self._fs + if kwargs.get("filesystem") is None: + unload_location = self._unload_location + else: + # pyarrow takes the path without the scheme with a filesystem. + bucket, key = parse_output_location(self._unload_location) + unload_location = f"{bucket}/{key}" else: raise ProgrammingError("Engine must be `pyarrow`.") - kwargs.update(self._kwargs) try: - return pd.read_parquet( - unload_location, - engine=self._engine, - filesystem=self._fs, - **kwargs, - ) + return pd.read_parquet(unload_location, engine=self._engine, **kwargs) except Exception as e: _logger.exception(f"Failed to read {self.output_location}.") raise OperationalError(*e.args) from e diff --git a/pyathena/polars/async_cursor.py b/pyathena/polars/async_cursor.py index cf63e6ad8..96bb6c600 100644 --- a/pyathena/polars/async_cursor.py +++ b/pyathena/polars/async_cursor.py @@ -183,10 +183,10 @@ def _collect_result_set( retry_config=self._retry_config, unload=self._unload, unload_location=unload_location, - block_size=self._block_size, - cache_type=self._cache_type, - max_workers=self._max_workers, - chunksize=self._chunksize, + block_size=kwargs.pop("block_size", self._block_size), + cache_type=kwargs.pop("cache_type", self._cache_type), + max_workers=kwargs.pop("max_workers", self._max_workers), + chunksize=kwargs.pop("chunksize", self._chunksize), result_set_type_hints=result_set_type_hints, **kwargs, ) @@ -230,6 +230,11 @@ def execute( :class:`~pyathena.options.ExecuteOptions` instance. Individual keyword arguments take precedence over ``options`` fields. **kwargs: Additional execution parameters passed to Polars read functions. + ``block_size``, ``cache_type``, ``max_workers``, and ``chunksize`` + override the cursor's values for this query. + Read function arguments replace the ones the result set chooses, such as + ``separator``, ``has_header``, ``schema_overrides``, and ``storage_options`` + (see :class:`~pyathena.polars.result_set.AthenaPolarsResultSet`). Returns: Tuple of (query_id, future) where future resolves to AthenaPolarsResultSet. diff --git a/pyathena/polars/cursor.py b/pyathena/polars/cursor.py index 095d1ce7a..05c91b307 100644 --- a/pyathena/polars/cursor.py +++ b/pyathena/polars/cursor.py @@ -189,6 +189,11 @@ def execute( :class:`~pyathena.options.ExecuteOptions` instance. Individual keyword arguments take precedence over ``options`` fields. **kwargs: Additional execution parameters passed to Polars read functions. + ``block_size``, ``cache_type``, ``max_workers``, and ``chunksize`` + override the cursor's values for this query. + Read function arguments replace the ones the result set chooses, such as + ``separator``, ``has_header``, ``schema_overrides``, and ``storage_options`` + (see :class:`~pyathena.polars.result_set.AthenaPolarsResultSet`). Returns: Self reference for method chaining. @@ -229,10 +234,10 @@ def execute( retry_config=self._retry_config, unload=self._unload, unload_location=unload_location, - block_size=self._block_size, - cache_type=self._cache_type, - max_workers=self._max_workers, - chunksize=self._chunksize, + block_size=kwargs.pop("block_size", self._block_size), + cache_type=kwargs.pop("cache_type", self._cache_type), + max_workers=kwargs.pop("max_workers", self._max_workers), + chunksize=kwargs.pop("chunksize", self._chunksize), result_set_type_hints=options.result_set_type_hints, **kwargs, ) diff --git a/pyathena/polars/result_set.py b/pyathena/polars/result_set.py index 34dbebe40..4641614ab 100644 --- a/pyathena/polars/result_set.py +++ b/pyathena/polars/result_set.py @@ -237,6 +237,11 @@ def __init__( result_set_type_hints: Athena type signatures for complex-type columns, keyed by column name (case-insensitive) or zero-based column index. **kwargs: Additional arguments passed to Polars read functions. + They replace the arguments the result set chooses, such as ``separator``, + ``has_header``, ``schema_overrides``, and ``storage_options``. A given + ``storage_options`` replaces PyAthena's S3 settings as a whole: non-chunked + CSV results are read through fsspec, and chunked CSV and UNLOAD results + through Polars' native object store. """ super().__init__( connection=connection, @@ -289,6 +294,42 @@ def __init__( self._column_names_cache: list[str] = self._get_column_names() self._iterrows = self._df_iter.iterrows() + def _storage_options(self, default: Callable[[], dict[str, Any]]) -> Any: + """Get the storage options for a Polars read function. + + Args: + default: Returns the storage options that the result set chooses. It is called + only when ``execute()`` was not given ``storage_options``, so replaced + options do not fetch credentials. + + Returns: + The ``storage_options`` given to ``execute()``, or else the default. + """ + if "storage_options" in self._kwargs: + return self._kwargs["storage_options"] + return default() + + def _read_kwargs( + self, storage_options: Callable[[], dict[str, Any]], **defaults: Any + ) -> dict[str, Any]: + """Combine the arguments of a Polars read function with the ones given to ``execute()``. + + Args: + storage_options: Returns the storage options that the result set chooses; + see ``_storage_options()``. + **defaults: The other arguments that the result set chooses, such as + ``separator``. + + Returns: + The arguments for the read function. A value given to ``execute()`` replaces + the one the result set chose, including the whole ``storage_options``. + """ + return { + **defaults, + **self._kwargs, + "storage_options": self._storage_options(storage_options), + } + @property def _csv_storage_options(self) -> dict[str, Any]: """Get storage options for Polars CSV reading via fsspec. @@ -456,11 +497,12 @@ def _read_csv(self) -> pl.DataFrame: try: df = pl.read_csv( self.output_location, - separator=separator, - has_header=has_header, - schema_overrides=self.dtypes, - storage_options=self._csv_storage_options, - **self._kwargs, + **self._read_kwargs( + lambda: self._csv_storage_options, + separator=separator, + has_header=has_header, + schema_overrides=self.dtypes, + ), ) if new_columns: df.columns = new_columns @@ -489,8 +531,7 @@ def _read_parquet(self) -> pl.DataFrame: try: return pl.read_parquet( self._unload_location, - storage_options=self._parquet_storage_options, - **self._kwargs, + **self._read_kwargs(lambda: self._parquet_storage_options), ) except Exception as e: _logger.exception(f"Failed to read {self._unload_location}.") @@ -507,7 +548,7 @@ def _read_parquet_schema(self) -> tuple[dict[str, Any], ...]: # Use scan_parquet to get schema without reading all data lazy_df = pl.scan_parquet( self._unload_location, - storage_options=self._parquet_storage_options, + storage_options=self._storage_options(lambda: self._parquet_storage_options), ) schema = lazy_df.collect_schema() return to_column_info(schema) @@ -649,11 +690,12 @@ def _iter_csv_chunks(self) -> Iterator[pl.DataFrame]: # not fsspec, so we use the same storage options as Parquet lazy_df = pl.scan_csv( self.output_location, - separator=separator, - has_header=has_header, - schema_overrides=self.dtypes, - storage_options=self._parquet_storage_options, - **self._kwargs, + **self._read_kwargs( + lambda: self._parquet_storage_options, + separator=separator, + has_header=has_header, + schema_overrides=self.dtypes, + ), ) for batch in lazy_df.collect_batches(chunk_size=self._chunksize): if new_columns: @@ -683,8 +725,7 @@ def _iter_parquet_chunks(self) -> Iterator[pl.DataFrame]: try: lazy_df = pl.scan_parquet( self._unload_location, - storage_options=self._parquet_storage_options, - **self._kwargs, + **self._read_kwargs(lambda: self._parquet_storage_options), ) yield from lazy_df.collect_batches(chunk_size=self._chunksize) except Exception as e: diff --git a/pyathena/s3fs/async_cursor.py b/pyathena/s3fs/async_cursor.py index f2ec8a7dd..02b7e61eb 100644 --- a/pyathena/s3fs/async_cursor.py +++ b/pyathena/s3fs/async_cursor.py @@ -172,7 +172,7 @@ def _collect_result_set( query_execution=query_execution, arraysize=self._arraysize, retry_config=self._retry_config, - csv_reader=self._csv_reader, + csv_reader=kwargs.pop("csv_reader", self._csv_reader), result_set_type_hints=result_set_type_hints, **kwargs, ) @@ -215,6 +215,8 @@ def execute( :class:`~pyathena.options.ExecuteOptions` instance. Individual keyword arguments take precedence over ``options`` fields. **kwargs: Additional execution parameters. + ``block_size`` sets the read block size for this query, and + ``csv_reader`` overrides the cursor's value. Returns: Tuple of (query_id, Future[AthenaS3FSResultSet]). diff --git a/pyathena/s3fs/cursor.py b/pyathena/s3fs/cursor.py index e78b51ac2..96698fb4f 100644 --- a/pyathena/s3fs/cursor.py +++ b/pyathena/s3fs/cursor.py @@ -165,6 +165,8 @@ def execute( :class:`~pyathena.options.ExecuteOptions` instance. Individual keyword arguments take precedence over ``options`` fields. **kwargs: Additional execution parameters. + ``block_size`` sets the read block size for this query, and + ``csv_reader`` overrides the cursor's value. Returns: Self reference for method chaining. @@ -203,7 +205,7 @@ def execute( query_execution=query_execution, arraysize=self.arraysize, retry_config=self._retry_config, - csv_reader=self._csv_reader, + csv_reader=kwargs.pop("csv_reader", self._csv_reader), result_set_type_hints=options.result_set_type_hints, **kwargs, ) diff --git a/tests/pyathena/aio/arrow/test_cursor.py b/tests/pyathena/aio/arrow/test_cursor.py index 9346d3bd2..fa4e14e11 100644 --- a/tests/pyathena/aio/arrow/test_cursor.py +++ b/tests/pyathena/aio/arrow/test_cursor.py @@ -5,10 +5,15 @@ # # SPDX-License-Identifier: MIT +from unittest.mock import MagicMock, patch + import pytest +from pyathena.aio.arrow.cursor import AioArrowCursor from pyathena.arrow.result_set import AthenaArrowResultSet from pyathena.error import ProgrammingError +from pyathena.model import AthenaQueryExecution +from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.aio.conftest import _aio_connect @@ -143,3 +148,30 @@ async def test_as_arrow_unload(self, aio_arrow_cursor): await aio_arrow_cursor.execute("SELECT * FROM one_row") table = aio_arrow_cursor.as_arrow() assert table.num_rows == 1 + + @pytest.mark.parametrize( + "execute_kwargs", [{}, {"connect_timeout": 3.0, "request_timeout": 4.0}] + ) + async def test_read_options(self, execute_kwargs): + """The cursor's read options reach the result set, and execute() overrides them. + + No AWS calls; the query and its result set are mocked. + """ + cursor_kwargs = {"connect_timeout": 1.0, "request_timeout": 2.0} + query_execution = MagicMock(state=AthenaQueryExecution.STATE_SUCCEEDED) + cursor = AioArrowCursor( + connection=MagicMock(), + converter=MagicMock(), + formatter=MagicMock(), + retry_config=RetryConfig(), + **cursor_kwargs, + ) + with ( + patch.object(AioArrowCursor, "_execute", return_value="query_id"), + patch.object(AioArrowCursor, "_poll", return_value=query_execution), + patch("pyathena.aio.arrow.cursor.AthenaArrowResultSet") as result_set_class, + ): + await cursor.execute("SELECT 1", **execute_kwargs) + kwargs = result_set_class.call_args.kwargs + expected = {**cursor_kwargs, **execute_kwargs} + assert {key: kwargs[key] for key in expected} == expected diff --git a/tests/pyathena/aio/pandas/test_cursor.py b/tests/pyathena/aio/pandas/test_cursor.py index 4d04c9aeb..ecb123a34 100644 --- a/tests/pyathena/aio/pandas/test_cursor.py +++ b/tests/pyathena/aio/pandas/test_cursor.py @@ -5,10 +5,15 @@ # # SPDX-License-Identifier: MIT +from unittest.mock import MagicMock, patch + import pytest +from pyathena.aio.pandas.cursor import AioPandasCursor from pyathena.error import ProgrammingError +from pyathena.model import AthenaQueryExecution from pyathena.pandas.result_set import AthenaPandasResultSet +from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.aio.conftest import _aio_connect @@ -134,3 +139,44 @@ async def test_as_pandas_unload(self, aio_pandas_cursor): await aio_pandas_cursor.execute("SELECT * FROM one_row") df = aio_pandas_cursor.as_pandas() assert len(df) == 1 + + @pytest.mark.parametrize( + "execute_kwargs", + [ + {}, + { + "block_size": 2048, + "cache_type": "none", + "max_workers": 3, + "auto_optimize_chunksize": False, + }, + ], + ) + async def test_read_options(self, execute_kwargs): + """The cursor's read options reach the result set, and execute() overrides them. + + No AWS calls; the query and its result set are mocked. + """ + cursor_kwargs = { + "block_size": 1024, + "cache_type": "bytes", + "max_workers": 2, + "auto_optimize_chunksize": True, + } + cursor = AioPandasCursor( + connection=MagicMock(), + converter=MagicMock(), + formatter=MagicMock(), + retry_config=RetryConfig(), + **cursor_kwargs, + ) + query_execution = MagicMock(state=AthenaQueryExecution.STATE_SUCCEEDED) + with ( + patch.object(AioPandasCursor, "_execute", return_value="query_id"), + patch.object(AioPandasCursor, "_poll", return_value=query_execution), + patch("pyathena.aio.pandas.cursor.AthenaPandasResultSet") as result_set_class, + ): + await cursor.execute("SELECT 1", **execute_kwargs) + kwargs = result_set_class.call_args.kwargs + expected = {**cursor_kwargs, **execute_kwargs} + assert {key: kwargs[key] for key in expected} == expected diff --git a/tests/pyathena/aio/polars/test_cursor.py b/tests/pyathena/aio/polars/test_cursor.py index 41b556a37..051d078dc 100644 --- a/tests/pyathena/aio/polars/test_cursor.py +++ b/tests/pyathena/aio/polars/test_cursor.py @@ -5,10 +5,15 @@ # # SPDX-License-Identifier: MIT +from unittest.mock import MagicMock, patch + import pytest +from pyathena.aio.polars.cursor import AioPolarsCursor from pyathena.error import ProgrammingError +from pyathena.model import AthenaQueryExecution from pyathena.polars.result_set import AthenaPolarsResultSet +from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.aio.conftest import _aio_connect @@ -128,3 +133,36 @@ async def test_as_polars_unload(self, aio_polars_cursor): await aio_polars_cursor.execute("SELECT * FROM one_row") df = aio_polars_cursor.as_polars() assert df.height == 1 + + @pytest.mark.parametrize( + "execute_kwargs", + [{}, {"block_size": 2048, "cache_type": "none", "max_workers": 3, "chunksize": 20}], + ) + async def test_read_options(self, execute_kwargs): + """The cursor's read options reach the result set, and execute() overrides them. + + No AWS calls; the query and its result set are mocked. + """ + cursor_kwargs = { + "block_size": 1024, + "cache_type": "bytes", + "max_workers": 2, + "chunksize": 10, + } + cursor = AioPolarsCursor( + connection=MagicMock(), + converter=MagicMock(), + formatter=MagicMock(), + retry_config=RetryConfig(), + **cursor_kwargs, + ) + query_execution = MagicMock(state=AthenaQueryExecution.STATE_SUCCEEDED) + with ( + patch.object(AioPolarsCursor, "_execute", return_value="query_id"), + patch.object(AioPolarsCursor, "_poll", return_value=query_execution), + patch("pyathena.aio.polars.cursor.AthenaPolarsResultSet") as result_set_class, + ): + await cursor.execute("SELECT 1", **execute_kwargs) + kwargs = result_set_class.call_args.kwargs + expected = {**cursor_kwargs, **execute_kwargs} + assert {key: kwargs[key] for key in expected} == expected diff --git a/tests/pyathena/aio/s3fs/test_cursor.py b/tests/pyathena/aio/s3fs/test_cursor.py index 4005da897..42f29ab1c 100644 --- a/tests/pyathena/aio/s3fs/test_cursor.py +++ b/tests/pyathena/aio/s3fs/test_cursor.py @@ -5,11 +5,16 @@ # # SPDX-License-Identifier: MIT +from unittest.mock import MagicMock, patch + import pytest from pyathena.aio.s3fs.cursor import AioS3FSCursor from pyathena.error import ProgrammingError +from pyathena.model import AthenaQueryExecution +from pyathena.s3fs.reader import AthenaCSVReader, DefaultCSVReader from pyathena.s3fs.result_set import AthenaS3FSResultSet +from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.aio.conftest import _aio_connect @@ -82,3 +87,28 @@ async def test_context_manager(self): async def test_execute_returns_self(self, aio_s3fs_cursor): result = await aio_s3fs_cursor.execute("SELECT * FROM one_row") assert result is aio_s3fs_cursor + + @pytest.mark.parametrize("execute_kwargs", [{}, {"csv_reader": AthenaCSVReader}]) + async def test_read_options(self, execute_kwargs): + """The cursor's read options reach the result set, and execute() overrides them. + + No AWS calls; the query and its result set are mocked. + """ + cursor_kwargs = {"csv_reader": DefaultCSVReader} + query_execution = MagicMock(state=AthenaQueryExecution.STATE_SUCCEEDED) + cursor = AioS3FSCursor( + connection=MagicMock(), + converter=MagicMock(), + formatter=MagicMock(), + retry_config=RetryConfig(), + **cursor_kwargs, + ) + with ( + patch.object(AioS3FSCursor, "_execute", return_value="query_id"), + patch.object(AioS3FSCursor, "_poll", return_value=query_execution), + patch("pyathena.aio.s3fs.cursor.AthenaS3FSResultSet") as result_set_class, + ): + await cursor.execute("SELECT 1", **execute_kwargs) + kwargs = result_set_class.call_args.kwargs + expected = {**cursor_kwargs, **execute_kwargs} + assert {key: kwargs[key] for key in expected} == expected diff --git a/tests/pyathena/arrow/test_async_cursor.py b/tests/pyathena/arrow/test_async_cursor.py index 734961e31..30ae8f5ce 100644 --- a/tests/pyathena/arrow/test_async_cursor.py +++ b/tests/pyathena/arrow/test_async_cursor.py @@ -4,6 +4,7 @@ import time from datetime import datetime from random import randint +from unittest.mock import MagicMock, patch import pytest @@ -11,6 +12,7 @@ from pyathena.error import NotSupportedError, ProgrammingError from pyathena.model import AthenaQueryExecution from pyathena.result_set import AthenaResultSet +from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect @@ -323,3 +325,31 @@ def test_empty_result_unload(self, async_arrow_cursor): table = future.result().as_arrow() assert table.shape[0] == 0 assert table.shape[1] == 0 + + @pytest.mark.parametrize( + "execute_kwargs", [{}, {"connect_timeout": 3.0, "request_timeout": 4.0}] + ) + def test_read_options(self, execute_kwargs): + """The cursor's read options reach the result set, and execute() overrides them. + + No AWS calls; the query and its result set are mocked. + """ + cursor_kwargs = {"connect_timeout": 1.0, "request_timeout": 2.0} + query_execution = MagicMock(state=AthenaQueryExecution.STATE_SUCCEEDED) + with ( + AsyncArrowCursor( + connection=MagicMock(), + converter=MagicMock(), + formatter=MagicMock(), + retry_config=RetryConfig(), + **cursor_kwargs, + ) as cursor, + patch.object(AsyncArrowCursor, "_execute", return_value="query_id"), + patch.object(AsyncArrowCursor, "_poll", return_value=query_execution), + patch("pyathena.arrow.async_cursor.AthenaArrowResultSet") as result_set_class, + ): + _, future = cursor.execute("SELECT 1", **execute_kwargs) + future.result() + kwargs = result_set_class.call_args.kwargs + expected = {**cursor_kwargs, **execute_kwargs} + assert {key: kwargs[key] for key in expected} == expected diff --git a/tests/pyathena/arrow/test_cursor.py b/tests/pyathena/arrow/test_cursor.py index 8b93ee698..fd3f04b42 100644 --- a/tests/pyathena/arrow/test_cursor.py +++ b/tests/pyathena/arrow/test_cursor.py @@ -12,6 +12,7 @@ from concurrent.futures import ThreadPoolExecutor from datetime import datetime from decimal import Decimal +from unittest.mock import MagicMock, patch import pandas as pd import polars as pl @@ -21,6 +22,8 @@ from pyathena.arrow.cursor import ArrowCursor from pyathena.arrow.result_set import AthenaArrowResultSet from pyathena.error import DatabaseError, ProgrammingError +from pyathena.model import AthenaQueryExecution +from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect @@ -1005,3 +1008,30 @@ def test_fetch_all_rows(self, arrow_cursor): (2, datetime(2017, 1, 1, 12, 34, 56).time(), b"\x01\x02", [1, "x"], "s", None), ] assert arrow_cursor.as_arrow().schema.field("col_json").type == pa.string() + + @pytest.mark.parametrize( + "execute_kwargs", [{}, {"connect_timeout": 3.0, "request_timeout": 4.0}] + ) + def test_read_options(self, execute_kwargs): + """The cursor's read options reach the result set, and execute() overrides them. + + No AWS calls; the query and its result set are mocked. + """ + cursor_kwargs = {"connect_timeout": 1.0, "request_timeout": 2.0} + query_execution = MagicMock(state=AthenaQueryExecution.STATE_SUCCEEDED) + cursor = ArrowCursor( + connection=MagicMock(), + converter=MagicMock(), + formatter=MagicMock(), + retry_config=RetryConfig(), + **cursor_kwargs, + ) + with ( + patch.object(ArrowCursor, "_execute", return_value="query_id"), + patch.object(ArrowCursor, "_poll", return_value=query_execution), + patch("pyathena.arrow.cursor.AthenaArrowResultSet") as result_set_class, + ): + cursor.execute("SELECT 1", **execute_kwargs) + kwargs = result_set_class.call_args.kwargs + expected = {**cursor_kwargs, **execute_kwargs} + assert {key: kwargs[key] for key in expected} == expected diff --git a/tests/pyathena/pandas/test_async_cursor.py b/tests/pyathena/pandas/test_async_cursor.py index 013136275..590d9d8b8 100644 --- a/tests/pyathena/pandas/test_async_cursor.py +++ b/tests/pyathena/pandas/test_async_cursor.py @@ -5,6 +5,7 @@ import time from datetime import datetime from random import randint +from unittest.mock import MagicMock, patch import numpy as np import pandas as pd @@ -13,7 +14,9 @@ from pyathena.error import NotSupportedError, ProgrammingError from pyathena.model import AthenaQueryExecution from pyathena.pandas.async_cursor import AsyncPandasCursor +from pyathena.pandas.result_set import AthenaPandasResultSet from pyathena.result_set import AthenaResultSet +from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect @@ -626,3 +629,58 @@ def test_null_decimal_value(self, async_pandas_cursor, parquet_engine): ) result_set = future.result() assert result_set.fetchall() == [(None,)] + + @pytest.mark.parametrize( + "execute_kwargs", + [ + {}, + { + "block_size": 2048, + "cache_type": "none", + "max_workers": 3, + "auto_optimize_chunksize": False, + }, + ], + ) + def test_read_options(self, execute_kwargs): + """The cursor's read options reach the result set, and execute() overrides them. + + No AWS calls; the query and its result set are mocked. + """ + cursor_kwargs = {"block_size": 1024, "cache_type": "bytes", "auto_optimize_chunksize": True} + query_execution = MagicMock(state=AthenaQueryExecution.STATE_SUCCEEDED) + with ( + AsyncPandasCursor( + connection=MagicMock(), + converter=MagicMock(), + formatter=MagicMock(), + retry_config=RetryConfig(), + **cursor_kwargs, + ) as cursor, + patch.object(AsyncPandasCursor, "_execute", return_value="query_id"), + patch.object(AsyncPandasCursor, "_poll", return_value=query_execution), + patch("pyathena.pandas.async_cursor.AthenaPandasResultSet") as result_set_class, + ): + _, future = cursor.execute("SELECT 1", **execute_kwargs) + future.result() + kwargs = result_set_class.call_args.kwargs + expected = {**cursor_kwargs, **execute_kwargs} + assert {key: kwargs[key] for key in expected} == expected + + @pytest.mark.parametrize( + "async_pandas_cursor", + [{"cursor_kwargs": {"auto_optimize_chunksize": True}}], + indirect=True, + ) + def test_auto_optimize_chunksize(self, async_pandas_cursor, monkeypatch): + """auto_optimize_chunksize given to the cursor chunks the CSV result.""" + # Make the five-row result exceed the threshold and read it two rows at a time. + monkeypatch.setattr(AthenaPandasResultSet, "LARGE_FILE_THRESHOLD_BYTES", 0) + monkeypatch.setattr(AthenaPandasResultSet, "ESTIMATED_BYTES_PER_ROW", 1) + monkeypatch.setattr(AthenaPandasResultSet, "AUTO_CHUNK_THRESHOLD_MEDIUM", 0) + monkeypatch.setattr(AthenaPandasResultSet, "AUTO_CHUNK_SIZE_MEDIUM", 2) + _, future = async_pandas_cursor.execute( + "SELECT number FROM (VALUES (1), (2), (3), (4), (5)) AS t(number)" + ) + result_set = future.result() + assert [df["number"].tolist() for df in result_set.iter_chunks()] == [[1, 2], [3, 4], [5]] diff --git a/tests/pyathena/pandas/test_cursor.py b/tests/pyathena/pandas/test_cursor.py index db5c55c70..7c646ded6 100644 --- a/tests/pyathena/pandas/test_cursor.py +++ b/tests/pyathena/pandas/test_cursor.py @@ -7,7 +7,7 @@ from concurrent.futures import ThreadPoolExecutor from datetime import datetime from decimal import Decimal -from unittest.mock import PropertyMock, patch +from unittest.mock import MagicMock, PropertyMock, patch import numpy as np import pandas as pd @@ -16,9 +16,11 @@ from pyathena.error import DatabaseError, ProgrammingError from pyathena.filesystem.s3 import S3FileSystem +from pyathena.model import AthenaQueryExecution from pyathena.pandas.converter import DefaultPandasTypeConverter from pyathena.pandas.cursor import PandasCursor from pyathena.pandas.result_set import AthenaPandasResultSet, PandasDataFrameIterator +from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect from tests.pyathena.util import cached_file_systems @@ -1399,6 +1401,19 @@ def test_callback(query_id: str): assert callback_results[0] == pandas_cursor.query_id assert pandas_cursor.query_id is not None + @pytest.mark.parametrize("pandas_cursor", [{"cursor_kwargs": {"unload": True}}], indirect=True) + @pytest.mark.parametrize("option", ["filesystem", "storage_options"]) + def test_unload_with_filesystem_options(self, pandas_cursor, option): + """filesystem or storage_options given to execute() read the UNLOAD result.""" + fs_kwargs = {"connection": pandas_cursor.connection, "skip_instance_cache": True} + kwargs = ( + {"filesystem": S3FileSystem(**fs_kwargs)} + if option == "filesystem" + else {"storage_options": fs_kwargs} + ) + df = pandas_cursor.execute("SELECT * FROM one_row", **kwargs).as_pandas() + assert df.to_dict("records") == [{"number_of_rows": 1}] + def test_pandas_cursor_iter_chunks_with_chunksize(self, pandas_cursor): """Test PandasCursor iter_chunks method with chunksize set.""" cursor = pandas_cursor @@ -1601,3 +1616,44 @@ def test_fetch_all_rows(self, pandas_cursor): (1, datetime(2017, 1, 1, 12, 34, 56).time(), b"\x00\x01", {"a": 1}, "[1, 2]", None), (2, datetime(2017, 1, 1, 12, 34, 56).time(), b"\x00\x01", [1, "x"], "s", None), ] + + @pytest.mark.parametrize( + "execute_kwargs", + [ + {}, + { + "block_size": 2048, + "cache_type": "none", + "max_workers": 3, + "auto_optimize_chunksize": False, + }, + ], + ) + def test_read_options(self, execute_kwargs): + """The cursor's read options reach the result set, and execute() overrides them. + + No AWS calls; the query and its result set are mocked. + """ + cursor_kwargs = { + "block_size": 1024, + "cache_type": "bytes", + "max_workers": 2, + "auto_optimize_chunksize": True, + } + cursor = PandasCursor( + connection=MagicMock(), + converter=MagicMock(), + formatter=MagicMock(), + retry_config=RetryConfig(), + **cursor_kwargs, + ) + query_execution = MagicMock(state=AthenaQueryExecution.STATE_SUCCEEDED) + with ( + patch.object(PandasCursor, "_execute", return_value="query_id"), + patch.object(PandasCursor, "_poll", return_value=query_execution), + patch("pyathena.pandas.cursor.AthenaPandasResultSet") as result_set_class, + ): + cursor.execute("SELECT 1", **execute_kwargs) + kwargs = result_set_class.call_args.kwargs + expected = {**cursor_kwargs, **execute_kwargs} + assert {key: kwargs[key] for key in expected} == expected diff --git a/tests/pyathena/pandas/test_result_set.py b/tests/pyathena/pandas/test_result_set.py index 11b4f7e3f..84ee98922 100644 --- a/tests/pyathena/pandas/test_result_set.py +++ b/tests/pyathena/pandas/test_result_set.py @@ -6,11 +6,16 @@ # SPDX-License-Identifier: MIT import io +from unittest.mock import MagicMock, patch import pandas as pd import pytest -from pyathena.pandas.result_set import PandasDataFrameIterator, _no_trunc_date +from pyathena.pandas.result_set import ( + AthenaPandasResultSet, + PandasDataFrameIterator, + _no_trunc_date, +) class TestPandasDataFrameIterator: @@ -64,3 +69,47 @@ def test_as_pandas_single_dataframe(self): df_iter = PandasDataFrameIterator(df, _no_trunc_date) assert df_iter.as_pandas() is df + + +_FS = MagicMock(name="pyathena_fs") +_USER_FS = MagicMock(name="user_fs") + + +class TestAthenaPandasResultSet: + @pytest.mark.parametrize( + ("execute_kwargs", "path", "filesystem_kwargs"), + [ + ({}, "bucket/unload/", {"filesystem": _FS}), + ({"filesystem": _USER_FS}, "bucket/unload/", {"filesystem": _USER_FS}), + ({"filesystem": None}, "s3://bucket/unload/", {"filesystem": None}), + ({"storage_options": {"anon": True}}, "s3://bucket/unload/", {}), + ({"storage_options": None}, "s3://bucket/unload/", {}), + ], + ids=["default", "filesystem", "filesystem-none", "storage-options", "storage-options-none"], + ) + def test_read_parquet_filesystem(self, execute_kwargs, path, filesystem_kwargs): + """filesystem or storage_options given to execute() replace PyAthena's filesystem. + + No AWS calls; the manifest and pandas.read_parquet are mocked. + """ + result_set = AthenaPandasResultSet.__new__(AthenaPandasResultSet) # bypass __init__ + result_set._unload_location = None + result_set._engine = "pyarrow" + result_set._fs = _FS + result_set._kwargs = dict(execute_kwargs) + with ( + patch.object( + AthenaPandasResultSet, + "_read_data_manifest", + return_value=["s3://bucket/unload/0.parquet"], + ), + patch("pandas.read_parquet") as read_parquet, + ): + result_set._read_parquet("pyarrow") + assert read_parquet.call_args.args == (path,) + assert read_parquet.call_args.kwargs == { + "engine": "pyarrow", + "use_threads": True, + **execute_kwargs, + **filesystem_kwargs, + } diff --git a/tests/pyathena/polars/test_async_cursor.py b/tests/pyathena/polars/test_async_cursor.py index dde6c4318..ac4a761d9 100644 --- a/tests/pyathena/polars/test_async_cursor.py +++ b/tests/pyathena/polars/test_async_cursor.py @@ -4,6 +4,7 @@ import time from datetime import datetime from random import randint +from unittest.mock import MagicMock, patch import polars as pl import pytest @@ -12,6 +13,7 @@ from pyathena.model import AthenaQueryExecution from pyathena.polars.async_cursor import AsyncPolarsCursor from pyathena.result_set import AthenaResultSet +from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect @@ -323,3 +325,37 @@ def test_empty_result_unload(self, async_polars_cursor): df = future.result().as_polars() assert df.height == 0 assert df.width == 0 + + @pytest.mark.parametrize( + "execute_kwargs", + [{}, {"block_size": 2048, "cache_type": "none", "max_workers": 3, "chunksize": 20}], + ) + def test_read_options(self, execute_kwargs): + """The cursor's read options reach the result set, and execute() overrides them. + + No AWS calls; the query and its result set are mocked. + """ + cursor_kwargs = { + "block_size": 1024, + "cache_type": "bytes", + "max_workers": 2, + "chunksize": 10, + } + query_execution = MagicMock(state=AthenaQueryExecution.STATE_SUCCEEDED) + with ( + AsyncPolarsCursor( + connection=MagicMock(), + converter=MagicMock(), + formatter=MagicMock(), + retry_config=RetryConfig(), + **cursor_kwargs, + ) as cursor, + patch.object(AsyncPolarsCursor, "_execute", return_value="query_id"), + patch.object(AsyncPolarsCursor, "_poll", return_value=query_execution), + patch("pyathena.polars.async_cursor.AthenaPolarsResultSet") as result_set_class, + ): + _, future = cursor.execute("SELECT 1", **execute_kwargs) + future.result() + kwargs = result_set_class.call_args.kwargs + expected = {**cursor_kwargs, **execute_kwargs} + assert {key: kwargs[key] for key in expected} == expected diff --git a/tests/pyathena/polars/test_cursor.py b/tests/pyathena/polars/test_cursor.py index c8727f03d..7988a29c5 100644 --- a/tests/pyathena/polars/test_cursor.py +++ b/tests/pyathena/polars/test_cursor.py @@ -12,13 +12,16 @@ from concurrent.futures import ThreadPoolExecutor from datetime import datetime from decimal import Decimal +from unittest.mock import MagicMock, patch import polars as pl import pytest from pyathena.error import DatabaseError, ProgrammingError +from pyathena.model import AthenaQueryExecution 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 cached_file_systems @@ -115,6 +118,18 @@ def test_as_polars(self, polars_cursor): assert df.width == 1 assert df.to_dicts() == [{"number_of_rows": 1}] + def test_as_polars_with_read_kwargs(self, polars_cursor): + """Read arguments given to execute() replace the ones the result set chooses.""" + df = polars_cursor.execute( + "SELECT * FROM one_row", + schema_overrides={"number_of_rows": pl.Utf8}, + storage_options={ + "connection": polars_cursor.connection, + "skip_instance_cache": True, + }, + ).as_polars() + assert df.to_dicts() == [{"number_of_rows": "1"}] + @pytest.mark.parametrize( "polars_cursor", [{"cursor_kwargs": {"unload": False}}, {"cursor_kwargs": {"unload": True}}], @@ -763,3 +778,36 @@ def test_fetch_all_rows(self, polars_cursor): (2, datetime(2017, 1, 1, 12, 34, 56).time(), b"\x00\x01", [1, "x"], "s", None), ] assert polars_cursor.as_polars()["col_json"].dtype == pl.String + + @pytest.mark.parametrize( + "execute_kwargs", + [{}, {"block_size": 2048, "cache_type": "none", "max_workers": 3, "chunksize": 20}], + ) + def test_read_options(self, execute_kwargs): + """The cursor's read options reach the result set, and execute() overrides them. + + No AWS calls; the query and its result set are mocked. + """ + cursor_kwargs = { + "block_size": 1024, + "cache_type": "bytes", + "max_workers": 2, + "chunksize": 10, + } + cursor = PolarsCursor( + connection=MagicMock(), + converter=MagicMock(), + formatter=MagicMock(), + retry_config=RetryConfig(), + **cursor_kwargs, + ) + query_execution = MagicMock(state=AthenaQueryExecution.STATE_SUCCEEDED) + with ( + patch.object(PolarsCursor, "_execute", return_value="query_id"), + patch.object(PolarsCursor, "_poll", return_value=query_execution), + patch("pyathena.polars.cursor.AthenaPolarsResultSet") as result_set_class, + ): + cursor.execute("SELECT 1", **execute_kwargs) + kwargs = result_set_class.call_args.kwargs + expected = {**cursor_kwargs, **execute_kwargs} + assert {key: kwargs[key] for key in expected} == expected diff --git a/tests/pyathena/polars/test_result_set.py b/tests/pyathena/polars/test_result_set.py index e6828b54b..c16318ecf 100644 --- a/tests/pyathena/polars/test_result_set.py +++ b/tests/pyathena/polars/test_result_set.py @@ -81,6 +81,95 @@ def test_iter_parquet_chunks_raises_when_read_fails_partway(self, tmp_path): ): list(result_set._iter_parquet_chunks()) + @pytest.mark.parametrize("reader", ["_read_csv", "_iter_csv_chunks"]) + def test_csv_read_kwargs_replace_defaults(self, tmp_path, reader): + """Read arguments given to execute() replace the ones the result set chooses.""" + path = tmp_path / "result.csv" + path.write_text("1;x\n2;y\n") + result_set = _chunked_result_set() + result_set._kwargs = { + "separator": ";", + "has_header": False, + "schema_overrides": {"column_1": pl.Utf8}, + } + with ( + patch.object( + AthenaPolarsResultSet, + "output_location", + new_callable=PropertyMock, + return_value=str(path), + ), + patch.object( + AthenaPolarsResultSet, + "dtypes", + new_callable=PropertyMock, + return_value={"1;x": pl.Int64}, + ), + patch.object( + AthenaPolarsResultSet, + "_csv_storage_options", + new_callable=PropertyMock, + return_value={}, + ), + patch.object( + AthenaPolarsResultSet, + "_parquet_storage_options", + new_callable=PropertyMock, + return_value={}, + ), + patch.object(AthenaPolarsResultSet, "_is_csv_readable", return_value=True), + ): + result = getattr(result_set, reader)() + df = result if isinstance(result, pl.DataFrame) else pl.concat(list(result)) + assert df.to_dict(as_series=False) == {"column_1": ["1", "2"], "column_2": ["x", "y"]} + + @pytest.mark.parametrize( + ("reader", "function"), + [ + ("_read_csv", "read_csv"), + ("_iter_csv_chunks", "scan_csv"), + ("_read_parquet", "read_parquet"), + ("_iter_parquet_chunks", "scan_parquet"), + ("_read_parquet_schema", "scan_parquet"), + ], + ) + def test_storage_options_replace_defaults(self, reader, function): + """storage_options given to execute() replace PyAthena's without computing them.""" + result_set = _chunked_result_set() + result_set._unload_location = "s3://bucket/unload/" + result_set._kwargs = {"storage_options": {"anon": True}} + with ( + patch.object( + AthenaPolarsResultSet, + "output_location", + new_callable=PropertyMock, + return_value="s3://bucket/result.csv", + ), + patch.object( + AthenaPolarsResultSet, "dtypes", new_callable=PropertyMock, return_value={} + ), + patch.object( + AthenaPolarsResultSet, + "_csv_storage_options", + new_callable=PropertyMock, + side_effect=AssertionError("replaced storage options were computed"), + ), + patch.object( + AthenaPolarsResultSet, + "_parquet_storage_options", + new_callable=PropertyMock, + side_effect=AssertionError("replaced storage options were computed"), + ), + patch.object(AthenaPolarsResultSet, "_is_csv_readable", return_value=True), + patch.object(AthenaPolarsResultSet, "_prepare_parquet_location", return_value=True), + patch(f"polars.{function}") as read, + patch("pyathena.polars.result_set.to_column_info"), + ): + result = getattr(result_set, reader)() + if not isinstance(result, (pl.DataFrame, tuple)): + list(result) + assert read.call_args.kwargs["storage_options"] == {"anon": True} + class TestPolarsDataFrameIterator: @pytest.mark.parametrize( diff --git a/tests/pyathena/s3fs/test_async_cursor.py b/tests/pyathena/s3fs/test_async_cursor.py index ffc96dfdf..ba28c93b6 100644 --- a/tests/pyathena/s3fs/test_async_cursor.py +++ b/tests/pyathena/s3fs/test_async_cursor.py @@ -10,13 +10,16 @@ import time from datetime import datetime from decimal import Decimal +from unittest.mock import MagicMock, patch import pytest from pyathena.error import ProgrammingError from pyathena.model import AthenaQueryExecution from pyathena.s3fs.async_cursor import AsyncS3FSCursor +from pyathena.s3fs.reader import AthenaCSVReader, DefaultCSVReader from pyathena.s3fs.result_set import AthenaS3FSResultSet +from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect @@ -179,3 +182,29 @@ def test_empty_result(self, async_s3fs_cursor): assert result_set.fetchmany() == [] assert result_set.fetchmany(10) == [] assert result_set.fetchall() == [] + + @pytest.mark.parametrize("execute_kwargs", [{}, {"csv_reader": AthenaCSVReader}]) + def test_read_options(self, execute_kwargs): + """The cursor's read options reach the result set, and execute() overrides them. + + No AWS calls; the query and its result set are mocked. + """ + cursor_kwargs = {"csv_reader": DefaultCSVReader} + query_execution = MagicMock(state=AthenaQueryExecution.STATE_SUCCEEDED) + with ( + AsyncS3FSCursor( + connection=MagicMock(), + converter=MagicMock(), + formatter=MagicMock(), + retry_config=RetryConfig(), + **cursor_kwargs, + ) as cursor, + patch.object(AsyncS3FSCursor, "_execute", return_value="query_id"), + patch.object(AsyncS3FSCursor, "_poll", return_value=query_execution), + patch("pyathena.s3fs.async_cursor.AthenaS3FSResultSet") as result_set_class, + ): + _, future = cursor.execute("SELECT 1", **execute_kwargs) + future.result() + kwargs = result_set_class.call_args.kwargs + expected = {**cursor_kwargs, **execute_kwargs} + assert {key: kwargs[key] for key in expected} == expected diff --git a/tests/pyathena/s3fs/test_cursor.py b/tests/pyathena/s3fs/test_cursor.py index 1d12d7041..4ebaa4130 100644 --- a/tests/pyathena/s3fs/test_cursor.py +++ b/tests/pyathena/s3fs/test_cursor.py @@ -5,15 +5,18 @@ from concurrent.futures import ThreadPoolExecutor from datetime import datetime from decimal import Decimal +from unittest.mock import MagicMock, patch import pytest from pyathena.converter import _to_default from pyathena.error import DatabaseError, ProgrammingError +from pyathena.model import AthenaQueryExecution from pyathena.s3fs.converter import DefaultS3FSTypeConverter from pyathena.s3fs.cursor import S3FSCursor from pyathena.s3fs.reader import AthenaCSVReader, DefaultCSVReader from pyathena.s3fs.result_set import AthenaS3FSResultSet +from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect from tests.pyathena.util import cached_file_systems @@ -559,3 +562,28 @@ def test_fetch_all_rows(self, s3fs_cursor): def test_fetch_all_rows_custom_converter(self, s3fs_cursor): s3fs_cursor.execute("SELECT 1 AS col, json_parse('{\"a\": 1}') AS col_json") assert s3fs_cursor.fetchall() == [(1, '{"a":1}')] + + @pytest.mark.parametrize("execute_kwargs", [{}, {"csv_reader": AthenaCSVReader}]) + def test_read_options(self, execute_kwargs): + """The cursor's read options reach the result set, and execute() overrides them. + + No AWS calls; the query and its result set are mocked. + """ + cursor_kwargs = {"csv_reader": DefaultCSVReader} + query_execution = MagicMock(state=AthenaQueryExecution.STATE_SUCCEEDED) + cursor = S3FSCursor( + connection=MagicMock(), + converter=MagicMock(), + formatter=MagicMock(), + retry_config=RetryConfig(), + **cursor_kwargs, + ) + with ( + patch.object(S3FSCursor, "_execute", return_value="query_id"), + patch.object(S3FSCursor, "_poll", return_value=query_execution), + patch("pyathena.s3fs.cursor.AthenaS3FSResultSet") as result_set_class, + ): + cursor.execute("SELECT 1", **execute_kwargs) + kwargs = result_set_class.call_args.kwargs + expected = {**cursor_kwargs, **execute_kwargs} + assert {key: kwargs[key] for key in expected} == expected