diff --git a/pyathena/sqlalchemy/base.py b/pyathena/sqlalchemy/base.py index 8eada2056..5295d08e8 100644 --- a/pyathena/sqlalchemy/base.py +++ b/pyathena/sqlalchemy/base.py @@ -8,6 +8,7 @@ from typing import ( TYPE_CHECKING, Any, + ClassVar, cast, ) @@ -143,7 +144,7 @@ class AthenaDialect(DefaultDialect): preparer: type[IdentifierPreparer] = AthenaDMLIdentifierPreparer statement_compiler: type[SQLCompiler] = AthenaStatementCompiler ddl_compiler: type[DDLCompiler] = AthenaDDLCompiler - type_compiler: type[GenericTypeCompiler] = AthenaTypeCompiler + type_compiler_cls: ClassVar[type[GenericTypeCompiler]] = AthenaTypeCompiler default_paramstyle: str = pyathena.paramstyle max_identifier_length: int = 255 cte_follows_insert: bool = True diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index 4211643b5..8687224e0 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -1182,7 +1182,7 @@ def get_column_specification(self, column: Column[Any], **kwargs) -> str: type_ = "INT" else: # type_expression marks column DDL so STRUCT and MAP use Hive syntax. - type_ = self.dialect.type_compiler.process(column.type, type_expression=column) + type_ = self.dialect.type_compiler_instance.process(column.type, type_expression=column) text = [f"{self.preparer.format_column(column)} {type_}"] if column.comment: text.append(f"{self._get_comment_specification(column.comment)}") diff --git a/tests/pyathena/sqlalchemy/test_base.py b/tests/pyathena/sqlalchemy/test_base.py index bf010643f..d956232c9 100644 --- a/tests/pyathena/sqlalchemy/test_base.py +++ b/tests/pyathena/sqlalchemy/test_base.py @@ -19,11 +19,13 @@ from sqlalchemy.sql.schema import Column, MetaData, Table from sqlalchemy.sql.selectable import TextualSelect +from pyathena.aio.sqlalchemy.base import AthenaAioDialect from pyathena.converter import DefaultTypeConverter from pyathena.cursor import Cursor from pyathena.error import DatabaseError, OperationalError from pyathena.formatter import DefaultParameterFormatter from pyathena.sqlalchemy.base import AthenaDialect +from pyathena.sqlalchemy.compiler import AthenaTypeCompiler from pyathena.sqlalchemy.types import ( TINYINT, AthenaArray, @@ -108,6 +110,14 @@ def close(self): class TestAthenaDialect: + @pytest.mark.parametrize("dialect_class", [AthenaDialect, AthenaAioDialect]) + def test_type_compiler(self, dialect_class): + # SQLAlchemy 2.0 builds the type compiler from type_compiler_cls. A legacy + # type_compiler class attribute would take precedence over it. + assert not hasattr(dialect_class, "type_compiler") + assert dialect_class.type_compiler_cls is AthenaTypeCompiler + assert isinstance(dialect_class().type_compiler_instance, AthenaTypeCompiler) + def test_columns_from_information_schema(self): # Rows arrive unordered, and Athena reports a missing comment as NULL. # The API cursor this path pins hands that over as None or as an empty