From a9d42e29bf735ae2909b2606509d9802bc46d359 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Fri, 2 Oct 2026 09:28:36 +0900 Subject: [PATCH 01/13] Check docstrings in pyathena/ with ruff's pydocstyle rules Enable ruff's D rules with the Google convention for pyathena/. Methods decorated with pyathena.util.override are exempt, and magic methods need no docstring. Files outside the package are not checked. per-file-ignores lists, per file, the codes that still have findings. A gap of another code, or in a file without an entry, fails the check; a new gap of a listed code in a listed file is not reported until that file's entry is removed. Fix the findings outside that list: complete Cursor.execute's Args, describe *args of connect() and aio_connect(), end the first lines of DBAPITypeObject, setinputsizes and setoutputsize with a period, and add docstrings to S3File, AthenaDictResultSet and the SQLAlchemy test-suite Requirements. Correct cache_size in the AsyncCursor and AioCursor execute docstrings, which described it as a cache size rather than the number of queries to check. Document the docstring rules in the contributing guide. Co-Authored-By: Claude Opus 5.5 --- AGENTS.md | 3 +- docs/contributing.md | 19 ++++++ pyathena/__init__.py | 6 +- pyathena/aio/cursor.py | 2 +- pyathena/async_cursor.py | 2 +- pyathena/common.py | 14 +++- pyathena/cursor.py | 7 ++ pyathena/filesystem/s3.py | 5 ++ pyathena/result_set.py | 2 + pyathena/sqlalchemy/requirements.py | 2 + pyproject.toml | 100 ++++++++++++++++++++++++++++ 11 files changed, 156 insertions(+), 6 deletions(-) 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..252efa5cf 100644 --- a/docs/contributing.md +++ b/docs/contributing.md @@ -71,6 +71,25 @@ 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, although ruff does not check them. + +`per-file-ignores` in `pyproject.toml` lists the files that still lack docstrings. +Remove a file's entry when its docstrings are complete. + ## Open a pull request Open a draft pull request with the repository's template completed. diff --git a/pyathena/__init__.py b/pyathena/__init__.py index 1e73d195d..ee80666d9 100644 --- a/pyathena/__init__.py +++ b/pyathena/__init__.py @@ -30,7 +30,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 +89,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 +146,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/cursor.py b/pyathena/aio/cursor.py index 38ef885fd..addfc6a54 100644 --- a/pyathena/aio/cursor.py +++ b/pyathena/aio/cursor.py @@ -104,7 +104,7 @@ 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). diff --git a/pyathena/async_cursor.py b/pyathena/async_cursor.py index 620df9d68..98e351f82 100644 --- a/pyathena/async_cursor.py +++ b/pyathena/async_cursor.py @@ -190,7 +190,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). diff --git a/pyathena/common.py b/pyathena/common.py index 8797bfb9d..ad5e422c6 100644 --- a/pyathena/common.py +++ b/pyathena/common.py @@ -1351,10 +1351,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/cursor.py b/pyathena/cursor.py index 2e789520a..8b91413f0 100644 --- a/pyathena/cursor.py +++ b/pyathena/cursor.py @@ -107,6 +107,13 @@ def execute( Args: operation: SQL query string to execute. parameters: Query parameters (optional). + 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 function called immediately after start_query_execution API is called. Function signature: (query_id: str) -> None diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index abbe270f4..4d29fa14d 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -1964,6 +1964,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__( diff --git a/pyathena/result_set.py b/pyathena/result_set.py index 8b6f0e3f1..2e719aa82 100644 --- a/pyathena/result_set.py +++ b/pyathena/result_set.py @@ -753,6 +753,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 diff --git a/pyathena/sqlalchemy/requirements.py b/pyathena/sqlalchemy/requirements.py index 048e92294..a4d9595b4 100644 --- a/pyathena/sqlalchemy/requirements.py +++ b/pyathena/sqlalchemy/requirements.py @@ -15,6 +15,8 @@ class Requirements(SuiteRequirements): + """Features of the Athena dialect for the SQLAlchemy test suite.""" + @property @override def comment_reflection(self): diff --git a/pyproject.toml b/pyproject.toml index 414ff78ff..8ad67713c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -159,22 +159,122 @@ 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 + # Missing docstrings (#882) + "D100", + "D102", + "D107", ] "pyathena/aio/sqlalchemy/base.py" = [ "N801", # SQLAlchemy async adapter naming convention (AsyncAdapt_*) + # Missing docstrings (#882) + "D100", + "D102", + "D107", ] +# Missing docstrings (#882). Remove an entry once its file is documented. +"pyathena/__init__.py" = ["D104"] +"pyathena/aio/__init__.py" = ["D104"] +"pyathena/aio/arrow/__init__.py" = ["D104"] +"pyathena/aio/arrow/cursor.py" = ["D100", "D107"] +"pyathena/aio/common.py" = ["D100"] +"pyathena/aio/connection.py" = ["D100", "D107"] +"pyathena/aio/cursor.py" = ["D100", "D107"] +"pyathena/aio/pandas/__init__.py" = ["D104"] +"pyathena/aio/pandas/cursor.py" = ["D100", "D107"] +"pyathena/aio/polars/__init__.py" = ["D104"] +"pyathena/aio/polars/cursor.py" = ["D100", "D107"] +"pyathena/aio/result_set.py" = ["D100", "D107"] +"pyathena/aio/s3fs/__init__.py" = ["D104"] +"pyathena/aio/s3fs/cursor.py" = ["D100", "D107"] +"pyathena/aio/spark/__init__.py" = ["D104"] +"pyathena/aio/spark/cursor.py" = ["D100"] +"pyathena/aio/sqlalchemy/__init__.py" = ["D104"] +"pyathena/aio/sqlalchemy/arrow.py" = ["D100"] +"pyathena/aio/sqlalchemy/pandas.py" = ["D100"] +"pyathena/aio/sqlalchemy/polars.py" = ["D100"] +"pyathena/aio/sqlalchemy/rest.py" = ["D100"] +"pyathena/aio/sqlalchemy/s3fs.py" = ["D100"] +"pyathena/aio/util.py" = ["D100"] +"pyathena/arrow/__init__.py" = ["D104"] +"pyathena/arrow/async_cursor.py" = ["D100"] +"pyathena/arrow/converter.py" = ["D100", "D107"] +"pyathena/arrow/cursor.py" = ["D100"] +"pyathena/arrow/result_set.py" = ["D100", "D102", "D107"] +"pyathena/async_cursor.py" = ["D100", "D102", "D107"] +"pyathena/common.py" = ["D100", "D102", "D107"] +"pyathena/connection.py" = ["D100"] +"pyathena/converter.py" = ["D100", "D102", "D107"] +"pyathena/cursor.py" = ["D100", "D107"] +"pyathena/error.py" = ["D100"] +"pyathena/filesystem/__init__.py" = ["D104"] +"pyathena/filesystem/s3.py" = ["D100", "D102", "D107"] +"pyathena/filesystem/s3_async.py" = ["D100", "D102", "D107"] +"pyathena/filesystem/s3_errors.py" = ["D107"] +"pyathena/filesystem/s3_executor.py" = ["D100", "D107"] +"pyathena/filesystem/s3_object.py" = ["D100", "D102", "D107"] +"pyathena/formatter.py" = ["D100", "D102", "D107"] +"pyathena/glue.py" = ["D107"] +"pyathena/model.py" = ["D100", "D102", "D107"] +"pyathena/options.py" = ["D100"] +"pyathena/pandas/__init__.py" = ["D104"] +"pyathena/pandas/async_cursor.py" = ["D100", "D107"] +"pyathena/pandas/converter.py" = ["D100", "D107"] +"pyathena/pandas/cursor.py" = ["D100"] +"pyathena/pandas/reader.py" = ["D100", "D107"] +"pyathena/pandas/result_set.py" = ["D100", "D102"] +"pyathena/pandas/util.py" = ["D100"] +"pyathena/parser.py" = ["D100", "D107"] +"pyathena/polars/__init__.py" = ["D104"] +"pyathena/polars/async_cursor.py" = ["D100"] +"pyathena/polars/converter.py" = ["D100", "D107"] +"pyathena/polars/cursor.py" = ["D100"] +"pyathena/polars/result_set.py" = ["D100"] +"pyathena/result_set.py" = ["D100", "D102", "D107"] +"pyathena/s3fs/__init__.py" = ["D104"] +"pyathena/s3fs/async_cursor.py" = ["D100"] +"pyathena/s3fs/converter.py" = ["D100", "D107"] +"pyathena/s3fs/cursor.py" = ["D100"] +"pyathena/s3fs/reader.py" = ["D100"] +"pyathena/s3fs/result_set.py" = ["D100", "D107"] +"pyathena/spark/__init__.py" = ["D104"] +"pyathena/spark/async_cursor.py" = ["D100", "D102"] +"pyathena/spark/common.py" = ["D100", "D102", "D107"] +"pyathena/spark/cursor.py" = ["D100", "D102"] +"pyathena/sqlalchemy/__init__.py" = ["D104"] +"pyathena/sqlalchemy/array.py" = ["D107"] +"pyathena/sqlalchemy/arrow.py" = ["D100"] +"pyathena/sqlalchemy/base.py" = ["D100", "D107"] +"pyathena/sqlalchemy/map.py" = ["D107"] +"pyathena/sqlalchemy/pandas.py" = ["D100"] +"pyathena/sqlalchemy/polars.py" = ["D100"] +"pyathena/sqlalchemy/preparer.py" = ["D100", "D107"] +"pyathena/sqlalchemy/requirements.py" = ["D100"] +"pyathena/sqlalchemy/rest.py" = ["D100"] +"pyathena/sqlalchemy/s3fs.py" = ["D100"] +"pyathena/sqlalchemy/struct.py" = ["D107"] +"pyathena/util.py" = ["D100", "D107"] [tool.mypy] follow_imports = "silent" From 13e43c3875f691e9b693157b46c328d6aa8fb5b6 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Fri, 2 Oct 2026 09:34:14 +0900 Subject: [PATCH 02/13] Keep other codes when removing a file's docstring ignores The compiler and asyncio dialect entries also hold naming exemptions, so completing a file removes its D codes, not the whole entry. ruff also checks the format of private docstrings that exist, so the guide says it only does not require them. Co-Authored-By: Claude Opus 5.5 --- docs/contributing.md | 7 ++++--- pyproject.toml | 2 +- 2 files changed, 5 insertions(+), 4 deletions(-) diff --git a/docs/contributing.md b/docs/contributing.md index 252efa5cf..8f031e762 100644 --- a/docs/contributing.md +++ b/docs/contributing.md @@ -85,10 +85,11 @@ Code under `pyathena/` uses [Google-style docstrings](https://google.github.io/s 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, although ruff does not check them. +- New or changed private functions and methods use the same style. + ruff does not require them to have a docstring, but checks the format of the docstrings they have. -`per-file-ignores` in `pyproject.toml` lists the files that still lack docstrings. -Remove a file's entry when its docstrings are complete. +`per-file-ignores` in `pyproject.toml` lists the `D` codes that each file still has findings for. +Remove a file's `D` codes when its docstrings are complete, and keep its other codes. ## Open a pull request diff --git a/pyproject.toml b/pyproject.toml index 8ad67713c..40a3530b9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -193,7 +193,7 @@ ignore-decorators = ["pyathena.util.override"] "D102", "D107", ] -# Missing docstrings (#882). Remove an entry once its file is documented. +# Missing docstrings (#882). Remove a file's D codes once it is documented. "pyathena/__init__.py" = ["D104"] "pyathena/aio/__init__.py" = ["D104"] "pyathena/aio/arrow/__init__.py" = ["D104"] From 919d238b702222ac72b8a3832eadc7e02af540c8 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Fri, 2 Oct 2026 09:34:51 +0900 Subject: [PATCH 03/13] Say that ruff checks existing private docstrings, not only their format Co-Authored-By: Claude Opus 5.5 --- docs/contributing.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/contributing.md b/docs/contributing.md index 8f031e762..5fc5a779e 100644 --- a/docs/contributing.md +++ b/docs/contributing.md @@ -86,7 +86,7 @@ Code under `pyathena/` uses [Google-style docstrings](https://google.github.io/s 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 format of the docstrings they have. + ruff does not require them to have a docstring, but checks the docstrings they have. `per-file-ignores` in `pyproject.toml` lists the `D` codes that each file still has findings for. Remove a file's `D` codes when its docstrings are complete, and keep its other codes. From 05a4fac865d73f36ec23a8041a9e532e5f00a6b4 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 11:25:24 +0900 Subject: [PATCH 04/13] Document the model classes Co-Authored-By: Claude Opus 5.5 --- pyathena/model.py | 206 ++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 206 insertions(+) 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, From 66b3ace48511d668d1d57ac1546fb78798279787 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 11:25:25 +0900 Subject: [PATCH 05/13] Document the result sets Co-Authored-By: Claude Opus 5.5 --- pyathena/aio/result_set.py | 17 +++++++ pyathena/arrow/result_set.py | 35 ++++++++++++++ pyathena/pandas/result_set.py | 11 +++++ pyathena/polars/result_set.py | 2 + pyathena/result_set.py | 91 +++++++++++++++++++++++++++++++++++ pyathena/s3fs/result_set.py | 25 ++++++++++ 6 files changed, 181 insertions(+) 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/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/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/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 2e719aa82..9467aeb65 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,7 @@ 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, or None without metadata or for DML.""" if self._metadata is None or ( self.substatement_type and self.substatement_type.upper() in self._DML_SUBSTATEMENT_TYPES @@ -338,6 +392,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 +787,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 @@ -864,24 +921,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 @@ -913,180 +974,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/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, From e6b43863878ef0e81cdb5df7fe4da8fad9ef5435 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 11:25:26 +0900 Subject: [PATCH 06/13] Document the S3 filesystem Co-Authored-By: Claude Opus 5.5 --- pyathena/filesystem/__init__.py | 2 + pyathena/filesystem/s3.py | 164 ++++++++++++++++++++++++++ pyathena/filesystem/s3_async.py | 182 +++++++++++++++++++++++++++++ pyathena/filesystem/s3_errors.py | 5 + pyathena/filesystem/s3_executor.py | 13 +++ pyathena/filesystem/s3_object.py | 176 ++++++++++++++++++++++++++++ 6 files changed, 542 insertions(+) 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 4d29fa14d..325ddf2dd 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,27 @@ def _list_object_versions_pages( break def info(self, path: str, **kwargs) -> S3Object: + """Return information about an S3 path. + + Returns a matching entry from the directory cache when one exists. + 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 +817,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 +834,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 +1066,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 +1349,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 +1843,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: @@ -1987,6 +2104,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, and a larger one is copied as the first + parts of the multipart upload. + + 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 {} @@ -2050,6 +2201,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() @@ -2155,6 +2307,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() @@ -2191,6 +2354,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__) From 1191251adafff616fb5a0d1e7ca202d11976e4ae Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 11:25:26 +0900 Subject: [PATCH 07/13] Document the SQLAlchemy dialects Co-Authored-By: Claude Opus 5.5 --- pyathena/aio/sqlalchemy/__init__.py | 8 ++ pyathena/aio/sqlalchemy/arrow.py | 2 + pyathena/aio/sqlalchemy/base.py | 122 ++++++++++++++++++++++- pyathena/aio/sqlalchemy/pandas.py | 2 + pyathena/aio/sqlalchemy/polars.py | 2 + pyathena/aio/sqlalchemy/rest.py | 2 + pyathena/aio/sqlalchemy/s3fs.py | 2 + pyathena/sqlalchemy/__init__.py | 8 ++ pyathena/sqlalchemy/array.py | 14 +++ pyathena/sqlalchemy/arrow.py | 2 + pyathena/sqlalchemy/base.py | 13 +++ pyathena/sqlalchemy/compiler.py | 149 ++++++++++++++++++++++++++++ pyathena/sqlalchemy/map.py | 8 ++ pyathena/sqlalchemy/pandas.py | 2 + pyathena/sqlalchemy/polars.py | 2 + pyathena/sqlalchemy/preparer.py | 13 +++ pyathena/sqlalchemy/requirements.py | 2 + pyathena/sqlalchemy/rest.py | 2 + pyathena/sqlalchemy/s3fs.py | 2 + pyathena/sqlalchemy/struct.py | 11 ++ 20 files changed, 367 insertions(+), 1 deletion(-) 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/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 a4d9595b4..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 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: From f841ea6e0a76b2a5af49bb29c5a368d6de2893b8 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 11:25:27 +0900 Subject: [PATCH 08/13] Document the connections, base cursors and shared modules Co-Authored-By: Claude Opus 5.5 --- pyathena/__init__.py | 2 ++ pyathena/aio/__init__.py | 8 ++++++ pyathena/aio/common.py | 2 ++ pyathena/aio/connection.py | 8 ++++++ pyathena/aio/cursor.py | 27 ++++++++++++++++++ pyathena/aio/util.py | 2 ++ pyathena/async_cursor.py | 44 +++++++++++++++++++++++++++++ pyathena/common.py | 57 ++++++++++++++++++++++++++++++++++++++ pyathena/connection.py | 2 ++ pyathena/converter.py | 22 +++++++++++++++ pyathena/cursor.py | 27 ++++++++++++++++++ pyathena/error.py | 2 ++ pyathena/formatter.py | 34 +++++++++++++++++++++++ pyathena/glue.py | 9 ++++++ pyathena/options.py | 2 ++ pyathena/parser.py | 9 ++++++ pyathena/util.py | 14 ++++++++++ 17 files changed, 271 insertions(+) diff --git a/pyathena/__init__.py b/pyathena/__init__.py index ee80666d9..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 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/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 addfc6a54..f924ef013 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: Stop the running query when polling is cancelled + with ``asyncio.CancelledError``. + 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, @@ -225,6 +245,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/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/async_cursor.py b/pyathena/async_cursor.py index 98e351f82..b99068454 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,30 @@ 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: Stop the running query on ``KeyboardInterrupt`` + while polling. + 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 +114,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 +139,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]: @@ -298,6 +335,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 ad5e422c6..438638744 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,29 @@ 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: Stop the running query when polling is interrupted. + 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 the query ID, 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 +272,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( @@ -1314,6 +1356,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 +1372,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: diff --git a/pyathena/connection.py b/pyathena/connection.py index 6f27c1d5e..3607276df 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 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 8b91413f0..6fdffb72c 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: Stop the running query on ``KeyboardInterrupt`` + while polling. + 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, @@ -196,6 +216,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/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/options.py b/pyathena/options.py index adcb9ec7a..92ce52f38 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 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/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 From 748a6aca62a10eeb6da091e983d9105456f03f96 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 11:25:27 +0900 Subject: [PATCH 09/13] Document the Arrow, pandas, Polars, S3FS and Spark cursors Co-Authored-By: Claude Opus 5.5 --- pyathena/aio/arrow/__init__.py | 8 ++++++ pyathena/aio/arrow/cursor.py | 24 +++++++++++++++++ pyathena/aio/pandas/__init__.py | 8 ++++++ pyathena/aio/pandas/cursor.py | 28 ++++++++++++++++++++ pyathena/aio/polars/__init__.py | 8 ++++++ pyathena/aio/polars/cursor.py | 25 ++++++++++++++++++ pyathena/aio/s3fs/__init__.py | 8 ++++++ pyathena/aio/s3fs/cursor.py | 22 +++++++++++++++ pyathena/aio/spark/__init__.py | 8 ++++++ pyathena/aio/spark/cursor.py | 2 ++ pyathena/arrow/__init__.py | 8 ++++++ pyathena/arrow/async_cursor.py | 2 ++ pyathena/arrow/converter.py | 4 +++ pyathena/arrow/cursor.py | 2 ++ pyathena/pandas/__init__.py | 2 ++ pyathena/pandas/async_cursor.py | 23 ++++++++++++++++ pyathena/pandas/converter.py | 4 +++ pyathena/pandas/cursor.py | 2 ++ pyathena/pandas/reader.py | 9 +++++++ pyathena/pandas/util.py | 2 ++ pyathena/polars/__init__.py | 2 ++ pyathena/polars/async_cursor.py | 2 ++ pyathena/polars/converter.py | 4 +++ pyathena/polars/cursor.py | 2 ++ pyathena/s3fs/__init__.py | 8 ++++++ pyathena/s3fs/async_cursor.py | 2 ++ pyathena/s3fs/converter.py | 3 +++ pyathena/s3fs/cursor.py | 2 ++ pyathena/s3fs/reader.py | 2 ++ pyathena/spark/__init__.py | 8 ++++++ pyathena/spark/async_cursor.py | 47 +++++++++++++++++++++++++++++++++ pyathena/spark/common.py | 26 ++++++++++++++++++ pyathena/spark/cursor.py | 8 ++++++ 33 files changed, 315 insertions(+) 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..7a2576c9d 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: Whether to stop the running query when the waiting + task is cancelled. + 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, 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..ccb6e29df 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: Whether to stop the running query when the waiting + task is cancelled. + 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. 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, 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..66ca2d814 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: Whether to stop the running query when the waiting + task is cancelled. + 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, results 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, 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..a485d59ce 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: Whether to stop the running query when the waiting + task is cancelled. + 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, 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/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..c968337f4 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 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..63c225292 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,27 @@ 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: Whether to stop the running query on ``KeyboardInterrupt``. + 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. + 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..8f8913c53 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 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/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/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..fdf8cd201 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 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..cef1eca90 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 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/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) From 24213c77180b3078953c5dfe34cb62e2c79554e1 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 11:25:31 +0900 Subject: [PATCH 10/13] Check every docstring in pyathena/ without per-file exceptions All public modules, packages, classes, functions, methods, property getters and __init__ methods in pyathena/ now have docstrings, so the per-file ignores for the D rules are no longer needed. Co-Authored-By: Claude Opus 5.5 --- docs/contributing.md | 3 -- pyproject.toml | 90 -------------------------------------------- 2 files changed, 93 deletions(-) diff --git a/docs/contributing.md b/docs/contributing.md index 5fc5a779e..320147260 100644 --- a/docs/contributing.md +++ b/docs/contributing.md @@ -88,9 +88,6 @@ Code under `pyathena/` uses [Google-style docstrings](https://google.github.io/s - 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. -`per-file-ignores` in `pyproject.toml` lists the `D` codes that each file still has findings for. -Remove a file's `D` codes when its docstrings are complete, and keep its other codes. - ## Open a pull request Open a draft pull request with the repository's template completed. diff --git a/pyproject.toml b/pyproject.toml index 40a3530b9..f56cbae18 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -181,100 +181,10 @@ ignore-decorators = ["pyathena.util.override"] ] "pyathena/sqlalchemy/compiler.py" = [ "N802", # SQLAlchemy TypeCompiler requires visit_UPPERCASE method names - # Missing docstrings (#882) - "D100", - "D102", - "D107", ] "pyathena/aio/sqlalchemy/base.py" = [ "N801", # SQLAlchemy async adapter naming convention (AsyncAdapt_*) - # Missing docstrings (#882) - "D100", - "D102", - "D107", ] -# Missing docstrings (#882). Remove a file's D codes once it is documented. -"pyathena/__init__.py" = ["D104"] -"pyathena/aio/__init__.py" = ["D104"] -"pyathena/aio/arrow/__init__.py" = ["D104"] -"pyathena/aio/arrow/cursor.py" = ["D100", "D107"] -"pyathena/aio/common.py" = ["D100"] -"pyathena/aio/connection.py" = ["D100", "D107"] -"pyathena/aio/cursor.py" = ["D100", "D107"] -"pyathena/aio/pandas/__init__.py" = ["D104"] -"pyathena/aio/pandas/cursor.py" = ["D100", "D107"] -"pyathena/aio/polars/__init__.py" = ["D104"] -"pyathena/aio/polars/cursor.py" = ["D100", "D107"] -"pyathena/aio/result_set.py" = ["D100", "D107"] -"pyathena/aio/s3fs/__init__.py" = ["D104"] -"pyathena/aio/s3fs/cursor.py" = ["D100", "D107"] -"pyathena/aio/spark/__init__.py" = ["D104"] -"pyathena/aio/spark/cursor.py" = ["D100"] -"pyathena/aio/sqlalchemy/__init__.py" = ["D104"] -"pyathena/aio/sqlalchemy/arrow.py" = ["D100"] -"pyathena/aio/sqlalchemy/pandas.py" = ["D100"] -"pyathena/aio/sqlalchemy/polars.py" = ["D100"] -"pyathena/aio/sqlalchemy/rest.py" = ["D100"] -"pyathena/aio/sqlalchemy/s3fs.py" = ["D100"] -"pyathena/aio/util.py" = ["D100"] -"pyathena/arrow/__init__.py" = ["D104"] -"pyathena/arrow/async_cursor.py" = ["D100"] -"pyathena/arrow/converter.py" = ["D100", "D107"] -"pyathena/arrow/cursor.py" = ["D100"] -"pyathena/arrow/result_set.py" = ["D100", "D102", "D107"] -"pyathena/async_cursor.py" = ["D100", "D102", "D107"] -"pyathena/common.py" = ["D100", "D102", "D107"] -"pyathena/connection.py" = ["D100"] -"pyathena/converter.py" = ["D100", "D102", "D107"] -"pyathena/cursor.py" = ["D100", "D107"] -"pyathena/error.py" = ["D100"] -"pyathena/filesystem/__init__.py" = ["D104"] -"pyathena/filesystem/s3.py" = ["D100", "D102", "D107"] -"pyathena/filesystem/s3_async.py" = ["D100", "D102", "D107"] -"pyathena/filesystem/s3_errors.py" = ["D107"] -"pyathena/filesystem/s3_executor.py" = ["D100", "D107"] -"pyathena/filesystem/s3_object.py" = ["D100", "D102", "D107"] -"pyathena/formatter.py" = ["D100", "D102", "D107"] -"pyathena/glue.py" = ["D107"] -"pyathena/model.py" = ["D100", "D102", "D107"] -"pyathena/options.py" = ["D100"] -"pyathena/pandas/__init__.py" = ["D104"] -"pyathena/pandas/async_cursor.py" = ["D100", "D107"] -"pyathena/pandas/converter.py" = ["D100", "D107"] -"pyathena/pandas/cursor.py" = ["D100"] -"pyathena/pandas/reader.py" = ["D100", "D107"] -"pyathena/pandas/result_set.py" = ["D100", "D102"] -"pyathena/pandas/util.py" = ["D100"] -"pyathena/parser.py" = ["D100", "D107"] -"pyathena/polars/__init__.py" = ["D104"] -"pyathena/polars/async_cursor.py" = ["D100"] -"pyathena/polars/converter.py" = ["D100", "D107"] -"pyathena/polars/cursor.py" = ["D100"] -"pyathena/polars/result_set.py" = ["D100"] -"pyathena/result_set.py" = ["D100", "D102", "D107"] -"pyathena/s3fs/__init__.py" = ["D104"] -"pyathena/s3fs/async_cursor.py" = ["D100"] -"pyathena/s3fs/converter.py" = ["D100", "D107"] -"pyathena/s3fs/cursor.py" = ["D100"] -"pyathena/s3fs/reader.py" = ["D100"] -"pyathena/s3fs/result_set.py" = ["D100", "D107"] -"pyathena/spark/__init__.py" = ["D104"] -"pyathena/spark/async_cursor.py" = ["D100", "D102"] -"pyathena/spark/common.py" = ["D100", "D102", "D107"] -"pyathena/spark/cursor.py" = ["D100", "D102"] -"pyathena/sqlalchemy/__init__.py" = ["D104"] -"pyathena/sqlalchemy/array.py" = ["D107"] -"pyathena/sqlalchemy/arrow.py" = ["D100"] -"pyathena/sqlalchemy/base.py" = ["D100", "D107"] -"pyathena/sqlalchemy/map.py" = ["D107"] -"pyathena/sqlalchemy/pandas.py" = ["D100"] -"pyathena/sqlalchemy/polars.py" = ["D100"] -"pyathena/sqlalchemy/preparer.py" = ["D100", "D107"] -"pyathena/sqlalchemy/requirements.py" = ["D100"] -"pyathena/sqlalchemy/rest.py" = ["D100"] -"pyathena/sqlalchemy/s3fs.py" = ["D100"] -"pyathena/sqlalchemy/struct.py" = ["D107"] -"pyathena/util.py" = ["D100", "D107"] [tool.mypy] follow_imports = "silent" From 8753ca8e19d8bc276ef98a5e523369d68fc46bc9 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 11:27:29 +0900 Subject: [PATCH 11/13] Say when on_start_query_execution and kill_on_interrupt apply The cursors call on_start_query_execution once _execute() returns a query ID, which is either a new query from StartQueryExecution or a reusable query found through cache_size. The docstrings and the usage guide said it runs right after StartQueryExecution or "when the query starts", which does not hold for a cache hit. Since #853, kill_on_interrupt also cancels a query whose start is interrupted. AsyncCursor waits on worker threads, which do not receive KeyboardInterrupt, so it only covers the start there. Co-Authored-By: Claude Opus 5.5 --- docs/usage.md | 1 + pyathena/aio/arrow/cursor.py | 4 +++- pyathena/aio/cursor.py | 4 +++- pyathena/aio/pandas/cursor.py | 4 +++- pyathena/aio/polars/cursor.py | 4 +++- pyathena/aio/s3fs/cursor.py | 4 +++- pyathena/arrow/cursor.py | 4 +++- pyathena/async_cursor.py | 5 +++-- pyathena/common.py | 10 ++++++---- pyathena/connection.py | 3 ++- pyathena/cursor.py | 5 +++-- pyathena/options.py | 5 +++-- pyathena/pandas/cursor.py | 4 +++- pyathena/polars/cursor.py | 4 +++- pyathena/s3fs/cursor.py | 4 +++- 15 files changed, 45 insertions(+), 20 deletions(-) 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/aio/arrow/cursor.py b/pyathena/aio/arrow/cursor.py index 7a2576c9d..304a2c4e9 100644 --- a/pyathena/aio/arrow/cursor.py +++ b/pyathena/aio/arrow/cursor.py @@ -135,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/cursor.py b/pyathena/aio/cursor.py index f924ef013..2274bcf9c 100644 --- a/pyathena/aio/cursor.py +++ b/pyathena/aio/cursor.py @@ -129,7 +129,9 @@ async def execute( 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. diff --git a/pyathena/aio/pandas/cursor.py b/pyathena/aio/pandas/cursor.py index ccb6e29df..89e4bea64 100644 --- a/pyathena/aio/pandas/cursor.py +++ b/pyathena/aio/pandas/cursor.py @@ -158,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/cursor.py b/pyathena/aio/polars/cursor.py index 66ca2d814..e3c72c3d0 100644 --- a/pyathena/aio/polars/cursor.py +++ b/pyathena/aio/polars/cursor.py @@ -142,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/s3fs/cursor.py b/pyathena/aio/s3fs/cursor.py index a485d59ce..4e3197ac3 100644 --- a/pyathena/aio/s3fs/cursor.py +++ b/pyathena/aio/s3fs/cursor.py @@ -138,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/arrow/cursor.py b/pyathena/arrow/cursor.py index c968337f4..ded70ee77 100644 --- a/pyathena/arrow/cursor.py +++ b/pyathena/arrow/cursor.py @@ -161,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/async_cursor.py b/pyathena/async_cursor.py index b99068454..3363901c7 100644 --- a/pyathena/async_cursor.py +++ b/pyathena/async_cursor.py @@ -80,8 +80,9 @@ def __init__( 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: Stop the running query on ``KeyboardInterrupt`` - while polling. + 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. diff --git a/pyathena/common.py b/pyathena/common.py index 438638744..31c445db0 100644 --- a/pyathena/common.py +++ b/pyathena/common.py @@ -227,11 +227,12 @@ def __init__( 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: Stop the running query when polling is interrupted. + 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 the query ID, by cursors - whose ``execute()`` supports it. + 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. @@ -1225,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) diff --git a/pyathena/connection.py b/pyathena/connection.py index 3607276df..826ff69d3 100644 --- a/pyathena/connection.py +++ b/pyathena/connection.py @@ -233,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/cursor.py b/pyathena/cursor.py index 6fdffb72c..bc4f531e7 100644 --- a/pyathena/cursor.py +++ b/pyathena/cursor.py @@ -134,8 +134,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 function called immediately after - start_query_execution API is called. + 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. diff --git a/pyathena/options.py b/pyathena/options.py index 92ce52f38..33d95bf9e 100644 --- a/pyathena/options.py +++ b/pyathena/options.py @@ -56,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/cursor.py b/pyathena/pandas/cursor.py index 8f8913c53..81738492a 100644 --- a/pyathena/pandas/cursor.py +++ b/pyathena/pandas/cursor.py @@ -180,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/polars/cursor.py b/pyathena/polars/cursor.py index fdf8cd201..53288bf91 100644 --- a/pyathena/polars/cursor.py +++ b/pyathena/polars/cursor.py @@ -180,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/s3fs/cursor.py b/pyathena/s3fs/cursor.py index cef1eca90..839b8c128 100644 --- a/pyathena/s3fs/cursor.py +++ b/pyathena/s3fs/cursor.py @@ -156,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. From 35fad823de412d8d5bb5d06bb067cf2535bd18e5 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 11:45:16 +0900 Subject: [PATCH 12/13] Align the new docstrings with the current interrupt and cache behavior The new kill_on_interrupt descriptions predate #853: the SQL cursors also cancel a query whose start is interrupted, the asyncio cursors react to task cancellation while starting or waiting, and the AsyncCursor variants only cover the start. S3FileSystem.info() also treats a cached listing of the path as a directory and a cached parent listing without it as missing, and S3File copies a large existing object for append only once a multipart upload starts. Co-Authored-By: Claude Opus 5.5 --- pyathena/aio/arrow/cursor.py | 4 ++-- pyathena/aio/cursor.py | 4 ++-- pyathena/aio/pandas/cursor.py | 4 ++-- pyathena/aio/polars/cursor.py | 4 ++-- pyathena/aio/s3fs/cursor.py | 4 ++-- pyathena/cursor.py | 4 ++-- pyathena/filesystem/s3.py | 8 +++++--- pyathena/pandas/async_cursor.py | 4 +++- 8 files changed, 20 insertions(+), 16 deletions(-) diff --git a/pyathena/aio/arrow/cursor.py b/pyathena/aio/arrow/cursor.py index 304a2c4e9..f1b4aa71d 100644 --- a/pyathena/aio/arrow/cursor.py +++ b/pyathena/aio/arrow/cursor.py @@ -66,8 +66,8 @@ def __init__( 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: Whether to stop the running query when the waiting - task is cancelled. + 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. diff --git a/pyathena/aio/cursor.py b/pyathena/aio/cursor.py index 2274bcf9c..89e7c5aa6 100644 --- a/pyathena/aio/cursor.py +++ b/pyathena/aio/cursor.py @@ -62,8 +62,8 @@ def __init__( 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: Stop the running query when polling is cancelled - with ``asyncio.CancelledError``. + 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 diff --git a/pyathena/aio/pandas/cursor.py b/pyathena/aio/pandas/cursor.py index 89e4bea64..413c291d0 100644 --- a/pyathena/aio/pandas/cursor.py +++ b/pyathena/aio/pandas/cursor.py @@ -75,8 +75,8 @@ def __init__( 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: Whether to stop the running query when the waiting - task is cancelled. + 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. If set, it takes precedence diff --git a/pyathena/aio/polars/cursor.py b/pyathena/aio/polars/cursor.py index e3c72c3d0..e7ca9d718 100644 --- a/pyathena/aio/polars/cursor.py +++ b/pyathena/aio/polars/cursor.py @@ -70,8 +70,8 @@ def __init__( 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: Whether to stop the running query when the waiting - task is cancelled. + 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. diff --git a/pyathena/aio/s3fs/cursor.py b/pyathena/aio/s3fs/cursor.py index 4e3197ac3..91a176170 100644 --- a/pyathena/aio/s3fs/cursor.py +++ b/pyathena/aio/s3fs/cursor.py @@ -67,8 +67,8 @@ def __init__( 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: Whether to stop the running query when the waiting - task is cancelled. + 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, diff --git a/pyathena/cursor.py b/pyathena/cursor.py index bc4f531e7..d0b6de4f3 100644 --- a/pyathena/cursor.py +++ b/pyathena/cursor.py @@ -68,8 +68,8 @@ def __init__( 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: Stop the running query on ``KeyboardInterrupt`` - while polling. + 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 diff --git a/pyathena/filesystem/s3.py b/pyathena/filesystem/s3.py index 325ddf2dd..4104f9354 100644 --- a/pyathena/filesystem/s3.py +++ b/pyathena/filesystem/s3.py @@ -552,7 +552,9 @@ def _list_object_versions_pages( def info(self, path: str, **kwargs) -> S3Object: """Return information about an S3 path. - Returns a matching entry from the directory cache when one exists. + 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 @@ -2109,8 +2111,8 @@ def __init__( 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, and a larger one is copied as the first - parts of the multipart upload. + 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. diff --git a/pyathena/pandas/async_cursor.py b/pyathena/pandas/async_cursor.py index 63c225292..2c66aa7a6 100644 --- a/pyathena/pandas/async_cursor.py +++ b/pyathena/pandas/async_cursor.py @@ -92,7 +92,9 @@ def __init__( 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: Whether to stop the running query on ``KeyboardInterrupt``. + 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. From 6be315ec4c87758f993778b0b2167696e848390a Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 11:53:16 +0900 Subject: [PATCH 13/13] Limit the chunking and description docstrings to where they apply Pandas chunksize only applies to CSV results (UNLOAD reads the whole Parquet output), Polars reads lazily only from result files in S3, and description is None only for INSERT, UPDATE, DELETE and MERGE, not for every DML statement type. Co-Authored-By: Claude Opus 5.5 --- pyathena/aio/pandas/cursor.py | 4 ++-- pyathena/aio/polars/cursor.py | 4 ++-- pyathena/pandas/async_cursor.py | 2 +- pyathena/result_set.py | 5 ++++- 4 files changed, 9 insertions(+), 6 deletions(-) diff --git a/pyathena/aio/pandas/cursor.py b/pyathena/aio/pandas/cursor.py index 413c291d0..967ec3e4f 100644 --- a/pyathena/aio/pandas/cursor.py +++ b/pyathena/aio/pandas/cursor.py @@ -79,8 +79,8 @@ def __init__( ``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. If set, it takes precedence - over ``auto_optimize_chunksize``. + 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. diff --git a/pyathena/aio/polars/cursor.py b/pyathena/aio/polars/cursor.py index e7ca9d718..8338c9d61 100644 --- a/pyathena/aio/polars/cursor.py +++ b/pyathena/aio/polars/cursor.py @@ -78,8 +78,8 @@ def __init__( 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, results are read lazily - in chunks of this size. + 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__``. """ diff --git a/pyathena/pandas/async_cursor.py b/pyathena/pandas/async_cursor.py index 2c66aa7a6..bd2279ea6 100644 --- a/pyathena/pandas/async_cursor.py +++ b/pyathena/pandas/async_cursor.py @@ -99,7 +99,7 @@ def __init__( 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. + 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``, diff --git a/pyathena/result_set.py b/pyathena/result_set.py index 9467aeb65..152715fdb 100644 --- a/pyathena/result_set.py +++ b/pyathena/result_set.py @@ -371,7 +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, or None without metadata or for DML.""" + """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