diff --git a/pyathena/aio/cursor.py b/pyathena/aio/cursor.py index 89e7c5aa6..00a57be0b 100644 --- a/pyathena/aio/cursor.py +++ b/pyathena/aio/cursor.py @@ -85,6 +85,7 @@ def __init__( ) self._result_set: AthenaAioResultSet | None = None self._result_set_class = AthenaAioResultSet + self._result_set_kwargs: dict[str, Any] = {} @property # type: ignore[explicit-override] # python/mypy#15900 @override @@ -174,6 +175,7 @@ async def execute( self.arraysize, self._retry_config, result_set_type_hints=options.result_set_type_hints, + **self._result_set_kwargs, ) else: raise OperationalError(query_execution.state_change_reason) @@ -246,15 +248,15 @@ class AioDictCursor(AioCursor): ... print(row["name"]) """ - def __init__(self, **kwargs) -> None: + def __init__(self, dict_type: type[Any] | None = None, **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. + dict_type: The type used to build each row of this cursor's result + sets. If None, the result set class's ``dict_type`` is used. + **kwargs: Arguments forwarded to ``AioCursor.__init__``. """ super().__init__(**kwargs) self._result_set_class = AthenaAioDictResultSet - if "dict_type" in kwargs: - AthenaAioDictResultSet.dict_type = kwargs["dict_type"] + if dict_type is not None: + self._result_set_kwargs = {"dict_type": dict_type} diff --git a/pyathena/aio/result_set.py b/pyathena/aio/result_set.py index 5a628da7e..59299f101 100644 --- a/pyathena/aio/result_set.py +++ b/pyathena/aio/result_set.py @@ -75,6 +75,7 @@ async def create( arraysize: int, retry_config: RetryConfig, result_set_type_hints: dict[str | int, str] | None = None, + **kwargs: Any, ) -> AthenaAioResultSet: """Async factory method. @@ -88,6 +89,8 @@ async def create( retry_config: Retry configuration for API calls. result_set_type_hints: Optional dictionary mapping column names to Athena DDL type signatures for precise type conversion. + **kwargs: Additional arguments passed to the constructor of ``cls``, + such as ``dict_type`` for ``AthenaAioDictResultSet``. Returns: A fully initialized ``AthenaAioResultSet``. @@ -99,6 +102,7 @@ async def create( arraysize, retry_config, result_set_type_hints=result_set_type_hints, + **kwargs, ) if result_set.state == AthenaQueryExecution.STATE_SUCCEEDED: await result_set._async_pre_fetch() diff --git a/pyathena/async_cursor.py b/pyathena/async_cursor.py index 3363901c7..2963e5f40 100644 --- a/pyathena/async_cursor.py +++ b/pyathena/async_cursor.py @@ -112,6 +112,7 @@ def __init__( self._max_workers = max_workers self._executor = ThreadPoolExecutor(max_workers=max_workers) self._result_set_class = AthenaResultSet + self._result_set_kwargs: dict[str, Any] = {} @property def arraysize(self) -> int: @@ -198,6 +199,7 @@ def _collect_result_set( arraysize=self._arraysize, retry_config=self._retry_config, result_set_type_hints=result_set_type_hints, + **self._result_set_kwargs, ) @override @@ -335,15 +337,15 @@ class AsyncDictCursor(AsyncCursor): >>> print(f"User: {row['name']} ({row['email']})") """ - def __init__(self, **kwargs) -> None: + def __init__(self, dict_type: type[Any] | None = None, **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. + dict_type: The type used to build each row of this cursor's result + sets. If None, the result set class's ``dict_type`` is used. + **kwargs: Arguments forwarded to ``AsyncCursor.__init__``. """ super().__init__(**kwargs) self._result_set_class = AthenaDictResultSet - if "dict_type" in kwargs: - AthenaDictResultSet.dict_type = kwargs["dict_type"] + if dict_type is not None: + self._result_set_kwargs = {"dict_type": dict_type} diff --git a/pyathena/cursor.py b/pyathena/cursor.py index d0b6de4f3..099f8e8eb 100644 --- a/pyathena/cursor.py +++ b/pyathena/cursor.py @@ -90,6 +90,7 @@ def __init__( **kwargs, ) self._result_set_class = AthenaResultSet + self._result_set_kwargs: dict[str, Any] = {} @property # type: ignore[explicit-override] # python/mypy#15900 @override @@ -192,6 +193,7 @@ def execute( self.arraysize, self._retry_config, result_set_type_hints=options.result_set_type_hints, + **self._result_set_kwargs, ) else: raise OperationalError(query_execution.state_change_reason) @@ -216,15 +218,15 @@ class DictCursor(Cursor): ... print(f"Product {row['id']}: {row['name']} - ${row['price']}") """ - def __init__(self, **kwargs) -> None: + def __init__(self, dict_type: type[Any] | None = None, **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. + dict_type: The type used to build each row of this cursor's result + sets. If None, the result set class's ``dict_type`` is used. + **kwargs: Arguments forwarded to ``Cursor.__init__``. """ super().__init__(**kwargs) self._result_set_class = AthenaDictResultSet - if "dict_type" in kwargs: - AthenaDictResultSet.dict_type = kwargs["dict_type"] + if dict_type is not None: + self._result_set_kwargs = {"dict_type": dict_type} diff --git a/pyathena/result_set.py b/pyathena/result_set.py index 152715fdb..adbb36c96 100644 --- a/pyathena/result_set.py +++ b/pyathena/result_set.py @@ -818,6 +818,19 @@ class AthenaDictResultSet(AthenaResultSet): # You can override this to use OrderedDict or other dict-like types. dict_type: type[Any] = dict + def __init__(self, *args: Any, dict_type: type[Any] | None = None, **kwargs: Any) -> None: + """Initialize the result set with an optional row type for this instance. + + Args: + *args: Positional arguments passed to the next ``__init__`` in the MRO. + dict_type: The type used to build each row of this result set. If + None, the class attribute ``dict_type`` is used. + **kwargs: Keyword arguments passed to the next ``__init__`` in the MRO. + """ + if dict_type is not None: + self.dict_type = dict_type + super().__init__(*args, **kwargs) + @override def _get_rows( self, diff --git a/tests/pyathena/aio/test_cursor.py b/tests/pyathena/aio/test_cursor.py index 568449568..4561b4613 100644 --- a/tests/pyathena/aio/test_cursor.py +++ b/tests/pyathena/aio/test_cursor.py @@ -1,6 +1,7 @@ import asyncio import re import threading +from collections import OrderedDict from datetime import UTC, datetime, timedelta from unittest.mock import AsyncMock, MagicMock, call, patch @@ -8,7 +9,8 @@ from botocore.exceptions import ClientError from pyathena import BINARY, Binary, ExecuteOptions -from pyathena.aio.cursor import AioCursor +from pyathena.aio.cursor import AioCursor, AioDictCursor +from pyathena.aio.result_set import AthenaAioDictResultSet from pyathena.error import DatabaseError, OperationalError, ProgrammingError from pyathena.glue import GlueMetadataClient from pyathena.model import AthenaQueryExecution @@ -896,3 +898,27 @@ async def test_fetchall(self, aio_dict_cursor): assert await aio_dict_cursor.fetchall() == [{"number_of_rows": 1}] await aio_dict_cursor.execute("SELECT a FROM many_rows ORDER BY a") assert await aio_dict_cursor.fetchall() == [{"a": i} for i in range(10000)] + + async def test_dict_type(self, aio_dict_cursor): + async with aio_dict_cursor.connection.cursor(dict_type=OrderedDict) as ordered_cursor: + await ordered_cursor.execute("SELECT * FROM one_row") + assert type(await ordered_cursor.fetchone()) is OrderedDict + # dict_type of another cursor does not change the row type of this one. + await aio_dict_cursor.execute("SELECT * FROM one_row") + assert type(await aio_dict_cursor.fetchone()) is dict + + async def test_dict_type_custom_result_set(self, aio_dict_cursor): + class CustomResultSet(AthenaAioDictResultSet): + pass + + class CustomDictCursor(AioDictCursor): + def __init__(self, **kwargs): + super().__init__(**kwargs) + self._result_set_class = CustomResultSet + + async with aio_dict_cursor.connection.cursor( + CustomDictCursor, dict_type=OrderedDict + ) as cursor: + await cursor.execute("SELECT * FROM one_row") + assert isinstance(cursor.result_set, CustomResultSet) + assert type(await cursor.fetchone()) is OrderedDict diff --git a/tests/pyathena/test_async_cursor.py b/tests/pyathena/test_async_cursor.py index 1cf4d0d91..d08ed5801 100644 --- a/tests/pyathena/test_async_cursor.py +++ b/tests/pyathena/test_async_cursor.py @@ -1,14 +1,15 @@ import contextlib import time +from collections import OrderedDict from datetime import datetime from random import randint import pytest -from pyathena.async_cursor import AsyncCursor +from pyathena.async_cursor import AsyncCursor, AsyncDictCursor from pyathena.error import NotSupportedError, ProgrammingError from pyathena.model import AthenaQueryExecution -from pyathena.result_set import AthenaResultSet +from pyathena.result_set import AthenaDictResultSet, AthenaResultSet from tests import ENV from tests.pyathena.conftest import connect @@ -246,3 +247,26 @@ def test_fetchall(self, async_dict_cursor): query_id, future = async_dict_cursor.execute("SELECT a FROM many_rows ORDER BY a") result_set = future.result() assert result_set.fetchall() == [{"a": i} for i in range(10000)] + + def test_dict_type(self, async_dict_cursor): + with async_dict_cursor.connection.cursor(dict_type=OrderedDict) as ordered_cursor: + _, future = ordered_cursor.execute("SELECT * FROM one_row") + assert type(future.result().fetchone()) is OrderedDict + # dict_type of another cursor does not change the row type of this one. + _, future = async_dict_cursor.execute("SELECT * FROM one_row") + assert type(future.result().fetchone()) is dict + + def test_dict_type_custom_result_set(self, async_dict_cursor): + class CustomResultSet(AthenaDictResultSet): + pass + + class CustomDictCursor(AsyncDictCursor): + def __init__(self, **kwargs): + super().__init__(**kwargs) + self._result_set_class = CustomResultSet + + with async_dict_cursor.connection.cursor(CustomDictCursor, dict_type=OrderedDict) as cursor: + _, future = cursor.execute("SELECT * FROM one_row") + result_set = future.result() + assert isinstance(result_set, CustomResultSet) + assert type(result_set.fetchone()) is OrderedDict diff --git a/tests/pyathena/test_cursor.py b/tests/pyathena/test_cursor.py index 720a5b230..22afd5781 100644 --- a/tests/pyathena/test_cursor.py +++ b/tests/pyathena/test_cursor.py @@ -7,6 +7,7 @@ import threading import time import uuid +from collections import OrderedDict from concurrent import futures from concurrent.futures.thread import ThreadPoolExecutor from datetime import UTC, date, datetime, timedelta @@ -31,9 +32,10 @@ ) from pyathena.async_cursor import AsyncCursor from pyathena.converter import _to_array, _to_map, _to_struct -from pyathena.cursor import Cursor +from pyathena.cursor import Cursor, DictCursor from pyathena.error import DatabaseError, NotSupportedError, OperationalError, ProgrammingError from pyathena.model import AthenaQueryExecution +from pyathena.result_set import AthenaDictResultSet from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect @@ -1876,6 +1878,28 @@ def test_fetchall(self, dict_cursor): dict_cursor.execute("SELECT a FROM many_rows ORDER BY a") assert dict_cursor.fetchall() == [{"a": i} for i in range(10000)] + def test_dict_type(self, dict_cursor): + with dict_cursor.connection.cursor(dict_type=OrderedDict) as ordered_cursor: + ordered_cursor.execute("SELECT * FROM one_row") + assert type(ordered_cursor.fetchone()) is OrderedDict + # dict_type of another cursor does not change the row type of this one. + dict_cursor.execute("SELECT * FROM one_row") + assert type(dict_cursor.fetchone()) is dict + + def test_dict_type_custom_result_set(self, dict_cursor): + class CustomResultSet(AthenaDictResultSet): + pass + + class CustomDictCursor(DictCursor): + def __init__(self, **kwargs): + super().__init__(**kwargs) + self._result_set_class = CustomResultSet + + with dict_cursor.connection.cursor(CustomDictCursor, dict_type=OrderedDict) as cursor: + cursor.execute("SELECT * FROM one_row") + assert isinstance(cursor.result_set, CustomResultSet) + assert type(cursor.fetchone()) is OrderedDict + def test_null_vs_empty_string(self, dict_cursor): """ DictCursor should properly distinguish NULL from empty string.