diff --git a/docs/aio.md b/docs/aio.md index 060bc190..1682d337 100644 --- a/docs/aio.md +++ b/docs/aio.md @@ -191,6 +191,7 @@ df = cursor.as_pandas() # In-memory conversion, no await needed The `as_pandas()`, `as_arrow()`, and `as_polars()` convenience methods operate on already-loaded data and remain synchronous. +With a chunk size chosen by `auto_optimize_chunksize`, `as_pandas()` reads every remaining chunk. See each cursor's documentation page for detailed usage examples. diff --git a/docs/pandas.md b/docs/pandas.md index 97bbf570..5b924b16 100644 --- a/docs/pandas.md +++ b/docs/pandas.md @@ -475,6 +475,10 @@ for chunk in cursor.iter_chunks(): process_chunk(chunk) ``` +Without an explicit `chunksize`, `as_pandas()` returns a single DataFrame even when a chunk size was chosen automatically. +It reads every chunk and joins them, so the whole result is held in memory. +Use `iter_chunks()`, `fetchone()`, or `fetchmany()` to read a large result one chunk at a time. + **Priority of chunksize settings:** 1. **Explicit chunksize** (highest priority): Always respected @@ -551,6 +555,8 @@ Common performance options: - `dtype`: Explicit column data types - `parse_dates`: Columns to parse as dates +With `engine="pyarrow"`, tab-separated `.txt` results from DDL statements such as `SHOW TABLES`, `SHOW COLUMNS`, and `DESCRIBE` use the C engine to preserve leading zeros, exponent notation, and padding in string values. + ### Unload options PandasCursor also supports the unload option, as does {ref}`arrow-cursor`. diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index 60fd1916..d153851e 100644 --- a/docs/sqlalchemy.md +++ b/docs/sqlalchemy.md @@ -719,6 +719,10 @@ That includes top-level columns, fields of a STRUCT, STRUCT values inside MAP, a Integer fields, and integer MAP keys and values, use `INT` in that DDL. `CAST` and other SQL expressions keep `ROW(...)`, `MAP(...)`, and `ARRAY(...)`, and spell integers as `INTEGER`. +Reflected STRUCT and ROW columns use `AthenaStruct` with their field names and types. +A field type that the dialect does not recognize is reflected as `NullType` with a warning. +Selecting such a column returns the value from the cursor, as described under Data format support below; SQLAlchemy does not convert it to the reflected field types. + #### Querying STRUCT data PyAthena automatically converts STRUCT data between different formats: @@ -868,6 +872,10 @@ CREATE TABLE products ( `CREATE TABLE` renders integer MAP keys and values as `INT`. `CAST` still spells those integers as `INTEGER`. +Reflected MAP columns use `AthenaMap` with their key and value types. +A key or value type that the dialect does not recognize is reflected as `NullType` with a warning. +Selecting such a column returns the value from the cursor, as described under Data format support below; SQLAlchemy does not convert it to the reflected key and value types. + #### Querying MAP data PyAthena automatically converts MAP data between different formats: diff --git a/docs/usage.md b/docs/usage.md index d71e0004..322fc73f 100644 --- a/docs/usage.md +++ b/docs/usage.md @@ -221,6 +221,8 @@ See the [Athena documentation](https://docs.aws.amazon.com/athena/latest/ug/reus You can attempt to re-use the results from a previously executed query to help save time and money in the cases where your underlying data isn't changing. Set the `cache_size` or `cache_expiration_time` parameter of `cursor.execute()` to a number larger than 0 to enable caching. +`cache_size` is the number of the most recent query executions in the work group to search, including executions by other clients of the same work group. +In a busy work group, a previous execution may no longer be among them. ```python from pyathena import connect @@ -256,6 +258,7 @@ cursor.execute("SELECT * FROM one_row", cache_size=100, cache_expiration_time=36 Results will only be re-used if the query strings match *exactly*, and the query was a DML statement (the assumption being that you always want to re-run queries like `CREATE TABLE` and `DROP TABLE`). +The cache is not used for a `qmark` query with parameters. The S3 staging directory is not checked, so it's possible that the location of the results is not in your provided `s3_staging_dir`. diff --git a/pyathena/aio/common.py b/pyathena/aio/common.py index e490a544..9faac348 100644 --- a/pyathena/aio/common.py +++ b/pyathena/aio/common.py @@ -4,7 +4,7 @@ import logging import sys from datetime import datetime, timedelta, timezone -from typing import Any, cast +from typing import Any, NoReturn, cast from pyathena.aio.util import async_retry_api_call from pyathena.common import BaseCursor, CursorIterator @@ -60,12 +60,16 @@ async def _execute( # type: ignore[override] result_reuse_minutes=options.result_reuse_minutes, execution_parameters=execution_parameters, ) - query_id = await self._find_previous_query_id( - query, - options.work_group, - cache_size=options.cache_size, - cache_expiration_time=options.cache_expiration_time, - ) + query_id = None + # Athena does not return the ExecutionParameters of earlier executions, + # so the cache cannot tell which parameters an execution ran with (#941). + if not request.get("ExecutionParameters"): + query_id = await self._find_previous_query_id( + query, + options.work_group, + cache_size=options.cache_size, + cache_expiration_time=options.cache_expiration_time, + ) if query_id is None: try: response = await async_retry_api_call( @@ -376,6 +380,7 @@ class WithAsyncFetch(AioBaseCursor, CursorIterator, WithResultSet): ``rownumber``, ``rowcount``), lifecycle methods (``close``, ``executemany``, ``cancel``), default sync fetch (for cursors whose result sets load all data eagerly in ``__init__``), and the async iteration protocol. + Synchronous iteration raises ``TypeError``. Subclasses override ``execute()`` and optionally ``__init__`` and format-specific helpers. @@ -504,6 +509,14 @@ def fetchall( result_set = cast(AthenaResultSet, self.result_set) return result_set.fetchall() + def __iter__(self) -> NoReturn: + """Reject synchronous iteration; use ``async for`` instead. + + Raises: + TypeError: Always, because the fetch methods are coroutines. + """ + raise TypeError(f"'{type(self).__name__}' object is not iterable; use 'async for' instead.") + def __aiter__(self): return self diff --git a/pyathena/aio/result_set.py b/pyathena/aio/result_set.py index 47fe6508..624ab2b5 100644 --- a/pyathena/aio/result_set.py +++ b/pyathena/aio/result_set.py @@ -4,6 +4,7 @@ from typing import ( TYPE_CHECKING, Any, + NoReturn, cast, ) @@ -25,7 +26,8 @@ class AthenaAioResultSet(AthenaResultSet): Skips the synchronous ``_pre_fetch`` by passing ``_pre_fetch=False`` to the parent ``__init__`` and provides an ``async create()`` classmethod - factory instead. + factory instead. Synchronous iteration raises ``TypeError``; use + ``async for`` instead. """ def __init__( @@ -191,6 +193,14 @@ async def fetchall( # type: ignore[override] break return rows + def __iter__(self) -> NoReturn: + """Reject synchronous iteration; use ``async for`` instead. + + Raises: + TypeError: Always, because the fetch methods are coroutines. + """ + raise TypeError(f"'{type(self).__name__}' object is not iterable; use 'async for' instead.") + def __aiter__(self): return self diff --git a/pyathena/arrow/result_set.py b/pyathena/arrow/result_set.py index 2ff47683..d5783820 100644 --- a/pyathena/arrow/result_set.py +++ b/pyathena/arrow/result_set.py @@ -119,6 +119,9 @@ def __init__( import pyarrow as pa self._table = pa.Table.from_pydict({}) + # The fetch methods convert only the values read from a result file. + # GetQueryResults values are already converted. + self._convert_rows = bool(self.output_location) self._batches = iter(self._table.to_batches(arraysize)) def __s3_file_system(self): @@ -215,11 +218,15 @@ def _fetch(self) -> None: return else: dict_rows = rows.to_pydict() - column_names = dict_rows.keys() - processed_rows = [ - tuple(self.converters[k](v) for k, v in zip(column_names, row, strict=False)) - for row in zip(*dict_rows.values(), strict=False) - ] + if self._convert_rows: + converters = self.converters + column_names = dict_rows.keys() + processed_rows = [ + tuple(converters[k](v) for k, v in zip(column_names, row, strict=False)) + for row in zip(*dict_rows.values(), strict=False) + ] + else: + processed_rows = list(zip(*dict_rows.values(), strict=False)) self._rows.extend(processed_rows) def fetchone( @@ -297,9 +304,13 @@ def _read_csv(self) -> Table: parse_opts = csv.ParseOptions( delimiter=",", quote_char='"', - ignore_empty_lines=not binary_columns, + # Athena writes a single-column row with a NULL value as an empty line. + ignore_empty_lines=False, double_quote=True, escape_char=False, + # A quoted value can contain a newline, so the reader must not split + # blocks inside quotes. + newlines_in_values=True, ) else: return pa.Table.from_pydict({}) diff --git a/pyathena/arrow/util.py b/pyathena/arrow/util.py index 9497858d..256d79f1 100644 --- a/pyathena/arrow/util.py +++ b/pyathena/arrow/util.py @@ -94,7 +94,7 @@ def get_athena_type(type_: DataType) -> tuple[str, int, int]: return "date", 0, 0 if type_.id == types.Type_TIMESTAMP: # 18 return "timestamp", 3, 0 - if type_.id in [types.Type_DECIMAL128, types.Decimal256Type]: # 23, 24 + if type_.id in [types.Type_DECIMAL128, types.Type_DECIMAL256]: # 23, 24 type_ = cast(types.Decimal128Type, type_) return "decimal", type_.precision, type_.scale if type_.id in [ diff --git a/pyathena/common.py b/pyathena/common.py index 3a517694..187bd48b 100644 --- a/pyathena/common.py +++ b/pyathena/common.py @@ -750,12 +750,16 @@ def _execute( result_reuse_minutes=options.result_reuse_minutes, execution_parameters=execution_parameters, ) - query_id = self._find_previous_query_id( - query, - options.work_group, - cache_size=options.cache_size, - cache_expiration_time=options.cache_expiration_time, - ) + query_id = None + # Athena does not return the ExecutionParameters of earlier executions, + # so the cache cannot tell which parameters an execution ran with (#941). + if not request.get("ExecutionParameters"): + query_id = self._find_previous_query_id( + query, + options.work_group, + cache_size=options.cache_size, + cache_expiration_time=options.cache_expiration_time, + ) if query_id is None: try: query_id = retry_api_call( diff --git a/pyathena/converter.py b/pyathena/converter.py index 8ae5cfdf..d37d3288 100644 --- a/pyathena/converter.py +++ b/pyathena/converter.py @@ -7,7 +7,7 @@ from abc import ABCMeta, abstractmethod from collections.abc import Callable from copy import deepcopy -from datetime import date, datetime, time +from datetime import date, datetime, time, timedelta, timezone from decimal import Decimal from typing import Any, ClassVar @@ -40,11 +40,33 @@ def _to_datetime(varchar_value: str | None) -> datetime | None: return datetime.strptime(varchar_value, "%Y-%m-%d %H:%M:%S.%f") +_UTC_OFFSET_PATTERN: re.Pattern[str] = re.compile(r"([+-])(\d{2}):(\d{2})") + + +def _parse_utc_offset(value: str) -> timezone | None: + """Parse a ``+HH:MM`` or ``-HH:MM`` UTC offset. + + Args: + value: The text to parse. + + Returns: + The fixed-offset time zone, or None if the text is not an offset. + """ + match = _UTC_OFFSET_PATTERN.fullmatch(value) + if not match: + return None + sign, hours, minutes = match.groups() + offset = timedelta(hours=int(hours), minutes=int(minutes)) + return timezone(-offset if sign == "-" else offset) + + def _to_datetime_with_tz(varchar_value: str | None) -> datetime | None: if varchar_value is None: return None datetime_, _, tz = varchar_value.rpartition(" ") - return datetime.strptime(datetime_, "%Y-%m-%d %H:%M:%S.%f").replace(tzinfo=gettz(tz)) + return datetime.strptime(datetime_, "%Y-%m-%d %H:%M:%S.%f").replace( + tzinfo=_parse_utc_offset(tz) or gettz(tz) + ) def _to_time(varchar_value: str | None) -> time | None: diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index ab941678..1c4d4f0c 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -271,6 +271,11 @@ def _head_bucket(self, bucket, refresh: bool = False) -> S3Object | None: Bucket=bucket, ) except FileNotFoundError: + self.dircache.pop(bucket, None) + # Evict the cached bucket listing only if it still lists the bucket. + buckets = self.dircache.get("") + if buckets and any(b.name == bucket for b in buckets): + self.dircache.pop("", None) return None file = S3Object( init={ @@ -308,6 +313,7 @@ def _head_object( **request, ) except FileNotFoundError: + self.dircache.pop(path, None) return None if self.version_aware and not version_id: # Pin the version of the object so that subsequent reads see @@ -360,13 +366,34 @@ def _ls_dirs( max_keys: int | None = None, refresh: bool = False, ) -> list[S3Object]: + """List the objects and common prefixes under a path. + + A complete, non-empty listing of the path is cached under + ``(path, delimiter)``, and an empty one evicts it. + ``invalidate_cache`` drops it when the path or a path under it is + invalidated. + + Args: + path: The bucket or directory path to list. + prefix: Key prefix to filter by, relative to the path. A prefixed + listing is neither read from nor written to the cache. + delimiter: Delimiter to group keys by; ``""`` lists recursively. + next_token: Continuation token to start listing from. A listing + that starts from a token is neither read from nor written to + the cache. + max_keys: Maximum number of keys per ListObjectsV2 request. + refresh: If True, bypass the cache and list from S3. + + Returns: + The listed directories and files. + """ bucket, key, version_id = self.parse_path(path) + use_cache = not prefix and not next_token if key: prefix = f"{key}/{prefix if prefix else ''}" - # Create a cache key that includes the delimiter cache_key = (path, delimiter) - if cache_key in self.dircache and not refresh: + if use_cache and cache_key in self.dircache and not refresh: return cast(list[S3Object], self.dircache[cache_key]) files: list[S3Object] = [] @@ -400,8 +427,11 @@ def _ls_dirs( next_token = response.get("NextContinuationToken") if not next_token: break - if files: - self.dircache[cache_key] = files + if use_cache: + if files: + self.dircache[cache_key] = files + else: + self.dircache.pop(cache_key, None) return files def ls( @@ -629,13 +659,15 @@ def _find( raise ValueError("Cannot traverse all files in S3.") bucket, key, _ = self.parse_path(path) prefix = kwargs.pop("prefix", "") + # Keep refresh in kwargs so that the recursive calls also refresh. + refresh = kwargs.get("refresh", False) # When maxdepth is specified, use a recursive approach with delimiter if maxdepth is not None: result: list[S3Object] = [] # List files and directories at current level - current_items = self._ls_dirs(path, prefix=prefix, delimiter="/") + current_items = self._ls_dirs(path, prefix=prefix, delimiter="/", refresh=refresh) for item in current_items: if item.type == S3ObjectType.S3_OBJECT_TYPE_FILE: @@ -657,16 +689,17 @@ def _find( return result # For unlimited depth, use the original approach (get all files at once) - files = self._ls_dirs(path, prefix=prefix, delimiter="") + files = self._ls_dirs(path, prefix=prefix, delimiter="", refresh=refresh) if not files and key: try: - files = [self.info(path)] + files = [self.info(path, refresh=refresh)] except FileNotFoundError: files = [] # If withdirs is True, we need to derive directories from file paths if withdirs: - files.extend(self._extract_parent_directories(files, bucket, key)) + # Build a new list; files may be the cached listing. + files = files + self._extract_parent_directories(files, bucket, key) # Filter directories if withdirs is False (default) if withdirs is False or withdirs is None: @@ -693,7 +726,11 @@ def find( maxdepth: Maximum depth to recurse (None for unlimited). withdirs: Whether to include directories in results (None = default behavior). detail: If True, return dict of {path: S3Object}; if False, return list of paths. - **kwargs: Additional arguments. + **kwargs: Additional arguments including: + prefix: Key prefix, relative to the path, to filter the listed keys + by. Without maxdepth, if nothing is listed and the path itself is + an object, that object is returned regardless of the prefix. + refresh: If True, bypass the cache and list from S3. Returns: Dictionary mapping paths to S3Objects (if detail=True) or @@ -717,7 +754,8 @@ def exists(self, path: str, **kwargs) -> bool: Args: path: S3 path to check (e.g., "s3://bucket" or "s3://bucket/key"). - **kwargs: Additional arguments (unused). + **kwargs: Additional arguments including: + refresh: If True, bypass the cache and query S3. Returns: True if the path exists, False otherwise. @@ -727,6 +765,7 @@ def exists(self, path: str, **kwargs) -> bool: >>> fs.exists("s3://my-bucket/file.txt") >>> fs.exists("s3://my-bucket/") """ + refresh = kwargs.pop("refresh", False) path = self._strip_protocol(path) if path in ["", "/"]: # The root always exists. @@ -734,22 +773,22 @@ def exists(self, path: str, **kwargs) -> bool: bucket, key, _ = self.parse_path(path) if key: try: - if self._ls_from_cache(path): + if not refresh and self._ls_from_cache(path): return True - info = self.info(path) + info = self.info(path, refresh=refresh) return bool(info) except FileNotFoundError: return False - elif self.dircache.get(bucket, False): - return True - else: + if not refresh: + if self.dircache.get(bucket, False): + return True try: if self._ls_from_cache(bucket): return True except FileNotFoundError: pass - file = self._head_bucket(bucket) - return bool(file) + file = self._head_bucket(bucket, refresh=refresh) + return bool(file) def rm_file(self, path: str, **kwargs) -> None: bucket, key, version_id = self.parse_path(path) @@ -1104,12 +1143,7 @@ def _copy_object_with_multipart_upload( if version_id1: copy_source.update({"VersionId": version_id1}) - ranges = S3File._get_ranges( - 0, - size1, - max_workers, - block_size, - ) + ranges = self._get_copy_ranges(size1, block_size) multipart_upload = self._create_multipart_upload( bucket=bucket2, key=key2, @@ -1135,6 +1169,37 @@ def _copy_object_with_multipart_upload( futures=futures, ) + def _get_copy_ranges(self, size: int, block_size: int) -> list[tuple[int, int]]: + """Split an object into the source ranges of a multipart copy. + + The object is split into ranges of ``block_size`` bytes, whatever the + number of workers. A last range shorter than + ``MULTIPART_UPLOAD_MIN_PART_SIZE`` is merged into the previous one, + which is split in half if the result exceeds + ``MULTIPART_UPLOAD_MAX_PART_SIZE``. Every range is then within the + S3 part size limits, including the last one unless the whole object + is smaller than the minimum part size, so that more parts can follow + the copied ones, as in an append. + + Args: + size: The size of the source object in bytes. + block_size: The size in bytes to split the object by, between + ``MULTIPART_UPLOAD_MIN_PART_SIZE`` and + ``MULTIPART_UPLOAD_MAX_PART_SIZE``. The range that a short + last range is merged into can be longer, up to + ``MULTIPART_UPLOAD_MAX_PART_SIZE``. + + Returns: + The ``(start, end)`` byte ranges, with an exclusive end, that + cover the whole object in order. + """ + starts = list(range(0, size, block_size)) + if len(starts) > 1 and size - starts[-1] < self.MULTIPART_UPLOAD_MIN_PART_SIZE: + starts.pop() + if size - starts[-1] > self.MULTIPART_UPLOAD_MAX_PART_SIZE: + starts.append(starts[-1] + (size - starts[-1]) // 2) + return list(zip(starts, [*starts[1:], size], strict=True)) + def pipe_file( self, path: str, value: bytes | bytearray | memoryview, mode: str = "overwrite", **kwargs ) -> None: @@ -1737,6 +1802,9 @@ def invalidate_cache(self, path: str | None = None) -> None: path = self._strip_protocol(path) while path: self.dircache.pop(path, None) + # _ls_dirs caches listings under (path, delimiter). + for delimiter in ("/", ""): + self.dircache.pop((path, delimiter), None) path = self._parent(path) def _ls_from_cache(self, path: str) -> list[S3Object] | S3Object | None: @@ -2004,11 +2072,15 @@ def __init__( self.s3_additional_kwargs.update({"IfMatch": etag}) self._details = info elif "a" in mode and self.fs.exists(path): - self.append_block = True info = self.fs.info(self.path, version_id=self.version_id) loc = info.get("size", 0) if loc < self.fs.MULTIPART_UPLOAD_MIN_PART_SIZE: + # Too small to be a part of a multipart upload: rewrite it + # from the buffer. self.write(self.fs.cat(self.path)) + else: + # Copied with UploadPartCopy as the leading part(s). + self.append_block = True self.loc = loc self.s3_additional_kwargs.update(info.to_api_repr()) self._details = info @@ -2023,8 +2095,10 @@ def close(self) -> None: self._executor.shutdown() def _initiate_upload(self) -> None: - if self.tell() < self.blocksize: + if not self.append_block and self.tell() < self.blocksize: # Files smaller than block size in size cannot be multipart uploaded. + # An append to an object copied with UploadPartCopy always uses + # a multipart upload, whatever the block size. return self.multipart_upload = self.fs._create_multipart_upload( @@ -2033,14 +2107,12 @@ def _initiate_upload(self) -> None: **self.s3_additional_kwargs, ) if self.append_block: - if self.tell() > S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE: + if self.tell() > self.fs.MULTIPART_UPLOAD_MAX_PART_SIZE: info = self.fs.info(self.path, version_id=self.version_id) - ranges = self._get_ranges( - 0, + ranges = self.fs._get_copy_ranges( # Set copy source file byte size info.get("size", 0), - self.max_workers, - S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE, + self.fs.MULTIPART_UPLOAD_MAX_PART_SIZE, ) for i, range_ in enumerate(ranges): self.multipart_upload_parts.append( @@ -2074,7 +2146,7 @@ def _upload_chunk(self, final: bool = False) -> bool: # can still read the bytes; resetting it there would upload an empty # object for small files. Mid-stream chunks (final=False) return True so # fsspec clears the already-uploaded buffer between parts. - if self.tell() < self.blocksize: + if not self.append_block and self.tell() < self.blocksize: # Files smaller than block size in size cannot be multipart uploaded. if self.autocommit and final: self.commit() @@ -2085,9 +2157,12 @@ def _upload_chunk(self, final: bool = False) -> bool: part_number = len(self.multipart_upload_parts) self.buffer.seek(0) - while data := self.buffer.read(self.blocksize): - # The last part of a multipart request should be adjusted - # to be larger than the minimum part size. + data = self.buffer.read(self.blocksize) + while data: + # Only the last part of a multipart upload may be smaller than the + # minimum part size, and more data may follow a mid-stream chunk. + # A single write() can leave several blocks in the buffer, so look + # ahead one block and merge a short last block into this one. next_data = self.buffer.read(self.blocksize) next_data_size = len(next_data) if 0 < next_data_size < self.fs.MULTIPART_UPLOAD_MIN_PART_SIZE: @@ -2098,10 +2173,9 @@ def _upload_chunk(self, final: bool = False) -> bool: else: split_size = upload_data_size // 2 uploads = [upload_data[:split_size], upload_data[split_size:]] + next_data = b"" else: uploads = [data] - if next_data: - uploads.append(next_data) for upload in uploads: part_number += 1 @@ -2116,8 +2190,7 @@ def _upload_chunk(self, final: bool = False) -> bool: ) ) - if not next_data: - break + data = next_data if self.autocommit and final: self.commit() @@ -2163,12 +2236,19 @@ def discard(self) -> None: if self.multipart_upload: for f in self.multipart_upload_parts: f.cancel() + # s3_additional_kwargs also holds object parameters (e.g., the + # existing object's metadata in append mode) that + # AbortMultipartUpload rejects. self.fs._call( "abort_multipart_upload", Bucket=self.bucket, Key=self.key, UploadId=self.multipart_upload.upload_id, - **self.s3_additional_kwargs, + **{ + k: v + for k, v in self.s3_additional_kwargs.items() + if k in ("RequestPayer", "ExpectedBucketOwner") + }, ) self.multipart_upload = None diff --git a/pyathena/filesystem/s3_async.py b/pyathena/filesystem/s3_async.py index 834ffd92..c34975b1 100644 --- a/pyathena/filesystem/s3_async.py +++ b/pyathena/filesystem/s3_async.py @@ -250,12 +250,7 @@ async def _copy_object_with_multipart_upload( if version_id1: copy_source["VersionId"] = version_id1 - ranges = S3File._get_ranges( - 0, - size1, - self._sync_fs.max_workers, - block_size, - ) + ranges = self._sync_fs._get_copy_ranges(size1, block_size) multipart_upload = await asyncio.to_thread( self._sync_fs._create_multipart_upload, bucket=bucket2, diff --git a/pyathena/options.py b/pyathena/options.py index bb058201..bf6e8e89 100644 --- a/pyathena/options.py +++ b/pyathena/options.py @@ -38,6 +38,7 @@ class ExecuteOptions: caching. 0 (default) disables the cache lookup, unless ``cache_expiration_time`` is set to a positive value, in which case all queries within the expiration window are scanned. + A ``qmark`` query with parameters is never looked up. cache_expiration_time: Maximum age in seconds of a cached query result to consider for reuse. 0 (default) means no age limit. result_reuse_enable: Enable Athena server-side result reuse for this diff --git a/pyathena/pandas/result_set.py b/pyathena/pandas/result_set.py index 4eba6d60..dc2ffbf4 100644 --- a/pyathena/pandas/result_set.py +++ b/pyathena/pandas/result_set.py @@ -166,7 +166,15 @@ def get_chunk(self, size: int | None = None) -> DataFrame: raise def as_pandas(self) -> DataFrame: - """Collect all chunks into a single DataFrame. + """Collect all remaining chunks into a single DataFrame. + + The chunks keep their index, so the result has the row numbers or the + ``index_col`` values of the CSV file. Categorical columns and a categorical + index stay categorical. Categories given in the dtype keep their order; + when the chunks inferred different categories, they are inferred again from + all chunks in sorted order. A whole-file read of a large file can order + inferred categories differently, because pandas joins its internal parser + blocks in the order they were read. Returns: Single pandas DataFrame containing all data. @@ -178,7 +186,20 @@ def as_pandas(self) -> DataFrame: return pd.DataFrame() if len(dfs) == 1: return dfs[0] - return pd.concat(dfs, ignore_index=True) + df = pd.concat(dfs) + # Each chunk infers its own categories, and concat turns categorical columns + # and indexes whose categories differ into object or string ones. + for column, dtype in dfs[0].dtypes.items(): + if isinstance(dtype, pd.CategoricalDtype) and not isinstance( + df[column].dtype, pd.CategoricalDtype + ): + df[column] = df[column].astype(pd.CategoricalDtype(ordered=dtype.ordered)) + index_dtype = dfs[0].index.dtype + if isinstance(index_dtype, pd.CategoricalDtype) and not isinstance( + df.index.dtype, pd.CategoricalDtype + ): + df.index = df.index.astype(pd.CategoricalDtype(ordered=index_dtype.ordered)) + return df class AthenaPandasResultSet(AthenaResultSet): @@ -358,6 +379,10 @@ def _get_csv_engine( # Use PyArrow only when explicitly requested and all compatibility # checks pass; otherwise fall through to the C engine default. if self._engine == "pyarrow": + # Header-less DDL results lose numeric-looking strings and their padding + # when PyArrow infers types before applying the string dtypes. + if self.output_location and self.output_location.endswith(".txt"): + return "c" effective_chunksize = chunksize if chunksize is not None else self._chunksize is_compatible = ( effective_chunksize is None @@ -828,20 +853,20 @@ def _as_pandas_from_api(self, converter: Converter | None = None) -> DataFrame: def as_pandas(self) -> PandasDataFrameIterator | DataFrame: if self._chunksize is None: - return next(self._df_iter) + return self._df_iter.as_pandas() return self._df_iter def iter_chunks(self) -> PandasDataFrameIterator: """Iterate over result chunks as pandas DataFrames. This method provides an iterator interface for processing large result sets. - When chunksize is specified, it yields DataFrames in chunks for memory-efficient - processing. When chunksize is not specified, it yields the entire result as a - single DataFrame. + When chunksize is specified, or ``auto_optimize_chunksize`` chose a chunk size + for a large CSV result, it yields DataFrames in chunks for memory-efficient + processing. Otherwise, it yields the entire result as a single DataFrame. Returns: PandasDataFrameIterator that yields pandas DataFrames for each chunk - of rows, or the entire DataFrame if chunksize was not specified. + of rows, or the entire DataFrame if the result was not read in chunks. Example: >>> # With chunking for large datasets diff --git a/pyathena/pandas/util.py b/pyathena/pandas/util.py index b2b96612..9337ce32 100644 --- a/pyathena/pandas/util.py +++ b/pyathena/pandas/util.py @@ -261,13 +261,17 @@ def to_sql( ).Bucket(bucket_name) cursor = conn.cursor() + # Athena stores identifiers in lowercase and information_schema reports them + # that way, so compare lowercase literals. + schema_literal = schema.lower().replace("'", "''") + name_literal = name.lower().replace("'", "''") table = cursor.execute( textwrap.dedent( f""" SELECT table_name FROM information_schema.tables - WHERE table_schema = '{schema}' - AND table_name = '{name}' + WHERE table_schema = '{schema_literal}' + AND table_name = '{name_literal}' """ ) ).fetchall() diff --git a/pyathena/polars/result_set.py b/pyathena/polars/result_set.py index 988e2d52..7ff84bf2 100644 --- a/pyathena/polars/result_set.py +++ b/pyathena/polars/result_set.py @@ -110,8 +110,10 @@ def close(self) -> None: """Close the iterator and release resources.""" from types import GeneratorType - if isinstance(self._reader, GeneratorType): - self._reader.close() + reader = self._reader + self._reader = iter(()) + if isinstance(reader, GeneratorType): + reader.close() def iterrows(self) -> Iterator[tuple[int, dict[str, Any]]]: """Iterate over rows as (index, row_dict) tuples. diff --git a/pyathena/spark/async_cursor.py b/pyathena/spark/async_cursor.py index 1a476f20..17aa2a3e 100644 --- a/pyathena/spark/async_cursor.py +++ b/pyathena/spark/async_cursor.py @@ -75,6 +75,27 @@ def __init__( max_workers: int = (cpu_count() or 1) * 5, **kwargs, ): + """Initialize the cursor and start or attach to a Spark session. + + Args: + session_id: ID of an existing session to use. If omitted, a new + session is started. + description: Description of a new session. + engine_configuration: Engine configuration of a new session. + notebook_version: Notebook version of a new session. + session_idle_timeout_minutes: Idle timeout of a new session in minutes. + max_workers: Maximum number of threads for asynchronous operations. + **kwargs: Arguments passed to ``SparkBaseCursor``. + + Raises: + ValueError: If ``max_workers`` is not greater than 0. + OperationalError: If the supplied session does not exist, or the + session cannot be started or does not become idle. + """ + # Created before the session so that an invalid max_workers cannot leave + # a newly started session behind; the executor starts no threads until used. + self._max_workers = max_workers + self._executor = ThreadPoolExecutor(max_workers=max_workers) super().__init__( session_id=session_id, description=description, @@ -83,12 +104,24 @@ def __init__( session_idle_timeout_minutes=session_idle_timeout_minutes, **kwargs, ) - self._max_workers = max_workers - self._executor = ThreadPoolExecutor(max_workers=max_workers) def close(self, wait: bool = False) -> None: - super().close() - self._executor.shutdown(wait=wait) + """Terminate the Spark session, then shut down the executor. + + The executor is shut down even if terminating the session fails. + If termination fails, calling this method again retries it. + + Args: + wait: Whether to wait for submitted futures to finish before returning + or raising. + + Raises: + OperationalError: If terminating the session fails. + """ + try: + super().close() + finally: + self._executor.shutdown(wait=wait) def calculation_execution(self, query_id: str) -> "Future[AthenaCalculationExecution]": return self._executor.submit(self._get_calculation_execution, query_id) diff --git a/pyathena/spark/common.py b/pyathena/spark/common.py index d25c0b0f..de56273c 100644 --- a/pyathena/spark/common.py +++ b/pyathena/spark/common.py @@ -1,5 +1,6 @@ from __future__ import annotations +import contextlib import logging import time from abc import ABCMeta, abstractmethod @@ -56,6 +57,22 @@ def __init__( session_idle_timeout_minutes: int | None = None, **kwargs, ) -> None: + """Initialize the cursor and start or attach to a Spark session. + + Args: + session_id: ID of an existing session to use. If omitted, a new + session is started. + description: Description of a new session. + engine_configuration: Engine configuration of a new session. + Defaults to ``get_default_engine_configuration()``. + notebook_version: Notebook version of a new session. + session_idle_timeout_minutes: Idle timeout of a new session in minutes. + **kwargs: Arguments passed to ``BaseCursor``. + + Raises: + OperationalError: If the supplied session does not exist, or the + session cannot be started or does not become idle. + """ super().__init__(**kwargs) self._engine_configuration = ( engine_configuration @@ -65,18 +82,11 @@ def __init__( self._notebook_version = notebook_version self._session_description = description self._session_idle_timeout_minutes = session_idle_timeout_minutes - - if session_id: - if self._exists_session(session_id): - self._session_id = session_id - else: - raise OperationalError(f"Session: {session_id} not found.") - else: - self._session_id = self._start_session() - self._calculation_id: str | None = None self._calculation_execution: AthenaCalculationExecution | None = None + # Created before the session so that a local failure cannot leave + # a newly started session behind. self._client = self.connection.session.client( "s3", region_name=self.connection.region_name, @@ -84,6 +94,14 @@ def __init__( **self.connection._client_kwargs, ) + if session_id: + if self._exists_session(session_id): + self._session_id = session_id + else: + raise OperationalError(f"Session: {session_id} not found.") + else: + self._session_id = self._start_session() + @property def session_id(self) -> str: return self._session_id @@ -126,17 +144,29 @@ def _get_session_status(self, session_id: str): else: return AthenaSessionStatus(response) - def _wait_for_idle_session(self, session_id: str): + def _wait_for_idle_session(self, session_id: str) -> None: + """Poll a Spark session with ``GetSessionStatus`` until it is idle. + + Args: + session_id: The session ID. + + Raises: + OperationalError: If the session is terminated, degraded, or failed, + or if the request fails. + """ while True: session_status = self._get_session_status(session_id) if session_status.state in [AthenaSessionStatus.STATE_IDLE]: break - if session_status in [ + if session_status.state in [ AthenaSessionStatus.STATE_TERMINATED, AthenaSessionStatus.STATE_DEGRADED, AthenaSessionStatus.STATE_FAILED, ]: - raise OperationalError(session_status.state_change_reason) + message = f"Session: {session_id} is {session_status.state}." + if session_status.state_change_reason: + message += f" {session_status.state_change_reason}" + raise OperationalError(message) time.sleep(self._poll_interval) def _exists_session(self, session_id: str) -> bool: @@ -161,6 +191,19 @@ def _exists_session(self, session_id: str) -> bool: return True def _start_session(self) -> str: + """Start a Spark session with ``StartSession`` and wait until it is idle. + + If waiting for the new session raises, including ``KeyboardInterrupt``, + the session is terminated on a best-effort basis before the exception is + re-raised. + + Returns: + The ID of the new session. + + Raises: + OperationalError: If the session cannot be started or does not + become idle. + """ request: dict[str, Any] = { "WorkGroup": self._work_group, "EngineConfiguration": self._engine_configuration, @@ -181,12 +224,33 @@ def _start_session(self) -> str: except Exception as e: _logger.exception("Failed to start session.") raise OperationalError(*e.args) from e - else: + + try: self._wait_for_idle_session(session_id) - return session_id + except BaseException: + # The caller receives no cursor to close, so the session is released here. + with contextlib.suppress(OperationalError): + # Already logged with the session ID; the original error takes precedence. + self.__terminate_session(session_id) + raise + return session_id def _terminate_session(self) -> None: - request = {"SessionId": self._session_id} + self.__terminate_session(self._session_id) + + def __terminate_session(self, session_id: str) -> None: + """Terminate a Spark session with ``TerminateSession``. + + Session startup calls this synchronously in every cursor variant, + including those that override ``_terminate_session`` with a coroutine. + + Args: + session_id: The session ID. + + Raises: + OperationalError: If the request fails. + """ + request = {"SessionId": session_id} try: retry_api_call( self._connection.client.terminate_session, @@ -195,7 +259,7 @@ def _terminate_session(self) -> None: **request, ) except Exception as e: - _logger.exception("Failed to terminate session.") + _logger.exception(f"Failed to terminate session: {session_id}.") raise OperationalError(*e.args) from e def __poll(self, query_id: str) -> AthenaQueryExecution | AthenaCalculationExecution: diff --git a/pyathena/sqlalchemy/base.py b/pyathena/sqlalchemy/base.py index 1c2fce4b..5ceb0233 100644 --- a/pyathena/sqlalchemy/base.py +++ b/pyathena/sqlalchemy/base.py @@ -384,6 +384,34 @@ def get_columns(self, connection: Connection, table_name: str, schema: str | Non return columns def _get_column_type(self, type_: str, _nested: bool = False): + """Map an Athena column type string to a SQLAlchemy type. + + Accepts both the Hive (``struct``, ``map``) and the + Trino (``row(a integer)``, ``map(integer, integer)``) spellings, and + parses the element, key, value, and field types of ARRAY, MAP, and + STRUCT/ROW types. + + Args: + type_: The column type reported by Athena. + _nested: Whether ``type_`` is nested in another type. A nested MAP + or STRUCT/ROW that cannot be parsed raises, so that the + enclosing type is reported as unrecognized. + + Returns: + The SQLAlchemy type. A type name that is not recognized, such as + ``foo`` in ``struct``, becomes ``NullType`` in place with a + warning. A type that cannot be parsed, such as ``map`` or + ``varchar(x)``, makes its innermost enclosing ARRAY ``NullType`` + with a warning; without an enclosing ARRAY, a top-level MAP or + STRUCT/ROW becomes ``NullType`` instead. + + Raises: + ValueError: If a type cannot be parsed and neither an enclosing + ARRAY nor a top-level MAP or STRUCT/ROW handles it, for example + a top-level ``varchar(x)`` or a nested ``map``. + TypeError: In the same case, for a DECIMAL type with more + arguments than SQLAlchemy's ``DECIMAL`` accepts. + """ type_ = type_.strip() match = self._pattern_column_type.match(type_) if match: @@ -399,6 +427,12 @@ def _get_column_type(self, type_: str, _nested: bool = False): except (TypeError, ValueError): util.warn(f"Did not recognize type '{type_}'") return types.NullType() + if not _nested and name in ("map", "row", "struct") and length: + try: + return self._get_column_type(type_, _nested=True) + except (TypeError, ValueError): + util.warn(f"Did not recognize type '{type_}'") + return types.NullType() if _nested and name == "map" and length: key, value = _split_type_arguments(length) return AthenaMap( diff --git a/tests/pyathena/aio/test_common.py b/tests/pyathena/aio/test_common.py new file mode 100644 index 00000000..57d3ddd0 --- /dev/null +++ b/tests/pyathena/aio/test_common.py @@ -0,0 +1,37 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT +from unittest.mock import MagicMock + +import pytest + +from pyathena.aio.arrow.cursor import AioArrowCursor +from pyathena.aio.cursor import AioCursor, AioDictCursor +from pyathena.aio.pandas.cursor import AioPandasCursor +from pyathena.aio.polars.cursor import AioPolarsCursor +from pyathena.aio.s3fs.cursor import AioS3FSCursor + + +class TestWithAsyncFetch: + @pytest.mark.parametrize( + "cursor_class", + [ + AioCursor, + AioDictCursor, + AioArrowCursor, + AioPandasCursor, + AioPolarsCursor, + AioS3FSCursor, + ], + ) + def test_sync_iteration_raises(self, cursor_class): + cursor = cursor_class( + connection=MagicMock(), converter=None, formatter=None, retry_config=None + ) + with pytest.raises( + TypeError, match=rf"'{cursor_class.__name__}' object is not iterable; use 'async for'" + ): + iter(cursor) diff --git a/tests/pyathena/aio/test_cursor.py b/tests/pyathena/aio/test_cursor.py index 26733c90..2b777793 100644 --- a/tests/pyathena/aio/test_cursor.py +++ b/tests/pyathena/aio/test_cursor.py @@ -148,6 +148,37 @@ async def test_execute_internal_legacy_kwargs_passthrough(self): cache_expiration_time=100, ) + async def test_execute_qmark_parameters_skip_cache(self): + """A qmark query with parameters never searches the cache (no AWS, #941). + + Mirrors the synchronous cursor test. + """ + cursor = AioCursor.__new__(AioCursor) # bypass __init__ to avoid AWS calls + cursor._connection = MagicMock() + cursor._connection.client.start_query_execution.return_value = { + "QueryExecutionId": "test_query_id" + } + cursor._retry_config = RetryConfig() + cursor._kill_on_interrupt = True + + with ( + patch.object( + AioCursor, + "_build_start_query_execution_request", + return_value={"ExecutionParameters": ["'1'"]}, + ) as request_mock, + patch.object( + AioCursor, "_find_previous_query_id", new_callable=AsyncMock, return_value="cached" + ) as cache_mock, + ): + query_id = await cursor._execute( + "SELECT ?", ["'1'"], paramstyle="qmark", cache_size=10, cache_expiration_time=100 + ) + + assert query_id == "test_query_id" + assert request_mock.call_args.kwargs["execution_parameters"] == ["'1'"] + cache_mock.assert_not_awaited() + async def test_cache_size_different_schema(self): """A cached result is only reused when it ran against the same schema (#739). diff --git a/tests/pyathena/aio/test_result_set.py b/tests/pyathena/aio/test_result_set.py new file mode 100644 index 00000000..b34910af --- /dev/null +++ b/tests/pyathena/aio/test_result_set.py @@ -0,0 +1,63 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT +from unittest.mock import MagicMock + +import pytest + +from pyathena.aio.result_set import AthenaAioDictResultSet, AthenaAioResultSet +from pyathena.converter import DefaultTypeConverter +from pyathena.model import AthenaQueryExecution +from pyathena.util import RetryConfig + + +async def _create_result_set(result_set_class): + """Create a result set of a succeeded query from a stubbed ``GetQueryResults``. + + Args: + result_set_class: The asyncio result set class to create. + + Returns: + The result set, holding the rows 1 and 2 of an integer column ``a``. + """ + connection = MagicMock() + connection.client.get_query_results.return_value = { + "ResultSet": { + "ResultSetMetadata": {"ColumnInfo": [{"Name": "a", "Type": "integer"}]}, + "Rows": [{"Data": [{"VarCharValue": "1"}]}, {"Data": [{"VarCharValue": "2"}]}], + } + } + query_execution = AthenaQueryExecution( + { + "QueryExecution": { + "QueryExecutionId": "test_query_id", + "Query": "SELECT a", + "Status": {"State": AthenaQueryExecution.STATE_SUCCEEDED}, + } + } + ) + return await result_set_class.create( + connection, DefaultTypeConverter(), query_execution, 1000, RetryConfig() + ) + + +class TestAthenaAioResultSet: + @pytest.mark.parametrize("result_set_class", [AthenaAioResultSet, AthenaAioDictResultSet]) + async def test_sync_iteration_raises(self, result_set_class): + result_set = await _create_result_set(result_set_class) + with pytest.raises( + TypeError, + match=rf"'{result_set_class.__name__}' object is not iterable; use 'async for'", + ): + iter(result_set) + + @pytest.mark.parametrize( + ("result_set_class", "expected"), + [(AthenaAioResultSet, [(1,), (2,)]), (AthenaAioDictResultSet, [{"a": 1}, {"a": 2}])], + ) + async def test_async_iteration(self, result_set_class, expected): + result_set = await _create_result_set(result_set_class) + assert [row async for row in result_set] == expected diff --git a/tests/pyathena/arrow/test_cursor.py b/tests/pyathena/arrow/test_cursor.py index 11a3dabb..7f8fed60 100644 --- a/tests/pyathena/arrow/test_cursor.py +++ b/tests/pyathena/arrow/test_cursor.py @@ -35,10 +35,38 @@ def test_binary_null_vs_empty(self, arrow_cursor): ] assert [row[3] for row in rows] == ["", "", "NULL"] + def test_multiline_values_across_blocks(self, arrow_cursor): + # The 50 two-line values of 301 bytes span several 1024-byte blocks. + arrow_cursor.execute( + """ + SELECT array_join(repeat('x', 150), '') || chr(10) || array_join(repeat('y', 150), '') + AS v + FROM UNNEST(sequence(1, 50)) AS t(i) + """, + block_size=1024, + ) + assert arrow_cursor.fetchall() == [("x" * 150 + "\n" + "y" * 150,)] * 50 + def test_binary_single_null(self, arrow_cursor): arrow_cursor.execute("SELECT CAST(NULL AS VARBINARY) AS value") assert arrow_cursor.fetchall() == [(None,)] + @pytest.mark.parametrize( + ("query", "expected"), + [ + ( + "SELECT x FROM (VALUES 1, NULL, 2) AS t(x) ORDER BY x NULLS FIRST", + [(None,), (1,), (2,)], + ), + # Arrow reads a NULL string from a CSV result as an empty string. + ("SELECT CAST(NULL AS VARCHAR) AS v", [("",)]), + ], + ) + def test_single_column_null(self, arrow_cursor, query, expected): + arrow_cursor.execute(query) + assert arrow_cursor.as_arrow().num_rows == len(expected) + assert arrow_cursor.fetchall() == expected + @pytest.mark.parametrize( "arrow_cursor", [{"cursor_kwargs": {"unload": False}}, {"cursor_kwargs": {"unload": True}}], @@ -962,5 +990,16 @@ def test_null_vs_empty_string(self, arrow_cursor): indirect=["arrow_cursor"], ) def test_fetch_all_rows(self, arrow_cursor): - arrow_cursor.execute("SELECT 1 AS col") - assert arrow_cursor.fetchall() == [(1,)] + arrow_cursor.execute( + """ + SELECT + 1 AS col + ,CAST('12:34:56' AS TIME) AS col_time + ,X'0102' AS col_varbinary + ,json_parse('{"a": 1}') AS col_json + ,CAST('{"a": 1}' AS JSON) AS col_json_string + """ + ) + assert arrow_cursor.fetchall() == [ + (1, datetime(2017, 1, 1, 12, 34, 56).time(), b"\x01\x02", {"a": 1}, '{"a": 1}') + ] diff --git a/tests/pyathena/arrow/test_util.py b/tests/pyathena/arrow/test_util.py index 1060ecb5..06d326f2 100644 --- a/tests/pyathena/arrow/test_util.py +++ b/tests/pyathena/arrow/test_util.py @@ -1,6 +1,7 @@ import pyarrow as pa +import pytest -from pyathena.arrow.util import to_column_info +from pyathena.arrow.util import get_athena_type, to_column_info def test_to_column_info(): @@ -141,3 +142,14 @@ def test_to_column_info(): "Type": "decimal", }, ) + + +@pytest.mark.parametrize( + ("type_", "expected"), + [ + (pa.decimal128(38, 4), ("decimal", 38, 4)), + (pa.decimal256(40, 5), ("decimal", 40, 5)), + ], +) +def test_get_athena_type_decimal(type_, expected): + assert get_athena_type(type_) == expected diff --git a/tests/pyathena/filesystem/test_s3.py b/tests/pyathena/filesystem/test_s3.py index 14ae956b..0498ec29 100644 --- a/tests/pyathena/filesystem/test_s3.py +++ b/tests/pyathena/filesystem/test_s3.py @@ -1,3 +1,4 @@ +import functools import io import os import tempfile @@ -148,6 +149,16 @@ def _make_fs(): fs.version_aware = False return fs + @staticmethod + def _file_object(key): + # Build a listed file entry in the bucket named "bucket". + return S3Object( + init={"Key": key}, + type=S3ObjectType.S3_OBJECT_TYPE_FILE, + bucket="bucket", + key=key, + ) + def test_get_client_compatible_with_s3fs(self): # Only constructs a boto3 client; no AWS access. fs = S3FileSystem( @@ -192,6 +203,137 @@ def test_ls_from_cache_with_cached_object(self): # every cache value is a listing and raises TypeError here). assert fs._ls_from_cache("bucket/key/child") is None + def test_invalidate_cache_drops_listings_of_path_and_parents(self): + fs = self._make_fs() + invalidated = [ + "bucket/a/b/c.txt", + ("bucket/a/b", "/"), + ("bucket/a/b", ""), + ("bucket/a", "/"), + ("bucket/a", ""), + ("bucket", "/"), + ("bucket", ""), + ] + kept = ["", ("bucket/a/x", "/")] + for cache_key in invalidated + kept: + fs.dircache[cache_key] = [] + + fs.invalidate_cache("s3://bucket/a/b/c.txt") + assert list(fs.dircache) == kept + + @pytest.mark.parametrize( + ("prefix", "next_token"), + [ + ("test_", None), + ("", "token"), + ], + ) + def test_ls_dirs_partial_listing_bypasses_cache(self, prefix, next_token): + fs = self._make_fs() + cached = self._file_object("dir/cached") + fs.dircache[("bucket/dir", "")] = [cached] + fs._call.return_value = {"Contents": [{"Key": "dir/test_1"}]} + + files = fs._ls_dirs("bucket/dir", prefix=prefix, delimiter="", next_token=next_token) + assert [f.name for f in files] == ["bucket/dir/test_1"] + assert fs.dircache[("bucket/dir", "")] == [cached] + + # A complete listing of the path is still served from the cache. + fs._call.reset_mock() + assert fs._ls_dirs("bucket/dir", delimiter="") == [cached] + fs._call.assert_not_called() + + def test_ls_dirs_empty_refresh_evicts_cached_listing(self): + fs = self._make_fs() + fs.dircache[("bucket/dir", "/")] = [self._file_object("dir/deleted")] + fs._call.return_value = {} + + assert fs._ls_dirs("bucket/dir", refresh=True) == [] + # The next listing must not return the deleted object from the cache. + fs._call.reset_mock() + assert fs._ls_dirs("bucket/dir") == [] + fs._call.assert_called_once() + + def test_find_withdirs_does_not_modify_cached_listing(self): + fs = self._make_fs() + fs.dircache[("bucket/dir", "")] = [self._file_object("dir/sub/file")] + + expected = ["bucket/dir/sub", "bucket/dir/sub/file"] + assert sorted(fs.find("s3://bucket/dir", withdirs=True)) == expected + assert sorted(fs.find("s3://bucket/dir", withdirs=True)) == expected + assert fs.find("s3://bucket/dir") == ["bucket/dir/sub/file"] + fs._call.assert_not_called() + + def test_find_refresh_bypasses_cached_listings(self): + fs = self._make_fs() + fs.dircache[("bucket/dir", "")] = [self._file_object("dir/old")] + fs.dircache[("bucket/dir", "/")] = [self._file_object("dir/old")] + fs.dircache[("bucket/dir/sub", "/")] = [self._file_object("dir/sub/old")] + responses = { + ("dir/", ""): {"Contents": [{"Key": "dir/sub/new"}]}, + ("dir/", "/"): {"CommonPrefixes": [{"Prefix": "dir/sub/"}]}, + ("dir/sub/", "/"): {"Contents": [{"Key": "dir/sub/new"}]}, + } + fs._call.side_effect = lambda method, **kwargs: responses[ + (kwargs["Prefix"], kwargs["Delimiter"]) + ] + + assert fs.find("s3://bucket/dir", refresh=True) == ["bucket/dir/sub/new"] + # The subdirectory listings of maxdepth are refreshed as well. + assert fs.find("s3://bucket/dir", maxdepth=1, refresh=True) == ["bucket/dir/sub/new"] + + def test_refresh_evicts_cached_object_and_bucket_not_found(self): + fs = self._make_fs() + fs.dircache["bucket/key"] = self._file_object("key") + fs.dircache["bucket"] = fs._directory_object("bucket", None) + fs.dircache[""] = [fs._directory_object("bucket", None)] + + def call(method, **kwargs): + if method in (fs._client.head_object, fs._client.head_bucket): + raise FileNotFoundError + return {} + + fs._call.side_effect = call + + assert fs.ls("s3://bucket/key", refresh=True) == [] + # The next lookups must not return the deleted object and bucket from the cache. + assert fs.ls("s3://bucket/key") == [] + assert not fs.exists("s3://bucket/key") + with pytest.raises(FileNotFoundError): + fs.info("s3://bucket", refresh=True) + assert not fs.exists("s3://bucket") + + def test_exists_refresh_bypasses_cache(self): + fs = self._make_fs() + fs.dircache["bucket/key"] = self._file_object("key") + fs.dircache["bucket"] = fs._directory_object("bucket", None) + fs.dircache[""] = [fs._directory_object("bucket", None)] + + def call(method, **kwargs): + if method in (fs._client.head_object, fs._client.head_bucket): + raise FileNotFoundError + return {} + + fs._call.side_effect = call + + assert fs.exists("s3://bucket/key") + assert fs.exists("s3://bucket") + fs._call.assert_not_called() + + assert not fs.exists("s3://bucket/key", refresh=True) + assert not fs.exists("s3://bucket", refresh=True) + + def test_missing_bucket_keeps_bucket_listing_without_it(self): + fs = self._make_fs() + fs.dircache[""] = [fs._directory_object("bucket", None)] + fs._call.side_effect = FileNotFoundError + + assert not fs.exists("s3://missing") + # Other buckets are still answered from the cached bucket listing. + fs._call.reset_mock() + assert fs.exists("s3://bucket") + fs._call.assert_not_called() + def test_mkdir_creates_bucket(self): fs = self._make_fs() fs.allow_bucket_creation = True @@ -336,6 +478,63 @@ def test_finish_multipart_upload_abort_failure_does_not_mask_the_original_error( bucket="bucket", key="key", upload_id="uploadid", futures=[future] ) + @pytest.mark.parametrize( + ("size", "block_size", "ranges"), + [ + # A single range. + (5 * 2**20, 5 * 2**20, [(0, 5 * 2**20)]), + # The size is an exact multiple of the block size. + (10 * 2**30, 5 * 2**30, [(0, 5 * 2**30), (5 * 2**30, 10 * 2**30)]), + # A last range of the minimum part size is kept. + ( + 5 * 2**30 + 5 * 2**20, + 5 * 2**30, + [(0, 5 * 2**30), (5 * 2**30, 5 * 2**30 + 5 * 2**20)], + ), + # GH-951: a last range shorter than the minimum part size is + # merged into the previous one, + (15 * 2**20 - 1, 5 * 2**20, [(0, 5 * 2**20), (5 * 2**20, 15 * 2**20 - 1)]), + # which is split in half if it exceeds the maximum part size. + ( + 5 * 2**30 + 2**20, + 5 * 2**30, + [(0, 5 * 2**29 + 2**19), (5 * 2**29 + 2**19, 5 * 2**30 + 2**20)], + ), + ], + ) + def test_get_copy_ranges(self, size, block_size, ranges): + assert self._make_fs()._get_copy_ranges(size, block_size) == ranges + + @pytest.mark.parametrize("max_workers", [1, 4]) + def test_copy_object_with_multipart_upload_part_sizes(self, max_workers): + # GH-951: the parts are within the S3 part size limits whatever the + # number of workers; a single worker used to copy the whole object + # as one part larger than 5 GiB. + fs = self._make_fs() + fs._create_multipart_upload = mock.MagicMock( + return_value=SimpleNamespace(upload_id="uploadid") + ) + fs._upload_part_copy = mock.MagicMock() + fs._finish_multipart_upload = mock.MagicMock() + + fs._copy_object_with_multipart_upload( + bucket1="bucket", + key1="src", + size1=5 * 2**30 + 2**20, + bucket2="bucket", + key2="dst", + max_workers=max_workers, + ) + + parts = sorted( + (c.kwargs["part_number"], c.kwargs["copy_source_ranges"]) + for c in fs._upload_part_copy.call_args_list + ) + assert parts == [ + (1, (0, 5 * 2**29 + 2**19)), + (2, (5 * 2**29 + 2**19, 5 * 2**30 + 2**20)), + ] + def test_head_object_version_aware(self): fs = self._make_fs() fs._call.return_value = {"ContentLength": 4, "ETag": '"etag"', "VersionId": "v1"} @@ -591,6 +790,20 @@ def test_write(self, fs, base, exp): assert len(actual) == len(data) assert actual == data + def test_write_multiple_blocks_then_more(self, fs): + # GH-942: a single write() of more than two blocks with a short tail, + # followed by more data, must not leave a part smaller than the + # minimum part size before the last part. + path = ( + f"s3://{ENV.s3_staging_bucket}/{ENV.s3_staging_key}{ENV.schema}/" + f"filesystem/test_write_multiple_blocks_then_more/{uuid.uuid4()}" + ) + size = 2 * fs.default_block_size + 2**20 + with fs.open(path, "wb") as f: + f.write(b"a" * size) + f.write(b"b") + assert fs.info(path).get("size") == size + 1 + @pytest.mark.parametrize( "size", [ @@ -676,6 +889,65 @@ def test_append(self, fs, base, exp): assert len(actual) == len(data + extra) assert actual == data + extra + @pytest.mark.parametrize( + ("size", "extra_size", "block_size"), + [ + # GH-921: an existing object of at least 5 MiB, appended within a + # larger block size, is copied with UploadPartCopy. + (6 * 2**20, 5, 16 * 2**20), + # An existing object smaller than 5 MiB is rewritten from the + # buffer, not copied as well, when the append crosses the block size. + (2**10, 5 * 2**20, None), + ], + ) + def test_append_with_block_size(self, fs, size, extra_size, block_size): + data = b"a" * size + extra = b"b" * extra_size + path = ( + f"s3://{ENV.s3_staging_bucket}/{ENV.s3_staging_key}{ENV.schema}/" + f"filesystem/test_append_with_block_size/{uuid.uuid4()}" + ) + fs.pipe_file(path, data) + with fs.open(path, "ab", block_size=block_size) as f: + f.write(extra) + # Check the size and the bytes at the ends and around the boundary + # instead of reading the whole object back, to keep the transfer small. + assert fs.info(path, refresh=True).size == size + extra_size + assert fs.cat_file(path, start=0, end=1) == b"a" + assert fs.cat_file(path, start=size - 1, end=size + 1) == b"ab" + assert fs.cat_file(path, start=-1) == b"b" + + @pytest.mark.parametrize("block_size", [None, 16 * 2**20]) + def test_append_transaction_rollback(self, fs, block_size): + # Raising inside the transaction aborts the multipart upload that + # copies the existing object and leaves the object unchanged. + data = b"a" * (6 * 2**20) + path = ( + f"s3://{ENV.s3_staging_bucket}/{ENV.s3_staging_key}{ENV.schema}/" + f"filesystem/test_append_transaction_rollback/{uuid.uuid4()}" + ) + fs.pipe_file(path, data) + before = fs.info(path, refresh=True) + + def append_then_fail(): + with fs.transaction: + f = fs.open(path, "ab", block_size=block_size) + f.write(b"b" * 5) + f.close() + raise RuntimeError("rollback") + + with pytest.raises(RuntimeError): + append_then_fail() + # A committed append (a multipart upload, or the appended bytes alone) + # would change the ETag and the size, so the object is not read back. + after = fs.info(path, refresh=True) + assert (after.etag, after.last_modified, after.size) == ( + before.etag, + before.last_modified, + before.size, + ) + assert not fs.list_multipart_uploads(path) + def test_ls_buckets(self, fs): fs.invalidate_cache() actual = fs.ls("s3://") @@ -719,6 +991,29 @@ def test_ls_dirs(self, fs): assert test_1_detail[0].name == fs._strip_protocol(f"{dir_}/prefix/test_1") assert test_1_detail[0].size == 1 + def test_ls_and_find_reflect_changes_through_the_filesystem(self, fs): + dir_ = ( + f"s3://{ENV.s3_staging_bucket}/{ENV.s3_staging_key}{ENV.schema}/" + f"filesystem/test_ls_and_find_reflect_changes/{uuid.uuid4()}" + ) + path = fs._strip_protocol(dir_) + fs.touch(f"{dir_}/a.txt") + fs.touch(f"{dir_}/b.txt") + assert sorted(fs.ls(dir_)) == [f"{path}/a.txt", f"{path}/b.txt"] + assert sorted(fs.find(dir_)) == [f"{path}/a.txt", f"{path}/b.txt"] + + fs.rm(f"{dir_}/a.txt") + assert fs.ls(dir_) == [f"{path}/b.txt"] + assert fs.find(dir_) == [f"{path}/b.txt"] + + fs.touch(f"{dir_}/c.txt") + assert sorted(fs.ls(dir_)) == [f"{path}/b.txt", f"{path}/c.txt"] + assert sorted(fs.find(dir_)) == [f"{path}/b.txt", f"{path}/c.txt"] + # A prefixed find must not be served from the unprefixed listing. + assert fs.find(dir_, prefix="c") == [f"{path}/c.txt"] + + fs.rm(dir_, recursive=True) + def test_info_bucket(self, fs): dir_ = f"s3://{ENV.s3_staging_bucket}" bucket, key, version_id = fs.parse_path(dir_) @@ -1511,6 +1806,7 @@ def _make_write_file(data: bytes, autocommit: bool): file.s3_additional_kwargs = {} file.autocommit = autocommit file.blocksize = S3FileSystem.MULTIPART_UPLOAD_MIN_PART_SIZE + file.append_block = False file.multipart_upload = None file.multipart_upload_parts = [] file.buffer = io.BytesIO(data) @@ -1534,6 +1830,156 @@ def _make_multipart_write_file(data: bytes, autocommit: bool): ) return file + @staticmethod + def _make_append_fs(existing: bytes): + # A mocked filesystem holding an existing object, with a minimum part + # size of 4 bytes so that the write and append paths can be exercised + # with tiny data and no AWS access. + fs = mock.MagicMock(spec=S3FileSystem) + fs.MULTIPART_UPLOAD_MIN_PART_SIZE = 4 + fs.MULTIPART_UPLOAD_MAX_PART_SIZE = 64 + fs.exists.return_value = True + fs.info.return_value = S3Object( + init={"ContentLength": len(existing)}, + type=S3ObjectType.S3_OBJECT_TYPE_FILE, + bucket="bucket", + key="key.txt", + ) + fs.cat.return_value = existing + fs._create_multipart_upload.return_value = SimpleNamespace(upload_id="uploadid") + + def part(**kw): + return SimpleNamespace(etag=f'"e{kw["part_number"]}"', part_number=kw["part_number"]) + + fs._upload_part.side_effect = part + fs._upload_part_copy.side_effect = part + fs._get_copy_ranges.side_effect = functools.partial(S3FileSystem._get_copy_ranges, fs) + return fs + + @staticmethod + def _uploaded_object(fs, existing: bytes) -> bytes: + # Rebuild the object S3 would store from the mocked upload calls. + # A part copy without a range copies the whole existing object. + if fs._put_object.called: + fs._create_multipart_upload.assert_not_called() + return fs._put_object.call_args.kwargs["body"] + fs._finish_multipart_upload.assert_called_once() + parts = [] + for c in fs._upload_part_copy.call_args_list: + start, end = c.kwargs.get("copy_source_ranges", (0, len(existing))) + parts.append((c.kwargs["part_number"], existing[start:end])) + parts += [ + (c.kwargs["part_number"], c.kwargs["body"]) for c in fs._upload_part.call_args_list + ] + part_numbers = sorted(n for n, _ in parts) + assert part_numbers == list(range(1, len(parts) + 1)) + return b"".join(body for _, body in sorted(parts)) + + @pytest.mark.parametrize( + ("existing", "appended", "multipart", "part_copy"), + [ + # Smaller than the minimum part size: read into the buffer. + (b"aa", b"bb", False, False), + # GH-921: an existing object of at least the minimum part size is + # copied with UploadPartCopy even when the block size is larger + # than the whole object. + (b"a" * 6, b"bb", True, True), + (b"a" * 6, b"", True, True), + # An existing object read into the buffer is not copied again + # when the append crosses the block size. + (b"aa", b"b" * 16, True, False), + ], + ) + def test_append(self, existing, appended, multipart, part_copy): + fs = self._make_append_fs(existing) + + with S3File(fs, "s3://bucket/key.txt", mode="ab", block_size=16) as f: + f.write(appended) + + assert self._uploaded_object(fs, existing) == existing + appended + assert fs._create_multipart_upload.called is multipart + assert fs._upload_part_copy.called is part_copy + fs.touch.assert_not_called() + + @pytest.mark.parametrize("max_workers", [1, 4]) + def test_append_part_copy_ranges(self, max_workers): + # GH-951: an existing object larger than the maximum part size is + # copied in parts within the part size limits whatever the number of + # workers. A short remainder used to be copied as its own part, + # which is not the last one when data is appended. + existing = b"a" * 129 + fs = self._make_append_fs(existing) + + with S3File( + fs, "s3://bucket/key.txt", mode="ab", block_size=16, max_workers=max_workers + ) as f: + f.write(b"b") + + assert self._uploaded_object(fs, existing) == existing + b"b" + ranges = sorted( + (c.kwargs["part_number"], c.kwargs["copy_source_ranges"]) + for c in fs._upload_part_copy.call_args_list + ) + assert ranges == [(1, (0, 64)), (2, (64, 96)), (3, (96, 129))] + + @pytest.mark.parametrize( + ("writes", "block_size"), + [ + # GH-942: a single write() that leaves more than two blocks with a + # short tail in the buffer, followed by more data. + ([b"a" * 11, b"b"], 4), + ([b"a" * 9, b"b"], 4), + ([b"a" * 11], 4), + ([b"a" * 12, b"b" * 3], 4), + ([b"a" * 3, b"b" * 10, b"c" * 2, b"d"], 4), + # A merged tail that reaches the maximum part size is split in half. + ([b"a" * 127, b"b"], 62), + ], + ) + def test_write_part_sizes(self, writes, block_size): + fs = self._make_append_fs(b"") + + with S3File(fs, "s3://bucket/key.txt", mode="wb", block_size=block_size) as f: + for data in writes: + f.write(data) + + assert self._uploaded_object(fs, b"") == b"".join(writes) + parts = sorted( + (c.kwargs["part_number"], len(c.kwargs["body"])) for c in fs._upload_part.call_args_list + ) + sizes = [size for _, size in parts] + assert all(size >= fs.MULTIPART_UPLOAD_MIN_PART_SIZE for size in sizes[:-1]) + assert all(size <= fs.MULTIPART_UPLOAD_MAX_PART_SIZE for size in sizes) + + def test_append_discard(self): + # Rolling back an append aborts its multipart upload without the + # existing object's metadata, which AbortMultipartUpload rejects, + # but with the request parameters it accepts. + fs = self._make_append_fs(b"a" * 6) + f = S3File( + fs, + "s3://bucket/key.txt", + mode="ab", + block_size=16, + autocommit=False, + s3_additional_kwargs={"RequestPayer": "requester", "ExpectedBucketOwner": "123"}, + ) + f.write(b"bb") + f.close() + + f.discard() + + fs._call.assert_called_once_with( + "abort_multipart_upload", + Bucket="bucket", + Key="key.txt", + UploadId="uploadid", + RequestPayer="requester", + ExpectedBucketOwner="123", + ) + fs._finish_multipart_upload.assert_not_called() + fs._put_object.assert_not_called() + @pytest.mark.parametrize( ("objects", "target"), [ diff --git a/tests/pyathena/filesystem/test_s3_async.py b/tests/pyathena/filesystem/test_s3_async.py index aae3e1af..f609b8a8 100644 --- a/tests/pyathena/filesystem/test_s3_async.py +++ b/tests/pyathena/filesystem/test_s3_async.py @@ -7,6 +7,8 @@ from datetime import datetime, timezone from itertools import chain from pathlib import Path +from types import SimpleNamespace +from unittest import mock import fsspec import pytest @@ -129,6 +131,41 @@ def test_parse_path_invalid(self): with pytest.raises(ValueError, match="Invalid S3 path format"): AioS3FileSystem.parse_path("s3a://bucket/path/to/obj?foo=bar") + @pytest.mark.parametrize("max_workers", [1, 4]) + @pytest.mark.asyncio + async def test_copy_object_with_multipart_upload_part_sizes(self, max_workers): + # GH-951: the parts are within the S3 part size limits whatever the + # number of workers; a single worker used to copy the whole object + # as one part larger than 5 GiB. + fs = AioS3FileSystem( + connection=mock.MagicMock(), max_workers=max_workers, skip_instance_cache=True + ) + sync_fs = fs._sync_fs + sync_fs._create_multipart_upload = mock.MagicMock( + return_value=SimpleNamespace(upload_id="uploadid") + ) + sync_fs._upload_part_copy = mock.MagicMock( + side_effect=lambda **kw: SimpleNamespace(etag='"e"', part_number=kw["part_number"]) + ) + sync_fs._complete_multipart_upload = mock.MagicMock() + + await fs._copy_object_with_multipart_upload( + bucket1="bucket", + key1="src", + size1=5 * 2**30 + 2**20, + bucket2="bucket", + key2="dst", + ) + + parts = sorted( + (c.kwargs["part_number"], c.kwargs["copy_source_ranges"]) + for c in sync_fs._upload_part_copy.call_args_list + ) + assert parts == [ + (1, (0, 5 * 2**29 + 2**19)), + (2, (5 * 2**29 + 2**19, 5 * 2**30 + 2**20)), + ] + @pytest.fixture(scope="class") def fs(self, request): if not hasattr(request, "param"): diff --git a/tests/pyathena/pandas/test_cursor.py b/tests/pyathena/pandas/test_cursor.py index 00bd26fb..31ebba39 100644 --- a/tests/pyathena/pandas/test_cursor.py +++ b/tests/pyathena/pandas/test_cursor.py @@ -12,6 +12,7 @@ import numpy as np import pandas as pd import pytest +from pandas.io.parsers import TextFileReader from pyathena.error import DatabaseError, ProgrammingError from pyathena.pandas.converter import DefaultPandasTypeConverter @@ -897,6 +898,7 @@ def test_get_csv_engine_explicit_specification(self): result_set = AthenaPandasResultSet.__new__(AthenaPandasResultSet) result_set._chunksize = None # Default values result_set._quoting = 1 + result_set._query_execution = None # Test C engine specification result_set._engine = "c" @@ -1353,17 +1355,32 @@ def test_pandas_cursor_iter_chunks_with_chunksize(self, pandas_cursor): assert chunk_count >= 1 - def test_pandas_cursor_auto_optimize_chunksize_enabled(self, pandas_cursor): - """Test PandasCursor with auto_optimize_chunksize enabled.""" - cursor = pandas_cursor - cursor._chunksize = None # No explicit chunksize - cursor._auto_optimize_chunksize = True - cursor.execute("SELECT number FROM (VALUES (1), (2), (3), (4), (5)) as t(number)") - - # Should work without error (auto-optimization for small files may not trigger chunking) - result = cursor.as_pandas() - # Small test data likely won't trigger chunking, so expect DataFrame - assert isinstance(result, (pd.DataFrame, PandasDataFrameIterator)) + @pytest.mark.parametrize( + ("pandas_cursor", "chunked"), + [ + ({"cursor_kwargs": {"auto_optimize_chunksize": True}}, False), + ({"cursor_kwargs": {"auto_optimize_chunksize": True}}, True), + ], + indirect=["pandas_cursor"], + ) + def test_pandas_cursor_auto_optimize_chunksize_enabled( + self, pandas_cursor, chunked, monkeypatch + ): + """Test that as_pandas() returns the whole result with auto_optimize_chunksize.""" + if chunked: + # 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) + pandas_cursor.execute("SELECT number FROM (VALUES (1), (2), (3), (4), (5)) AS t(number)") + reader = pandas_cursor.result_set._df_iter._reader + assert isinstance(reader, TextFileReader) is chunked + + df = pandas_cursor.as_pandas() + assert isinstance(df, pd.DataFrame) + assert df["number"].tolist() == [1, 2, 3, 4, 5] + assert df.index.tolist() == [0, 1, 2, 3, 4] def test_pandas_cursor_auto_optimize_chunksize_disabled(self, pandas_cursor): """Test PandasCursor with auto_optimize_chunksize disabled (default).""" diff --git a/tests/pyathena/pandas/test_result_set.py b/tests/pyathena/pandas/test_result_set.py new file mode 100644 index 00000000..67dbb5ba --- /dev/null +++ b/tests/pyathena/pandas/test_result_set.py @@ -0,0 +1,106 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT + +import csv +import io +from unittest.mock import MagicMock + +import pandas as pd +import pytest + +from pyathena.pandas.converter import DefaultPandasTypeConverter +from pyathena.pandas.result_set import ( + AthenaPandasResultSet, + PandasDataFrameIterator, + _no_trunc_date, +) + + +class TestPandasDataFrameIterator: + @pytest.mark.parametrize( + ("csv", "read_csv_kwargs"), + [ + ("id,kind\n10,a\n11,a\n12,b\n13,b\n14,c\n", {}), + ("id,kind\n10,a\n11,a\n12,b\n13,b\n14,c\n", {"index_col": "id"}), + ("id,kind\n10,c\n11,c\n12,a\n13,a\n14,b\n", {"dtype": {"kind": "category"}}), + ("id,kind\n10,\n11,\n12,b\n13,a\n14,\n", {"dtype": {"kind": "category"}}), + ( + "id,kind\n10,c\n11,c\n12,a\n13,a\n14,b\n", + {"dtype": {"kind": pd.CategoricalDtype(ordered=True)}}, + ), + ( + "id,kind\n10,c\n11,c\n12,a\n13,a\n14,b\n", + {"index_col": "kind", "dtype": {"kind": "category"}}, + ), + ], + ids=[ + "default_index", + "index_col", + "category", + "category_null_chunk", + "ordered_category", + "category_index", + ], + ) + def test_as_pandas_matches_whole_read(self, csv, read_csv_kwargs): + """Joining the chunks gives the DataFrame that reading the whole file gives.""" + expected = pd.read_csv(io.StringIO(csv), **read_csv_kwargs) + reader = pd.read_csv(io.StringIO(csv), chunksize=2, **read_csv_kwargs) + df_iter = PandasDataFrameIterator(reader, _no_trunc_date) + + pd.testing.assert_frame_equal(df_iter.as_pandas(), expected) + + def test_as_pandas_remaining_chunks(self): + """After a chunk was read, the remaining chunks keep their row numbers.""" + reader = pd.read_csv(io.StringIO("n\n1\n2\n3\n4\n5\n"), chunksize=2) + df_iter = PandasDataFrameIterator(reader, _no_trunc_date) + next(df_iter) + + df = df_iter.as_pandas() + assert df["n"].tolist() == [3, 4, 5] + assert df.index.tolist() == [2, 3, 4] + assert df_iter.as_pandas().empty + + def test_as_pandas_single_dataframe(self): + """A DataFrame that was read at once is returned as is.""" + df = pd.DataFrame({"n": [1, 2]}) + df_iter = PandasDataFrameIterator(df, _no_trunc_date) + + assert df_iter.as_pandas() is df + + +class TestAthenaPandasResultSet: + @pytest.mark.parametrize("engine", ["auto", "c", "python", "pyarrow"]) + @pytest.mark.parametrize( + ("output_location", "file_size_bytes", "pyarrow_engine"), + [ + ("s3://bucket/result.txt", None, "c"), + ("s3://bucket/result.txt", 99, "c"), + ("s3://bucket/result.txt", 100, "c"), + ("s3://bucket/result.txt", 101, "c"), + ("s3://bucket/result.csv", None, "pyarrow"), + ("s3://bucket/result.csv", 99, "c"), + ("s3://bucket/result.csv", 100, "pyarrow"), + ("s3://bucket/result.csv", 101, "pyarrow"), + (None, None, "pyarrow"), + ], + ) + def test_get_csv_engine_result_format( + self, engine, output_location, file_size_bytes, pyarrow_engine + ): + """PyArrow falls back for DDL text while C and Python choices are preserved.""" + result_set = AthenaPandasResultSet.__new__(AthenaPandasResultSet) + result_set._query_execution = MagicMock(output_location=output_location) + result_set._metadata = None + result_set._converter = DefaultPandasTypeConverter() + result_set._engine = engine + result_set._chunksize = None + result_set._quoting = csv.QUOTE_ALL + result_set._kwargs = {} + + expected = {"auto": "c", "c": "c", "python": "python", "pyarrow": pyarrow_engine}[engine] + assert result_set._get_csv_engine(file_size_bytes) == expected diff --git a/tests/pyathena/pandas/test_util.py b/tests/pyathena/pandas/test_util.py index 788308da..1acbc07e 100644 --- a/tests/pyathena/pandas/test_util.py +++ b/tests/pyathena/pandas/test_util.py @@ -342,7 +342,9 @@ def test_to_sql(cursor): "col_binary", ] ] - table_name = f"""to_sql_{str(uuid.uuid4()).replace("-", "")}""" + # Uppercase letters: Athena reports names in lowercase, and the existence + # check behind if_exists has to find the table anyway. + table_name = f"To_Sql_{uuid.uuid4().hex}" location = f"{ENV.s3_staging_dir}{ENV.schema}/{table_name}/" to_sql( df, diff --git a/tests/pyathena/polars/test_result_set.py b/tests/pyathena/polars/test_result_set.py new file mode 100644 index 00000000..c7040332 --- /dev/null +++ b/tests/pyathena/polars/test_result_set.py @@ -0,0 +1,24 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT + +import polars as pl +import pytest + +from pyathena.polars.result_set import PolarsDataFrameIterator + + +class TestPolarsDataFrameIterator: + @pytest.mark.parametrize( + "reader", + [pl.DataFrame({"a": [1, 2]}), (df for df in [pl.DataFrame({"a": [1]})] * 2)], + ids=["dataframe", "generator"], + ) + def test_close_stops_iteration(self, reader): + """A closed iterator yields nothing for either reader kind.""" + df_iter = PolarsDataFrameIterator(reader, {}, ["a"]) + df_iter.close() + assert list(df_iter) == [] diff --git a/tests/pyathena/spark/test_async_cursor.py b/tests/pyathena/spark/test_async_cursor.py index e921bdc3..6f9f2d25 100644 --- a/tests/pyathena/spark/test_async_cursor.py +++ b/tests/pyathena/spark/test_async_cursor.py @@ -1,10 +1,20 @@ import textwrap +import threading import time +from concurrent.futures import ThreadPoolExecutor from random import randint +from unittest.mock import MagicMock +import pytest + +from pyathena import OperationalError from pyathena.model import AthenaCalculationExecutionStatus +from pyathena.spark.async_cursor import AsyncSparkCursor from tests import ENV +# Bounds how long the executor test task blocks when nothing releases it. +_TIMEOUT = 10 + class TestAsyncSparkCursor: def test_spark_dataframe(self, async_spark_cursor): @@ -121,3 +131,87 @@ def test_cancel(self, async_spark_cursor): calculation_execution = future.result() assert calculation_execution.state == AthenaCalculationExecutionStatus.STATE_CANCELED + + @staticmethod + def _cursor_with_submitted_work(): + """Build a cursor whose executor holds one running and one queued future. + + Returns: + A tuple of the cursor, the event that releases the running future, + the running future, and the queued future. + """ + cursor = AsyncSparkCursor.__new__(AsyncSparkCursor) # bypass __init__ to avoid AWS calls + cursor._executor = MagicMock(wraps=ThreadPoolExecutor(max_workers=1)) + started = threading.Event() + release = threading.Event() + + def run(): + started.set() + release.wait(_TIMEOUT) + return "running" + + running = cursor._executor.submit(run) + queued = cursor._executor.submit(lambda: "queued") + assert started.wait(_TIMEOUT) + return cursor, release, running, queued + + @pytest.mark.parametrize("fails", [False, True]) + def test_close_wait_shuts_down_executor(self, fails): + cursor, release, running, queued = self._cursor_with_submitted_work() + + def terminate(): + # Let the running future finish so that shutdown(wait=True) returns. + release.set() + if fails: + raise OperationalError("termination failed") + + cursor._terminate_session = MagicMock(side_effect=terminate) + if fails: + with pytest.raises(OperationalError, match="termination failed"): + cursor.close(wait=True) + else: + cursor.close(wait=True) + + cursor._terminate_session.assert_called_once_with() + cursor._executor.shutdown.assert_called_once_with(wait=True) + assert running.done() + assert running.result() == "running" + assert queued.done() + assert queued.result() == "queued" + with pytest.raises(RuntimeError): + cursor._executor.submit(lambda: None) + + @pytest.mark.parametrize("fails", [False, True]) + def test_close_no_wait_shuts_down_executor(self, fails): + cursor, release, running, queued = self._cursor_with_submitted_work() + cursor._terminate_session = MagicMock( + side_effect=OperationalError("termination failed") if fails else None + ) + try: + if fails: + with pytest.raises(OperationalError, match="termination failed"): + cursor.close(wait=False) + else: + cursor.close(wait=False) + + cursor._terminate_session.assert_called_once_with() + cursor._executor.shutdown.assert_called_once_with(wait=False) + with pytest.raises(RuntimeError): + cursor._executor.submit(lambda: None) + finally: + release.set() + assert running.result(_TIMEOUT) == "running" + assert queued.result(_TIMEOUT) == "queued" + + def test_close_retries_termination_after_failure(self): + cursor = AsyncSparkCursor.__new__(AsyncSparkCursor) # bypass __init__ to avoid AWS calls + cursor._terminate_session = MagicMock( + side_effect=[OperationalError("termination failed"), None] + ) + cursor._executor = ThreadPoolExecutor(max_workers=1) + + with pytest.raises(OperationalError, match="termination failed"): + cursor.close() + cursor.close() + + assert cursor._terminate_session.call_count == 2 diff --git a/tests/pyathena/spark/test_common.py b/tests/pyathena/spark/test_common.py new file mode 100644 index 00000000..9dbfd27f --- /dev/null +++ b/tests/pyathena/spark/test_common.py @@ -0,0 +1,238 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT + +import logging +from unittest.mock import MagicMock, patch + +import pytest +from botocore.exceptions import ClientError + +from pyathena import OperationalError +from pyathena.aio.spark.cursor import AioSparkCursor +from pyathena.model import AthenaSessionStatus +from pyathena.spark.async_cursor import AsyncSparkCursor +from pyathena.spark.common import SparkBaseCursor +from pyathena.spark.cursor import SparkCursor +from pyathena.util import RetryConfig + +SPARK_CURSOR_CLASSES = [SparkCursor, AsyncSparkCursor, AioSparkCursor] + + +def _session_status(state: str, reason: str | None = None) -> AthenaSessionStatus: + return AthenaSessionStatus( + {"SessionId": "session_id", "Status": {"State": state, "StateChangeReason": reason}} + ) + + +def _cursor() -> SparkCursor: + cursor = SparkCursor.__new__(SparkCursor) # bypass __init__ to avoid AWS calls + cursor._connection = MagicMock() + cursor._retry_config = RetryConfig() + cursor._poll_interval = 0 + return cursor + + +def _connection(): + connection = MagicMock() + connection.client.start_session.return_value = {"SessionId": "new-session"} + connection.client.get_session.return_value = {"SessionId": "supplied-session"} + connection.client.get_session_status.return_value = { + "Status": {"State": AthenaSessionStatus.STATE_IDLE} + } + return connection + + +def _init_cursor(cursor_class, connection, **kwargs): + return cursor_class( + connection=connection, + converter=MagicMock(), + formatter=MagicMock(), + retry_config=RetryConfig(), + s3_staging_dir=None, + schema_name=None, + catalog_name=None, + work_group="spark", + poll_interval=0, + encryption_option=None, + kms_key=None, + kill_on_interrupt=False, + result_reuse_enable=False, + result_reuse_minutes=60, + **kwargs, + ) + + +class TestSparkBaseCursor: + @pytest.mark.parametrize( + "state", + [ + AthenaSessionStatus.STATE_TERMINATED, + AthenaSessionStatus.STATE_DEGRADED, + AthenaSessionStatus.STATE_FAILED, + ], + ) + def test_wait_for_idle_session_raises_on_failure_state(self, state): + cursor = _cursor() + with ( + patch.object( + SparkCursor, + "_get_session_status", + return_value=_session_status(state, "session failure reason"), + ), + patch("pyathena.spark.common.time.sleep", side_effect=AssertionError("slept")), + pytest.raises( + OperationalError, + match=rf"^Session: session_id is {state}\. session failure reason$", + ), + ): + cursor._wait_for_idle_session("session_id") + + def test_wait_for_idle_session_raises_without_reason(self): + cursor = _cursor() + with ( + patch.object( + SparkCursor, + "_get_session_status", + return_value=_session_status(AthenaSessionStatus.STATE_TERMINATED), + ), + patch("pyathena.spark.common.time.sleep", side_effect=AssertionError("slept")), + pytest.raises(OperationalError, match=r"^Session: session_id is TERMINATED\.$"), + ): + cursor._wait_for_idle_session("session_id") + + def test_wait_for_idle_session_waits_until_idle(self): + cursor = _cursor() + statuses = [ + _session_status(AthenaSessionStatus.STATE_CREATING), + _session_status(AthenaSessionStatus.STATE_BUSY), + _session_status(AthenaSessionStatus.STATE_IDLE), + ] + with ( + patch.object(SparkCursor, "_get_session_status", side_effect=statuses) as get_status, + patch("pyathena.spark.common.time.sleep") as sleep, + ): + cursor._wait_for_idle_session("session_id") + assert get_status.call_count == 3 + assert sleep.call_count == 2 + + def test_exists_session_raises_on_failure_state(self): + cursor = _cursor() + with ( + patch.object( + SparkCursor, + "_get_session_status", + return_value=_session_status( + AthenaSessionStatus.STATE_TERMINATED, "session failure reason" + ), + ), + patch("pyathena.spark.common.time.sleep", side_effect=AssertionError("slept")), + pytest.raises(OperationalError, match="session failure reason"), + ): + cursor._exists_session("session_id") + cursor._connection.client.get_session.assert_called_once_with(SessionId="session_id") + + @pytest.mark.parametrize("cursor_class", SPARK_CURSOR_CLASSES) + def test_init_starts_session(self, cursor_class): + connection = _connection() + with patch.object(SparkBaseCursor, "_wait_for_idle_session"): + cursor = _init_cursor(cursor_class, connection) + + assert cursor.session_id == "new-session" + connection.client.terminate_session.assert_not_called() + + @pytest.mark.parametrize("cursor_class", SPARK_CURSOR_CLASSES) + @pytest.mark.parametrize( + "error", + [OperationalError("Session did not become idle."), KeyboardInterrupt()], + ) + def test_init_terminates_new_session_that_does_not_become_idle(self, cursor_class, error): + connection = _connection() + with ( + patch.object(SparkBaseCursor, "_wait_for_idle_session", side_effect=error), + pytest.raises(type(error)) as exc_info, + ): + _init_cursor(cursor_class, connection) + + assert exc_info.value is error + connection.client.terminate_session.assert_called_once_with(SessionId="new-session") + + @pytest.mark.parametrize("cursor_class", SPARK_CURSOR_CLASSES) + @pytest.mark.parametrize( + "state", + [ + AthenaSessionStatus.STATE_TERMINATED, + AthenaSessionStatus.STATE_DEGRADED, + AthenaSessionStatus.STATE_FAILED, + ], + ) + def test_init_terminates_new_session_in_failure_state(self, cursor_class, state): + connection = _connection() + connection.client.get_session_status.return_value = { + "SessionId": "new-session", + "Status": {"State": state, "StateChangeReason": "session failure reason"}, + } + with ( + patch("pyathena.spark.common.time.sleep", side_effect=AssertionError("slept")), + pytest.raises( + OperationalError, + match=rf"^Session: new-session is {state}\. session failure reason$", + ), + ): + _init_cursor(cursor_class, connection) + + connection.client.get_session_status.assert_called_once_with(SessionId="new-session") + connection.client.terminate_session.assert_called_once_with(SessionId="new-session") + + @pytest.mark.parametrize("cursor_class", SPARK_CURSOR_CLASSES) + def test_init_keeps_original_error_when_cleanup_fails(self, cursor_class, caplog): + connection = _connection() + connection.client.terminate_session.side_effect = ClientError( + {"Error": {"Code": "InternalServerException", "Message": "Cleanup failed."}}, + "TerminateSession", + ) + error = OperationalError("Session did not become idle.") + with ( + caplog.at_level(logging.ERROR, logger="pyathena.spark.common"), + patch.object(SparkBaseCursor, "_wait_for_idle_session", side_effect=error), + pytest.raises(OperationalError) as exc_info, + ): + _init_cursor(cursor_class, connection) + + assert exc_info.value is error + connection.client.terminate_session.assert_called_once_with(SessionId="new-session") + assert "Failed to terminate session: new-session." in caplog.text + + @pytest.mark.parametrize("cursor_class", SPARK_CURSOR_CLASSES) + def test_init_does_not_terminate_supplied_session(self, cursor_class): + connection = _connection() + error = OperationalError("Session did not become idle.") + with ( + patch.object(SparkBaseCursor, "_wait_for_idle_session", side_effect=error), + pytest.raises(OperationalError) as exc_info, + ): + _init_cursor(cursor_class, connection, session_id="supplied-session") + + assert exc_info.value is error + connection.client.start_session.assert_not_called() + connection.client.terminate_session.assert_not_called() + + @pytest.mark.parametrize("cursor_class", SPARK_CURSOR_CLASSES) + def test_init_does_not_start_session_when_s3_client_fails(self, cursor_class): + connection = _connection() + connection.session.client.side_effect = ValueError("Invalid S3 client configuration.") + with pytest.raises(ValueError, match=r"^Invalid S3 client configuration\.$"): + _init_cursor(cursor_class, connection) + + connection.client.start_session.assert_not_called() + connection.client.terminate_session.assert_not_called() + + def test_async_init_does_not_start_session_when_executor_fails(self): + connection = _connection() + with pytest.raises(ValueError, match="max_workers must be greater than 0"): + _init_cursor(AsyncSparkCursor, connection, max_workers=0) + + connection.client.start_session.assert_not_called() diff --git a/tests/pyathena/sqlalchemy/test_base.py b/tests/pyathena/sqlalchemy/test_base.py index 2863fb77..8aec1fa4 100644 --- a/tests/pyathena/sqlalchemy/test_base.py +++ b/tests/pyathena/sqlalchemy/test_base.py @@ -20,6 +20,7 @@ from sqlalchemy.sql.selectable import TextualSelect from sqlalchemy.util import PluginLoader +from pyathena.sqlalchemy.base import AthenaDialect from pyathena.sqlalchemy.rest import AthenaRestDialect from pyathena.sqlalchemy.types import ( TINYINT, @@ -81,6 +82,54 @@ def test_compliance_suite_registry_matches_entry_points(self, monkeypatch): assert entry_points assert {name: load() for name, load in loader.impls.items()} == entry_points + @pytest.mark.parametrize( + ("map_type", "struct_type"), + [ + # Metadata API (Hive) spellings. + ("map", "struct>>"), + # information_schema (Trino) spellings. + ("map(integer, varchar)", 'row(a integer, "b c" array(row(x integer)))'), + ], + ) + def test_top_level_map_and_struct_reflect_their_types(self, map_type, struct_type): + dialect = AthenaDialect() + map_ = dialect._get_column_type(map_type) + struct = dialect._get_column_type(struct_type) + + assert isinstance(map_, AthenaMap) + assert isinstance(map_.key_type, types.INTEGER) + assert isinstance(map_.value_type, (types.String, types.VARCHAR)) + assert isinstance(struct, AthenaStruct) + assert list(struct.fields) == ["a", "b c"] + assert isinstance(struct.fields["a"], types.INTEGER) + assert isinstance(struct.fields["b c"], AthenaArray) + assert list(struct.fields["b c"].item_type.fields) == ["x"] + + table = Table( + "t", + MetaData(), + Column("m", map_), + Column("s", struct), + awsathena_location="s3://bucket/path/", + ) + ddl = str(CreateTable(table).compile(dialect=dialect)) + assert "\tm MAP,\n" in ddl + assert "\ts STRUCT>>\n" in ddl + + def test_unrecognized_field_type_reflects_null_type(self): + dialect = AthenaDialect() + with pytest.warns(sqlalchemy.exc.SAWarning, match="Did not recognize type"): + struct = dialect._get_column_type("struct") + assert isinstance(struct.fields["a"], types.INTEGER) + assert isinstance(struct.fields["b"], types.NullType) + + @pytest.mark.parametrize( + "type_", ["map", "struct", "row(a)", "struct>", "map>"] + ) + def test_unrecognized_top_level_map_or_struct_reflects_null_type(self, type_): + with pytest.warns(sqlalchemy.exc.SAWarning, match="Did not recognize type"): + assert isinstance(AthenaDialect()._get_column_type(type_), types.NullType) + class TestSQLAlchemyAthena: @pytest.mark.parametrize( @@ -616,10 +665,15 @@ def test_reflect_select(self, engine): assert isinstance(one_row_complex.c.col_binary.type, types.BINARY) assert isinstance(one_row_complex.c.col_array.type, AthenaArray) assert isinstance(one_row_complex.c.col_array.type.item_type, types.INTEGER) - assert isinstance(one_row_complex.c.col_map.type, types.String) - # With struct support, col_struct should now be recognized as AthenaStruct - + assert isinstance(one_row_complex.c.col_map.type, AthenaMap) + assert isinstance(one_row_complex.c.col_map.type.key_type, types.INTEGER) + assert isinstance(one_row_complex.c.col_map.type.value_type, types.INTEGER) assert isinstance(one_row_complex.c.col_struct.type, AthenaStruct) + assert list(one_row_complex.c.col_struct.type.fields) == ["a", "b"] + assert isinstance(one_row_complex.c.col_struct.type.fields["a"], types.INTEGER) + ddl = str(CreateTable(one_row_complex).compile(dialect=engine.dialect)) + assert "\tcol_map MAP,\n" in ddl + assert "\tcol_struct STRUCT,\n" in ddl assert isinstance( one_row_complex.c.col_decimal.type, types.DECIMAL, @@ -669,11 +723,13 @@ def test_get_column_type(self, engine): assert isinstance(dialect._get_column_type("date"), types.DATE) assert isinstance(dialect._get_column_type("binary"), types.BINARY) assert isinstance(dialect._get_column_type("array"), AthenaArray) - assert isinstance(dialect._get_column_type("map"), types.String) - # With struct support, struct types should be recognized as AthenaStruct - - assert isinstance(dialect._get_column_type("struct"), AthenaStruct) - assert isinstance(dialect._get_column_type("row"), AthenaStruct) + assert isinstance(dialect._get_column_type("map"), AthenaMap) + struct = dialect._get_column_type("struct") + assert isinstance(struct, AthenaStruct) + assert list(struct.fields) == ["a", "b"] + row = dialect._get_column_type("row") + assert isinstance(row, AthenaStruct) + assert list(row.fields) == ["name", "age"] decimal_with_args = dialect._get_column_type("decimal(10,1)") assert isinstance(decimal_with_args, types.DECIMAL) assert decimal_with_args.precision == 10 diff --git a/tests/pyathena/sqlalchemy/test_compiler.py b/tests/pyathena/sqlalchemy/test_compiler.py index db9f5baa..88af4baa 100644 --- a/tests/pyathena/sqlalchemy/test_compiler.py +++ b/tests/pyathena/sqlalchemy/test_compiler.py @@ -214,6 +214,11 @@ def test_floating_point_types(self, type_, ddl, cast_type): assert dialect.type_compiler_instance.process(type_) == ddl assert str(cast(column("x"), type_).compile(dialect=dialect)) == f"CAST(x AS {cast_type})" + def test_null_type_cast_is_unchanged(self): + assert str(cast(column("x"), types.NullType()).compile(dialect=AthenaDialect())) == ( + "CAST(x AS NULL)" + ) + class TestAthenaStatementCompiler: """Test cases for Athena statement compiler functionality.""" diff --git a/tests/pyathena/test_converter.py b/tests/pyathena/test_converter.py index 667d9509..74fb9d6f 100644 --- a/tests/pyathena/test_converter.py +++ b/tests/pyathena/test_converter.py @@ -1,8 +1,12 @@ +from datetime import datetime, timedelta, timezone + import pytest +from dateutil.tz import gettz from pyathena.converter import ( DefaultTypeConverter, _to_array, + _to_datetime_with_tz, _to_map, _to_struct, ) @@ -533,3 +537,36 @@ def test_normalize_hive_syntax_mixed(self): type_hint="array", ) assert result == [{"a": 1, "b": "hello"}] + + +@pytest.mark.parametrize( + ("input_value", "expected"), + [ + (None, None), + ( + "2024-02-29 23:59:58.123 +05:30", + datetime( + 2024, 2, 29, 23, 59, 58, 123000, tzinfo=timezone(timedelta(hours=5, minutes=30)) + ), + ), + ( + "2024-02-29 23:59:58.123456 -08:00", + datetime(2024, 2, 29, 23, 59, 58, 123456, tzinfo=timezone(-timedelta(hours=8))), + ), + ( + "2024-02-29 23:59:58.123 UTC", + datetime(2024, 2, 29, 23, 59, 58, 123000, tzinfo=gettz("UTC")), + ), + ( + "2024-02-29 23:59:58.123 America/New_York", + datetime(2024, 2, 29, 23, 59, 58, 123000, tzinfo=gettz("America/New_York")), + ), + ], +) +def test_to_datetime_with_tz_offsets_and_zone_names(input_value, expected): + """Numeric UTC offsets give fixed-offset time zones; zone names keep their zone.""" + result = _to_datetime_with_tz(input_value) + assert result == expected + if expected is not None: + assert result.utcoffset() == expected.utcoffset() + assert result.tzinfo is not None diff --git a/tests/pyathena/test_cursor.py b/tests/pyathena/test_cursor.py index cea162fd..0ff324ac 100644 --- a/tests/pyathena/test_cursor.py +++ b/tests/pyathena/test_cursor.py @@ -92,6 +92,9 @@ def test_iterator(self, cursor): assert list(cursor) == [(1,)] pytest.raises(StopIteration, cursor.__next__) + # Cache hits are asserted in ENV.work_group: the default work group runs most test + # queries, which can push earlier executions out of the cache_size window. + @pytest.mark.parametrize("cursor", [{"work_group": ENV.work_group}], indirect=["cursor"]) def test_cache_size(self, cursor): # To test caching, we need to make sure the query is unique, otherwise # we might accidentally pick up the cache results from another CI run. @@ -127,6 +130,24 @@ def test_cache_size_with_work_group(self, cursor): assert first_query_id != second_query_id assert third_query_id in [first_query_id, second_query_id] + @pytest.mark.parametrize("cursor", [{"work_group": ENV.work_group}], indirect=["cursor"]) + def test_cache_size_with_qmark_parameters(self, cursor): + query = f"SELECT ? AS v -- {datetime.now(timezone.utc)!s}" + + cursor.execute(query, ["'1'"], paramstyle="qmark") + first_query_id = cursor.query_id + + # Different parameters must not reuse the earlier execution (#941). + cursor.execute(query, ["'2'"], paramstyle="qmark", cache_size=100) + assert cursor.query_id != first_query_id + assert cursor.fetchall() == [("2",)] + + # Athena does not return the parameters of earlier executions, + # so even the same parameters run again. + cursor.execute(query, ["'1'"], paramstyle="qmark", cache_size=100) + assert cursor.query_id != first_query_id + assert cursor.fetchall() == [("1",)] + def test_cache_expiration_time(self, cursor): query = f"SELECT * FROM one_row -- {datetime.now(timezone.utc)!s}" @@ -142,6 +163,7 @@ def test_cache_expiration_time(self, cursor): assert query_id_1 != query_id_2 assert query_id_3 in [query_id_1, query_id_2] + @pytest.mark.parametrize("cursor", [{"work_group": ENV.work_group}], indirect=["cursor"]) def test_cache_expiration_time_with_cache_size(self, cursor): # Cache miss query = f"SELECT * FROM one_row -- {datetime.now(timezone.utc)!s}" @@ -1064,6 +1086,32 @@ def test_execute_internal_legacy_kwargs_passthrough(self): cache_expiration_time=100, ) + def test_execute_qmark_parameters_skip_cache(self): + """A qmark query with parameters never searches the cache (no AWS, #941).""" + cursor = Cursor.__new__(Cursor) # bypass __init__ to avoid AWS calls + cursor._connection = MagicMock() + cursor._connection.client.start_query_execution.return_value = { + "QueryExecutionId": "test_query_id" + } + cursor._retry_config = RetryConfig() + cursor._kill_on_interrupt = True + + with ( + patch.object( + Cursor, + "_build_start_query_execution_request", + return_value={"ExecutionParameters": ["'1'"]}, + ) as request_mock, + patch.object(Cursor, "_find_previous_query_id", return_value="cached") as cache_mock, + ): + query_id = cursor._execute( + "SELECT ?", ["'1'"], paramstyle="qmark", cache_size=10, cache_expiration_time=100 + ) + + assert query_id == "test_query_id" + assert request_mock.call_args.kwargs["execution_parameters"] == ["'1'"] + cache_mock.assert_not_called() + def test_connection_level_callback(self): """Test connection-level default callback.""" callback_results = []