diff --git a/AGENTS.md b/AGENTS.md index d87ab8051..51c300485 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -134,4 +134,5 @@ Versions are derived from git tags via `hatch-vcs` — never edit `pyathena/_ver ### Google-style Docstrings -Use Google-style docstrings for public methods. See existing code for examples. +Follow the [docstring rules](docs/contributing.md#write-docstrings), which `just lint` checks with ruff's pydocstyle rules for `pyathena/`. +Mark overrides with `pyathena.util.override` instead of repeating the base method's docstring. diff --git a/docs/contributing.md b/docs/contributing.md index c0fbac149..320147260 100644 --- a/docs/contributing.md +++ b/docs/contributing.md @@ -71,6 +71,23 @@ External-fork pull requests must not be run in the project's AWS integration CI. Do not submit an untested change expecting a maintainer to approve an AWS CI run to validate it. Checks that need no AWS access may still run, but their success does not establish integration coverage. +## Write docstrings + +Code under `pyathena/` uses [Google-style docstrings](https://google.github.io/styleguide/pyguide.html#38-comments-and-docstrings). +`just lint` checks them with ruff's pydocstyle rules. + +- Modules, packages, public classes, and public functions and methods have a docstring. + Describe arguments, return values, and raised exceptions in `Args:`, `Returns:`, and `Raises:` sections. +- `__init__` describes the constructor arguments in its `Args:` section. +- A property getter has a one-line docstring; its setter needs none. +- Magic methods such as `__enter__` and `__iter__` need no docstring. +- A method that overrides a base class method is decorated with `override` from `pyathena.util`. + It needs a docstring only when its behavior differs from the base method; otherwise the API reference shows the base method's docstring. + mypy reports a missing decorator, except on an unannotated property. +- Overrides of fsspec methods are not decorated, because fsspec has no type information, so they need a docstring. +- New or changed private functions and methods use the same style. + ruff does not require them to have a docstring, but checks the docstrings they have. + ## Open a pull request Open a draft pull request with the repository's template completed. diff --git a/docs/usage.md b/docs/usage.md index 885ef3e1e..f62a83575 100644 --- a/docs/usage.md +++ b/docs/usage.md @@ -320,6 +320,7 @@ Result columns from a passthrough query come back with the source system's types PyAthena provides a callback mechanism that allows you to get immediate access to the query ID as soon as the `start_query_execution` API call is made, before waiting for query completion. This is useful for monitoring, logging, or cancelling long-running queries from another thread. +When `cache_size` finds a reusable query, no new query starts, and the callback receives the reused query's ID. The `on_start_query_execution` callback can be configured at both the connection level and the execute level. When both are set, both callbacks will be invoked. diff --git a/pyathena/__init__.py b/pyathena/__init__.py index 1e73d195d..504105113 100644 --- a/pyathena/__init__.py +++ b/pyathena/__init__.py @@ -1,3 +1,5 @@ +"""DB API 2.0 interface to Amazon Athena: ``connect()``, ``aio_connect()``, and type objects.""" + from __future__ import annotations import datetime @@ -30,7 +32,7 @@ class DBAPITypeObject(frozenset[str]): - """Type Objects and Constructors + """A DB API type object that compares equal to each of its Athena type names. https://www.python.org/dev/peps/pep-0249/#type-objects-and-constructors """ @@ -89,6 +91,8 @@ def connect(*args, **kwargs) -> Connection[Any]: SQL queries. Args: + *args: Positional arguments passed to the Connection constructor, in the + order of its parameters (``s3_staging_dir``, ``region_name``, ...). s3_staging_dir: S3 location to store query results. Required if not using workgroups or if the workgroup doesn't have a result location. Pass an empty string to explicitly disable S3 staging and skip @@ -144,6 +148,8 @@ async def aio_connect(*args, **kwargs) -> AioConnection: and API calls, keeping the event loop free. Args: + *args: Forwarded to ``AioConnection.create()``, which accepts keyword + arguments only. **kwargs: Arguments forwarded to ``AioConnection.create()``. See :func:`connect` for the full list of supported arguments. diff --git a/pyathena/aio/__init__.py b/pyathena/aio/__init__.py index e69de29bb..8a96ff3fa 100644 --- a/pyathena/aio/__init__.py +++ b/pyathena/aio/__init__.py @@ -0,0 +1,8 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT + +"""Native asyncio connections and cursors for Amazon Athena.""" diff --git a/pyathena/aio/arrow/__init__.py b/pyathena/aio/arrow/__init__.py index e69de29bb..46ae24237 100644 --- a/pyathena/aio/arrow/__init__.py +++ b/pyathena/aio/arrow/__init__.py @@ -0,0 +1,8 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT + +"""Native asyncio cursor that returns Athena query results as Apache Arrow tables.""" diff --git a/pyathena/aio/arrow/cursor.py b/pyathena/aio/arrow/cursor.py index 2932b24fd..f1b4aa71d 100644 --- a/pyathena/aio/arrow/cursor.py +++ b/pyathena/aio/arrow/cursor.py @@ -1,3 +1,5 @@ +"""Native asyncio cursor that returns Athena query results as Apache Arrow tables.""" + from __future__ import annotations import asyncio @@ -54,6 +56,28 @@ def __init__( request_timeout: float | None = None, **kwargs, ) -> None: + """Initialize an AioArrowCursor. + + Args: + s3_staging_dir: S3 location for query results. + schema_name: Default schema name. + catalog_name: Default catalog name. + work_group: Athena workgroup name. + poll_interval: Query status polling interval in seconds. + encryption_option: S3 encryption option for query results. + kms_key: KMS key for encrypting query results. + kill_on_interrupt: Cancel the query when the task is cancelled while + ``execute()`` starts or waits for the query. + unload: Whether to wrap queries in ``UNLOAD`` and read the Parquet output. + result_reuse_enable: Whether to enable Athena query result reuse. + result_reuse_minutes: Maximum age of a reused query result in minutes. + connect_timeout: Connection timeout in seconds of the pyarrow S3 filesystem + that reads the results. If None, the pyarrow default is used. + request_timeout: Request timeout in seconds of the pyarrow S3 filesystem + that reads the results. If None, the pyarrow default is used. + **kwargs: Other cursor arguments, such as ``connection`` and ``arraysize``, + passed to the parent ``__init__``. + """ super().__init__( s3_staging_dir=s3_staging_dir, schema_name=schema_name, @@ -111,7 +135,9 @@ async def execute( result_reuse_enable: Enable Athena result reuse for this query. result_reuse_minutes: Minutes to reuse cached results. paramstyle: Parameter style ('qmark' or 'pyformat'). - on_start_query_execution: Callback called when query starts. + on_start_query_execution: Callback invoked with the query ID before ``execute()`` + waits for the query: after the ``StartQueryExecution`` call, or after a + reusable query ID is found through ``cache_size``. result_set_type_hints: Optional dictionary mapping column names to Athena DDL type signatures for precise type conversion within complex types. diff --git a/pyathena/aio/common.py b/pyathena/aio/common.py index d143bcd0a..7e581ab7f 100644 --- a/pyathena/aio/common.py +++ b/pyathena/aio/common.py @@ -1,3 +1,5 @@ +"""Asyncio base cursor and the fetch mixin shared by the asyncio SQL cursors.""" + from __future__ import annotations import asyncio diff --git a/pyathena/aio/connection.py b/pyathena/aio/connection.py index b775b6901..91fa6e15d 100644 --- a/pyathena/aio/connection.py +++ b/pyathena/aio/connection.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""Asyncio-aware connection to Amazon Athena.""" + from __future__ import annotations import asyncio @@ -31,6 +33,12 @@ class AioConnection(Connection[AioCursor]): """ def __init__(self, **kwargs: Any) -> None: + """Initialize the connection with ``AioCursor`` as the default cursor class. + + Args: + **kwargs: Arguments forwarded to ``Connection.__init__``. If they do not + include ``cursor_class``, it is set to ``AioCursor``. + """ if "cursor_class" not in kwargs: kwargs["cursor_class"] = AioCursor super().__init__(**kwargs) diff --git a/pyathena/aio/cursor.py b/pyathena/aio/cursor.py index 38ef885fd..89e7c5aa6 100644 --- a/pyathena/aio/cursor.py +++ b/pyathena/aio/cursor.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""Native asyncio cursors that return rows as tuples or dictionaries.""" + from __future__ import annotations import logging @@ -50,6 +52,24 @@ def __init__( result_reuse_minutes: int = CursorIterator.DEFAULT_RESULT_REUSE_MINUTES, **kwargs, ) -> None: + """Initialize an AioCursor. + + Args: + s3_staging_dir: S3 location for query results. + schema_name: Default schema name. + catalog_name: Default catalog name. + work_group: Athena workgroup name. + poll_interval: Query status polling interval in seconds. + encryption_option: S3 encryption option (SSE_S3, SSE_KMS, CSE_KMS). + kms_key: KMS key for encryption. + kill_on_interrupt: Cancel the query when the task is cancelled while + ``execute()`` starts or waits for the query. + result_reuse_enable: Enable Athena query result reuse. + result_reuse_minutes: Maximum age in minutes of a reused result. + **kwargs: Arguments forwarded to ``WithResultSet.__init__`` and + ``AioBaseCursor.__init__``, such as ``arraysize``, ``connection``, + ``converter``, ``formatter``, and ``retry_config``. + """ super().__init__( s3_staging_dir=s3_staging_dir, schema_name=schema_name, @@ -104,12 +124,14 @@ async def execute( parameters: Query parameters (optional). work_group: Athena workgroup to use (optional). s3_staging_dir: S3 location for query results (optional). - cache_size: Query result cache size (optional). + cache_size: Number of queries to check for result caching (optional). cache_expiration_time: Cache expiration time in seconds (optional). result_reuse_enable: Enable result reuse (optional). result_reuse_minutes: Result reuse duration in minutes (optional). paramstyle: Parameter style to use (optional). - on_start_query_execution: Callback called when query starts. + on_start_query_execution: Callback invoked with the query ID before ``execute()`` + waits for the query: after the ``StartQueryExecution`` call, or after a + reusable query ID is found through ``cache_size``. result_set_type_hints: Optional dictionary mapping column names to Athena DDL type signatures for precise type conversion within complex types. @@ -225,6 +247,13 @@ class AioDictCursor(AioCursor): """ def __init__(self, **kwargs) -> None: + """Initialize an AioDictCursor. + + Args: + **kwargs: Arguments forwarded to ``AioCursor.__init__``. If they include + ``dict_type``, it is also assigned to the class attribute + ``AthenaAioDictResultSet.dict_type``, the type used to build each row. + """ super().__init__(**kwargs) self._result_set_class = AthenaAioDictResultSet if "dict_type" in kwargs: diff --git a/pyathena/aio/pandas/__init__.py b/pyathena/aio/pandas/__init__.py index e69de29bb..b5af65223 100644 --- a/pyathena/aio/pandas/__init__.py +++ b/pyathena/aio/pandas/__init__.py @@ -0,0 +1,8 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT + +"""Native asyncio cursor that returns Athena query results as pandas DataFrames.""" diff --git a/pyathena/aio/pandas/cursor.py b/pyathena/aio/pandas/cursor.py index 863fe5168..967ec3e4f 100644 --- a/pyathena/aio/pandas/cursor.py +++ b/pyathena/aio/pandas/cursor.py @@ -1,3 +1,5 @@ +"""Native asyncio cursor that returns Athena query results as pandas DataFrames.""" + from __future__ import annotations import asyncio @@ -63,6 +65,32 @@ def __init__( auto_optimize_chunksize: bool = False, **kwargs, ) -> None: + """Initialize an AioPandasCursor. + + Args: + s3_staging_dir: S3 location for query results. + schema_name: Default schema name. + catalog_name: Default catalog name. + work_group: Athena workgroup name. + poll_interval: Query status polling interval in seconds. + encryption_option: S3 encryption option for query results. + kms_key: KMS key for encrypting query results. + kill_on_interrupt: Cancel the query when the task is cancelled while + ``execute()`` starts or waits for the query. + unload: Whether to wrap queries in ``UNLOAD`` and read the Parquet output. + engine: Parsing engine (``auto``, ``c``, ``python``, or ``pyarrow``). + chunksize: Number of rows per DataFrame chunk when reading CSV results. If set, + it takes precedence over ``auto_optimize_chunksize``. + block_size: Default block size of the S3 filesystem that reads the results. + cache_type: Default cache type of the S3 filesystem that reads the results. + max_workers: Maximum number of workers of the S3 filesystem. + result_reuse_enable: Whether to enable Athena query result reuse. + result_reuse_minutes: Maximum age of a reused query result in minutes. + auto_optimize_chunksize: Whether to choose a chunk size from the size of the + CSV result file when ``chunksize`` is None. + **kwargs: Other cursor arguments, such as ``connection`` and ``arraysize``, + passed to the parent ``__init__``. + """ super().__init__( s3_staging_dir=s3_staging_dir, schema_name=schema_name, @@ -130,7 +158,9 @@ async def execute( keep_default_na: Whether to keep default pandas NA values. na_values: Additional values to treat as NA. quoting: CSV quoting behavior (pandas csv.QUOTE_* constants). - on_start_query_execution: Callback called when query starts. + on_start_query_execution: Callback invoked with the query ID before ``execute()`` + waits for the query: after the ``StartQueryExecution`` call, or after a + reusable query ID is found through ``cache_size``. result_set_type_hints: Optional dictionary mapping column names to Athena DDL type signatures for precise type conversion within complex types. diff --git a/pyathena/aio/polars/__init__.py b/pyathena/aio/polars/__init__.py index e69de29bb..f194bd646 100644 --- a/pyathena/aio/polars/__init__.py +++ b/pyathena/aio/polars/__init__.py @@ -0,0 +1,8 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT + +"""Native asyncio cursor that returns Athena query results as Polars DataFrames.""" diff --git a/pyathena/aio/polars/cursor.py b/pyathena/aio/polars/cursor.py index ffe330187..8338c9d61 100644 --- a/pyathena/aio/polars/cursor.py +++ b/pyathena/aio/polars/cursor.py @@ -1,3 +1,5 @@ +"""Native asyncio cursor that returns Athena query results as Polars DataFrames.""" + from __future__ import annotations import asyncio @@ -58,6 +60,29 @@ def __init__( chunksize: int | None = None, **kwargs, ) -> None: + """Initialize an AioPolarsCursor. + + Args: + s3_staging_dir: S3 location for query results. + schema_name: Default schema name. + catalog_name: Default catalog name. + work_group: Athena workgroup name. + poll_interval: Query status polling interval in seconds. + encryption_option: S3 encryption option for query results. + kms_key: KMS key for encrypting query results. + kill_on_interrupt: Cancel the query when the task is cancelled while + ``execute()`` starts or waits for the query. + unload: Whether to wrap queries in ``UNLOAD`` and read the Parquet output. + result_reuse_enable: Whether to enable Athena query result reuse. + result_reuse_minutes: Maximum age of a reused query result in minutes. + block_size: Default block size of the S3 filesystem that reads the results. + cache_type: Default cache type of the S3 filesystem that reads the results. + max_workers: Maximum number of workers of the S3 filesystem. + chunksize: Number of rows per chunk. If set, result files in S3 are read + lazily in chunks of this size. + **kwargs: Other cursor arguments, such as ``connection`` and ``arraysize``, + passed to the parent ``__init__``. + """ super().__init__( s3_staging_dir=s3_staging_dir, schema_name=schema_name, @@ -117,7 +142,9 @@ async def execute( result_reuse_enable: Enable Athena result reuse for this query. result_reuse_minutes: Minutes to reuse cached results. paramstyle: Parameter style ('qmark' or 'pyformat'). - on_start_query_execution: Callback called when query starts. + on_start_query_execution: Callback invoked with the query ID before ``execute()`` + waits for the query: after the ``StartQueryExecution`` call, or after a + reusable query ID is found through ``cache_size``. result_set_type_hints: Optional dictionary mapping column names to Athena DDL type signatures for precise type conversion within complex types. diff --git a/pyathena/aio/result_set.py b/pyathena/aio/result_set.py index 54c4df592..5a628da7e 100644 --- a/pyathena/aio/result_set.py +++ b/pyathena/aio/result_set.py @@ -1,3 +1,5 @@ +"""Asyncio result sets that fetch Athena query results with ``GetQueryResults``.""" + from __future__ import annotations import logging @@ -39,6 +41,21 @@ def __init__( retry_config: RetryConfig, result_set_type_hints: dict[str | int, str] | None = None, ) -> None: + """Initialize the result set without fetching rows; ``create()`` fetches the first page. + + Args: + connection: The connection that ran the query. + converter: The converter for result values. + query_execution: The query execution whose results to read. + arraysize: The number of rows per ``GetQueryResults`` page and the default + ``fetchmany()`` size. + retry_config: The retry configuration for API calls. + result_set_type_hints: Athena type signatures for complex-type columns, + keyed by column name (case-insensitive) or zero-based column index. + + Raises: + ProgrammingError: If ``query_execution`` is not given. + """ super().__init__( connection=connection, converter=converter, diff --git a/pyathena/aio/s3fs/__init__.py b/pyathena/aio/s3fs/__init__.py index e69de29bb..795ab0a09 100644 --- a/pyathena/aio/s3fs/__init__.py +++ b/pyathena/aio/s3fs/__init__.py @@ -0,0 +1,8 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT + +"""Native asyncio cursor that reads Athena CSV query results through ``AioS3FileSystem``.""" diff --git a/pyathena/aio/s3fs/cursor.py b/pyathena/aio/s3fs/cursor.py index 5850ac460..91a176170 100644 --- a/pyathena/aio/s3fs/cursor.py +++ b/pyathena/aio/s3fs/cursor.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""Native asyncio cursor that reads Athena CSV query results through ``AioS3FileSystem``.""" + from __future__ import annotations import asyncio @@ -55,6 +57,26 @@ def __init__( csv_reader: CSVReaderType | None = None, **kwargs, ) -> None: + """Initialize an AioS3FSCursor. + + Args: + s3_staging_dir: S3 location for query results. + schema_name: Default schema name. + catalog_name: Default catalog name. + work_group: Athena workgroup name. + poll_interval: Query status polling interval in seconds. + encryption_option: S3 encryption option for query results. + kms_key: KMS key for encrypting query results. + kill_on_interrupt: Cancel the query when the task is cancelled while + ``execute()`` starts or waits for the query. + result_reuse_enable: Whether to enable Athena query result reuse. + result_reuse_minutes: Maximum age of a reused query result in minutes. + csv_reader: CSV reader class for parsing the result files. If None, + ``AthenaCSVReader`` is used, which distinguishes NULL from empty + strings. ``DefaultCSVReader`` reads both as empty strings. + **kwargs: Other cursor arguments, such as ``connection`` and ``arraysize``, + passed to the parent ``__init__``. + """ super().__init__( s3_staging_dir=s3_staging_dir, schema_name=schema_name, @@ -116,7 +138,9 @@ async def execute( result_reuse_enable: Enable Athena result reuse for this query. result_reuse_minutes: Minutes to reuse cached results. paramstyle: Parameter style ('qmark' or 'pyformat'). - on_start_query_execution: Callback called when query starts. + on_start_query_execution: Callback invoked with the query ID before ``execute()`` + waits for the query: after the ``StartQueryExecution`` call, or after a + reusable query ID is found through ``cache_size``. result_set_type_hints: Optional dictionary mapping column names to Athena DDL type signatures for precise type conversion within complex types. diff --git a/pyathena/aio/spark/__init__.py b/pyathena/aio/spark/__init__.py index e69de29bb..2813e88e0 100644 --- a/pyathena/aio/spark/__init__.py +++ b/pyathena/aio/spark/__init__.py @@ -0,0 +1,8 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT + +"""Native asyncio cursor that runs PySpark code in Athena for Apache Spark sessions.""" diff --git a/pyathena/aio/spark/cursor.py b/pyathena/aio/spark/cursor.py index b6091666a..effeceb02 100644 --- a/pyathena/aio/spark/cursor.py +++ b/pyathena/aio/spark/cursor.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""Native asyncio cursor that runs PySpark code in an Athena for Apache Spark session.""" + from __future__ import annotations import asyncio diff --git a/pyathena/aio/sqlalchemy/__init__.py b/pyathena/aio/sqlalchemy/__init__.py index e69de29bb..253b7ce16 100644 --- a/pyathena/aio/sqlalchemy/__init__.py +++ b/pyathena/aio/sqlalchemy/__init__.py @@ -0,0 +1,8 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT + +"""Async SQLAlchemy dialects for Amazon Athena.""" diff --git a/pyathena/aio/sqlalchemy/arrow.py b/pyathena/aio/sqlalchemy/arrow.py index f65671137..1dd574b0e 100644 --- a/pyathena/aio/sqlalchemy/arrow.py +++ b/pyathena/aio/sqlalchemy/arrow.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""Async SQLAlchemy dialect for Athena that uses ``AioArrowCursor``.""" + from typing import TYPE_CHECKING from pyathena.aio.sqlalchemy.base import AthenaAioDialect diff --git a/pyathena/aio/sqlalchemy/base.py b/pyathena/aio/sqlalchemy/base.py index 6c65dc3b6..abb600c53 100644 --- a/pyathena/aio/sqlalchemy/base.py +++ b/pyathena/aio/sqlalchemy/base.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""Async SQLAlchemy dialect base and DBAPI adapters for PyAthena asyncio cursors.""" + from __future__ import annotations from collections import deque @@ -55,22 +57,45 @@ class AsyncAdapt_pyathena_cursor: __slots__ = ("_cursor", "_rows") def __init__(self, cursor: Any) -> None: + """Initialize the adapter around an async cursor. + + Args: + cursor: The async PyAthena cursor to wrap. + """ self._cursor = cursor self._rows: deque[Any] = deque() @property def description(self) -> Any: + """The ``description`` of the wrapped cursor.""" return self._cursor.description @property def rowcount(self) -> int: + """The ``rowcount`` of the wrapped cursor.""" return self._cursor.rowcount # type: ignore[no-any-return] def close(self) -> None: + """Close the wrapped cursor and discard any buffered rows.""" self._cursor.close() self._rows.clear() def execute(self, operation: str, parameters: Any = None, **kwargs: Any) -> Any: + """Execute a statement and buffer all of its result rows. + + When the statement produces a result set (the cursor has a + ``description``), every row is fetched and buffered so that the fetch + methods can return rows without awaiting. + + Args: + operation: The SQL statement to execute. + parameters: Parameters to bind to the statement. + **kwargs: Additional keyword arguments forwarded to the wrapped + cursor's ``execute()``. + + Returns: + The value returned by the wrapped cursor's ``execute()``. + """ result = await_only(self._cursor.execute(operation, parameters, **kwargs)) if self._cursor.description: self._rows = deque(await_only(self._cursor.fetchall())) @@ -84,25 +109,59 @@ def executemany( seq_of_parameters: list[dict[str, Any] | list[str] | None], **kwargs: Any, ) -> None: + """Execute a statement once for each parameter set. + + Any buffered rows are discarded first. + + Args: + operation: The SQL statement to execute. + seq_of_parameters: The parameter sets to bind, one per execution. + **kwargs: Additional keyword arguments forwarded to the wrapped + cursor's ``executemany()``. + """ self._rows.clear() await_only(self._cursor.executemany(operation, seq_of_parameters, **kwargs)) def fetchone(self) -> Any: + """Fetch the next buffered row. + + Returns: + The next row, or ``None`` when no rows remain. + """ if self._rows: return self._rows.popleft() return None def fetchmany(self, size: int | None = None) -> Any: + """Fetch up to ``size`` buffered rows. + + Args: + size: Maximum number of rows to fetch. If ``None``, the wrapped + cursor's ``arraysize`` is used, or 1 if it has none. + + Returns: + A list of rows, empty when no rows remain. + """ if size is None: size = self._cursor.arraysize if hasattr(self._cursor, "arraysize") else 1 return [self._rows.popleft() for _ in range(min(size, len(self._rows)))] def fetchall(self) -> Any: + """Fetch all remaining buffered rows. + + Returns: + A list of the remaining rows. + """ items = list(self._rows) self._rows.clear() return items def setinputsizes(self, sizes: Any) -> None: + """Forward ``sizes`` to the wrapped cursor's ``setinputsizes()``. + + Args: + sizes: Sequence of parameter types or sizes. + """ self._cursor.setinputsizes(sizes) async def _async_soft_close(self) -> None: @@ -110,12 +169,39 @@ async def _async_soft_close(self) -> None: # PyAthena-specific methods used by AthenaDialect reflection def list_databases(self, *args: Any, **kwargs: Any) -> Any: + """Await the wrapped cursor's ``list_databases()`` and return its result. + + Args: + *args: Positional arguments forwarded to ``list_databases()``. + **kwargs: Keyword arguments forwarded to ``list_databases()``. + + Returns: + The result of the wrapped cursor's ``list_databases()``. + """ return await_only(self._cursor.list_databases(*args, **kwargs)) def get_table_metadata(self, *args: Any, **kwargs: Any) -> Any: + """Await the wrapped cursor's ``get_table_metadata()`` and return its result. + + Args: + *args: Positional arguments forwarded to ``get_table_metadata()``. + **kwargs: Keyword arguments forwarded to ``get_table_metadata()``. + + Returns: + The result of the wrapped cursor's ``get_table_metadata()``. + """ return await_only(self._cursor.get_table_metadata(*args, **kwargs)) def list_table_metadata(self, *args: Any, **kwargs: Any) -> Any: + """Await the wrapped cursor's ``list_table_metadata()`` and return its result. + + Args: + *args: Positional arguments forwarded to ``list_table_metadata()``. + **kwargs: Keyword arguments forwarded to ``list_table_metadata()``. + + Returns: + The result of the wrapped cursor's ``list_table_metadata()``. + """ return await_only(self._cursor.list_table_metadata(*args, **kwargs)) def __enter__(self) -> AsyncAdapt_pyathena_cursor: @@ -136,6 +222,12 @@ class AsyncAdapt_pyathena_connection(AdaptedConnection): __slots__ = ("_connection", "dbapi") def __init__(self, dbapi: AsyncAdapt_pyathena_dbapi, connection: AioConnection) -> None: + """Initialize the adapted connection. + + Args: + dbapi: The adapted DBAPI module that created this connection. + connection: The ``AioConnection`` to wrap. + """ self.dbapi = dbapi self._connection = connection # type: ignore[assignment] @@ -146,34 +238,54 @@ def driver_connection(self) -> AioConnection: @property def catalog_name(self) -> str | None: + """The catalog name of the wrapped connection.""" return self._connection.catalog_name # type: ignore[no-any-return] @property def schema_name(self) -> str | None: + """The schema name of the wrapped connection.""" return self._connection.schema_name # type: ignore[no-any-return] @property def cursor_kwargs(self) -> dict[str, Any]: + """The default cursor keyword arguments of the wrapped connection.""" return self._connection.cursor_kwargs # type: ignore[no-any-return] @property def retry_config(self) -> RetryConfig: + """The retry configuration of the wrapped connection.""" return self._connection.retry_config # type: ignore[no-any-return] def cursor(self, cursor: Any = None, **kwargs: Any) -> AsyncAdapt_pyathena_cursor: + """Create an async cursor on the wrapped connection and adapt it. + + A synchronous cursor class that has a registered async counterpart is + replaced with that counterpart; any other value is passed through. + + Args: + cursor: The cursor class to create, or ``None`` for the + connection's default cursor class. + **kwargs: Keyword arguments forwarded to the wrapped connection's + ``cursor()``. + + Returns: + The created cursor wrapped in ``AsyncAdapt_pyathena_cursor``. + """ # The shared dialect names a cursor class in its synchronous form; this # connection can only drive the async counterpart. raw_cursor = self._connection.cursor(_ASYNC_CURSOR_CLASSES.get(cursor, cursor), **kwargs) return AsyncAdapt_pyathena_cursor(raw_cursor) def close(self) -> None: + """Close the wrapped connection.""" self._connection.close() def commit(self) -> None: + """Call ``commit()`` on the wrapped connection.""" self._connection.commit() # type: ignore[unused-coroutine] def rollback(self) -> None: - pass + """Do nothing, because Athena does not support transactions.""" class AsyncAdapt_pyathena_dbapi: @@ -202,6 +314,14 @@ class AsyncAdapt_pyathena_dbapi: NotSupportedError = NotSupportedError def connect(self, **kwargs: Any) -> AsyncAdapt_pyathena_connection: + """Create an ``AioConnection`` and wrap it in an adapted connection. + + Args: + **kwargs: Keyword arguments forwarded to ``AioConnection.create()``. + + Returns: + The new connection wrapped in ``AsyncAdapt_pyathena_connection``. + """ connection = await_only(AioConnection.create(**kwargs)) return AsyncAdapt_pyathena_connection(self, connection) diff --git a/pyathena/aio/sqlalchemy/pandas.py b/pyathena/aio/sqlalchemy/pandas.py index 6f10eaf6d..f5ebbb2da 100644 --- a/pyathena/aio/sqlalchemy/pandas.py +++ b/pyathena/aio/sqlalchemy/pandas.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""Async SQLAlchemy dialect for Athena that uses ``AioPandasCursor``.""" + from typing import TYPE_CHECKING from pyathena.aio.sqlalchemy.base import AthenaAioDialect diff --git a/pyathena/aio/sqlalchemy/polars.py b/pyathena/aio/sqlalchemy/polars.py index e4f0218f7..7008a1be8 100644 --- a/pyathena/aio/sqlalchemy/polars.py +++ b/pyathena/aio/sqlalchemy/polars.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""Async SQLAlchemy dialect for Athena that uses ``AioPolarsCursor``.""" + from typing import TYPE_CHECKING from pyathena.aio.sqlalchemy.base import AthenaAioDialect diff --git a/pyathena/aio/sqlalchemy/rest.py b/pyathena/aio/sqlalchemy/rest.py index 74d3adb8f..8e7eb7363 100644 --- a/pyathena/aio/sqlalchemy/rest.py +++ b/pyathena/aio/sqlalchemy/rest.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""Async SQLAlchemy dialect for Athena that uses ``AioCursor``.""" + from typing import TYPE_CHECKING from pyathena.aio.sqlalchemy.base import AthenaAioDialect diff --git a/pyathena/aio/sqlalchemy/s3fs.py b/pyathena/aio/sqlalchemy/s3fs.py index 39d8e034b..bd55d59b1 100644 --- a/pyathena/aio/sqlalchemy/s3fs.py +++ b/pyathena/aio/sqlalchemy/s3fs.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""Async SQLAlchemy dialect for Athena that uses ``AioS3FSCursor``.""" + from typing import TYPE_CHECKING from pyathena.aio.sqlalchemy.base import AthenaAioDialect diff --git a/pyathena/aio/util.py b/pyathena/aio/util.py index 1e058fbbf..c9b6a69af 100644 --- a/pyathena/aio/util.py +++ b/pyathena/aio/util.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""Asyncio helpers for retrying AWS API calls.""" + from __future__ import annotations import asyncio diff --git a/pyathena/arrow/__init__.py b/pyathena/arrow/__init__.py index e69de29bb..5685cc76e 100644 --- a/pyathena/arrow/__init__.py +++ b/pyathena/arrow/__init__.py @@ -0,0 +1,8 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT + +"""Cursors that return Athena query results as Apache Arrow tables.""" diff --git a/pyathena/arrow/async_cursor.py b/pyathena/arrow/async_cursor.py index ab2a18eef..a9f5d6766 100644 --- a/pyathena/arrow/async_cursor.py +++ b/pyathena/arrow/async_cursor.py @@ -1,3 +1,5 @@ +"""Asynchronous cursor that returns Athena query results as Apache Arrow tables.""" + from __future__ import annotations import logging diff --git a/pyathena/arrow/converter.py b/pyathena/arrow/converter.py index cdb475eb2..ea6482557 100644 --- a/pyathena/arrow/converter.py +++ b/pyathena/arrow/converter.py @@ -1,3 +1,5 @@ +"""Type converters for Apache Arrow cursor results.""" + from __future__ import annotations import logging @@ -56,6 +58,7 @@ class DefaultArrowTypeConverter(Converter): """ def __init__(self) -> None: + """Initialize the converter with the default Arrow conversion functions and types.""" super().__init__( mappings=deepcopy(_DEFAULT_ARROW_CONVERTERS), default=_to_default, @@ -111,6 +114,7 @@ class DefaultArrowUnloadTypeConverter(Converter): """ def __init__(self) -> None: + """Initialize the converter with no type mappings.""" super().__init__( mappings={}, default=_to_default, diff --git a/pyathena/arrow/cursor.py b/pyathena/arrow/cursor.py index 757d98da8..ded70ee77 100644 --- a/pyathena/arrow/cursor.py +++ b/pyathena/arrow/cursor.py @@ -1,3 +1,5 @@ +"""Cursor that returns Athena query results as Apache Arrow tables.""" + from __future__ import annotations import logging @@ -159,7 +161,9 @@ def execute( result_reuse_enable: Enable Athena result reuse for this query. result_reuse_minutes: Minutes to reuse cached results. paramstyle: Parameter style ('qmark' or 'pyformat'). - on_start_query_execution: Callback called when query starts. + on_start_query_execution: Callback invoked with the query ID before ``execute()`` + waits for the query: after the ``StartQueryExecution`` call, or after a + reusable query ID is found through ``cache_size``. result_set_type_hints: Optional dictionary mapping column names to Athena DDL type signatures for precise type conversion within complex types. diff --git a/pyathena/arrow/result_set.py b/pyathena/arrow/result_set.py index 2b091fa6d..de9d4c66d 100644 --- a/pyathena/arrow/result_set.py +++ b/pyathena/arrow/result_set.py @@ -1,3 +1,5 @@ +"""Result set that reads Athena query results into Apache Arrow Tables.""" + from __future__ import annotations import logging @@ -94,6 +96,31 @@ def __init__( result_set_type_hints: dict[str | int, str] | None = None, **kwargs, ) -> None: + """Initialize the result set and load the query results into an Arrow Table. + + Args: + connection: The connection that ran the query. + converter: The converter for result values. + query_execution: The query execution whose results to read. + arraysize: The default ``fetchmany()`` size and the maximum number of rows per + record batch that the fetch methods read from the table. + retry_config: The retry configuration for API calls. + block_size: The block size in bytes for reading CSV results. If not set, + ``DEFAULT_BLOCK_SIZE`` is used. + unload: Whether the query is an ``UNLOAD`` whose Parquet output is read + instead of the CSV results. + unload_location: The S3 location of the ``UNLOAD`` output. If None, it is + derived from the first file in the data manifest. + connect_timeout: The connect timeout in seconds for the pyarrow S3 filesystem. + request_timeout: The request timeout in seconds for the pyarrow S3 filesystem. + result_set_type_hints: Athena type signatures for complex-type columns, + keyed by column name (case-insensitive) or zero-based column index. + **kwargs: Additional keyword arguments, stored but not used. + + Raises: + ProgrammingError: If ``query_execution`` is not given. + OperationalError: If reading the query results fails. + """ super().__init__( connection=connection, converter=converter, @@ -195,12 +222,14 @@ def _create_s3_file_system(self): @property def timestamp_parsers(self) -> list[str]: + """The timestamp formats for reading CSV results, starting with pyarrow's ``ISO8601``.""" from pyarrow.csv import ISO8601 return [ISO8601, *self._timestamp_parsers] @property def column_types(self) -> dict[str, type[Any]]: + """The converter's types for the result columns it maps, keyed by column name.""" description = self.description if self.description else [] return { d[0]: dtype @@ -210,6 +239,7 @@ def column_types(self) -> dict[str, type[Any]]: @property def converters(self) -> dict[str, Callable[[str | None], Any | None]]: + """The conversion functions for the result columns, keyed by column name.""" description = self.description if self.description else [] return {d[0]: self._converter.get(d[1]) for d in description} @@ -356,6 +386,11 @@ def _as_arrow_from_api(self, converter: Converter | None = None) -> Table: return pa.table(self._rows_to_columnar(rows, columns)) def as_arrow(self) -> Table: + """Return the query results as an Apache Arrow Table. + + Returns: + The Arrow Table that holds the query results. + """ return self._table def as_polars(self) -> pl.DataFrame: diff --git a/pyathena/async_cursor.py b/pyathena/async_cursor.py index 620df9d68..3363901c7 100644 --- a/pyathena/async_cursor.py +++ b/pyathena/async_cursor.py @@ -1,3 +1,5 @@ +"""Thread-pool cursors that run Athena queries concurrently and return futures.""" + from __future__ import annotations import logging @@ -68,6 +70,31 @@ def __init__( result_reuse_minutes: int = CursorIterator.DEFAULT_RESULT_REUSE_MINUTES, **kwargs, ) -> None: + """Initialize an AsyncCursor. + + Args: + s3_staging_dir: S3 location for query results. + schema_name: Default schema name. + catalog_name: Default catalog name. + work_group: Athena workgroup name. + poll_interval: Query status polling interval in seconds. + encryption_option: S3 encryption option (SSE_S3, SSE_KMS, CSE_KMS). + kms_key: KMS key for encryption. + kill_on_interrupt: Cancel a query whose start in ``execute()`` is interrupted by + ``KeyboardInterrupt``. Waiting runs on worker threads, which do not + receive the interrupt. + max_workers: Maximum number of threads in the cursor's thread pool. + arraysize: Default number of rows per ``fetchmany()`` call of the result + sets the cursor creates. + result_reuse_enable: Enable Athena query result reuse. + result_reuse_minutes: Maximum age in minutes of a reused result. + **kwargs: Arguments forwarded to ``BaseCursor.__init__``, such as + ``connection``, ``converter``, ``formatter``, and ``retry_config``. + + Raises: + ProgrammingError: If ``arraysize`` is not between 1 and + ``CursorIterator.DEFAULT_FETCH_SIZE``. + """ super().__init__( s3_staging_dir=s3_staging_dir, schema_name=schema_name, @@ -88,6 +115,7 @@ def __init__( @property def arraysize(self) -> int: + """The default number of rows per ``fetchmany()`` call of the result sets.""" return self._arraysize @arraysize.setter @@ -112,6 +140,16 @@ def _description( def description( self, query_id: str ) -> Future[list[tuple[str, str, None, None, int, int, str]] | None]: + """Get the column descriptions of a query's result set asynchronously. + + The future waits for the query to finish before it reads the result set. + + Args: + query_id: The Athena query execution ID. + + Returns: + Future object containing the DB API 2.0 column descriptions, or None. + """ return self._executor.submit(self._description, query_id) def query_execution(self, query_id: str) -> Future[AthenaQueryExecution]: @@ -190,7 +228,7 @@ def execute( parameters: Query parameters (optional). work_group: Athena workgroup to use (optional). s3_staging_dir: S3 location for query results (optional). - cache_size: Query result cache size in MB (optional). + cache_size: Number of queries to check for result caching (optional). cache_expiration_time: Cache expiration time in seconds (optional). result_reuse_enable: Enable result reuse for identical queries (optional). result_reuse_minutes: Result reuse duration in minutes (optional). @@ -298,6 +336,13 @@ class AsyncDictCursor(AsyncCursor): """ def __init__(self, **kwargs) -> None: + """Initialize an AsyncDictCursor. + + Args: + **kwargs: Arguments forwarded to ``AsyncCursor.__init__``. If they include + ``dict_type``, it is also assigned to the class attribute + ``AthenaDictResultSet.dict_type``, the type used to build each row. + """ super().__init__(**kwargs) self._result_set_class = AthenaDictResultSet if "dict_type" in kwargs: diff --git a/pyathena/common.py b/pyathena/common.py index 8797bfb9d..31c445db0 100644 --- a/pyathena/common.py +++ b/pyathena/common.py @@ -1,3 +1,5 @@ +"""Base classes and callback types shared by PyAthena cursors.""" + from __future__ import annotations import logging @@ -81,6 +83,16 @@ class CursorIterator(metaclass=ABCMeta): DEFAULT_RESULT_REUSE_MINUTES = 60 def __init__(self, **kwargs) -> None: + """Initialize the iterator with no current row and an unknown row count. + + Args: + **kwargs: Keyword arguments, of which only ``arraysize`` is used. If it is + absent, ``DEFAULT_FETCH_SIZE`` is used. + + Raises: + ProgrammingError: If ``arraysize`` is outside the range the + ``arraysize`` setter accepts. + """ super().__init__() self.arraysize: int = kwargs.get("arraysize", self.DEFAULT_FETCH_SIZE) self._rownumber: int | None = None @@ -88,6 +100,7 @@ def __init__(self, **kwargs) -> None: @property def arraysize(self) -> int: + """The default number of rows per ``fetchmany()`` call.""" return self._arraysize @arraysize.setter @@ -100,22 +113,27 @@ def arraysize(self, value: int) -> None: @property def rownumber(self) -> int | None: + """The zero-based index of the next row, or None if it is unknown.""" return self._rownumber @property def rowcount(self) -> int: + """The number of rows affected by the last operation, or -1 if it is unknown.""" return self._rowcount @abstractmethod def fetchone(self): + """Fetch the next row of the result.""" raise NotImplementedError # pragma: no cover @abstractmethod def fetchmany(self): + """Fetch the next set of rows of the result.""" raise NotImplementedError # pragma: no cover @abstractmethod def fetchall(self): + """Fetch all remaining rows of the result.""" raise NotImplementedError # pragma: no cover def __next__(self): @@ -195,6 +213,30 @@ def __init__( on_poll: OnPollCallback | None = None, **kwargs, ) -> None: + """Initialize the cursor with the settings it uses to run queries. + + Args: + connection: The connection that created the cursor. + converter: Converter for result values. + formatter: Formatter for query parameters. + retry_config: Retry configuration for API calls. + s3_staging_dir: S3 location for query results. + schema_name: Default schema name. + catalog_name: Default catalog name. + work_group: Athena workgroup name. + poll_interval: Query status polling interval in seconds. + encryption_option: S3 encryption option (SSE_S3, SSE_KMS, CSE_KMS). + kms_key: KMS key for encryption. + kill_on_interrupt: Cancel the execution when a ``KeyboardInterrupt`` interrupts + starting it or waiting for it. + result_reuse_enable: Enable Athena query result reuse. + result_reuse_minutes: Maximum age in minutes of a reused result. + on_start_query_execution: Callback invoked with each query ID before the cursor + waits for the query, by cursors whose ``execute()`` supports it. + on_poll: Callback invoked once per poll iteration with the current + execution object. + **kwargs: Ignored. + """ super().__init__() self._connection = connection self._converter = converter @@ -231,6 +273,7 @@ def get_default_converter(unload: bool = False) -> DefaultTypeConverter | Any: @property def connection(self) -> Connection[Any]: + """The connection that created this cursor.""" return self._connection def _build_start_query_execution_request( @@ -1183,7 +1226,8 @@ def _call_on_start_query_execution(self, query_id: str, options: ExecuteOptions) Both callbacks are invoked if set. Called by cursors whose execution model supports early access to the query ID (the synchronous and aio - cursors) immediately after the StartQueryExecution API call. + cursors) once ``_execute()`` returns it: after the StartQueryExecution + API call, or with a reusable query ID found through ``cache_size``. """ if self._on_start_query_execution: self._on_start_query_execution(query_id) @@ -1314,6 +1358,13 @@ def execute( parameters: dict[str, Any] | list[str] | None = None, **kwargs, ): + """Execute a SQL query. + + Args: + operation: SQL query string. + parameters: Query parameters. + **kwargs: Execution options defined by the cursor implementation. + """ raise NotImplementedError # pragma: no cover @abstractmethod @@ -1323,10 +1374,18 @@ def executemany( seq_of_parameters: list[dict[str, Any] | list[str] | None], **kwargs, ) -> None: + """Execute a SQL query once for each set of parameters. + + Args: + operation: SQL query string. + seq_of_parameters: Sequence of parameter sets. + **kwargs: Execution options defined by the cursor implementation. + """ raise NotImplementedError # pragma: no cover @abstractmethod def close(self) -> None: + """Close the cursor.""" raise NotImplementedError # pragma: no cover def _cancel(self, query_id: str) -> None: @@ -1351,10 +1410,20 @@ def _cancel(self, query_id: str) -> None: raise OperationalError(*e.args) from e def setinputsizes(self, sizes): # noqa: B027 - """Does nothing by default""" + """Accept input sizes as DB API 2.0 requires, and ignore them. + + Args: + sizes: Sequence of parameter types or sizes. + """ def setoutputsize(self, size, column=None): # noqa: B027 - """Does nothing by default""" + """Accept a column buffer size as DB API 2.0 requires, and ignore it. + + Args: + size: Buffer size for large columns. + column: Index of the column the size applies to, or None for all + large columns. + """ def __enter__(self): return self diff --git a/pyathena/connection.py b/pyathena/connection.py index 6f27c1d5e..826ff69d3 100644 --- a/pyathena/connection.py +++ b/pyathena/connection.py @@ -1,3 +1,5 @@ +"""DB API 2.0 connection to Amazon Athena.""" + from __future__ import annotations import logging @@ -231,7 +233,8 @@ def __init__( config: Boto3 Config object for client configuration. result_reuse_enable: Enable Athena query result reuse. Defaults to False. result_reuse_minutes: Minutes to reuse cached results. - on_start_query_execution: Callback function called when query starts. + on_start_query_execution: Callback invoked with each query ID before the cursor + waits for the query, as for the ``execute()`` argument of the same name. on_poll: Callback invoked once per poll iteration with the current execution object (``AthenaQueryExecution``, or ``AthenaCalculationExecutionStatus`` for Spark). Useful for diff --git a/pyathena/converter.py b/pyathena/converter.py index 7336bfaea..560549430 100644 --- a/pyathena/converter.py +++ b/pyathena/converter.py @@ -1,3 +1,5 @@ +"""Conversion of Athena result values to Python objects.""" + from __future__ import annotations import binascii @@ -509,6 +511,15 @@ def __init__( default: Callable[[str | None], Any | None] = _to_default, types: dict[str, type[Any]] | None = None, ) -> None: + """Initialize the converter. + + Args: + mappings: Conversion functions keyed by Athena type name. An empty + value is replaced with an empty dict. + default: Conversion function for types not in ``mappings``. + types: Python types keyed by Athena type name, returned by + ``get_dtype()``. None is replaced with an empty dict. + """ if mappings: self._mappings = mappings else: @@ -591,6 +602,16 @@ def update(self, mappings: dict[str, Callable[[str | None], Any | None]]) -> Non @abstractmethod def convert(self, type_: str, value: str | None, type_hint: str | None = None) -> Any | None: + """Convert a value returned by Athena to a Python object. + + Args: + type_: The Athena data type name. + value: The string value to convert, or None. + type_hint: Optional Athena DDL type signature of the value. + + Returns: + The converted value. + """ raise NotImplementedError # pragma: no cover @@ -629,6 +650,7 @@ class DefaultTypeConverter(Converter): _HIVE_REPLACEMENTS: ClassVar[dict[str, str]] = {"<": "(", ">": ")", ":": " "} def __init__(self) -> None: + """Initialize the converter with the default conversion functions.""" super().__init__(mappings=deepcopy(_DEFAULT_CONVERTERS), default=_to_default) self._parser = TypeSignatureParser() self._typed_converter = TypedValueConverter( diff --git a/pyathena/cursor.py b/pyathena/cursor.py index 2e789520a..d0b6de4f3 100644 --- a/pyathena/cursor.py +++ b/pyathena/cursor.py @@ -1,3 +1,5 @@ +"""DB API 2.0 cursors that return rows as tuples or dictionaries.""" + from __future__ import annotations import logging @@ -56,6 +58,24 @@ def __init__( result_reuse_minutes: int = CursorIterator.DEFAULT_RESULT_REUSE_MINUTES, **kwargs, ) -> None: + """Initialize a Cursor. + + Args: + s3_staging_dir: S3 location for query results. + schema_name: Default schema name. + catalog_name: Default catalog name. + work_group: Athena workgroup name. + poll_interval: Query status polling interval in seconds. + encryption_option: S3 encryption option (SSE_S3, SSE_KMS, CSE_KMS). + kms_key: KMS key for encryption. + kill_on_interrupt: Cancel the query when a ``KeyboardInterrupt`` interrupts + ``execute()`` while it starts or waits for the query. + result_reuse_enable: Enable Athena query result reuse. + result_reuse_minutes: Maximum age in minutes of a reused result. + **kwargs: Arguments forwarded to ``WithResultSet.__init__`` and + ``BaseCursor.__init__``, such as ``arraysize``, ``connection``, + ``converter``, ``formatter``, and ``retry_config``. + """ super().__init__( s3_staging_dir=s3_staging_dir, schema_name=schema_name, @@ -107,8 +127,16 @@ def execute( Args: operation: SQL query string to execute. parameters: Query parameters (optional). - on_start_query_execution: Callback function called immediately after - start_query_execution API is called. + work_group: Athena workgroup to use for this query. + s3_staging_dir: S3 location for query results. + cache_size: Number of queries to check for result caching. + cache_expiration_time: Cache expiration time in seconds. + result_reuse_enable: Enable Athena result reuse for this query. + result_reuse_minutes: Minutes to reuse cached results. + paramstyle: Parameter style ('qmark' or 'pyformat'). + on_start_query_execution: Callback invoked with the query ID before ``execute()`` + waits for the query: after the ``StartQueryExecution`` call, or after a + reusable query ID is found through ``cache_size``. Function signature: (query_id: str) -> None This allows early access to query_id for monitoring/cancellation. @@ -189,6 +217,13 @@ class DictCursor(Cursor): """ def __init__(self, **kwargs) -> None: + """Initialize a DictCursor. + + Args: + **kwargs: Arguments forwarded to ``Cursor.__init__``. If they include + ``dict_type``, it is also assigned to the class attribute + ``AthenaDictResultSet.dict_type``, the type used to build each row. + """ super().__init__(**kwargs) self._result_set_class = AthenaDictResultSet if "dict_type" in kwargs: diff --git a/pyathena/error.py b/pyathena/error.py index 666584f1e..7d15eb865 100644 --- a/pyathena/error.py +++ b/pyathena/error.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""DB API 2.0 exception hierarchy used by PyAthena.""" + __all__ = [ "DataError", "DatabaseError", diff --git a/pyathena/filesystem/__init__.py b/pyathena/filesystem/__init__.py index c65ed7373..65f3a088e 100644 --- a/pyathena/filesystem/__init__.py +++ b/pyathena/filesystem/__init__.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""fsspec filesystem implementations for Amazon S3.""" + import logging import fsspec diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index abbe270f4..4104f9354 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -1,3 +1,5 @@ +"""fsspec filesystem and file implementations for Amazon S3.""" + from __future__ import annotations import logging @@ -255,6 +257,22 @@ def _get_client_compatible_with_s3fs(self, **kwargs) -> BaseClient: @staticmethod def parse_path(path: str) -> tuple[str, str | None, str | None]: + """Parse an S3 path into its bucket, key and version ID. + + The path may have an ``s3://`` or ``s3a://`` scheme and a version ID + query (``?versionId=``, ``?versionID=``, ``?versionid=`` or + ``?version_id=``). + + Args: + path: The S3 path (e.g., "s3://bucket/key?versionId=..."). + + Returns: + Tuple of the bucket, the key (None for a bucket path) and the + version ID (None if the path has none). + + Raises: + ValueError: If the path is not a valid S3 path. + """ match = S3FileSystem.PATTERN_PATH.search(path) if match: return match.group("bucket"), match.group("key"), match.group("version_id") @@ -532,6 +550,29 @@ def _list_object_versions_pages( break def info(self, path: str, **kwargs) -> S3Object: + """Return information about an S3 path. + + Uses the directory cache first: a cached entry for the path is + returned, a cached listing of the path itself makes it a directory, + and a cached listing of its parent without it means it does not exist. + Otherwise, a key path is looked up with HeadObject and, if no object + exists, with a ListObjectsV2 request (``Delimiter="/"``, + ``MaxKeys=1``) that checks whether it is a key prefix; a bucket path + is looked up with HeadBucket. With ``version_aware``, a cached file + entry without a version ID is looked up again. + + Args: + path: S3 path (e.g., "s3://bucket" or "s3://bucket/key"). + **kwargs: Additional arguments including: + refresh: If True, bypass the cache and query S3. + version_id: The version ID to look up when the path has none. + + Returns: + S3Object describing the bucket, directory, or file. + + Raises: + FileNotFoundError: If the path does not exist. + """ refresh = kwargs.pop("refresh", False) path = self._strip_protocol(path) bucket, key, path_version_id = self.parse_path(path) @@ -778,6 +819,16 @@ def exists(self, path: str, **kwargs) -> bool: return bool(file) def rm_file(self, path: str, **kwargs) -> None: + """Delete an S3 object with DeleteObject. + + Does nothing for a bucket path. If the path has a version ID, that + version is deleted. + + Args: + path: S3 path (s3://bucket/key) of the object to delete. + **kwargs: Accepted for fsspec compatibility; not used in the + request. + """ bucket, key, version_id = self.parse_path(path) if not key: return @@ -785,6 +836,21 @@ def rm_file(self, path: str, **kwargs) -> None: self.invalidate_cache(path) def rm(self, path, recursive=False, maxdepth=None, **kwargs) -> None: + """Delete objects with DeleteObjects requests. + + Expands the path with ``expand_path`` and deletes the matched objects + in parallel requests of up to ``DELETE_OBJECTS_MAX_KEYS`` keys each. + + Args: + path: S3 path (s3://bucket/key) to delete. + recursive: Whether to delete all objects below the path. + maxdepth: Maximum depth to expand when ``recursive`` is True. + **kwargs: Additional parameters passed to the DeleteObjects API. + ``Quiet`` (default True) sets the quiet mode of the requests. + + Raises: + ValueError: If the path is a bucket. + """ bucket, key, version_id = self.parse_path(path) if not key: raise ValueError("Cannot delete the bucket.") @@ -1002,6 +1068,22 @@ def rmdir(self, path: str) -> None: self.dircache.pop("", None) def touch(self, path: str, truncate: bool = True, **kwargs) -> dict[str, Any]: + """Create an empty object with PutObject. + + Args: + path: S3 path (s3://bucket/key) of the object. + truncate: If True, replace an existing object with an empty one; + if False, raise if the object exists. + **kwargs: Additional parameters passed to the PutObject API. + + Returns: + The PutObject response as a dictionary (see + :meth:`S3PutObject.to_dict`). + + Raises: + ValueError: If the path has a version ID, is a bucket, or exists + while ``truncate`` is False. + """ bucket, key, version_id = self.parse_path(path) if version_id: raise ValueError("Cannot touch the file with the version specified.") @@ -1269,6 +1351,19 @@ def _finish_multipart_upload( def cat_file( self, path: str, start: int | None = None, end: int | None = None, **kwargs ) -> bytes: + """Read the contents of an S3 object with GetObject. + + Args: + path: S3 path (s3://bucket/key) of the object. + start: Byte offset to start reading at. A negative value counts + from the end of the object. + end: Byte offset to stop reading at (exclusive). A negative value + counts from the end of the object. + **kwargs: Additional parameters passed to the GetObject API. + + Returns: + The bytes read from the object. + """ bucket, key, version_id = self.parse_path(path) if start is not None or end is not None: size = self.info(path).get("size", 0) @@ -1750,13 +1845,37 @@ def clear_multipart_uploads(self, path: str) -> None: future.result() def created(self, path: str) -> datetime: + """Return the creation time of the path. + + Returns the same value as :meth:`modified`. + + Args: + path: S3 path (s3://bucket/key). + + Returns: + The last-modified time of the object. + """ return self.modified(path) def modified(self, path: str) -> datetime: + """Return the last-modified time of the path. + + Args: + path: S3 path (s3://bucket/key). + + Returns: + The ``last_modified`` field from :meth:`info`, which is None for + buckets and directories. + """ info = self.info(path) return cast(datetime, info.get("last_modified")) def invalidate_cache(self, path: str | None = None) -> None: + """Remove the cached entries of the path and its parent paths. + + Args: + path: The path to invalidate. If None, clear the whole cache. + """ if path is None: self.dircache.clear() else: @@ -1964,6 +2083,11 @@ def _call(self, method: str | Callable[..., Any], **kwargs) -> dict[str, Any]: class S3File(AbstractBufferedFile): + """A buffered file object for reading and writing an S3 object. + + Instances are returned by ``S3FileSystem.open()``. + """ + fs: S3FileSystem def __init__( @@ -1982,6 +2106,40 @@ def __init__( s3_additional_kwargs: dict[str, Any] | None = None, **kwargs, ) -> None: + """Initialize the file for the path and mode. + + In read mode, the object is looked up with ``info()`` and the reads + are made conditional on its ETag (``IfMatch``). In append mode, an + existing object smaller than ``MULTIPART_UPLOAD_MIN_PART_SIZE`` is + read into the write buffer; a larger one is copied as the first parts + once a multipart upload starts. + + Args: + fs: The filesystem that the file belongs to. + path: S3 path (s3://bucket/key) of the file. + mode: The file mode, such as ``rb``, ``wb`` or ``ab``. + version_id: The version ID to read. Must match the version ID in + the path if both are given. + max_workers: The number of parallel workers for range reads and + part copies. + executor: The executor for parallel operations. If None, a new + ``S3ThreadPoolExecutor`` is created. + block_size: The block size for reads and writes. Must be at least + ``MULTIPART_UPLOAD_MIN_PART_SIZE`` unless reading. + cache_type: The fsspec cache type for reads. + autocommit: Whether to commit the written data when the file is + closed. If False, :meth:`commit` must be called. + cache_options: Options for the fsspec cache. + size: The size of the object, if known. Passed to + ``fsspec.spec.AbstractBufferedFile``. + s3_additional_kwargs: Additional parameters for the object requests + of the file. + **kwargs: Accepted for compatibility; not used. + + Raises: + ValueError: If the path has no key, the version IDs do not match, + or the block size is too small for writing. + """ self.max_workers = max_workers self._executor: S3Executor = executor or S3ThreadPoolExecutor(max_workers=max_workers) self.s3_additional_kwargs = s3_additional_kwargs if s3_additional_kwargs else {} @@ -2045,6 +2203,7 @@ def __init__( self.multipart_upload_parts: list[Future[S3MultipartUploadPart]] = [] def close(self) -> None: + """Close the file, flushing any written data, and shut down its executor.""" super().close() self._executor.shutdown() @@ -2150,6 +2309,17 @@ def _upload_chunk(self, final: bool = False) -> bool: return not final def commit(self) -> None: + """Complete the upload of the written data. + + Creates an empty object if nothing was written, uploads the buffered + data with PutObject if no multipart upload part was submitted, and + otherwise completes the multipart upload, which is aborted if the + completion fails. Invalidates the cache of the path afterwards. + + Raises: + RuntimeError: If parts were submitted but no multipart upload is + initialized. + """ if self.tell() == 0: if self.buffer is not None: self.discard() @@ -2186,6 +2356,7 @@ def commit(self) -> None: self.fs.invalidate_cache(self.path) def discard(self) -> None: + """Cancel pending part uploads and abort the multipart upload, if any.""" if self.multipart_upload: for f in self.multipart_upload_parts: f.cancel() diff --git a/pyathena/filesystem/s3_async.py b/pyathena/filesystem/s3_async.py index fa9ba9962..3e8279728 100644 --- a/pyathena/filesystem/s3_async.py +++ b/pyathena/filesystem/s3_async.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""Asynchronous fsspec filesystem for Amazon S3 built on ``S3FileSystem``.""" + from __future__ import annotations import asyncio @@ -84,6 +86,23 @@ def __init__( batch_size: int | None = None, **kwargs, ) -> None: + """Initialize the filesystem and its internal ``S3FileSystem``. + + Args: + connection: Passed to the internal ``S3FileSystem``. + default_block_size: Passed to the internal ``S3FileSystem``. + default_cache_type: Passed to the internal ``S3FileSystem``. + max_workers: Passed to the internal ``S3FileSystem``. + s3_additional_kwargs: Passed to the internal ``S3FileSystem``. + allow_bucket_creation: Passed to the internal ``S3FileSystem``. + allow_bucket_deletion: Passed to the internal ``S3FileSystem``. + version_aware: Passed to the internal ``S3FileSystem``. + asynchronous: Passed to ``fsspec.asyn.AsyncFileSystem``. + loop: Passed to ``fsspec.asyn.AsyncFileSystem``. + batch_size: Passed to ``fsspec.asyn.AsyncFileSystem``. + **kwargs: Passed to both ``fsspec.asyn.AsyncFileSystem`` and the + internal ``S3FileSystem``. + """ super().__init__( asynchronous=asynchronous, loop=loop, @@ -106,6 +125,19 @@ def __init__( @staticmethod def parse_path(path: str) -> tuple[str, str | None, str | None]: + """Parse an S3 path into its bucket, key and version ID. + + See :meth:`S3FileSystem.parse_path`. + + Args: + path: The S3 path. + + Returns: + Tuple of the bucket, the key and the version ID. + + Raises: + ValueError: If the path is not a valid S3 path. + """ return S3FileSystem.parse_path(path) async def _info(self, path: str, **kwargs) -> S3Object: @@ -348,50 +380,200 @@ async def _rmdir(self, path: str) -> None: await asyncio.to_thread(self._sync_fs.rmdir, path) def rmdir(self, path: str) -> None: + """Remove an S3 bucket, which must be empty. + + See :meth:`S3FileSystem.rmdir`. + + Args: + path: S3 bucket path (e.g., "s3://bucket"). + """ self._sync_fs.rmdir(path) def sign(self, path: str, expiration: int = 3600, **kwargs) -> str: + """Generate a presigned URL for S3 object access. + + See :meth:`S3FileSystem.sign`. + + Args: + path: S3 path (s3://bucket/key) to generate the URL for. + expiration: URL expiration time in seconds. + **kwargs: Additional parameters passed to :meth:`S3FileSystem.sign`. + + Returns: + The presigned URL. + """ return cast(str, self._sync_fs.sign(path, expiration=expiration, **kwargs)) def metadata(self, path: str, **kwargs) -> S3Metadata: + """Return the metadata of the path. + + See :meth:`S3FileSystem.metadata`. + + Args: + path: S3 path (s3://bucket/key) to get metadata for. + **kwargs: Additional parameters passed to the HeadObject API. + + Returns: + S3Metadata of the object. + """ return self._sync_fs.metadata(path, **kwargs) def getxattr(self, path: str, attr_name: str, **kwargs) -> str | None: + """Get an attribute from the user-defined metadata of the path. + + See :meth:`S3FileSystem.getxattr`. + + Args: + path: S3 path (s3://bucket/key) to get the attribute for. + attr_name: The name of the attribute. + **kwargs: Additional parameters passed to the HeadObject API. + + Returns: + The value of the attribute, or None if the attribute is not set. + """ return self._sync_fs.getxattr(path, attr_name, **kwargs) def setxattr(self, path: str, copy_kwargs: dict[str, Any] | None = None, **kwargs) -> None: + """Set the user-defined metadata of the path. + + See :meth:`S3FileSystem.setxattr`. + + Args: + path: S3 path (s3://bucket/key) to set metadata for. + copy_kwargs: Additional parameters to use for the underlying + CopyObject API call. + **kwargs: Key-value pairs of metadata to set; a None value deletes the + key. + """ self._sync_fs.setxattr(path, copy_kwargs=copy_kwargs, **kwargs) def get_tags(self, path: str) -> dict[str, str]: + """Retrieve the tag key/values for the given path. + + See :meth:`S3FileSystem.get_tags`. + + Args: + path: S3 path (s3://bucket/key) to get tags for. + + Returns: + Dictionary mapping tag keys to tag values. + """ return self._sync_fs.get_tags(path) def put_tags(self, path: str, tags: dict[str, str], mode: str = "o") -> None: + """Set the tags for the given existing key. + + See :meth:`S3FileSystem.put_tags`. + + Args: + path: S3 path (s3://bucket/key) of the existing key. + tags: Tags to apply. + mode: ``o`` to overwrite the existing tags or ``m`` to merge with them. + """ self._sync_fs.put_tags(path, tags, mode=mode) def chmod(self, path: str, acl: str, recursive: bool = False, **kwargs) -> None: + """Set the canned ACL of a bucket or key. + + See :meth:`S3FileSystem.chmod`. + + Args: + path: S3 path (s3://bucket or s3://bucket/key) to set the ACL on. + acl: The canned ACL to apply. + recursive: Whether to apply the ACL to all keys below the path too. + **kwargs: Additional parameters passed to the PutObjectAcl or + PutBucketAcl API. + """ self._sync_fs.chmod(path, acl, recursive=recursive, **kwargs) def object_version_info( self, path: str, delete_markers: bool = False, **kwargs ) -> list[S3ObjectVersion]: + """List the versions of the objects under the path. + + See :meth:`S3FileSystem.object_version_info`. + + Args: + path: S3 path (s3://bucket/key or a key prefix) to list the versions for. + delete_markers: Whether to include delete markers in the result. + **kwargs: Additional parameters passed to the ListObjectVersions API. + + Returns: + List of S3ObjectVersion instances describing the versions. + """ return self._sync_fs.object_version_info(path, delete_markers=delete_markers, **kwargs) def list_multipart_uploads(self, path: str) -> list[S3MultipartUpload]: + """List in-progress (incomplete) multipart uploads in a bucket. + + See :meth:`S3FileSystem.list_multipart_uploads`. + + Args: + path: S3 bucket or prefix path (e.g., "s3://bucket" or "s3://bucket/prefix"). + + Returns: + List of S3MultipartUpload instances describing the uploads. + """ return self._sync_fs.list_multipart_uploads(path) def clear_multipart_uploads(self, path: str) -> None: + """Abort any incomplete multipart uploads in the bucket. + + See :meth:`S3FileSystem.clear_multipart_uploads`. + + Args: + path: S3 bucket or prefix path (e.g., "s3://bucket" or "s3://bucket/prefix"). + """ self._sync_fs.clear_multipart_uploads(path) def checksum(self, path: str, **kwargs) -> int: + """Get the checksum of an S3 object or directory. + + See :meth:`S3FileSystem.checksum`. + + Args: + path: S3 path (s3://bucket/key) to get the checksum for. + **kwargs: Additional arguments passed to :meth:`S3FileSystem.checksum`. + + Returns: + Integer checksum derived from the ETag or the directory token. + """ return cast(int, self._sync_fs.checksum(path, **kwargs)) def created(self, path: str) -> datetime: + """Return the creation time of the path. + + See :meth:`S3FileSystem.created`. + + Args: + path: S3 path (s3://bucket/key). + + Returns: + The last-modified time of the object. + """ return self._sync_fs.created(path) def modified(self, path: str) -> datetime: + """Return the last-modified time of the path. + + See :meth:`S3FileSystem.modified`. + + Args: + path: S3 path (s3://bucket/key). + + Returns: + The last-modified time of the object. + """ return self._sync_fs.modified(path) def invalidate_cache(self, path: str | None = None) -> None: + """Remove the cached entries of the path and its parent paths. + + See :meth:`S3FileSystem.invalidate_cache`. + + Args: + path: The path to invalidate. If None, clear the whole cache. + """ self._sync_fs.invalidate_cache(path) async def _touch(self, path: str, truncate: bool = True, **kwargs) -> None: diff --git a/pyathena/filesystem/s3_errors.py b/pyathena/filesystem/s3_errors.py index 78a9aedd6..0929ffc4c 100644 --- a/pyathena/filesystem/s3_errors.py +++ b/pyathena/filesystem/s3_errors.py @@ -97,6 +97,11 @@ class S3ClientError: } def __init__(self, error: botocore.exceptions.ClientError) -> None: + """Initialize the error from a botocore ``ClientError``. + + Args: + error: The ``ClientError`` raised by the S3 client. + """ error_info = error.response.get("Error", {}) self._code: str = str(error_info.get("Code", "")) self._message: str = str(error_info.get("Message", error)) diff --git a/pyathena/filesystem/s3_executor.py b/pyathena/filesystem/s3_executor.py index 603005555..4cf80093b 100644 --- a/pyathena/filesystem/s3_executor.py +++ b/pyathena/filesystem/s3_executor.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""Executors that run S3 filesystem operations in parallel.""" + from __future__ import annotations import asyncio @@ -53,6 +55,11 @@ class S3ThreadPoolExecutor(S3Executor): """ def __init__(self, max_workers: int) -> None: + """Initialize the executor with a new ``ThreadPoolExecutor``. + + Args: + max_workers: The maximum number of threads of the thread pool. + """ self._executor = ThreadPoolExecutor(max_workers=max_workers) @override @@ -83,6 +90,12 @@ class S3AioExecutor(S3Executor): """ def __init__(self, loop: asyncio.AbstractEventLoop | None = None) -> None: + """Initialize the executor with the event loop to schedule work on. + + Args: + loop: The asyncio event loop. ``submit`` raises ``RuntimeError`` + if it is None or not running. + """ self._loop = loop @override diff --git a/pyathena/filesystem/s3_object.py b/pyathena/filesystem/s3_object.py index 392fe0358..c28e45dd9 100644 --- a/pyathena/filesystem/s3_object.py +++ b/pyathena/filesystem/s3_object.py @@ -1,3 +1,5 @@ +"""Model classes for S3 objects and S3 API responses used by the S3 filesystem.""" + from __future__ import annotations import copy @@ -109,6 +111,22 @@ def __init__( init: dict[str, Any], **kwargs, ) -> None: + """Initialize the object from an S3 API response. + + Only the fields of ``init`` that have a property mapping are kept, + under their property names (e.g., ``ContentType`` -> ``content_type``). + ``storage_class`` defaults to ``STANDARD`` when ``init`` has no + ``StorageClass``, and ``size`` is taken from ``Size`` or + ``ContentLength``. ``name`` is set to ``bucket/key``, or to the bucket + when there is no key. + + Args: + init: An S3 API response or listing entry, such as a HeadObject + response or a ListObjectsV2 ``Contents`` entry. + **kwargs: Additional fields stored as-is, such as ``type``, + ``bucket``, ``key`` and ``version_id``. S3 API field names + are stored under their property names. + """ if init: filtered = {} for k, v in init.items(): @@ -181,6 +199,13 @@ def to_dict(self) -> dict[str, Any]: return copy.deepcopy(self.__dict__) def to_api_repr(self) -> dict[str, Any]: + """Convert the object metadata to S3 API request parameters. + + Returns: + Dictionary keyed by S3 API field names (e.g., ``ContentType``) of + the fields that are set, excluding ``ETag``, ``ContentLength`` and + ``LastModified``. + """ fields = {} for k, v in _API_FIELD_TO_S3_OBJECT_PROPERTY.items(): if k in ["ETag", "ContentLength", "LastModified"]: @@ -213,6 +238,11 @@ class S3Metadata(Mapping[str, str]): """ def __init__(self, response: dict[str, Any]) -> None: + """Initialize the metadata from a HeadObject response. + + Args: + response: The HeadObject response. + """ self._cache_control: str | None = response.get("CacheControl") self._content_disposition: str | None = response.get("ContentDisposition") self._content_encoding: str | None = response.get("ContentEncoding") @@ -255,70 +285,87 @@ def __repr__(self) -> str: @property def cache_control(self) -> str | None: + """The ``CacheControl`` header of the object.""" return self._cache_control @property def content_disposition(self) -> str | None: + """The ``ContentDisposition`` header of the object.""" return self._content_disposition @property def content_encoding(self) -> str | None: + """The ``ContentEncoding`` header of the object.""" return self._content_encoding @property def content_language(self) -> str | None: + """The ``ContentLanguage`` header of the object.""" return self._content_language @property def content_length(self) -> int | None: + """The ``ContentLength`` of the object in bytes.""" return self._content_length @property def content_type(self) -> str | None: + """The ``ContentType`` of the object.""" return self._content_type @property def etag(self) -> str | None: + """The ``ETag`` (entity tag) of the object.""" return self._etag @property def expiration(self) -> str | None: + """The ``Expiration`` header of the object.""" return self._expiration @property def expires(self) -> datetime | None: + """The ``Expires`` date of the object.""" return self._expires @property def last_modified(self) -> datetime | None: + """The ``LastModified`` time of the object.""" return self._last_modified @property def storage_class(self) -> str: + """The ``StorageClass`` of the object; ``STANDARD`` when S3 omits it.""" return self._storage_class @property def server_side_encryption(self) -> str | None: + """The ``ServerSideEncryption`` algorithm of the object.""" return self._server_side_encryption @property def sse_customer_algorithm(self) -> str | None: + """The ``SSECustomerAlgorithm`` of the object.""" return self._sse_customer_algorithm @property def sse_kms_key_id(self) -> str | None: + """The ``SSEKMSKeyId`` of the KMS key for the object.""" return self._sse_kms_key_id @property def bucket_key_enabled(self) -> bool | None: + """Whether the object uses an S3 Bucket Key (``BucketKeyEnabled``).""" return self._bucket_key_enabled @property def website_redirect_location(self) -> str | None: + """The ``WebsiteRedirectLocation`` of the object.""" return self._website_redirect_location @property def version_id(self) -> str | None: + """The ``VersionId`` of the object.""" return self._version_id @property @@ -349,6 +396,15 @@ class S3ObjectVersion: """ def __init__(self, bucket: str, is_delete_marker: bool, response: dict[str, Any]) -> None: + """Initialize the version from a ListObjectVersions entry. + + Args: + bucket: The name of the bucket that contains the version. + is_delete_marker: Whether the entry comes from ``DeleteMarkers`` + rather than ``Versions``. + response: A ``Versions`` or ``DeleteMarkers`` entry of the + ListObjectVersions response. + """ self._bucket = bucket self._is_delete_marker = is_delete_marker self._key: str = response["Key"] @@ -363,46 +419,57 @@ def __init__(self, bucket: str, is_delete_marker: bool, response: dict[str, Any] @property def bucket(self) -> str: + """The name of the bucket that contains the version.""" return self._bucket @property def key(self) -> str: + """The ``Key`` of the object.""" return self._key @property def name(self) -> str: + """The path of the version in ``bucket/key`` form.""" return f"{self._bucket}/{self._key}" @property def version_id(self) -> str | None: + """The ``VersionId`` of the version.""" return self._version_id @property def is_latest(self) -> bool: + """Whether the version is the latest version of the object (``IsLatest``).""" return self._is_latest @property def is_delete_marker(self) -> bool: + """Whether the version is a delete marker.""" return self._is_delete_marker @property def last_modified(self) -> datetime | None: + """The ``LastModified`` time of the version.""" return self._last_modified @property def etag(self) -> str | None: + """The ``ETag`` of the version; None for delete markers.""" return self._etag @property def size(self) -> int | None: + """The ``Size`` of the version in bytes; None for delete markers.""" return self._size @property def storage_class(self) -> str | None: + """The ``StorageClass`` of the version; None for delete markers.""" return self._storage_class @property def owner(self) -> S3Owner | None: + """The ``Owner`` of the version, or None if the response has none.""" return self._owner @@ -426,6 +493,11 @@ class S3PutObject: """ def __init__(self, response: dict[str, Any]) -> None: + """Initialize the result from a PutObject response. + + Args: + response: The PutObject response. + """ self._expiration: str | None = response.get("Expiration") self._version_id: str | None = response.get("VersionId") self._etag: str | None = response.get("ETag") @@ -443,61 +515,81 @@ def __init__(self, response: dict[str, Any]) -> None: @property def expiration(self) -> str | None: + """The ``Expiration`` header of the uploaded object.""" return self._expiration @property def version_id(self) -> str | None: + """The ``VersionId`` of the uploaded object.""" return self._version_id @property def etag(self) -> str | None: + """The ``ETag`` of the uploaded object.""" return self._etag @property def checksum_crc32(self) -> str | None: + """The ``ChecksumCRC32`` of the uploaded object.""" return self._checksum_crc32 @property def checksum_crc32c(self) -> str | None: + """The ``ChecksumCRC32C`` of the uploaded object.""" return self._checksum_crc32c @property def checksum_sha1(self) -> str | None: + """The ``ChecksumSHA1`` of the uploaded object.""" return self._checksum_sha1 @property def checksum_sha256(self) -> str | None: + """The ``ChecksumSHA256`` of the uploaded object.""" return self._checksum_sha256 @property def server_side_encryption(self) -> str | None: + """The ``ServerSideEncryption`` algorithm of the uploaded object.""" return self._server_side_encryption @property def sse_customer_algorithm(self) -> str | None: + """The ``SSECustomerAlgorithm`` of the uploaded object.""" return self._sse_customer_algorithm @property def sse_customer_key_md5(self) -> str | None: + """The ``SSECustomerKeyMD5`` of the customer-provided key for the uploaded object.""" return self._sse_customer_key_md5 @property def sse_kms_key_id(self) -> str | None: + """The ``SSEKMSKeyId`` of the KMS key for the uploaded object.""" return self._sse_kms_key_id @property def sse_kms_encryption_context(self) -> str | None: + """The ``SSEKMSEncryptionContext`` of the uploaded object.""" return self._sse_kms_encryption_context @property def bucket_key_enabled(self) -> bool | None: + """Whether the uploaded object uses an S3 Bucket Key (``BucketKeyEnabled``).""" return self._bucket_key_enabled @property def request_charged(self) -> str | None: + """The ``RequestCharged`` field of the response.""" return self._request_charged def to_dict(self) -> dict[str, Any]: + """Convert the response to a dictionary. + + Returns: + Deep copy of the instance attributes, keyed by their attribute + names (e.g., ``_etag``). + """ return copy.deepcopy(self.__dict__) @@ -510,15 +602,22 @@ class S3Owner: """ def __init__(self, response: dict[str, Any]) -> None: + """Initialize the owner from an ``Owner`` or ``Initiator`` response field. + + Args: + response: The ``Owner`` or ``Initiator`` field of an S3 API response. + """ self._display_name: str | None = response.get("DisplayName") self._id: str | None = response.get("ID") @property def display_name(self) -> str | None: + """The ``DisplayName`` of the owner.""" return self._display_name @property def id(self) -> str | None: + """The canonical user ``ID`` of the owner.""" return self._id @@ -544,6 +643,12 @@ class S3MultipartUpload: """ def __init__(self, response: dict[str, Any]) -> None: + """Initialize the upload from an S3 API response. + + Args: + response: A CreateMultipartUpload response or an ``Uploads`` entry + of a ListMultipartUploads response. + """ self._abort_date = response.get("AbortDate") self._abort_rule_id = response.get("AbortRuleId") self._bucket = response.get("Bucket") @@ -567,70 +672,87 @@ def __init__(self, response: dict[str, Any]) -> None: @property def abort_date(self) -> datetime | None: + """The ``AbortDate`` of the upload set by a lifecycle rule.""" return self._abort_date @property def abort_rule_id(self) -> str | None: + """The ``AbortRuleId`` of the lifecycle rule that applies to the upload.""" return self._abort_rule_id @property def bucket(self) -> str | None: + """The ``Bucket`` of the upload.""" return self._bucket @property def key(self) -> str | None: + """The ``Key`` of the object being uploaded.""" return self._key @property def upload_id(self) -> str | None: + """The ``UploadId`` of the multipart upload.""" return self._upload_id @property def server_side_encryption(self) -> str | None: + """The ``ServerSideEncryption`` algorithm of the upload.""" return self._server_side_encryption @property def sse_customer_algorithm(self) -> str | None: + """The ``SSECustomerAlgorithm`` of the upload.""" return self._sse_customer_algorithm @property def sse_customer_key_md5(self) -> str | None: + """The ``SSECustomerKeyMD5`` of the customer-provided key for the upload.""" return self._sse_customer_key_md5 @property def sse_kms_key_id(self) -> str | None: + """The ``SSEKMSKeyId`` of the KMS key for the upload.""" return self._sse_kms_key_id @property def sse_kms_encryption_context(self) -> str | None: + """The ``SSEKMSEncryptionContext`` of the upload.""" return self._sse_kms_encryption_context @property def bucket_key_enabled(self) -> bool | None: + """Whether the upload uses an S3 Bucket Key (``BucketKeyEnabled``).""" return self._bucket_key_enabled @property def request_charged(self) -> str | None: + """The ``RequestCharged`` field of the response.""" return self._request_charged @property def checksum_algorithm(self) -> str | None: + """The ``ChecksumAlgorithm`` of the upload.""" return self._checksum_algorithm @property def initiated(self) -> datetime | None: + """The ``Initiated`` time of the upload, returned by ListMultipartUploads.""" return self._initiated @property def storage_class(self) -> str | None: + """The ``StorageClass`` of the upload, returned by ListMultipartUploads.""" return self._storage_class @property def owner(self) -> S3Owner | None: + """The ``Owner`` of the upload, returned by ListMultipartUploads.""" return self._owner @property def initiator(self) -> S3Owner | None: + """The ``Initiator`` of the upload, returned by ListMultipartUploads.""" return self._initiator @@ -653,6 +775,15 @@ class S3MultipartUploadPart: """ def __init__(self, part_number: int, response: dict[str, Any]) -> None: + """Initialize the part from an UploadPart or UploadPartCopy response. + + For an UploadPartCopy response, the ``ETag``, ``LastModified`` and + checksums are read from its ``CopyPartResult``. + + Args: + part_number: The part number of the part. + response: The UploadPart or UploadPartCopy response. + """ self._part_number = part_number self._copy_source_version_id: str | None = response.get("CopySourceVersionId") copy_part_result = response.get("CopyPartResult") @@ -679,61 +810,81 @@ def __init__(self, part_number: int, response: dict[str, Any]) -> None: @property def part_number(self) -> int: + """The part number of the part.""" return self._part_number @property def copy_source_version_id(self) -> str | None: + """The ``CopySourceVersionId`` of the source object of a copied part.""" return self._copy_source_version_id @property def last_modified(self) -> datetime | None: + """The ``LastModified`` time from ``CopyPartResult``; None for uploaded parts.""" return self._last_modified @property def etag(self) -> str | None: + """The ``ETag`` of the part.""" return self._etag @property def checksum_crc32(self) -> str | None: + """The ``ChecksumCRC32`` of the part.""" return self._checksum_crc32 @property def checksum_crc32c(self) -> str | None: + """The ``ChecksumCRC32C`` of the part.""" return self._checksum_crc32c @property def checksum_sha1(self) -> str | None: + """The ``ChecksumSHA1`` of the part.""" return self._checksum_sha1 @property def checksum_sha256(self) -> str | None: + """The ``ChecksumSHA256`` of the part.""" return self._checksum_sha256 @property def server_side_encryption(self) -> str | None: + """The ``ServerSideEncryption`` algorithm of the part.""" return self._server_side_encryption @property def sse_customer_algorithm(self) -> str | None: + """The ``SSECustomerAlgorithm`` of the part.""" return self._sse_customer_algorithm @property def sse_customer_key_md5(self) -> str | None: + """The ``SSECustomerKeyMD5`` of the customer-provided key for the part.""" return self._sse_customer_key_md5 @property def sse_kms_key_id(self) -> str | None: + """The ``SSEKMSKeyId`` of the KMS key for the part.""" return self._sse_kms_key_id @property def bucket_key_enabled(self) -> bool | None: + """Whether the part uses an S3 Bucket Key (``BucketKeyEnabled``).""" return self._bucket_key_enabled @property def request_charged(self) -> str | None: + """The ``RequestCharged`` field of the response.""" return self._request_charged def to_api_repr(self) -> dict[str, Any]: + """Convert the part to a part entry of a CompleteMultipartUpload request. + + Returns: + Dictionary with the ``ETag``, checksum and ``PartNumber`` fields of + the part. + """ return { "ETag": self.etag, "ChecksumCRC32": self.checksum_crc32, @@ -765,6 +916,11 @@ class S3CompleteMultipartUpload: """ def __init__(self, response: dict[str, Any]) -> None: + """Initialize the result from a CompleteMultipartUpload response. + + Args: + response: The CompleteMultipartUpload response. + """ self._location: str | None = response.get("Location") self._bucket: str | None = response.get("Bucket") self._key: str | None = response.get("Key") @@ -782,59 +938,79 @@ def __init__(self, response: dict[str, Any]) -> None: @property def location(self) -> str | None: + """The ``Location`` URI of the completed object.""" return self._location @property def bucket(self) -> str | None: + """The ``Bucket`` of the completed object.""" return self._bucket @property def key(self) -> str | None: + """The ``Key`` of the completed object.""" return self._key @property def expiration(self) -> str | None: + """The ``Expiration`` header of the completed object.""" return self._expiration @property def version_id(self) -> str | None: + """The ``VersionId`` of the completed object.""" return self._version_id @property def etag(self) -> str | None: + """The ``ETag`` of the completed object.""" return self._etag @property def checksum_crc32(self) -> str | None: + """The ``ChecksumCRC32`` of the completed object.""" return self._checksum_crc32 @property def checksum_crc32c(self) -> str | None: + """The ``ChecksumCRC32C`` of the completed object.""" return self._checksum_crc32c @property def checksum_sha1(self) -> str | None: + """The ``ChecksumSHA1`` of the completed object.""" return self._checksum_sha1 @property def checksum_sha256(self) -> str | None: + """The ``ChecksumSHA256`` of the completed object.""" return self._checksum_sha256 @property def server_side_encryption(self) -> str | None: + """The ``ServerSideEncryption`` algorithm of the completed object.""" return self._server_side_encryption @property def sse_kms_key_id(self) -> str | None: + """The ``SSEKMSKeyId`` of the KMS key for the completed object.""" return self._sse_kms_key_id @property def bucket_key_enabled(self) -> bool | None: + """Whether the completed object uses an S3 Bucket Key (``BucketKeyEnabled``).""" return self._bucket_key_enabled @property def request_charged(self) -> str | None: + """The ``RequestCharged`` field of the response.""" return self._request_charged def to_dict(self): + """Convert the response to a dictionary. + + Returns: + Deep copy of the instance attributes, keyed by their attribute + names (e.g., ``_etag``). + """ return copy.deepcopy(self.__dict__) diff --git a/pyathena/formatter.py b/pyathena/formatter.py index b1c686e7a..632c8e676 100644 --- a/pyathena/formatter.py +++ b/pyathena/formatter.py @@ -1,3 +1,5 @@ +"""Formatting of query parameters as SQL literals, and wrapping of queries in UNLOAD.""" + from __future__ import annotations import logging @@ -47,6 +49,12 @@ def __init__( mappings: dict[type[Any], Callable[[Formatter, Callable[[str], str], Any], Any]], default: Callable[[Formatter, Callable[[str], str], Any], Any] | None = None, ) -> None: + """Initialize the formatter. + + Args: + mappings: Formatting functions keyed by Python type. + default: Formatting function for types not in ``mappings``. + """ self._mappings = mappings self._default = default @@ -77,18 +85,43 @@ def set( type_: type[Any], formatter: Callable[[Formatter, Callable[[str], str], Any], Any], ) -> None: + """Set the formatting function for a Python type. + + Args: + type_: The Python type. + formatter: The formatting function to use for this type. + """ self.mappings[type_] = formatter def remove(self, type_: type[Any]) -> None: + """Remove the formatting function for a Python type. + + Args: + type_: The Python type to remove. + """ self.mappings.pop(type_, None) def update( self, mappings: dict[type[Any], Callable[[Formatter, Callable[[str], str], Any], Any]] ) -> None: + """Update multiple formatting functions at once. + + Args: + mappings: Dictionary of Python types to formatting functions. + """ self.mappings.update(mappings) @abstractmethod def format(self, operation: str, parameters: dict[str, Any] | None = None) -> str: + """Format a query with its parameters. + + Args: + operation: SQL query string. + parameters: Query parameters. + + Returns: + The formatted query. + """ raise NotImplementedError # pragma: no cover @staticmethod @@ -395,6 +428,7 @@ class DefaultParameterFormatter(Formatter): """ def __init__(self) -> None: + """Initialize the formatter with the default formatting functions and no default.""" super().__init__(mappings=deepcopy(_DEFAULT_FORMATTERS), default=None) @override diff --git a/pyathena/glue.py b/pyathena/glue.py index 56147943d..7ac57e757 100644 --- a/pyathena/glue.py +++ b/pyathena/glue.py @@ -54,6 +54,15 @@ def __init__( config: Config | None, client_kwargs: Mapping[str, Any], ) -> None: + """Initialize the client without building the Glue client. + + Args: + session: The connection's boto3 session. + region_name: The connection's region. + config: The connection's botocore config. + client_kwargs: The connection's client arguments. Athena's + ``endpoint_url`` and ``api_version`` are not passed to Glue. + """ self._session = session self._region_name = region_name self._config = config diff --git a/pyathena/model.py b/pyathena/model.py index 6cbb335db..d76872703 100644 --- a/pyathena/model.py +++ b/pyathena/model.py @@ -1,3 +1,5 @@ +"""Model classes that wrap Amazon Athena API responses and table format constants.""" + from __future__ import annotations import logging @@ -66,6 +68,15 @@ class AthenaQueryExecution: S3_ACL_OPTION_BUCKET_OWNER_FULL_CONTROL = "BUCKET_OWNER_FULL_CONTROL" def __init__(self, response: dict[str, Any]) -> None: + """Initialize the query execution from a ``GetQueryExecution`` response. + + Args: + response: The API response containing a ``QueryExecution`` object. + + Raises: + DataError: If ``QueryExecution``, ``QueryExecutionId``, ``Query``, + or ``Status`` is missing from the response. + """ query_execution = response.get("QueryExecution") if not query_execution: raise DataError("KeyError `QueryExecution`") @@ -165,162 +176,202 @@ def __init__(self, response: dict[str, Any]) -> None: @property def database(self) -> str | None: + """The ``Database`` of the query execution context.""" return self._database @property def catalog(self) -> str | None: + """The ``Catalog`` of the query execution context.""" return self._catalog @property def query_id(self) -> str | None: + """The ``QueryExecutionId`` of the query.""" return self._query_id @property def query(self) -> str | None: + """The ``Query`` string that was executed.""" return self._query @property def statement_type(self) -> str | None: + """The ``StatementType`` of the query, such as ``DDL`` or ``DML``.""" return self._statement_type @property def substatement_type(self) -> str | None: + """The ``SubstatementType`` of the query.""" return self._substatement_type @property def work_group(self) -> str | None: + """The ``WorkGroup`` in which the query ran.""" return self._work_group @property def execution_parameters(self) -> list[str]: + """The ``ExecutionParameters`` of the query, or an empty list.""" return self._execution_parameters @property def state(self) -> str | None: + """The ``State`` of the query execution, such as ``RUNNING`` or ``SUCCEEDED``.""" return self._state @property def state_change_reason(self) -> str | None: + """The ``StateChangeReason`` of the query execution status.""" return self._state_change_reason @property def submission_date_time(self) -> datetime | None: + """The ``SubmissionDateTime`` of the query.""" return self._submission_date_time @property def completion_date_time(self) -> datetime | None: + """The ``CompletionDateTime`` of the query.""" return self._completion_date_time @property def error_category(self) -> int | None: + """The ``ErrorCategory`` of the ``AthenaError``.""" return self._error_category @property def error_type(self) -> int | None: + """The ``ErrorType`` of the ``AthenaError``.""" return self._error_type @property def retryable(self) -> bool | None: + """The ``Retryable`` flag of the ``AthenaError``.""" return self._retryable @property def error_message(self) -> str | None: + """The ``ErrorMessage`` of the ``AthenaError``.""" return self._error_message @property def data_scanned_in_bytes(self) -> int | None: + """The ``DataScannedInBytes`` statistic of the query.""" return self._data_scanned_in_bytes @property def engine_execution_time_in_millis(self) -> int | None: + """The ``EngineExecutionTimeInMillis`` statistic of the query.""" return self._engine_execution_time_in_millis @property def query_queue_time_in_millis(self) -> int | None: + """The ``QueryQueueTimeInMillis`` statistic of the query.""" return self._query_queue_time_in_millis @property def total_execution_time_in_millis(self) -> int | None: + """The ``TotalExecutionTimeInMillis`` statistic of the query.""" return self._total_execution_time_in_millis @property def query_planning_time_in_millis(self) -> int | None: + """The ``QueryPlanningTimeInMillis`` statistic of the query.""" return self._query_planning_time_in_millis @property def service_pre_processing_time_in_millis(self) -> int | None: + """The ``ServicePreProcessingTimeInMillis`` statistic of the query.""" return self._service_pre_processing_time_in_millis @property def service_processing_time_in_millis(self) -> int | None: + """The ``ServiceProcessingTimeInMillis`` statistic of the query.""" return self._service_processing_time_in_millis @property def dpu_count(self) -> float | None: + """The ``DpuCount`` statistic of the query.""" return self._dpu_count @property def output_location(self) -> str | None: + """The ``OutputLocation`` of the result configuration.""" return self._output_location @property def data_manifest_location(self) -> str | None: + """The ``DataManifestLocation`` statistic of the query.""" return self._data_manifest_location @property def reused_previous_result(self) -> bool | None: + """The ``ReusedPreviousResult`` flag of the result reuse information.""" return self._reused_previous_result @property def encryption_option(self) -> str | None: + """The ``EncryptionOption`` of the result encryption configuration.""" return self._encryption_option @property def kms_key(self) -> str | None: + """The ``KmsKey`` of the result encryption configuration.""" return self._kms_key @property def expected_bucket_owner(self) -> str | None: + """The ``ExpectedBucketOwner`` of the result configuration.""" return self._expected_bucket_owner @property def s3_acl_option(self) -> str | None: + """The ``S3AclOption`` of the result ACL configuration.""" return self._s3_acl_option @property def selected_engine_version(self) -> str | None: + """The ``SelectedEngineVersion`` of the query.""" return self._selected_engine_version @property def effective_engine_version(self) -> str | None: + """The ``EffectiveEngineVersion`` of the query.""" return self._effective_engine_version @property def result_reuse_enabled(self) -> bool | None: + """The ``Enabled`` flag of the result reuse by age configuration.""" return self._result_reuse_enabled @property def result_reuse_minutes(self) -> int | None: + """The ``MaxAgeInMinutes`` of the result reuse by age configuration.""" return self._result_reuse_minutes @property def managed_query_results_enabled(self) -> bool | None: + """The ``Enabled`` flag of the managed query results configuration.""" return self._managed_query_results_enabled @property def managed_query_results_kms_key(self) -> str | None: + """The ``KmsKey`` of the managed query results encryption configuration.""" return self._managed_query_results_kms_key @property def enable_s3_access_grants(self) -> bool | None: + """The ``EnableS3AccessGrants`` flag of the S3 Access Grants configuration.""" return self._enable_s3_access_grants @property def create_user_level_prefix(self) -> bool | None: + """The ``CreateUserLevelPrefix`` flag of the S3 Access Grants configuration.""" return self._create_user_level_prefix @property def s3_access_grants_authentication_type(self) -> str | None: + """The ``AuthenticationType`` of the S3 Access Grants configuration.""" return self._s3_access_grants_authentication_type @@ -357,6 +408,14 @@ class AthenaCalculationExecutionStatus: TERMINAL_STATES: tuple[str, ...] = (STATE_COMPLETED, STATE_FAILED, STATE_CANCELED) def __init__(self, response: dict[str, Any]) -> None: + """Initialize the calculation status from an Athena API response. + + Args: + response: The API response containing ``Status`` and ``Statistics`` objects. + + Raises: + DataError: If ``Status`` or ``Statistics`` is missing from the response. + """ status = response.get("Status") if not status: raise DataError("KeyError `Status`") @@ -373,26 +432,32 @@ def __init__(self, response: dict[str, Any]) -> None: @property def state(self) -> str | None: + """The ``State`` of the calculation, such as ``RUNNING`` or ``COMPLETED``.""" return self._state @property def state_change_reason(self) -> str | None: + """The ``StateChangeReason`` of the calculation status.""" return self._state_change_reason @property def submission_date_time(self) -> datetime | None: + """The ``SubmissionDateTime`` of the calculation.""" return self._submission_date_time @property def completion_date_time(self) -> datetime | None: + """The ``CompletionDateTime`` of the calculation.""" return self._completion_date_time @property def dpu_execution_in_millis(self) -> int | None: + """The ``DpuExecutionInMillis`` statistic of the calculation.""" return self._dpu_execution_in_millis @property def progress(self) -> str | None: + """The ``Progress`` statistic of the calculation.""" return self._progress @@ -412,6 +477,16 @@ class AthenaCalculationExecution(AthenaCalculationExecutionStatus): """ def __init__(self, response: dict[str, Any]) -> None: + """Initialize the calculation execution from a ``GetCalculationExecution`` response. + + Args: + response: The API response containing the calculation fields, ``Status``, + ``Statistics``, and an optional ``Result`` object. + + Raises: + DataError: If ``Status``, ``Statistics``, ``CalculationExecutionId``, + or ``SessionId`` is missing from the response. + """ super().__init__(response) self._calculation_id: str | None = response.get("CalculationExecutionId") @@ -432,34 +507,42 @@ def __init__(self, response: dict[str, Any]) -> None: @property def calculation_id(self) -> str | None: + """The ``CalculationExecutionId`` of the calculation.""" return self._calculation_id @property def session_id(self) -> str | None: + """The ``SessionId`` of the session that ran the calculation.""" return self._session_id @property def description(self) -> str | None: + """The ``Description`` of the calculation.""" return self._description @property def working_directory(self) -> str | None: + """The ``WorkingDirectory`` of the calculation.""" return self._working_directory @property def std_out_s3_uri(self) -> str | None: + """The ``StdOutS3Uri`` of the calculation result.""" return self._std_out_s3_uri @property def std_error_s3_uri(self) -> str | None: + """The ``StdErrorS3Uri`` of the calculation result.""" return self._std_error_s3_uri @property def result_s3_uri(self) -> str | None: + """The ``ResultS3Uri`` of the calculation result.""" return self._result_s3_uri @property def result_type(self) -> str | None: + """The ``ResultType`` of the calculation result.""" return self._result_type @@ -495,6 +578,14 @@ class AthenaSessionStatus: STATE_FAILED: str = "FAILED" def __init__(self, response: dict[str, Any]) -> None: + """Initialize the session status from an Athena API response. + + Args: + response: The API response containing ``SessionId`` and a ``Status`` object. + + Raises: + DataError: If ``Status`` is missing from the response. + """ self._session_id: str | None = response.get("SessionId") status = response.get("Status") @@ -509,30 +600,37 @@ def __init__(self, response: dict[str, Any]) -> None: @property def session_id(self) -> str | None: + """The ``SessionId`` of the session.""" return self._session_id @property def state(self) -> str | None: + """The ``State`` of the session, such as ``IDLE`` or ``BUSY``.""" return self._state @property def state_change_reason(self) -> str | None: + """The ``StateChangeReason`` of the session status.""" return self._state_change_reason @property def start_date_time(self) -> datetime | None: + """The ``StartDateTime`` of the session.""" return self._start_date_time @property def last_modified_date_time(self) -> datetime | None: + """The ``LastModifiedDateTime`` of the session.""" return self._last_modified_date_time @property def end_date_time(self) -> datetime | None: + """The ``EndDateTime`` of the session.""" return self._end_date_time @property def idle_since_date_time(self) -> datetime | None: + """The ``IdleSinceDateTime`` of the session.""" return self._idle_since_date_time @@ -549,6 +647,14 @@ class AthenaDatabase: """ def __init__(self, response): + """Initialize the database from an Athena API response. + + Args: + response: A dictionary containing a ``Database`` object. + + Raises: + DataError: If ``Database`` is missing from the response. + """ database = response.get("Database") if not database: raise DataError("KeyError `Database`") @@ -559,14 +665,17 @@ def __init__(self, response): @property def name(self) -> str | None: + """The ``Name`` of the database.""" return self._name @property def description(self) -> str | None: + """The ``Description`` of the database.""" return self._description @property def parameters(self) -> dict[str, str]: + """The ``Parameters`` of the database, or an empty dictionary.""" return self._parameters @@ -582,20 +691,28 @@ class AthenaTableMetadataColumn: """ def __init__(self, response): + """Initialize the column from an Athena ``Column`` object. + + Args: + response: The ``Column`` object with ``Name``, ``Type``, and ``Comment``. + """ self._name: str | None = response.get("Name") self._type: str | None = response.get("Type") self._comment: str | None = response.get("Comment") @property def name(self) -> str | None: + """The ``Name`` of the column.""" return self._name @property def type(self) -> str | None: + """The ``Type`` of the column.""" return self._type @property def comment(self) -> str | None: + """The ``Comment`` of the column.""" return self._comment @@ -612,20 +729,28 @@ class AthenaTableMetadataPartitionKey: """ def __init__(self, response): + """Initialize the partition key from an Athena ``Column`` object. + + Args: + response: The ``Column`` object with ``Name``, ``Type``, and ``Comment``. + """ self._name: str | None = response.get("Name") self._type: str | None = response.get("Type") self._comment: str | None = response.get("Comment") @property def name(self) -> str | None: + """The ``Name`` of the partition key.""" return self._name @property def type(self) -> str | None: + """The ``Type`` of the partition key.""" return self._type @property def comment(self) -> str | None: + """The ``Comment`` of the partition key.""" return self._comment @@ -645,6 +770,14 @@ class AthenaTableMetadata: """ def __init__(self, response): + """Initialize the table metadata from an Athena API response. + + Args: + response: A dictionary containing a ``TableMetadata`` object. + + Raises: + DataError: If ``TableMetadata`` is missing from the response. + """ table_metadata = response.get("TableMetadata") if not table_metadata: raise DataError("KeyError `TableMetadata`") @@ -668,50 +801,62 @@ def __init__(self, response): @property def name(self) -> str | None: + """The ``Name`` of the table.""" return self._name @property def create_time(self) -> datetime | None: + """The ``CreateTime`` of the table.""" return self._create_time @property def last_access_time(self) -> datetime | None: + """The ``LastAccessTime`` of the table.""" return self._last_access_time @property def table_type(self) -> str | None: + """The ``TableType`` of the table.""" return self._table_type @property def columns(self) -> list[AthenaTableMetadataColumn]: + """The ``Columns`` of the table.""" return self._columns @property def partition_keys(self) -> list[AthenaTableMetadataPartitionKey]: + """The ``PartitionKeys`` of the table.""" return self._partition_keys @property def parameters(self) -> dict[str, str]: + """The ``Parameters`` of the table, or an empty dictionary.""" return self._parameters @property def comment(self) -> str | None: + """The ``comment`` table parameter.""" return self._parameters.get("comment") @property def location(self) -> str | None: + """The ``location`` table parameter.""" return self._parameters.get("location") @property def input_format(self) -> str | None: + """The ``inputformat`` table parameter.""" return self._parameters.get("inputformat") @property def output_format(self) -> str | None: + """The ``outputformat`` table parameter.""" return self._parameters.get("outputformat") @property def row_format(self) -> str | None: + """The ``SERDE ''`` clause built from ``serde_serialization_lib``, or ``None``.""" serde = self.serde_serialization_lib if serde: return f"SERDE '{serde}'" @@ -719,6 +864,7 @@ def row_format(self) -> str | None: @property def file_format(self) -> str | None: + """The ``INPUTFORMAT '...' OUTPUTFORMAT '...'`` clause, or ``None`` unless both are set.""" input = self.input_format output = self.output_format if input and output: @@ -727,10 +873,16 @@ def file_format(self) -> str | None: @property def serde_serialization_lib(self) -> str | None: + """The ``serde.serialization.lib`` table parameter.""" return self._parameters.get("serde.serialization.lib") @property def compression(self) -> str | None: + """The compression codec from the table parameters, or ``None``. + + The first parameter present is used, in the order ``write.compression``, + ``serde.param.write.compression``, ``parquet.compress``, and ``orc.compress``. + """ if "write.compression" in self._parameters: # text or json return self._parameters["write.compression"] if "serde.param.write.compression" in self._parameters: # text or json @@ -743,6 +895,7 @@ def compression(self) -> str | None: @property def serde_properties(self) -> dict[str, str]: + """The ``serde.param.``-prefixed table parameters with the prefix removed.""" return { k.replace("serde.param.", ""): v for k, v in self._parameters.items() @@ -751,6 +904,7 @@ def serde_properties(self) -> dict[str, str]: @property def table_properties(self) -> dict[str, str]: + """The table parameters that do not start with ``serde.param.``.""" return {k: v for k, v in self._parameters.items() if not k.startswith("serde.param.")} @@ -797,10 +951,26 @@ class AthenaFileFormat: @staticmethod def is_parquet(value: str) -> bool: + """Check whether a file format name is ``PARQUET``, ignoring case. + + Args: + value: The file format name. + + Returns: + True if the value is ``PARQUET``, False otherwise. + """ return value.upper() == AthenaFileFormat.FILE_FORMAT_PARQUET @staticmethod def is_orc(value: str) -> bool: + """Check whether a file format name is ``ORC``, ignoring case. + + Args: + value: The file format name. + + Returns: + True if the value is ``ORC``, False otherwise. + """ return value.upper() == AthenaFileFormat.FILE_FORMAT_ORC @@ -846,6 +1016,15 @@ class AthenaRowFormatSerde: @staticmethod def is_parquet(value: str) -> bool: + """Check whether a ``SERDE ''`` row format uses the Parquet SerDe. + + Args: + value: The row format string, such as the value of + ``AthenaTableMetadata.row_format``. + + Returns: + True if the SerDe is ``ROW_FORMAT_SERDE_PARQUET``, False otherwise. + """ match = AthenaRowFormatSerde.PATTERN_ROW_FORMAT_SERDE.search(value) if match: serde = match.group("serde") @@ -855,6 +1034,15 @@ def is_parquet(value: str) -> bool: @staticmethod def is_orc(value: str) -> bool: + """Check whether a ``SERDE ''`` row format uses the ORC SerDe. + + Args: + value: The row format string, such as the value of + ``AthenaTableMetadata.row_format``. + + Returns: + True if the SerDe is ``ROW_FORMAT_SERDE_ORC``, False otherwise. + """ match = AthenaRowFormatSerde.PATTERN_ROW_FORMAT_SERDE.search(value) if match: serde = match.group("serde") @@ -911,6 +1099,15 @@ class AthenaCompression: @staticmethod def is_valid(value: str) -> bool: + """Check whether a value is a supported compression format, ignoring case. + + Args: + value: The compression format name. + + Returns: + True if the value matches one of the ``COMPRESSION_*`` constants, + False otherwise. + """ return value.upper() in [ AthenaCompression.COMPRESSION_BZIP2, AthenaCompression.COMPRESSION_DEFLATE, @@ -963,6 +1160,15 @@ class AthenaPartitionTransform: @staticmethod def is_valid(value: str) -> bool: + """Check whether a value is a supported partition transform, ignoring case. + + Args: + value: The partition transform name. + + Returns: + True if the value matches one of the ``PARTITION_TRANSFORM_*`` constants, + False otherwise. + """ return value.lower() in [ AthenaPartitionTransform.PARTITION_TRANSFORM_YEAR, AthenaPartitionTransform.PARTITION_TRANSFORM_MONTH, diff --git a/pyathena/options.py b/pyathena/options.py index adcb9ec7a..33d95bf9e 100644 --- a/pyathena/options.py +++ b/pyathena/options.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""Options shared by the ``execute()`` methods of the SQL cursors.""" + from __future__ import annotations from collections.abc import Callable @@ -54,8 +56,9 @@ class ExecuteOptions: falls back to the connection-level setting. paramstyle: Parameter style for this query ('qmark' or 'pyformat'). None (default) uses the module-level ``pyathena.paramstyle``. - on_start_query_execution: Callback invoked with the query ID - immediately after the StartQueryExecution API call. Invoked by + on_start_query_execution: Callback invoked with the query ID before + ``execute()`` waits for the query: after the StartQueryExecution API + call, or after a reusable query ID is found through ``cache_size``. Invoked by synchronous and aio cursors; ``AsyncCursor``-based cursors return the query ID directly through their execution model and do not invoke it. diff --git a/pyathena/pandas/__init__.py b/pyathena/pandas/__init__.py index 5c077197a..23392ba1a 100644 --- a/pyathena/pandas/__init__.py +++ b/pyathena/pandas/__init__.py @@ -1,3 +1,5 @@ +"""Cursors that return Athena query results as pandas DataFrames.""" + from pyathena.filesystem import register_s3_filesystem register_s3_filesystem() diff --git a/pyathena/pandas/async_cursor.py b/pyathena/pandas/async_cursor.py index a98927a83..bd2279ea6 100644 --- a/pyathena/pandas/async_cursor.py +++ b/pyathena/pandas/async_cursor.py @@ -1,3 +1,5 @@ +"""Asynchronous cursor that returns Athena query results as pandas DataFrames.""" + from __future__ import annotations import logging @@ -80,6 +82,29 @@ def __init__( result_reuse_minutes: int = CursorIterator.DEFAULT_RESULT_REUSE_MINUTES, **kwargs, ) -> None: + """Initialize an AsyncPandasCursor. + + Args: + s3_staging_dir: S3 location for query results. + schema_name: Default schema name. + catalog_name: Default catalog name. + work_group: Athena workgroup name. + poll_interval: Query status polling interval in seconds. + encryption_option: S3 encryption option for query results. + kms_key: KMS key for encrypting query results. + kill_on_interrupt: Cancel a query whose start in ``execute()`` is interrupted by + ``KeyboardInterrupt``. Waiting runs on worker threads, which do not + receive the interrupt. + max_workers: Maximum number of threads that run queries concurrently. + arraysize: Number of rows to fetch per batch. Must be a positive integer. + unload: Whether to wrap queries in ``UNLOAD`` and read the Parquet output. + engine: Parsing engine (``auto``, ``c``, ``python``, or ``pyarrow``). + chunksize: Number of rows per DataFrame chunk when reading CSV results. + result_reuse_enable: Whether to enable Athena query result reuse. + result_reuse_minutes: Maximum age of a reused query result in minutes. + **kwargs: Other cursor arguments, such as ``connection`` and ``converter``, + passed to ``AsyncCursor.__init__``. + """ super().__init__( s3_staging_dir=s3_staging_dir, schema_name=schema_name, diff --git a/pyathena/pandas/converter.py b/pyathena/pandas/converter.py index 54f97ca9b..a80975871 100644 --- a/pyathena/pandas/converter.py +++ b/pyathena/pandas/converter.py @@ -1,3 +1,5 @@ +"""Type converters for pandas cursor results.""" + from __future__ import annotations import logging @@ -53,6 +55,7 @@ class DefaultPandasTypeConverter(Converter): """ def __init__(self) -> None: + """Initialize the converter with the default pandas conversion functions and dtypes.""" super().__init__( mappings=deepcopy(_DEFAULT_PANDAS_CONVERTERS), default=_to_default, @@ -101,6 +104,7 @@ class DefaultPandasUnloadTypeConverter(Converter): """ def __init__(self) -> None: + """Initialize the converter with no type mappings.""" super().__init__( mappings={}, default=_to_default, diff --git a/pyathena/pandas/cursor.py b/pyathena/pandas/cursor.py index f574a2d30..81738492a 100644 --- a/pyathena/pandas/cursor.py +++ b/pyathena/pandas/cursor.py @@ -1,3 +1,5 @@ +"""Cursor that returns Athena query results as pandas DataFrames.""" + from __future__ import annotations import logging @@ -178,7 +180,9 @@ def execute( keep_default_na: Whether to keep default pandas NA values. na_values: Additional values to treat as NA. quoting: CSV quoting behavior (pandas csv.QUOTE_* constants). - on_start_query_execution: Callback called when query starts. + on_start_query_execution: Callback invoked with the query ID before ``execute()`` + waits for the query: after the ``StartQueryExecution`` call, or after a + reusable query ID is found through ``cache_size``. result_set_type_hints: Optional dictionary mapping column names to Athena DDL type signatures for precise type conversion within complex types. diff --git a/pyathena/pandas/reader.py b/pyathena/pandas/reader.py index bfdcc4777..da346b6d0 100644 --- a/pyathena/pandas/reader.py +++ b/pyathena/pandas/reader.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""Raw CSV stream that marks binary NULL fields for the pandas CSV reader.""" + from __future__ import annotations import re @@ -27,6 +29,13 @@ class BinaryCSVReader(RawIOBase): """ def __init__(self, stream: Any, binary_columns: set[int]) -> None: + """Initialize the reader. + + Args: + stream: Text stream of Athena CSV output, read with ``AthenaCSVReader``. + binary_columns: Zero-based indexes of the binary columns whose unquoted + empty fields are replaced with the binary NULL marker. + """ super().__init__() self._reader = AthenaCSVReader(stream) self._binary_columns = binary_columns diff --git a/pyathena/pandas/result_set.py b/pyathena/pandas/result_set.py index 1b36ad027..f7f488d1d 100644 --- a/pyathena/pandas/result_set.py +++ b/pyathena/pandas/result_set.py @@ -1,3 +1,5 @@ +"""Result set that reads Athena query results into pandas DataFrames.""" + from __future__ import annotations import csv @@ -463,6 +465,7 @@ def dtypes(self) -> dict[str, type[Any]]: def converters( self, ) -> dict[Any | None, Callable[[str | None], Any | None]]: + """The conversion functions for the result columns the converter maps, keyed by name.""" description = self.description if self.description else [] return { d[0]: self._converter.get(d[1]) for d in description if d[1] in self._converter.mappings @@ -470,6 +473,7 @@ def converters( @property def parse_dates(self) -> list[Any | None]: + """The names of the result columns with date, time, or timestamp types.""" description = self.description if self.description else [] return [d[0] for d in description if d[1] in self._PARSE_DATES] @@ -809,6 +813,13 @@ def _as_pandas_from_api(self, converter: Converter | None = None) -> DataFrame: return pd.DataFrame(self._rows_to_columnar(rows, columns)) def as_pandas(self) -> PandasDataFrameIterator | DataFrame: + """Return the query results as a DataFrame or an iterator of DataFrame chunks. + + Returns: + If ``chunksize`` is None, the next DataFrame from the result iterator, which + holds the whole result unless ``auto_optimize_chunksize`` chose a chunk size; + otherwise the ``PandasDataFrameIterator`` that yields DataFrame chunks. + """ if self._chunksize is None: return next(self._df_iter) return self._df_iter diff --git a/pyathena/pandas/util.py b/pyathena/pandas/util.py index 9337ce326..e12aef7f2 100644 --- a/pyathena/pandas/util.py +++ b/pyathena/pandas/util.py @@ -1,3 +1,5 @@ +"""Helpers that convert query results to pandas DataFrames and write DataFrames to Athena.""" + from __future__ import annotations import concurrent diff --git a/pyathena/parser.py b/pyathena/parser.py index 5d3635cf0..1694cebc3 100644 --- a/pyathena/parser.py +++ b/pyathena/parser.py @@ -1,3 +1,5 @@ +"""Parsing of Athena type signatures and conversion of values by parsed type.""" + from __future__ import annotations import json @@ -231,6 +233,13 @@ def __init__( default_converter: Callable[[str | None], Any | None], struct_parser: Callable[[str | None], dict[str, Any] | str | None], ) -> None: + """Initialize the converter. + + Args: + converters: Mapping of type names to conversion functions. + default_converter: Fallback conversion function for unknown types. + struct_parser: Function to parse untyped struct values. + """ self._converters = converters self._default_converter = default_converter self._struct_parser = struct_parser diff --git a/pyathena/polars/__init__.py b/pyathena/polars/__init__.py index 5c077197a..6f57d0aed 100644 --- a/pyathena/polars/__init__.py +++ b/pyathena/polars/__init__.py @@ -1,3 +1,5 @@ +"""Cursors that return Athena query results as Polars DataFrames.""" + from pyathena.filesystem import register_s3_filesystem register_s3_filesystem() diff --git a/pyathena/polars/async_cursor.py b/pyathena/polars/async_cursor.py index 719475547..9e42ca78d 100644 --- a/pyathena/polars/async_cursor.py +++ b/pyathena/polars/async_cursor.py @@ -1,3 +1,5 @@ +"""Asynchronous cursor that returns Athena query results as Polars DataFrames.""" + from __future__ import annotations import logging diff --git a/pyathena/polars/converter.py b/pyathena/polars/converter.py index 3b3835288..9cf87718a 100644 --- a/pyathena/polars/converter.py +++ b/pyathena/polars/converter.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""Type converters for Polars cursor results.""" + from __future__ import annotations import logging @@ -59,6 +61,7 @@ class DefaultPolarsTypeConverter(Converter): """ def __init__(self) -> None: + """Initialize the converter with the default Polars conversion functions and dtypes.""" super().__init__( mappings=deepcopy(_DEFAULT_POLARS_CONVERTERS), default=_to_default, @@ -132,6 +135,7 @@ class DefaultPolarsUnloadTypeConverter(Converter): """ def __init__(self) -> None: + """Initialize the converter with no type mappings.""" super().__init__( mappings={}, default=_to_default, diff --git a/pyathena/polars/cursor.py b/pyathena/polars/cursor.py index 233cdd7c7..53288bf91 100644 --- a/pyathena/polars/cursor.py +++ b/pyathena/polars/cursor.py @@ -1,3 +1,5 @@ +"""Cursor that returns Athena query results as Polars DataFrames.""" + from __future__ import annotations import logging @@ -178,7 +180,9 @@ def execute( result_reuse_enable: Enable Athena result reuse for this query. result_reuse_minutes: Minutes to reuse cached results. paramstyle: Parameter style ('qmark' or 'pyformat'). - on_start_query_execution: Callback called when query starts. + on_start_query_execution: Callback invoked with the query ID before ``execute()`` + waits for the query: after the ``StartQueryExecution`` call, or after a + reusable query ID is found through ``cache_size``. result_set_type_hints: Optional dictionary mapping column names to Athena DDL type signatures for precise type conversion within complex types. diff --git a/pyathena/polars/result_set.py b/pyathena/polars/result_set.py index 4fc9e1014..237b4ac6f 100644 --- a/pyathena/polars/result_set.py +++ b/pyathena/polars/result_set.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""Result set that reads Athena query results into Polars DataFrames.""" + from __future__ import annotations import logging diff --git a/pyathena/result_set.py b/pyathena/result_set.py index 8b6f0e3f1..152715fdb 100644 --- a/pyathena/result_set.py +++ b/pyathena/result_set.py @@ -1,3 +1,5 @@ +"""Result sets for ``GetQueryResults`` and the cursor mixins that expose them.""" + from __future__ import annotations import collections @@ -67,6 +69,24 @@ def __init__( _pre_fetch: bool = True, result_set_type_hints: dict[str | int, str] | None = None, ) -> None: + """Initialize the result set and fetch the first page if the query succeeded. + + Args: + connection: The connection that ran the query. + converter: The converter for result values. + query_execution: The query execution whose results to read. + arraysize: The number of rows per ``GetQueryResults`` page and the default + ``fetchmany()`` size. + retry_config: The retry configuration for API calls. + _pre_fetch: Whether to fetch the first page here when the query succeeded. + The async result set passes False and fetches it itself. + result_set_type_hints: Athena type signatures for complex-type columns, + keyed by column name (case-insensitive) or zero-based column index. + + Raises: + ProgrammingError: If ``query_execution`` is not given. + OperationalError: If fetching the first page fails. + """ super().__init__(arraysize=arraysize) self._connection: Connection[Any] | None = connection self._converter = converter @@ -105,150 +125,175 @@ def __init__( @property def database(self) -> str | None: + """The database in the ``QueryExecutionContext`` of the query.""" if not self._query_execution: return None return self._query_execution.database @property def catalog(self) -> str | None: + """The data catalog in the ``QueryExecutionContext`` of the query.""" if not self._query_execution: return None return self._query_execution.catalog @property def query_id(self) -> str | None: + """The ID of the query execution.""" if not self._query_execution: return None return self._query_execution.query_id @property def query(self) -> str | None: + """The SQL statement that the query execution ran.""" if not self._query_execution: return None return self._query_execution.query @property def statement_type(self) -> str | None: + """The ``StatementType`` of the query, such as ``DDL``, ``DML``, or ``UTILITY``.""" if not self._query_execution: return None return self._query_execution.statement_type @property def substatement_type(self) -> str | None: + """The ``SubstatementType`` of the query, such as ``INSERT`` or ``MERGE``.""" if not self._query_execution: return None return self._query_execution.substatement_type @property def work_group(self) -> str | None: + """The work group in which the query ran.""" if not self._query_execution: return None return self._query_execution.work_group @property def execution_parameters(self) -> list[str]: + """The ``ExecutionParameters`` values of the query.""" if not self._query_execution: return [] return self._query_execution.execution_parameters @property def state(self) -> str | None: + """The state of the query execution, such as ``RUNNING`` or ``SUCCEEDED``.""" if not self._query_execution: return None return self._query_execution.state @property def state_change_reason(self) -> str | None: + """The ``StateChangeReason`` that gives further detail about the state.""" if not self._query_execution: return None return self._query_execution.state_change_reason @property def submission_date_time(self) -> datetime | None: + """The date and time when the query was submitted.""" if not self._query_execution: return None return self._query_execution.submission_date_time @property def completion_date_time(self) -> datetime | None: + """The date and time when the query completed.""" if not self._query_execution: return None return self._query_execution.completion_date_time @property def error_category(self) -> int | None: + """The ``ErrorCategory`` of the failure: 1 for system, 2 for user, 3 for other.""" if not self._query_execution: return None return self._query_execution.error_category @property def error_type(self) -> int | None: + """The ``ErrorType`` code of the query failure.""" if not self._query_execution: return None return self._query_execution.error_type @property def retryable(self) -> bool | None: + """Whether Athena reports the query failure as retryable.""" if not self._query_execution: return None return self._query_execution.retryable @property def error_message(self) -> str | None: + """The ``ErrorMessage`` that describes the query failure.""" if not self._query_execution: return None return self._query_execution.error_message @property def data_scanned_in_bytes(self) -> int | None: + """The number of bytes that the query scanned.""" if not self._query_execution: return None return self._query_execution.data_scanned_in_bytes @property def engine_execution_time_in_millis(self) -> int | None: + """The time in milliseconds that the query engine took to run the query.""" if not self._query_execution: return None return self._query_execution.engine_execution_time_in_millis @property def query_queue_time_in_millis(self) -> int | None: + """The time in milliseconds that the query waited in the queue.""" if not self._query_execution: return None return self._query_execution.query_queue_time_in_millis @property def total_execution_time_in_millis(self) -> int | None: + """The total time in milliseconds that Athena took to run the query.""" if not self._query_execution: return None return self._query_execution.total_execution_time_in_millis @property def query_planning_time_in_millis(self) -> int | None: + """The time in milliseconds that Athena took to plan the query.""" if not self._query_execution: return None return self._query_execution.query_planning_time_in_millis @property def service_processing_time_in_millis(self) -> int | None: + """The time in milliseconds that Athena took to publish the query results.""" if not self._query_execution: return None return self._query_execution.service_processing_time_in_millis @property def output_location(self) -> str | None: + """The S3 location of the query results.""" if not self._query_execution: return None return self._query_execution.output_location @property def data_manifest_location(self) -> str | None: + """The S3 location of the data manifest that lists the files the query wrote.""" if not self._query_execution: return None return self._query_execution.data_manifest_location @property def reused_previous_result(self) -> bool | None: + """Whether Athena reused a previous query result instead of running the query.""" if not self._query_execution: return None return self._query_execution.reused_previous_result @@ -268,48 +313,56 @@ def is_unload(self) -> bool: @property def encryption_option(self) -> str | None: + """The ``EncryptionOption`` of the query results, such as ``SSE_S3`` or ``SSE_KMS``.""" if not self._query_execution: return None return self._query_execution.encryption_option @property def kms_key(self) -> str | None: + """The KMS key used to encrypt the query results.""" if not self._query_execution: return None return self._query_execution.kms_key @property def expected_bucket_owner(self) -> str | None: + """The AWS account ID expected to own the S3 bucket of the query results.""" if not self._query_execution: return None return self._query_execution.expected_bucket_owner @property def s3_acl_option(self) -> str | None: + """The ``S3AclOption`` of the query results, such as ``BUCKET_OWNER_FULL_CONTROL``.""" if not self._query_execution: return None return self._query_execution.s3_acl_option @property def selected_engine_version(self) -> str | None: + """The Athena engine version selected to run the query.""" if not self._query_execution: return None return self._query_execution.selected_engine_version @property def effective_engine_version(self) -> str | None: + """The Athena engine version that ran the query.""" if not self._query_execution: return None return self._query_execution.effective_engine_version @property def result_reuse_enabled(self) -> bool | None: + """Whether reuse of previous query results by age is enabled for the query.""" if not self._query_execution: return None return self._query_execution.result_reuse_enabled @property def result_reuse_minutes(self) -> int | None: + """The maximum age in minutes of a previous query result that Athena can reuse.""" if not self._query_execution: return None return self._query_execution.result_reuse_minutes @@ -318,6 +371,10 @@ def result_reuse_minutes(self) -> int | None: def description( self, ) -> list[tuple[str, str, None, None, int, int, str]] | None: + """The DB API 2.0 column descriptions. + + None without result metadata, or for ``INSERT``, ``UPDATE``, ``DELETE``, and ``MERGE``. + """ if self._metadata is None or ( self.substatement_type and self.substatement_type.upper() in self._DML_SUBSTATEMENT_TYPES @@ -338,6 +395,7 @@ def description( @property def connection(self) -> Connection[Any]: + """The connection of the result set; raises ``ProgrammingError`` if closed.""" if self.is_closed: raise ProgrammingError("AthenaResultSet is closed.") return cast("Connection[Any]", self._connection) @@ -732,9 +790,11 @@ def _read_data_manifest(self) -> list[str]: @property def is_closed(self) -> bool: + """Whether the result set is closed.""" return self._connection is None def close(self) -> None: + """Close the result set and discard its query execution, metadata, and rows.""" self._connection = None self._query_execution = None self._metadata = None @@ -753,6 +813,8 @@ def __exit__(self, exc_type, exc_val, exc_tb): class AthenaDictResultSet(AthenaResultSet): + """A result set that returns each row as a dictionary keyed by column name.""" + # You can override this to use OrderedDict or other dict-like types. dict_type: type[Any] = dict @@ -862,24 +924,28 @@ def result_set(self, val: AthenaResultSet | None) -> None: @property def has_result_set(self) -> bool: + """Whether the cursor has a result set.""" return self.result_set is not None @property def description( self, ) -> list[tuple[str, str, None, None, int, int, str]] | None: + """The DB API 2.0 column descriptions of the result set, or None without one.""" if not self.result_set: return None return self.result_set.description @property def database(self) -> str | None: + """The database in the ``QueryExecutionContext`` of the query.""" if not self.result_set: return None return self.result_set.database @property def catalog(self) -> str | None: + """The data catalog in the ``QueryExecutionContext`` of the query.""" if not self.result_set: return None return self.result_set.catalog @@ -911,180 +977,210 @@ def _set_interrupted_execution_id(self, execution_id: str) -> None: @property def query(self) -> str | None: + """The SQL statement that the query execution ran.""" if not self.result_set: return None return self.result_set.query @property def statement_type(self) -> str | None: + """The ``StatementType`` of the query, such as ``DDL``, ``DML``, or ``UTILITY``.""" if not self.result_set: return None return self.result_set.statement_type @property def substatement_type(self) -> str | None: + """The ``SubstatementType`` of the query, such as ``INSERT`` or ``MERGE``.""" if not self.result_set: return None return self.result_set.substatement_type @property def work_group(self) -> str | None: + """The work group in which the query ran.""" if not self.result_set: return None return self.result_set.work_group @property def execution_parameters(self) -> list[str]: + """The ``ExecutionParameters`` values of the query.""" if not self.result_set: return [] return self.result_set.execution_parameters @property def state(self) -> str | None: + """The state of the query execution, such as ``RUNNING`` or ``SUCCEEDED``.""" if not self.result_set: return None return self.result_set.state @property def state_change_reason(self) -> str | None: + """The ``StateChangeReason`` that gives further detail about the state.""" if not self.result_set: return None return self.result_set.state_change_reason @property def submission_date_time(self) -> datetime | None: + """The date and time when the query was submitted.""" if not self.result_set: return None return self.result_set.submission_date_time @property def completion_date_time(self) -> datetime | None: + """The date and time when the query completed.""" if not self.result_set: return None return self.result_set.completion_date_time @property def error_category(self) -> int | None: + """The ``ErrorCategory`` of the failure: 1 for system, 2 for user, 3 for other.""" if not self.result_set: return None return self.result_set.error_category @property def error_type(self) -> int | None: + """The ``ErrorType`` code of the query failure.""" if not self.result_set: return None return self.result_set.error_type @property def retryable(self) -> bool | None: + """Whether Athena reports the query failure as retryable.""" if not self.result_set: return None return self.result_set.retryable @property def error_message(self) -> str | None: + """The ``ErrorMessage`` that describes the query failure.""" if not self.result_set: return None return self.result_set.error_message @property def data_scanned_in_bytes(self) -> int | None: + """The number of bytes that the query scanned.""" if not self.result_set: return None return self.result_set.data_scanned_in_bytes @property def engine_execution_time_in_millis(self) -> int | None: + """The time in milliseconds that the query engine took to run the query.""" if not self.result_set: return None return self.result_set.engine_execution_time_in_millis @property def query_queue_time_in_millis(self) -> int | None: + """The time in milliseconds that the query waited in the queue.""" if not self.result_set: return None return self.result_set.query_queue_time_in_millis @property def total_execution_time_in_millis(self) -> int | None: + """The total time in milliseconds that Athena took to run the query.""" if not self.result_set: return None return self.result_set.total_execution_time_in_millis @property def query_planning_time_in_millis(self) -> int | None: + """The time in milliseconds that Athena took to plan the query.""" if not self.result_set: return None return self.result_set.query_planning_time_in_millis @property def service_processing_time_in_millis(self) -> int | None: + """The time in milliseconds that Athena took to publish the query results.""" if not self.result_set: return None return self.result_set.service_processing_time_in_millis @property def output_location(self) -> str | None: + """The S3 location of the query results.""" if not self.result_set: return None return self.result_set.output_location @property def data_manifest_location(self) -> str | None: + """The S3 location of the data manifest that lists the files the query wrote.""" if not self.result_set: return None return self.result_set.data_manifest_location @property def reused_previous_result(self) -> bool | None: + """Whether Athena reused a previous query result instead of running the query.""" if not self.result_set: return None return self.result_set.reused_previous_result @property def encryption_option(self) -> str | None: + """The ``EncryptionOption`` of the query results, such as ``SSE_S3`` or ``SSE_KMS``.""" if not self.result_set: return None return self.result_set.encryption_option @property def kms_key(self) -> str | None: + """The KMS key used to encrypt the query results.""" if not self.result_set: return None return self.result_set.kms_key @property def expected_bucket_owner(self) -> str | None: + """The AWS account ID expected to own the S3 bucket of the query results.""" if not self.result_set: return None return self.result_set.expected_bucket_owner @property def s3_acl_option(self) -> str | None: + """The ``S3AclOption`` of the query results, such as ``BUCKET_OWNER_FULL_CONTROL``.""" if not self.result_set: return None return self.result_set.s3_acl_option @property def selected_engine_version(self) -> str | None: + """The Athena engine version selected to run the query.""" if not self.result_set: return None return self.result_set.selected_engine_version @property def effective_engine_version(self) -> str | None: + """The Athena engine version that ran the query.""" if not self.result_set: return None return self.result_set.effective_engine_version @property def result_reuse_enabled(self) -> bool | None: + """Whether reuse of previous query results by age is enabled for the query.""" if not self.result_set: return None return self.result_set.result_reuse_enabled @property def result_reuse_minutes(self) -> int | None: + """The maximum age in minutes of a previous query result that Athena can reuse.""" if not self.result_set: return None return self.result_set.result_reuse_minutes diff --git a/pyathena/s3fs/__init__.py b/pyathena/s3fs/__init__.py index e69de29bb..9e25e3ab9 100644 --- a/pyathena/s3fs/__init__.py +++ b/pyathena/s3fs/__init__.py @@ -0,0 +1,8 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT + +"""Cursors that read Athena CSV query results through PyAthena's S3 filesystem.""" diff --git a/pyathena/s3fs/async_cursor.py b/pyathena/s3fs/async_cursor.py index 67ed69638..98bac4fc3 100644 --- a/pyathena/s3fs/async_cursor.py +++ b/pyathena/s3fs/async_cursor.py @@ -1,3 +1,5 @@ +"""Asynchronous cursor that reads Athena CSV query results through PyAthena's S3 filesystem.""" + from __future__ import annotations import logging diff --git a/pyathena/s3fs/converter.py b/pyathena/s3fs/converter.py index 168029609..4e39fd75c 100644 --- a/pyathena/s3fs/converter.py +++ b/pyathena/s3fs/converter.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""Type converter for S3FS cursor results.""" + from __future__ import annotations import logging @@ -50,6 +52,7 @@ class DefaultS3FSTypeConverter(Converter): """ def __init__(self) -> None: + """Initialize the converter with the standard Athena conversion functions.""" super().__init__( mappings=deepcopy(_DEFAULT_CONVERTERS), default=_to_default, diff --git a/pyathena/s3fs/cursor.py b/pyathena/s3fs/cursor.py index 1cb04d74a..839b8c128 100644 --- a/pyathena/s3fs/cursor.py +++ b/pyathena/s3fs/cursor.py @@ -1,3 +1,5 @@ +"""Cursor that reads Athena CSV query results through PyAthena's S3 filesystem.""" + from __future__ import annotations import logging @@ -154,7 +156,9 @@ def execute( result_reuse_enable: Enable Athena result reuse for this query. result_reuse_minutes: Minutes to reuse cached results. paramstyle: Parameter style ('qmark' or 'pyformat'). - on_start_query_execution: Callback called when query starts. + on_start_query_execution: Callback invoked with the query ID before ``execute()`` + waits for the query: after the ``StartQueryExecution`` call, or after a + reusable query ID is found through ``cache_size``. result_set_type_hints: Optional dictionary mapping column names to Athena DDL type signatures for precise type conversion within complex types. diff --git a/pyathena/s3fs/reader.py b/pyathena/s3fs/reader.py index 607cb9445..098857401 100644 --- a/pyathena/s3fs/reader.py +++ b/pyathena/s3fs/reader.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""CSV readers that parse Athena query result files for the S3FS cursor.""" + from __future__ import annotations import csv diff --git a/pyathena/s3fs/result_set.py b/pyathena/s3fs/result_set.py index e0ee25bf3..151bf82fc 100644 --- a/pyathena/s3fs/result_set.py +++ b/pyathena/s3fs/result_set.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""Result set that reads Athena CSV query results through an fsspec filesystem.""" + from __future__ import annotations import logging @@ -74,6 +76,29 @@ def __init__( result_set_type_hints: dict[str | int, str] | None = None, **kwargs, ) -> None: + """Initialize the result set and prepare to read the query results. + + Args: + connection: The connection that ran the query. + converter: The converter for result values. + query_execution: The query execution whose results to read. + arraysize: The number of rows read from the CSV results per fetch and the + default ``fetchmany()`` size. + retry_config: The retry configuration for API calls. + block_size: The default block size in bytes for the filesystem. If not set, + ``DEFAULT_BLOCK_SIZE`` is used. + csv_reader: The CSV reader class for the results. If None, + ``AthenaCSVReader`` is used. + filesystem_class: The filesystem class for reading the results. If None, + PyAthena's ``S3FileSystem`` is used. + result_set_type_hints: Athena type signatures for complex-type columns, + keyed by column name (case-insensitive) or zero-based column index. + **kwargs: Additional keyword arguments, which are ignored. + + Raises: + ProgrammingError: If ``query_execution`` is not given. + OperationalError: If reading the query results fails. + """ super().__init__( connection=connection, converter=converter, diff --git a/pyathena/spark/__init__.py b/pyathena/spark/__init__.py index e69de29bb..e9abb198c 100644 --- a/pyathena/spark/__init__.py +++ b/pyathena/spark/__init__.py @@ -0,0 +1,8 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT + +"""Cursors that run PySpark code in Athena for Apache Spark sessions.""" diff --git a/pyathena/spark/async_cursor.py b/pyathena/spark/async_cursor.py index 8d12fd56c..84f35ff42 100644 --- a/pyathena/spark/async_cursor.py +++ b/pyathena/spark/async_cursor.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""Asynchronous cursor that runs PySpark code in an Athena for Apache Spark session.""" + import logging from concurrent.futures import Future, ThreadPoolExecutor from multiprocessing import cpu_count @@ -138,11 +140,29 @@ def close(self, wait: bool = False) -> None: self._executor.shutdown(wait=wait) def calculation_execution(self, query_id: str) -> "Future[AthenaCalculationExecution]": + """Get calculation execution details asynchronously. + + Args: + query_id: The calculation execution ID. + + Returns: + Future object containing the ``AthenaCalculationExecution``. + """ return self._executor.submit(self._get_calculation_execution, query_id) def get_std_out( self, calculation_execution: AthenaCalculationExecution ) -> "Future[str] | None": + """Read the standard output of a calculation from S3 asynchronously. + + Args: + calculation_execution: The calculation execution whose + ``std_out_s3_uri`` is read. + + Returns: + Future object containing the output text with leading and trailing + whitespace removed, or None if the calculation has no ``std_out_s3_uri``. + """ if not calculation_execution.std_out_s3_uri: return None return self._executor.submit( @@ -152,6 +172,16 @@ def get_std_out( def get_std_error( self, calculation_execution: AthenaCalculationExecution ) -> "Future[str] | None": + """Read the standard error output of a calculation from S3 asynchronously. + + Args: + calculation_execution: The calculation execution whose + ``std_error_s3_uri`` is read. + + Returns: + Future object containing the error output text with leading and trailing + whitespace removed, or None if the calculation has no ``std_error_s3_uri``. + """ if not calculation_execution.std_error_s3_uri: return None return self._executor.submit( @@ -159,6 +189,14 @@ def get_std_error( ) def poll(self, query_id: str) -> "Future[AthenaCalculationExecution]": + """Wait for a calculation to reach a terminal state asynchronously. + + Args: + query_id: The calculation execution ID. + + Returns: + Future object containing the calculation execution in a terminal state. + """ return cast( "Future[AthenaCalculationExecution]", self._executor.submit(self._poll, query_id) ) @@ -183,4 +221,13 @@ def execute( return calculation_id, self._executor.submit(self._poll, calculation_id) def cancel(self, query_id: str) -> "Future[None]": + """Stop a calculation execution asynchronously. + + Args: + query_id: The calculation execution ID. + + Returns: + Future object that completes when the ``StopCalculationExecution`` + request has been sent. + """ return self._executor.submit(self._cancel, query_id) diff --git a/pyathena/spark/common.py b/pyathena/spark/common.py index 53a42e9e6..b7c3948f6 100644 --- a/pyathena/spark/common.py +++ b/pyathena/spark/common.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""Base classes shared by the Athena for Apache Spark cursors.""" + from __future__ import annotations import contextlib @@ -123,14 +125,22 @@ def __init__( @property def session_id(self) -> str: + """The ID of the Spark session that this cursor runs calculations in.""" return self._session_id @property def calculation_id(self) -> str | None: + """The ID of the calculation tracked by this cursor, or None if there is none.""" return self._calculation_id @staticmethod def get_default_engine_configuration() -> dict[str, Any]: + """Return the engine configuration used when none is given. + + Returns: + The ``EngineConfiguration`` of a new session: a coordinator DPU size of 1, + at most 2 concurrent DPUs, and a default executor DPU size of 1. + """ return { "CoordinatorDpuSize": 1, "MaxConcurrentDpus": 2, @@ -496,91 +506,107 @@ class WithCalculationExecution: """ def __init__(self): + """Initialize the mixin, which keeps no state of its own.""" super().__init__() @property @abstractmethod def calculation_execution(self) -> AthenaCalculationExecution | None: + """The calculation execution that the other properties read, or None.""" raise NotImplementedError # pragma: no cover @property @abstractmethod def session_id(self) -> str: + """The ID of the Spark session that runs the calculations.""" raise NotImplementedError # pragma: no cover @property @abstractmethod def calculation_id(self) -> str | None: + """The ID of the current calculation, or None.""" raise NotImplementedError # pragma: no cover @property def description(self) -> str | None: + """The ``Description`` of the calculation, or None if there is none.""" if not self.calculation_execution: return None return self.calculation_execution.description @property def working_directory(self) -> str | None: + """The ``WorkingDirectory`` of the calculation, or None if there is none.""" if not self.calculation_execution: return None return self.calculation_execution.working_directory @property def state(self) -> str | None: + """The ``State`` of the calculation, or None if there is none.""" if not self.calculation_execution: return None return self.calculation_execution.state @property def state_change_reason(self) -> str | None: + """The ``StateChangeReason`` of the calculation, or None if there is none.""" if not self.calculation_execution: return None return self.calculation_execution.state_change_reason @property def submission_date_time(self) -> datetime | None: + """The ``SubmissionDateTime`` of the calculation, or None if there is none.""" if not self.calculation_execution: return None return self.calculation_execution.submission_date_time @property def completion_date_time(self) -> datetime | None: + """The ``CompletionDateTime`` of the calculation, or None if there is none.""" if not self.calculation_execution: return None return self.calculation_execution.completion_date_time @property def dpu_execution_in_millis(self) -> int | None: + """The ``DpuExecutionInMillis`` statistic of the calculation, or None if there is none.""" if not self.calculation_execution: return None return self.calculation_execution.dpu_execution_in_millis @property def progress(self) -> str | None: + """The ``Progress`` statistic of the calculation, or None if there is none.""" if not self.calculation_execution: return None return self.calculation_execution.progress @property def std_out_s3_uri(self) -> str | None: + """The ``StdOutS3Uri`` of the calculation, or None if there is none.""" if not self.calculation_execution: return None return self.calculation_execution.std_out_s3_uri @property def std_error_s3_uri(self) -> str | None: + """The ``StdErrorS3Uri`` of the calculation, or None if there is none.""" if not self.calculation_execution: return None return self.calculation_execution.std_error_s3_uri @property def result_s3_uri(self) -> str | None: + """The ``ResultS3Uri`` of the calculation, or None if there is none.""" if not self.calculation_execution: return None return self.calculation_execution.result_s3_uri @property def result_type(self) -> str | None: + """The ``ResultType`` of the calculation, or None if there is none.""" if not self.calculation_execution: return None return self.calculation_execution.result_type diff --git a/pyathena/spark/cursor.py b/pyathena/spark/cursor.py index ad6e655da..3376454de 100644 --- a/pyathena/spark/cursor.py +++ b/pyathena/spark/cursor.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""Cursor that runs PySpark code in an Athena for Apache Spark session.""" + from __future__ import annotations import logging @@ -127,6 +129,12 @@ def execute( return self def cancel(self) -> None: + """Stop the calculation that ``execute()`` last started. + + Raises: + ProgrammingError: If no calculation ID is set. + OperationalError: If the ``StopCalculationExecution`` request fails. + """ if not self.calculation_id: raise ProgrammingError("CalculationExecutionId is none or empty.") self._cancel(self.calculation_id) diff --git a/pyathena/sqlalchemy/__init__.py b/pyathena/sqlalchemy/__init__.py index e69de29bb..9fff16782 100644 --- a/pyathena/sqlalchemy/__init__.py +++ b/pyathena/sqlalchemy/__init__.py @@ -0,0 +1,8 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT + +"""SQLAlchemy dialects and types for Amazon Athena.""" diff --git a/pyathena/sqlalchemy/array.py b/pyathena/sqlalchemy/array.py index 3a91a1b51..ec61f28d6 100644 --- a/pyathena/sqlalchemy/array.py +++ b/pyathena/sqlalchemy/array.py @@ -80,6 +80,20 @@ def __init__( dimensions: int | None = None, zero_indexes: bool = False, ) -> None: + """Initialize the ARRAY type. + + Args: + item_type: SQLAlchemy type or type class of the array elements. A type + class is instantiated. Defaults to ``String``. + as_tuple: Return tuples instead of lists. + dimensions: Fixed number of array dimensions, a positive integer. + zero_indexes: Translate zero-based SQLAlchemy indexes to one-based SQL + indexes. + + Raises: + ValueError: If ``dimensions`` is not a positive integer, or if both a + nested ARRAY ``item_type`` and ``dimensions`` are given. + """ if dimensions is not None and ( isinstance(dimensions, bool) or not isinstance(dimensions, int) or dimensions < 1 ): diff --git a/pyathena/sqlalchemy/arrow.py b/pyathena/sqlalchemy/arrow.py index ea31b0733..9b2b35fa0 100644 --- a/pyathena/sqlalchemy/arrow.py +++ b/pyathena/sqlalchemy/arrow.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""SQLAlchemy dialect for Athena that returns results through ``ArrowCursor``.""" + from typing import TYPE_CHECKING from pyathena.sqlalchemy.base import AthenaDialect diff --git a/pyathena/sqlalchemy/base.py b/pyathena/sqlalchemy/base.py index 876c0098c..d33331e6b 100644 --- a/pyathena/sqlalchemy/base.py +++ b/pyathena/sqlalchemy/base.py @@ -1,3 +1,5 @@ +"""Base SQLAlchemy dialect for Amazon Athena.""" + from __future__ import annotations import contextlib @@ -227,6 +229,17 @@ class AthenaDialect(DefaultDialect): _URL_ENGINE_OPTIONS: tuple[str, ...] = ("insertmanyvalues_page_size", "use_insertmanyvalues") def __init__(self, json_deserializer=None, json_serializer=None, **kwargs): + """Initialize the dialect. + + Args: + json_deserializer: Callable used to deserialize JSON values, or + ``None`` for the default. + json_serializer: Callable used to serialize JSON values, or ``None`` + for the default. + **kwargs: Keyword arguments forwarded to ``DefaultDialect.__init__``. + Those that are URL engine options take precedence over the same + options in the connection URL query. + """ DefaultDialect.__init__(self, **kwargs) self._json_deserializer = json_deserializer self._json_serializer = json_serializer diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index d40cff29a..2e254a745 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -1,3 +1,5 @@ +"""SQLAlchemy type, statement, and DDL compilers for Amazon Athena.""" + from __future__ import annotations import re @@ -128,6 +130,15 @@ def visit_DECIMAL(self, type_: types.DECIMAL[Any], **kw: Any) -> str: return f"DECIMAL({type_.precision}, {type_.scale})" def visit_TINYINT(self, type_: types.Integer, **kw: Any) -> str: + """Render a TINYINT type. + + Args: + type_: The type to render. + **kw: Type-compiler keyword arguments. + + Returns: + ``TINYINT``. + """ return "TINYINT" @override @@ -207,6 +218,15 @@ def visit_BOOLEAN(self, type_: types.Boolean, **kw: Any) -> str: return "BOOLEAN" def visit_JSON(self, type_: types.JSON, **kw: Any) -> str: + """Render a JSON type. + + Args: + type_: The type to render. + **kw: Type-compiler keyword arguments. + + Returns: + ``JSON``. + """ return "JSON" @override @@ -226,6 +246,15 @@ def visit_null(self, type_, **kw): return "NULL" def visit_tinyint(self, type_, **kw): + """Render a tinyint type through ``visit_TINYINT``. + + Args: + type_: The type to render. + **kw: Type-compiler keyword arguments. + + Returns: + ``TINYINT``. + """ return self.visit_TINYINT(type_, **kw) @override @@ -253,6 +282,20 @@ def _enable_hive_column_ddl(self, kw: dict[str, Any]) -> bool: return False def visit_struct(self, type_, **kw): + """Render a STRUCT type. + + CREATE TABLE column types and types nested in an ARRAY render Hive + ``STRUCT``; other contexts render ``ROW(name type, ...)``. + A type that is not an ``AthenaStruct``, or one without fields, renders + ``ROW()``. + + Args: + type_: The type to render. + **kw: Type-compiler keyword arguments. + + Returns: + The STRUCT or ROW type clause. + """ # Empty structs keep the existing ROW() rendering in every context. if not isinstance(type_, AthenaStruct) or not type_.fields: return "ROW()" @@ -272,9 +315,29 @@ def visit_struct(self, type_, **kw): return f"ROW({', '.join(field_specs)})" def visit_STRUCT(self, type_, **kw): + """Render a STRUCT type through ``visit_struct``. + + Args: + type_: The type to render. + **kw: Type-compiler keyword arguments. + + Returns: + The STRUCT or ROW type clause. + """ return self.visit_struct(type_, **kw) def visit_map(self, type_, **kw): + """Render a MAP type as ``MAP``. + + A type that is not an ``AthenaMap`` renders ``MAP``. + + Args: + type_: The type to render. + **kw: Type-compiler keyword arguments. + + Returns: + The MAP type clause. + """ if isinstance(type_, AthenaMap): self._enable_hive_column_ddl(kw) key_type_str = self.process(type_.key_type, **kw) @@ -283,9 +346,30 @@ def visit_map(self, type_, **kw): return "MAP" def visit_MAP(self, type_, **kw): + """Render a MAP type through ``visit_map``. + + Args: + type_: The type to render. + **kw: Type-compiler keyword arguments. + + Returns: + The MAP type clause. + """ return self.visit_map(type_, **kw) def visit_array(self, type_, **kw): + """Render an ARRAY type as ``ARRAY``. + + Nested types of an ARRAY use Hive DDL syntax. A type that is not an + ARRAY renders ``ARRAY``. + + Args: + type_: The type to render. + **kw: Type-compiler keyword arguments. + + Returns: + The ARRAY type clause. + """ if isinstance(type_, types.ARRAY): kw["_athena_hive_ddl"] = True item_type_str = self.process(_ArrayTypeInspector.item_type(type_), **kw) @@ -293,6 +377,15 @@ def visit_array(self, type_, **kw): return "ARRAY" def visit_ARRAY(self, type_, **kw): + """Render an ARRAY type through ``visit_array``. + + Args: + type_: The type to render. + **kw: Type-compiler keyword arguments. + + Returns: + The ARRAY type clause. + """ return self.visit_array(type_, **kw) @@ -321,6 +414,15 @@ def _array_type_inspector(self): return _ArrayTypeInspector(self.dialect) def visit_char_length_func(self, fn: Function[Any], **kw: Any) -> str: + """Render ``char_length()`` as Athena ``length()``. + + Args: + fn: The function expression. + **kw: Compiler keyword arguments. + + Returns: + The ``length()`` function call. + """ return f"length{self.function_argspec(fn, **kw)}" @staticmethod @@ -338,6 +440,15 @@ def visit_update(self, update_stmt, visiting_cte=None, **kw): ) def visit_athena_array_update(self, expression, **kw): + """Render the whole-column expression of a partial ARRAY assignment. + + Args: + expression: The ``_ArrayUpdate`` expression created by ``visit_update``. + **kw: Compiler keyword arguments. + + Returns: + The SQL expression assigned to the ARRAY column. + """ return _ArrayUpdateCompiler(self).process(expression, **kw) def _array_lambda_name(self): @@ -418,6 +529,23 @@ def visit_binary( return super().visit_binary(binary, override_operator=override_operator, **kw) def visit_getitem_binary(self, binary, operator, **kw): + """Render an ARRAY index as ``element_at()`` and an ARRAY slice as ``slice()``. + + Slice bounds are inclusive and clamped to the array. An index below 1 is + passed to ``element_at()`` as NULL. + + Args: + binary: The index or slice expression. + operator: The getitem operator. + **kw: Compiler keyword arguments. + + Returns: + The rendered index or slice expression. + + Raises: + CompileError: If the indexed expression is not an ARRAY, or if a slice + step is not None or 1. + """ array_type = self._array_type_inspector.array_type(binary.left.type) if array_type is None: raise exc.CompileError("Athena indexing requires an ARRAY expression") @@ -828,6 +956,15 @@ def _complex_dml_type(self, type_, *, require_precision=False, timestamp_precisi return self.dialect.type_compiler_instance.process(type_) def visit_athena_array_json_projection(self, expression, **kw): + """Render an ARRAY result column as a JSON envelope string. + + Args: + expression: The ``_ArrayJSONProjection`` expression. + **kw: Compiler keyword arguments. + + Returns: + The SQL expression that serializes the ARRAY value as JSON. + """ value = self.process(expression.element, **kw) encoded = self._array_json(value, expression.array_type) # An object envelope keeps SQL NULL and CSV null markers out of the transport. @@ -955,6 +1092,18 @@ def __init__( render_schema_translate: bool = False, compile_kwargs: dict[str, Any] | None = None, ): + """Initialize the DDL compiler with an ``AthenaDDLIdentifierPreparer``. + + Args: + dialect: The Athena dialect. + statement: The DDL statement to compile. + schema_translate_map: Schema translation map forwarded to + ``DDLCompiler``. + render_schema_translate: Whether to render schema translation, + forwarded to ``DDLCompiler``. + compile_kwargs: Compiler keyword arguments. ``None`` is replaced with + an empty mapping. + """ self._preparer = AthenaDDLIdentifierPreparer(dialect) super().__init__( dialect=dialect, diff --git a/pyathena/sqlalchemy/map.py b/pyathena/sqlalchemy/map.py index e2e828b00..9e823c616 100644 --- a/pyathena/sqlalchemy/map.py +++ b/pyathena/sqlalchemy/map.py @@ -43,6 +43,14 @@ class AthenaMap(TypeEngine[dict[str, Any]]): __visit_name__ = "map" def __init__(self, key_type: Any = None, value_type: Any = None) -> None: + """Initialize the MAP type. + + Args: + key_type: SQLAlchemy type or type class for map keys. A type class is + instantiated. Defaults to ``String``. + value_type: SQLAlchemy type or type class for map values. A type class + is instantiated. Defaults to ``String``. + """ if key_type is None: self.key_type: TypeEngine[Any] = sqltypes.String() elif isinstance(key_type, TypeEngine): diff --git a/pyathena/sqlalchemy/pandas.py b/pyathena/sqlalchemy/pandas.py index 0d9759344..9fb9c2ae5 100644 --- a/pyathena/sqlalchemy/pandas.py +++ b/pyathena/sqlalchemy/pandas.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""SQLAlchemy dialect for Athena that returns results through ``PandasCursor``.""" + from typing import TYPE_CHECKING from pyathena.sqlalchemy.base import AthenaDialect diff --git a/pyathena/sqlalchemy/polars.py b/pyathena/sqlalchemy/polars.py index a715bc8de..38c7064b2 100644 --- a/pyathena/sqlalchemy/polars.py +++ b/pyathena/sqlalchemy/polars.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""SQLAlchemy dialect for Athena that returns results through ``PolarsCursor``.""" + from typing import TYPE_CHECKING from pyathena.sqlalchemy.base import AthenaDialect diff --git a/pyathena/sqlalchemy/preparer.py b/pyathena/sqlalchemy/preparer.py index e4185d597..f1511ab40 100644 --- a/pyathena/sqlalchemy/preparer.py +++ b/pyathena/sqlalchemy/preparer.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""SQLAlchemy identifier preparers for Athena DML and DDL statements.""" + from __future__ import annotations from typing import TYPE_CHECKING @@ -66,6 +68,17 @@ def __init__( quote_case_sensitive_collations: bool = True, omit_schema: bool = False, ): + """Initialize the preparer with backtick quoting by default. + + Args: + dialect: The dialect that uses this preparer. + initial_quote: Character that begins a delimited identifier. + final_quote: Character that ends a delimited identifier. ``None`` + uses ``initial_quote``. + escape_quote: Character that escapes a quote inside an identifier. + quote_case_sensitive_collations: Forwarded to ``IdentifierPreparer``. + omit_schema: Do not prepend the schema name to identifiers. + """ super().__init__( dialect=dialect, initial_quote=initial_quote, diff --git a/pyathena/sqlalchemy/requirements.py b/pyathena/sqlalchemy/requirements.py index 048e92294..ebf98065d 100644 --- a/pyathena/sqlalchemy/requirements.py +++ b/pyathena/sqlalchemy/requirements.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""SQLAlchemy test suite requirements for the Athena dialect.""" + from sqlalchemy.testing import exclusions from sqlalchemy.testing.requirements import SuiteRequirements @@ -15,6 +17,8 @@ class Requirements(SuiteRequirements): + """Features of the Athena dialect for the SQLAlchemy test suite.""" + @property @override def comment_reflection(self): diff --git a/pyathena/sqlalchemy/rest.py b/pyathena/sqlalchemy/rest.py index eeea5e12f..ac8d8b8c6 100644 --- a/pyathena/sqlalchemy/rest.py +++ b/pyathena/sqlalchemy/rest.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""SQLAlchemy dialect for Athena that uses the default ``Cursor``.""" + from typing import TYPE_CHECKING from pyathena.sqlalchemy.base import AthenaDialect diff --git a/pyathena/sqlalchemy/s3fs.py b/pyathena/sqlalchemy/s3fs.py index 293855f5d..9ea596b40 100644 --- a/pyathena/sqlalchemy/s3fs.py +++ b/pyathena/sqlalchemy/s3fs.py @@ -5,6 +5,8 @@ # # SPDX-License-Identifier: MIT +"""SQLAlchemy dialect for Athena that returns results through ``S3FSCursor``.""" + from typing import TYPE_CHECKING from pyathena.sqlalchemy.base import AthenaDialect diff --git a/pyathena/sqlalchemy/struct.py b/pyathena/sqlalchemy/struct.py index 865be0830..a49e30cb2 100644 --- a/pyathena/sqlalchemy/struct.py +++ b/pyathena/sqlalchemy/struct.py @@ -49,6 +49,17 @@ class AthenaStruct(TypeEngine[dict[str, Any]]): __visit_name__ = "struct" def __init__(self, *fields: str | tuple[str, Any]) -> None: + """Initialize the STRUCT type. + + Args: + *fields: Field specifications. A string is a field name of type + ``String``. A ``(field_name, field_type)`` tuple names a field and + its SQLAlchemy type or type class; a type class is instantiated. + + Raises: + ValueError: If a field specification is neither a string nor a + two-element tuple. + """ self.fields: dict[str, TypeEngine[Any]] = {} for field in fields: diff --git a/pyathena/util.py b/pyathena/util.py index d98798bc3..b4ea89fd6 100644 --- a/pyathena/util.py +++ b/pyathena/util.py @@ -1,3 +1,5 @@ +"""Helpers for S3 output locations, truth-value parsing, and retrying AWS API calls.""" + from __future__ import annotations import logging @@ -168,6 +170,18 @@ def __init__( max_delay: int = 100, exponential_base: int = 2, ) -> None: + """Initialize the retry configuration. + + Args: + exceptions: AWS error code, or iterable of error codes, to retry on. + Stored as a tuple. + attempt: Maximum number of attempts, including the first call. + multiplier: Base multiplier for exponential backoff in seconds, and + the upper bound of the random jitter added to each wait. + max_delay: Maximum exponential delay between retries in seconds, + before jitter is added. + exponential_base: Base for exponential backoff calculation. + """ self.exceptions = (exceptions,) if isinstance(exceptions, str) else tuple(exceptions) self.attempt = attempt self.multiplier = multiplier diff --git a/pyproject.toml b/pyproject.toml index 414ff78ff..f56cbae18 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -159,16 +159,26 @@ select = [ "PGH", # pygrep-hooks "G", # flake8-logging-format "PT", # flake8-pytest-style + "D", # pydocstyle ] ignore = [ "RUF059", # unused-unpacked-variable (too noisy for interface-heavy code) "G004", # logging-f-string (f-strings are preferred for log messages) + "D105", # undocumented-magic-method (protocol methods need no docstring) ] +[tool.ruff.lint.pydocstyle] +convention = "google" +# Overrides marked with @override inherit the base method's documentation. +ignore-decorators = ["pyathena.util.override"] + [tool.ruff.lint.per-file-ignores] "tests/**" = [ "RUF012", # mutable-class-default (test classes often use mutable defaults) ] +"!pyathena/**" = [ + "D", # docstring rules apply to the package only +] "pyathena/sqlalchemy/compiler.py" = [ "N802", # SQLAlchemy TypeCompiler requires visit_UPPERCASE method names ]