From 268f8d6b959e19f01a7666f67afab8ed087ee990 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 15:15:55 +0900 Subject: [PATCH 1/6] Split SQLAlchemy type rendering into Hive DDL and Trino DML compilers Athena parses DDL with Hive type syntax and queries with Trino type syntax. AthenaTypeCompiler mixed both through the type_expression column check and the _athena_hive_ddl flag, so direct compilation produced hybrids such as ROW(name STRING) and invalid DDL for empty STRUCT, JSON, and CLOB columns. AthenaTypeCompiler now renders Hive DDL types only: INT, STRING, STRUCT, MAP, and ARRAY in every context. JSON and a STRUCT without fields raise CompileError, CLOB/NCLOB render STRING, and an unexpected type class in visit_struct/map/array raises instead of falling back to a placeholder. The new AthenaDMLTypeCompiler renders the Trino types of CAST expressions and replaces the type branches of visit_cast and _complex_dml_type. CAST output is unchanged except that an empty ROW now raises CompileError. Refs #886 Co-Authored-By: Claude Opus 5.5 --- docs/api/sqlalchemy.rst | 3 + docs/sqlalchemy.md | 27 +- pyathena/sqlalchemy/array.py | 31 +- pyathena/sqlalchemy/compiler.py | 551 +++++++++++++-------- tests/pyathena/sqlalchemy/test_compiler.py | 139 ++++-- 5 files changed, 479 insertions(+), 272 deletions(-) diff --git a/docs/api/sqlalchemy.rst b/docs/api/sqlalchemy.rst index ebd68c0d9..910951f62 100644 --- a/docs/api/sqlalchemy.rst +++ b/docs/api/sqlalchemy.rst @@ -55,6 +55,9 @@ Compilers .. autoclass:: pyathena.sqlalchemy.compiler.AthenaTypeCompiler :members: +.. autoclass:: pyathena.sqlalchemy.compiler.AthenaDMLTypeCompiler + :members: + .. autoclass:: pyathena.sqlalchemy.compiler.AthenaStatementCompiler :members: diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index 16e244188..179502aaa 100644 --- a/docs/sqlalchemy.md +++ b/docs/sqlalchemy.md @@ -737,6 +737,25 @@ engine_arrow = create_engine( ) ``` +## DDL and CAST types + +Athena parses DDL statements such as `CREATE TABLE` with Hive type syntax, and queries with Trino type syntax. +`CREATE TABLE` column types, and types compiled with `TypeEngine.compile()`, use the Hive syntax. +`CAST` target types use the Trino syntax. + +| SQLAlchemy type | Table DDL | CAST | +|---|---|---| +| `Integer`, `INTEGER` | `INT` | `INTEGER` | +| `String`, `Text`, `CLOB` | `STRING` | `VARCHAR` | +| `LargeBinary`, `BINARY`, `VARBINARY` | `BINARY` | `VARBINARY` | +| `JSON` | Raises `CompileError` | `JSON` | +| `AthenaStruct` | `STRUCT` | `ROW(name type, ...)` | +| `AthenaMap` | `MAP` | `MAP(key, value)` | +| `AthenaArray`, `ARRAY` | `ARRAY` | `ARRAY(item)` | + +Complex types apply the same syntax to their nested types. +An `AthenaStruct` without fields raises `CompileError` in both, because Athena has no empty STRUCT type. + ## Floating-point types | SQLAlchemy type | Table DDL | CAST | @@ -824,10 +843,8 @@ CREATE TABLE users ( ) ``` -`CREATE TABLE` renders `AthenaStruct` columns with Hive `STRUCT` syntax at every nesting depth. -That includes top-level columns, fields of a STRUCT, STRUCT values inside MAP, and STRUCT values inside ARRAY. -Integer fields, and integer MAP keys and values, use `INT` in that DDL. -`CAST` and other SQL expressions keep `ROW(...)`, `MAP(...)`, and `ARRAY(...)`, and spell integers as `INTEGER`. +`CREATE TABLE` renders `AthenaStruct` columns with Hive `STRUCT` syntax at every nesting depth, and `CAST` renders them as `ROW(name type, ...)`. +See [DDL and CAST types](#ddl-and-cast-types). #### Querying STRUCT data @@ -1367,7 +1384,7 @@ Athena's JSON type support has specific limitations: - **JSON objects and arrays are supported** - `CAST('...' AS JSON)` accepts an object or a top-level array such as `[1, 2, 3]` - **Arrays within objects are supported** - JSON objects can contain arrays as property values -- **DML only** - JSON type is supported for SELECT queries but not in CREATE TABLE statements +- **DML only** - JSON type is supported for SELECT queries but not in CREATE TABLE statements; compiling `CREATE TABLE` with a `JSON` column raises `CompileError` ```python # Supported: JSON object with nested array diff --git a/pyathena/sqlalchemy/array.py b/pyathena/sqlalchemy/array.py index ec61f28d6..9fbba051c 100644 --- a/pyathena/sqlalchemy/array.py +++ b/pyathena/sqlalchemy/array.py @@ -235,6 +235,27 @@ def decorator_impl(self, type_: types.TypeDecorator[Any]) -> TypeEngine[Any]: return variant return type_.load_dialect_impl(self.dialect) + def dialect_type(self, type_: TypeEngine[Any]) -> TypeEngine[Any]: + """Resolve the type this dialect uses for a SQLAlchemy type. + + Takes the Athena variant from ``with_variant()`` and the implementation + of a TypeDecorator until neither applies. + + Args: + type_: The declared type. + + Returns: + The resolved type. + """ + while True: + variant = self.variant(type_) + if variant is not None: + type_ = variant + elif isinstance(type_, types.TypeDecorator): + type_ = self.decorator_impl(type_) + else: + return type_ + @staticmethod def has_unknown_element(type_: TypeEngine[Any]) -> bool: if isinstance(type_, sqltypes.ARRAY): @@ -579,7 +600,9 @@ def process(self, expression, **kw): ): raise exc.CompileError("An ARRAY slice assignment requires a non-NULL array") rhs = compiler.process(value, **kw) - rhs_type = compiler._complex_dml_type(expression.value_type, require_precision=True) + rhs_type = compiler._dml_type_compiler.process( + expression.value_type, require_precision=True + ) rhs = f"CAST({rhs} AS {rhs_type})" if final_slice: # Reject SQL expressions that evaluate to NULL without issuing a second statement. @@ -626,7 +649,7 @@ def _rebuild(self, array, array_type, path, rhs, **kw): array_type = self._type_inspector.array_type(array_type) if array_type is None: raise exc.CompileError("Partial ARRAY updates require an ARRAY column type") - array_sql_type = compiler._complex_dml_type(array_type) + array_sql_type = compiler._dml_type_compiler.process(array_type) array = f"coalesce({array}, CAST(ARRAY[] AS {array_sql_type}))" bound = path[0] if isinstance(bound, Slice): @@ -635,7 +658,9 @@ def _rebuild(self, array, array_type, path, rhs, **kw): def _prefix_and_padding(self, array, start, array_type): prefix = f"slice({array}, 1, least({start} - 1, cardinality({array})))" - element_type = self.compiler._complex_dml_type(_ArrayTypeInspector.item_type(array_type)) + element_type = self.compiler._dml_type_compiler.process( + _ArrayTypeInspector.item_type(array_type) + ) padding = ( f"repeat(CAST(NULL AS {element_type}), " f"CAST(greatest({start} - 1 - cardinality({array}), 0) AS INTEGER))" diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index a68ad60ea..5e5c6797a 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -4,7 +4,6 @@ import re from collections.abc import Mapping -from functools import partial from itertools import product from typing import TYPE_CHECKING, Any, cast @@ -78,25 +77,20 @@ class AthenaTypeCompiler(GenericTypeCompiler): - """Type compiler for Amazon Athena SQL types. + """Type compiler for Amazon Athena DDL types. - This compiler translates SQLAlchemy type objects into Athena-compatible - SQL type strings for use in DDL statements. It handles the mapping between - SQLAlchemy's portable types and Athena's specific type syntax. + Athena parses DDL statements such as CREATE TABLE with Hive type syntax, + and queries with Trino type syntax. This compiler renders the Hive types. + It is the dialect's type compiler, so it renders CREATE TABLE column types + and ``TypeEngine.compile()``. ``AthenaStatementCompiler`` renders the + types of CAST expressions with ``AthenaDMLTypeCompiler``. - Athena has specific requirements for type names that differ from standard - SQL. For example, FLOAT and REAL render as FLOAT here, while the statement - compiler renders them as REAL in CAST expressions. TEXT, and CHAR, NCHAR, - VARCHAR, or NVARCHAR without a length, render as STRING; with a length, - they render as CHAR(n) or VARCHAR(n). - - The compiler also supports Athena-specific complex types: - - STRUCT/ROW: Nested record types with named fields - - MAP: Key-value pair collections - - ARRAY: Ordered collections of elements - - CREATE TABLE columns render STRUCT fields as Hive ``STRUCT``. - Compiling a type on its own renders ``ROW(...)``. + Integers render as INT, FLOAT and REAL as FLOAT, and binary types as + BINARY. TEXT, CLOB, and character types without a length render as + STRING; with a length, they render as CHAR(n) or VARCHAR(n). Complex + types render as ``STRUCT``, ``MAP``, and + ``ARRAY``. TIME, JSON, and a STRUCT without fields have no Athena + DDL type and raise ``CompileError``. See Also: AWS Athena Data Types: @@ -145,7 +139,7 @@ def visit_TINYINT(self, type_: types.Integer, **kw: Any) -> str: @override def visit_INTEGER(self, type_: types.Integer, **kw: Any) -> str: - return "INT" if kw.get("_athena_hive_ddl") else "INTEGER" + return "INT" @override def visit_SMALLINT(self, type_: types.SmallInteger, **kw: Any) -> str: @@ -173,11 +167,11 @@ def visit_TIME(self, type_: types.Time, **kw: Any) -> str: @override def visit_CLOB(self, type_: types.CLOB, **kw: Any) -> str: - return self.visit_BINARY(type_, **kw) # type: ignore[arg-type] + return self.visit_TEXT(type_, **kw) @override def visit_NCLOB(self, type_: types.Text, **kw: Any) -> str: - return self.visit_BINARY(type_, **kw) # type: ignore[arg-type] + return self.visit_TEXT(type_, **kw) @override def visit_CHAR(self, type_: types.CHAR, **kw: Any) -> str: @@ -220,16 +214,16 @@ 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. + """Reject a JSON type, which Athena DDL does not accept. Args: type_: The type to render. **kw: Type-compiler keyword arguments. - Returns: - ``JSON``. + Raises: + CompileError: Always. """ - return "JSON" + raise exc.CompileError(f"Data type `{type_}` is not supported in Athena DDL") @override def visit_string(self, type_, **kw): @@ -263,58 +257,305 @@ def visit_tinyint(self, type_, **kw): def visit_enum(self, type_, **kw): return self.visit_string(type_, **kw) - def _enable_hive_column_ddl(self, kw: dict[str, Any]) -> bool: - """Enable Hive spelling for a CREATE TABLE column type. + @util.memoized_property + def _preparer(self) -> AthenaDDLIdentifierPreparer: + """The DDL identifier preparer that quotes STRUCT field names.""" + return AthenaDDLIdentifierPreparer(self.dialect) - ``get_column_specification`` passes the column as ``type_expression``. - ARRAY compilation sets ``_athena_hive_ddl`` so nested fields use - ``STRUCT`` and ``INT``. STRUCT and MAP reuse that flag in - column DDL. Direct compilation and CAST leave it unset. + def visit_struct(self, type_, **kw): + """Render a STRUCT type as ``STRUCT``. Args: - kw: Type-compiler keyword arguments. When Hive spelling applies, - ``_athena_hive_ddl`` is set so nested types keep it. + type_: The type to render. + **kw: Type-compiler keyword arguments. Returns: - True when the type should use Hive DDL syntax. + The STRUCT type clause. + + Raises: + CompileError: If the type is not an ``AthenaStruct`` or has no fields. """ - if kw.get("_athena_hive_ddl") or isinstance(kw.get("type_expression"), Column): - kw["_athena_hive_ddl"] = True - return True - return False + if not isinstance(type_, AthenaStruct): + raise exc.CompileError(f"Cannot render `{type_!r}` as STRUCT") + if not type_.fields: + raise exc.CompileError("STRUCT requires at least one field") + fields = ", ".join( + f"{self._preparer.quote(name)}:{self.process(field_type, **kw)}" + for name, field_type in type_.fields.items() + ) + return f"STRUCT<{fields}>" - def visit_struct(self, type_, **kw): - """Render a STRUCT type. + 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 type clause. + """ + return self.visit_struct(type_, **kw) + + def visit_map(self, type_, **kw): + """Render a MAP type as ``MAP``. + + Args: + type_: The type to render. + **kw: Type-compiler keyword arguments. + + Returns: + The MAP type clause. + + Raises: + CompileError: If the type is not an ``AthenaMap``. + """ + if not isinstance(type_, AthenaMap): + raise exc.CompileError(f"Cannot render `{type_!r}` as MAP") + key_type_str = self.process(type_.key_type, **kw) + value_type_str = self.process(type_.value_type, **kw) + return f"MAP<{key_type_str}, {value_type_str}>" + + 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``. + + Args: + type_: The type to render. + **kw: Type-compiler keyword arguments. + + Returns: + The ARRAY type clause. + + Raises: + CompileError: If the type is not an ARRAY. + """ + if not isinstance(type_, types.ARRAY): + raise exc.CompileError(f"Cannot render `{type_!r}` as ARRAY") + return f"ARRAY<{self.process(_ArrayTypeInspector.item_type(type_), **kw)}>" + + 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) + + +class AthenaDMLTypeCompiler(GenericTypeCompiler): + """Type compiler for the Trino types of Amazon Athena queries. + + ``AthenaStatementCompiler`` renders the types of CAST expressions with + this compiler, while ``AthenaTypeCompiler`` renders the Hive types of + DDL statements. Character types render as VARCHAR, FLOAT and REAL as + REAL, binary types as VARBINARY, and DateTime types as ``TIMESTAMP(6)`` + or ``TIMESTAMP(precision)``. Complex types render as + ``ROW(name type)``, ``MAP(key, value)``, and ``ARRAY(item)``. + + A type is resolved through its ``with_variant()`` type for this dialect + and the implementation of a TypeDecorator before it is rendered. + + Two keyword arguments of ``process()`` adjust the rendering: + + - ``require_precision``: Reject a Numeric type without a precision. + - ``timestamp_precision``: When False, render DateTime types as a bare + ``TIMESTAMP``. A cast that only types an empty value uses it, because + Athena widens a bare ``TIMESTAMP`` to the other operand's precision + instead of widening that operand. + + See Also: + AWS Athena Data Types: + https://docs.aws.amazon.com/athena/latest/ug/data-types.html + """ + + @util.memoized_property + def _type_inspector(self) -> _ArrayTypeInspector: + """The inspector that resolves variants and TypeDecorators for this dialect.""" + return _ArrayTypeInspector(self.dialect) + + @override + def process(self, type_: TypeEngine[Any], **kw: Any) -> str: + return self._type_inspector.dialect_type(type_)._compiler_dispatch(self, **kw) + + def _process_element(self, type_: TypeEngine[Any], **kw: Any) -> str: + """Render the element type of an ARRAY, MAP, or ROW. + + Args: + type_: The element type. + **kw: Type-compiler keyword arguments. + + Returns: + The element type clause. - 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()``. + Raises: + CompileError: If the element type is unknown. + """ + if isinstance(self._type_inspector.dialect_type(type_), types.NullType): + raise exc.CompileError("Bound ARRAY values require an explicit element type") + return self.process(type_, **kw) + + @override + def visit_FLOAT(self, type_: types.Float[Any], **kw: Any) -> str: + return "REAL" + + @override + def visit_REAL(self, type_: types.REAL[Any], **kw: Any) -> str: + return "REAL" + + @override + def visit_DOUBLE_PRECISION(self, type_, **kw) -> str: + return "DOUBLE" + + @override + def visit_NUMERIC(self, type_: types.Numeric[Any], **kw: Any) -> str: + return self.visit_DECIMAL(type_, **kw) # type: ignore[arg-type] + + @override + def visit_DECIMAL(self, type_: types.DECIMAL[Any], **kw: Any) -> str: + if kw.get("require_precision") and type_.precision is None: + raise exc.CompileError( + "ARRAY decimal values require explicit Numeric precision; " + "specify precision and scale to avoid implicit rounding" + ) + return super().visit_DECIMAL(type_, **kw) + + def visit_tinyint(self, type_, **kw): + """Render a tinyint type. Args: type_: The type to render. **kw: Type-compiler keyword arguments. Returns: - The STRUCT or ROW type clause. + ``TINYINT``. """ - # Empty structs keep the existing ROW() rendering in every context. - if not isinstance(type_, AthenaStruct) or not type_.fields: - return "ROW()" - hive_ddl = self._enable_hive_column_ddl(kw) - preparer = ( - AthenaDDLIdentifierPreparer(self.dialect) - if hive_ddl - else self.dialect.identifier_preparer + return "TINYINT" + + def visit_TINYINT(self, type_, **kw): + """Render a TINYINT type. + + Args: + type_: The type to render. + **kw: Type-compiler keyword arguments. + + Returns: + ``TINYINT``. + """ + return "TINYINT" + + @override + def visit_TIMESTAMP(self, type_: types.TIMESTAMP, **kw: Any) -> str: + # A bare TIMESTAMP is timestamp(3) in Athena and truncates microseconds. + if not kw.get("timestamp_precision", True): + return "TIMESTAMP" + if isinstance(type_, AthenaTimestamp) and type_.precision is not None: + return f"TIMESTAMP({type_.precision})" + return "TIMESTAMP(6)" + + @override + def visit_DATETIME(self, type_: types.DateTime, **kw: Any) -> str: + return self.visit_TIMESTAMP(type_, **kw) # type: ignore[arg-type] + + @override + def visit_TIME(self, type_: types.Time, **kw: Any) -> str: + raise exc.CompileError(f"Data type `{type_}` is not supported") + + @override + def visit_CHAR(self, type_: types.CHAR, **kw: Any) -> str: + return "VARCHAR" + + @override + def visit_NCHAR(self, type_: types.NCHAR, **kw: Any) -> str: + return "VARCHAR" + + @override + def visit_VARCHAR(self, type_: types.String, **kw: Any) -> str: + return "VARCHAR" + + @override + def visit_NVARCHAR(self, type_: types.NVARCHAR, **kw: Any) -> str: + return "VARCHAR" + + @override + def visit_TEXT(self, type_: types.Text, **kw: Any) -> str: + return "VARCHAR" + + @override + def visit_CLOB(self, type_: types.CLOB, **kw: Any) -> str: + return "VARCHAR" + + @override + def visit_NCLOB(self, type_: types.Text, **kw: Any) -> str: + return "VARCHAR" + + @override + def visit_BLOB(self, type_: types.LargeBinary, **kw: Any) -> str: + return "VARBINARY" + + @override + def visit_BINARY(self, type_: types.BINARY, **kw: Any) -> str: + return "VARBINARY" + + @override + def visit_VARBINARY(self, type_: types.VARBINARY, **kw: Any) -> str: + return "VARBINARY" + + 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 + def visit_null(self, type_, **kw): + return "NULL" + + def visit_struct(self, type_, **kw): + """Render a STRUCT type as ``ROW(name type, ...)``. + + Args: + type_: The type to render. + **kw: Type-compiler keyword arguments. + + Returns: + The ROW type clause. + + Raises: + CompileError: If the type is not an ``AthenaStruct``, has no fields, + or has a field of unknown type. + """ + if not isinstance(type_, AthenaStruct): + raise exc.CompileError(f"Cannot render `{type_!r}` as ROW") + if not type_.fields: + raise exc.CompileError("ROW requires at least one field") + preparer = self.dialect.identifier_preparer + fields = ", ".join( + f"{preparer.quote(name)} {self._process_element(field_type, **kw)}" + for name, field_type in type_.fields.items() ) - separator = ":" if hive_ddl else " " - field_specs = [] - for field_name, field_type in type_.fields.items(): - field_type_str = self.process(field_type, **kw) - field_specs.append(f"{preparer.quote(field_name)}{separator}{field_type_str}") - if hive_ddl: - return f"STRUCT<{', '.join(field_specs)}>" - return f"ROW({', '.join(field_specs)})" + return f"ROW({fields})" def visit_STRUCT(self, type_, **kw): """Render a STRUCT type through ``visit_struct``. @@ -324,14 +565,12 @@ def visit_STRUCT(self, type_, **kw): **kw: Type-compiler keyword arguments. Returns: - The STRUCT or ROW type clause. + The 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``. + """Render a MAP type as ``MAP(key, value)``. Args: type_: The type to render. @@ -339,13 +578,16 @@ def visit_map(self, type_, **kw): Returns: The MAP type clause. + + Raises: + CompileError: If the type is not an ``AthenaMap`` or has a key or + value of unknown type. """ - if isinstance(type_, AthenaMap): - self._enable_hive_column_ddl(kw) - key_type_str = self.process(type_.key_type, **kw) - value_type_str = self.process(type_.value_type, **kw) - return f"MAP<{key_type_str}, {value_type_str}>" - return "MAP" + if not isinstance(type_, AthenaMap): + raise exc.CompileError(f"Cannot render `{type_!r}` as MAP") + key_type_str = self._process_element(type_.key_type, **kw) + value_type_str = self._process_element(type_.value_type, **kw) + return f"MAP({key_type_str}, {value_type_str})" def visit_MAP(self, type_, **kw): """Render a MAP type through ``visit_map``. @@ -360,10 +602,7 @@ def visit_MAP(self, type_, **kw): 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``. + """Render an ARRAY type as ``ARRAY(item)``. Args: type_: The type to render. @@ -371,12 +610,13 @@ def visit_array(self, type_, **kw): Returns: The ARRAY type clause. + + Raises: + CompileError: If the type is not an ARRAY or its item type is unknown. """ - if isinstance(type_, types.ARRAY): - kw["_athena_hive_ddl"] = True - item_type_str = self.process(_ArrayTypeInspector.item_type(type_), **kw) - return f"ARRAY<{item_type_str}>" - return "ARRAY" + if not isinstance(type_, types.ARRAY): + raise exc.CompileError(f"Cannot render `{type_!r}` as ARRAY") + return f"ARRAY({self._process_element(_ArrayTypeInspector.item_type(type_), **kw)})" def visit_ARRAY(self, type_, **kw): """Render an ARRAY type through ``visit_array``. @@ -415,6 +655,11 @@ class AthenaStatementCompiler(SQLCompiler): def _array_type_inspector(self): return _ArrayTypeInspector(self.dialect) + @util.memoized_property + def _dml_type_compiler(self) -> AthenaDMLTypeCompiler: + """The type compiler for the Trino types of CAST expressions.""" + return AthenaDMLTypeCompiler(self.dialect) + def visit_char_length_func(self, fn: Function[Any], **kw: Any) -> str: """Render ``char_length()`` as Athena ``length()``. @@ -590,11 +835,11 @@ def _array_slice_step(self, sql, step, array_type, **kw): "CAST(concat('Unsupported ARRAY slice step: ', " f"coalesce(CAST({step_sql} AS VARCHAR), 'NULL')) AS BIGINT)" ) - empty = ( - f"slice({sql}, 1, 0)" - if _ArrayTypeInspector.has_unknown_element(array_type) - else f"CAST(ARRAY[] AS {self._complex_dml_type(array_type, timestamp_precision=False)})" - ) + if _ArrayTypeInspector.has_unknown_element(array_type): + empty = f"slice({sql}, 1, 0)" + else: + empty_type = self._dml_type_compiler.process(array_type, timestamp_precision=False) + empty = f"CAST(ARRAY[] AS {empty_type})" return f"IF({step_sql} = 1, {sql}, slice({empty}, {failure}, 0))" @override @@ -833,129 +1078,12 @@ def visit_cast(self, cast: Cast[Any], **kwargs): The CAST SQL. Raises: - CompileError: For an ARRAY, MAP, or ROW type that cannot be cast. + CompileError: For a type that cannot be cast. """ - type_ = self._dialect_type(cast.type) - if isinstance(type_, (types.ARRAY, AthenaMap, AthenaStruct)): - type_clause = self._complex_dml_type( - type_, require_precision=cast._annotations.get("_pyathena_array_bind", False) - ) - return f"CAST({self.process(cast.clause, **kwargs)} AS {type_clause})" - if (isinstance(type_, types.VARCHAR) and type_.length is None) or isinstance( - type_, types.String - ): - type_clause = "VARCHAR" - elif isinstance(type_, types.CHAR) and type_.length is None: - type_clause = "CHAR" - elif isinstance(type_, (types.LargeBinary, types.BINARY, types.VARBINARY)): - type_clause = "VARBINARY" - elif isinstance(type_, types.Double): - type_clause = "DOUBLE" - elif isinstance(type_, (types.FLOAT, types.Float, types.REAL)): - # https://docs.aws.amazon.com/athena/latest/ug/data-types.html - # In Athena, use float in DDL statements like CREATE TABLE - # and real in SQL functions like SELECT CAST. - type_clause = "REAL" - elif (timestamp_type := self._timestamp_dml_type(type_)) is not None: - type_clause = timestamp_type - else: - type_clause = cast.typeclause._compiler_dispatch(self, **kwargs) - return f"CAST({cast.clause._compiler_dispatch(self, **kwargs)} AS {type_clause})" - - def _dialect_type(self, type_: TypeEngine[Any]) -> TypeEngine[Any]: - """Resolve the type this dialect uses for a SQLAlchemy type. - - Takes the Athena variant from ``with_variant()`` and the implementation - of a TypeDecorator until neither applies. - - Args: - type_: The declared type. - - Returns: - The resolved type. - """ - while True: - variant = self._array_type_inspector.variant(type_) - if variant is not None: - type_ = variant - elif isinstance(type_, types.TypeDecorator): - type_ = self._array_type_inspector.decorator_impl(type_) - else: - return type_ - - def _timestamp_dml_type(self, type_: TypeEngine[Any]) -> str | None: - """Return the DML type clause for a DateTime type. - - A bare ``TIMESTAMP`` is ``timestamp(3)`` in Athena and truncates - microseconds, so DML casts use ``TIMESTAMP(6)``. - - Args: - type_: The type to cast to, possibly a TypeDecorator or a type - with an Athena variant. - - Returns: - ``TIMESTAMP(precision)`` for an AthenaTimestamp with a precision, - ``TIMESTAMP(6)`` for any other DateTime type, otherwise None. - """ - type_ = self._dialect_type(type_) - if isinstance(type_, AthenaTimestamp) and type_.precision is not None: - return f"TIMESTAMP({type_.precision})" - if isinstance(type_, (types.DateTime, AthenaTimestamp)): - return "TIMESTAMP(6)" - return None - - def _complex_dml_type(self, type_, *, require_precision=False, timestamp_precision=True): - """Render a type for a DML cast. - - Args: - type_: The type to render. - require_precision: Reject a Numeric without an explicit precision. - timestamp_precision: Render DateTime types with their precision. - A cast that only types an empty value passes False to keep a - bare ``TIMESTAMP``, which Athena widens to the other operand's - precision instead of widening that operand. - - Returns: - The type clause. - - Raises: - CompileError: For an element type that cannot be cast. - """ - recurse = partial( - self._complex_dml_type, - require_precision=require_precision, - timestamp_precision=timestamp_precision, + type_clause = self._dml_type_compiler.process( + cast.type, require_precision=cast._annotations.get("_pyathena_array_bind", False) ) - type_ = self._dialect_type(type_) - if isinstance(type_, types.NullType): - raise exc.CompileError("Bound ARRAY values require an explicit element type") - if isinstance(type_, types.ARRAY): - return f"ARRAY({recurse(_ArrayTypeInspector.item_type(type_))})" - if isinstance(type_, AthenaMap): - return f"MAP({recurse(type_.key_type)}, {recurse(type_.value_type)})" - if isinstance(type_, AthenaStruct): - fields = ", ".join( - f"{self.preparer.quote(name)} {recurse(field_type)}" - for name, field_type in type_.fields.items() - ) - return f"ROW({fields})" - if isinstance(type_, types.String): - return "VARCHAR" - if isinstance(type_, (types.LargeBinary, types.BINARY, types.VARBINARY)): - return "VARBINARY" - if isinstance(type_, types.Double): - return "DOUBLE" - if isinstance(type_, types.Float): - return "REAL" - timestamp_type = self._timestamp_dml_type(type_) if timestamp_precision else None - if timestamp_type is not None: - return timestamp_type - if require_precision and isinstance(type_, types.Numeric) and type_.precision is None: - raise exc.CompileError( - "ARRAY decimal values require explicit Numeric precision; " - "specify precision and scale to avoid implicit rounding" - ) - return self.dialect.type_compiler_instance.process(type_) + return f"CAST({self.process(cast.clause, **kwargs)} AS {type_clause})" def visit_athena_array_json_projection(self, expression, **kw): """Render an ARRAY result column as a JSON envelope string. @@ -973,7 +1101,7 @@ def visit_athena_array_json_projection(self, expression, **kw): return f"json_format(CAST(MAP(ARRAY['_pyathena_array'], ARRAY[{encoded}]) AS JSON))" def _array_json(self, value, type_, depth=0): - type_ = self._dialect_type(type_) + type_ = self._array_type_inspector.dialect_type(type_) # Each recursive value becomes JSON, including map keys and typed scalar leaves. variable = f"_pyathena_array_{depth}" if isinstance(type_, types.ARRAY): @@ -1367,14 +1495,7 @@ def _get_table_properties_specification( @override def get_column_specification(self, column: Column[Any], **kwargs) -> str: - if type(column.type) in [types.Integer, types.INTEGER, types.INT]: - # https://docs.aws.amazon.com/athena/latest/ug/create-table.html - # In Data Definition Language (DDL) queries like CREATE TABLE, - # use the int keyword to represent an integer - type_ = "INT" - else: - # type_expression marks column DDL so STRUCT and MAP use Hive syntax. - type_ = self.dialect.type_compiler_instance.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_compiler.py b/tests/pyathena/sqlalchemy/test_compiler.py index c25e886eb..6c55c7623 100644 --- a/tests/pyathena/sqlalchemy/test_compiler.py +++ b/tests/pyathena/sqlalchemy/test_compiler.py @@ -88,68 +88,51 @@ def test_visit_struct_empty(self): dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) struct_type = AthenaStruct() - result = compiler.visit_struct(struct_type) - assert result == "ROW()" + with pytest.raises(exc.CompileError, match="STRUCT requires at least one field"): + compiler.visit_struct(struct_type) def test_visit_struct_with_fields(self): dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) struct_type = AthenaStruct(("name", String), ("age", Integer)) result = compiler.visit_struct(struct_type) - # The exact order might vary, so we check that both fields are present - assert "ROW(" in result - assert "name STRING" in result or "name VARCHAR" in result - assert "age INTEGER" in result - assert result.endswith(")") + assert result == "STRUCT" def test_visit_struct_uppercase(self): dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) struct_type = STRUCT(("id", Integer), ("title", String)) result = compiler.visit_STRUCT(struct_type) - assert "ROW(" in result - assert "id INTEGER" in result - assert "title STRING" in result or "title VARCHAR" in result - assert result.endswith(")") + assert result == "STRUCT" def test_visit_struct_no_fields_attribute(self): # Test struct type without fields attribute dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) struct_type = type("MockStruct", (), {})() - result = compiler.visit_struct(struct_type) - assert result == "ROW()" + with pytest.raises(exc.CompileError, match="as STRUCT"): + compiler.visit_struct(struct_type) def test_visit_struct_single_field(self): dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) struct_type = AthenaStruct(("name", String)) result = compiler.visit_struct(struct_type) - assert result == "ROW(name STRING)" or result == "ROW(name VARCHAR)" + assert result == "STRUCT" - def test_visit_struct_and_map_without_column_context_stay_row(self): - dialect = AthenaDialect() - compiler = AthenaTypeCompiler(dialect) + def test_complex_types_render_hive_syntax_without_column(self): + compiler = AthenaDialect().type_compiler_instance struct_type = AthenaStruct( ("profile", AthenaStruct(("name", String), ("age", Integer))), ("metrics", AthenaMap(String, Integer)), ) map_type = AthenaMap(Integer, AthenaStruct(("n", Integer))) - assert compiler.process(struct_type) == ( - "ROW(profile ROW(name STRING, age INTEGER), metrics MAP)" - ) - assert compiler.process(map_type) == "MAP" - assert compiler.process(AthenaStruct()) == "ROW()" - - def test_type_expression_column_selects_hive_syntax(self): - compiler = AthenaDialect().type_compiler_instance - struct_type = AthenaStruct(("name", String), ("age", Integer)) - map_type = AthenaMap(Integer, AthenaStruct(("n", Integer))) - assert compiler.process(struct_type, type_expression=Column("profile", struct_type)) == ( - "STRUCT" - ) - assert compiler.process(map_type, type_expression=Column("labels", map_type)) == ( - "MAP>" + expected_struct = "STRUCT, metrics:MAP>" + assert compiler.process(struct_type) == expected_struct + assert struct_type.compile(dialect=AthenaDialect()) == expected_struct + assert compiler.process(map_type) == "MAP>" + assert compiler.process(struct_type, type_expression=Column("c", struct_type)) == ( + expected_struct ) def test_visit_map_default(self): @@ -164,22 +147,22 @@ def test_visit_map_with_types(self): compiler = AthenaTypeCompiler(dialect) map_type = AthenaMap(String, Integer) result = compiler.visit_map(map_type) - assert result == "MAP" or result == "MAP" + assert result == "MAP" def test_visit_map_uppercase(self): dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) map_type = MAP(Integer, String) result = compiler.visit_MAP(map_type) - assert result == "MAP" or result == "MAP" + assert result == "MAP" def test_visit_map_no_attributes(self): # Test map type without key_type/value_type attributes dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) map_type = type("MockMap", (), {})() - result = compiler.visit_map(map_type) - assert result == "MAP" + with pytest.raises(exc.CompileError, match="as MAP"): + compiler.visit_map(map_type) def test_visit_array_default(self): dialect = AthenaDialect() @@ -207,18 +190,40 @@ def test_visit_array_no_attributes(self): dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) array_type = type("MockArray", (), {})() - result = compiler.visit_array(array_type) - assert result == "ARRAY" + with pytest.raises(exc.CompileError, match="as ARRAY"): + compiler.visit_array(array_type) def test_visit_json(self): - """Test JSON type compilation.""" - from sqlalchemy import types - dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) json_type = types.JSON() - result = compiler.visit_JSON(json_type) - assert result == "JSON" + with pytest.raises(exc.CompileError, match="not supported in Athena DDL"): + compiler.visit_JSON(json_type) + + @pytest.mark.parametrize( + ("type_", "ddl", "cast_type"), + [ + (Integer(), "INT", "INTEGER"), + (types.INTEGER(), "INT", "INTEGER"), + (types.CLOB(), "STRING", "VARCHAR"), + (types.Text(), "STRING", "VARCHAR"), + (types.VARCHAR(10), "VARCHAR(10)", "VARCHAR"), + (types.BINARY(), "BINARY", "VARBINARY"), + ], + ) + def test_ddl_and_cast_types(self, type_, ddl, cast_type): + dialect = AthenaDialect() + assert dialect.type_compiler_instance.process(type_) == ddl + assert str(cast(column("x"), type_).compile(dialect=dialect)) == f"CAST(x AS {cast_type})" + + @pytest.mark.parametrize( + "type_", [types.JSON(), AthenaArray(types.JSON), AthenaMap(String, types.JSON)] + ) + def test_json_is_rejected_in_ddl_and_kept_in_cast(self, type_): + dialect = AthenaDialect() + with pytest.raises(exc.CompileError, match="not supported in Athena DDL"): + dialect.type_compiler_instance.process(type_) + assert "JSON" in str(cast(column("x"), type_).compile(dialect=dialect)) @pytest.mark.parametrize( ("type_", "ddl", "cast_type"), @@ -668,6 +673,25 @@ def test_cast_resolves_variants_and_decorators(self, type_, expected): def test_complex_cast_keeps_dml_syntax(self, type_, expected): assert self._compile_sql(cast(column("col"), type_)) == f"CAST(col AS {expected})" + @pytest.mark.parametrize( + ("type_", "expected"), + [ + (types.JSON(), "JSON"), + (types.ARRAY(types.JSON), "ARRAY(JSON)"), + (AthenaMap(String, types.JSON), "MAP(VARCHAR, JSON)"), + (types.CLOB(), "VARCHAR"), + ], + ) + def test_cast_renders_types_rejected_or_respelled_in_ddl(self, type_, expected): + assert self._compile_sql(cast(column("col"), type_)) == f"CAST(col AS {expected})" + + @pytest.mark.parametrize( + "type_", [AthenaStruct(), types.ARRAY(AthenaStruct()), AthenaMap(String, AthenaStruct())] + ) + def test_cast_to_empty_struct_raises(self, type_): + with pytest.raises(exc.CompileError, match="ROW requires at least one field"): + self._compile_sql(cast(column("col"), type_)) + def test_timestamp_precision_applies_to_compared_values(self): col = column("col", AthenaTimestamp(precision=3)) value = datetime(2012, 10, 15, 12, 57, 18, 789999) @@ -960,14 +984,31 @@ def test_struct_field_quoting_follows_ddl_preparer(self): '"a`b" VARCHAR, "first name" VARCHAR, _hidden INTEGER))' ) - def test_empty_struct_column_stays_row(self): + @pytest.mark.parametrize( + "type_", [AthenaStruct(), AthenaArray(AthenaStruct()), AthenaMap(String, AthenaStruct())] + ) + def test_empty_struct_column_raises(self, type_): + with pytest.raises( + exc.CompileError, + match=r"column 'empty'.*STRUCT requires at least one field", + ): + self._ddl(Column("empty", type_)) + + def test_json_column_raises(self): + with pytest.raises( + exc.CompileError, match=r"column 'payload'.*not supported in Athena DDL" + ): + self._ddl(Column("payload", types.JSON)) + + def test_integer_subclass_and_decorator_columns_use_int(self): ddl = self._ddl( - Column("empty", AthenaStruct()), - Column("filled", AthenaStruct(("n", Integer))), + Column("decorated", decorated(Integer())), + Column("subclassed", type("MyInteger", (Integer,), {})()), + Column("text_value", types.CLOB), ) - assert "empty ROW()" in ddl - assert "filled STRUCT" in ddl - assert "STRUCT<>" not in ddl + assert "decorated INT" in ddl + assert "subclassed INT" in ddl + assert "text_value STRING" in ddl def test_unsupported_type_inside_struct_column_still_raises(self): with pytest.raises(exc.CompileError, match="not supported"): From e8812ad6900d3f9be77970ac89a83512000bbe4e Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 15:19:21 +0900 Subject: [PATCH 2/6] Render a type subclass with its own visit name as its base in CAST The previous visit_cast matched String, LargeBinary, Float, and DateTime with isinstance, so a subclass with its own __visit_name__ still cast as VARCHAR, VARBINARY, REAL, or TIMESTAMP(6). The DML type compiler dispatches by visit name, so it now falls back to the nearest base class that it renders instead of raising UnsupportedCompilationError. Also correct the AthenaTypeCompiler docstring: String with a length renders STRING, not VARCHAR(n). Co-Authored-By: Claude Opus 5.5 --- pyathena/sqlalchemy/compiler.py | 33 ++++++++++++++++++++-- tests/pyathena/sqlalchemy/test_compiler.py | 17 +++++++++++ 2 files changed, 47 insertions(+), 3 deletions(-) diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index 5e5c6797a..a676faf6c 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -85,9 +85,10 @@ class AthenaTypeCompiler(GenericTypeCompiler): and ``TypeEngine.compile()``. ``AthenaStatementCompiler`` renders the types of CAST expressions with ``AthenaDMLTypeCompiler``. - Integers render as INT, FLOAT and REAL as FLOAT, and binary types as - BINARY. TEXT, CLOB, and character types without a length render as - STRING; with a length, they render as CHAR(n) or VARCHAR(n). Complex + INTEGER renders as INT, FLOAT and REAL as FLOAT, and binary types as + BINARY. String, TEXT, CLOB, and CHAR, NCHAR, VARCHAR, or NVARCHAR + without a length render as STRING; with a length, CHAR, NCHAR, VARCHAR, + and NVARCHAR render as CHAR(n) or VARCHAR(n). Complex types render as ``STRUCT``, ``MAP``, and ``ARRAY``. TIME, JSON, and a STRUCT without fields have no Athena DDL type and raise ``CompileError``. @@ -410,6 +411,32 @@ def _process_element(self, type_: TypeEngine[Any], **kw: Any) -> str: raise exc.CompileError("Bound ARRAY values require an explicit element type") return self.process(type_, **kw) + @override + def visit_unsupported_compilation( # type: ignore[override] # base returns NoReturn + self, element: Any, err: Exception, **kw: Any + ) -> str: + """Render a type whose own visit name has no method as its nearest base type. + + A subclass of ``String`` with its own ``__visit_name__``, for example, + renders as VARCHAR. + + Args: + element: The type to render. + err: The error from the missing visit method. + **kw: Type-compiler keyword arguments. + + Returns: + The type clause of the nearest base type that this compiler renders. + + Raises: + UnsupportedCompilationError: If no base type has a visit method. + """ + for base in type(element).__mro__[1:]: + visit_name = base.__dict__.get("__visit_name__") + if isinstance(visit_name, str) and hasattr(self, f"visit_{visit_name}"): + return cast("str", getattr(self, f"visit_{visit_name}")(element, **kw)) + return super().visit_unsupported_compilation(element, err, **kw) + @override def visit_FLOAT(self, type_: types.Float[Any], **kw: Any) -> str: return "REAL" diff --git a/tests/pyathena/sqlalchemy/test_compiler.py b/tests/pyathena/sqlalchemy/test_compiler.py index 6c55c7623..0d0dbf25d 100644 --- a/tests/pyathena/sqlalchemy/test_compiler.py +++ b/tests/pyathena/sqlalchemy/test_compiler.py @@ -685,6 +685,23 @@ def test_complex_cast_keeps_dml_syntax(self, type_, expected): def test_cast_renders_types_rejected_or_respelled_in_ddl(self, type_, expected): assert self._compile_sql(cast(column("col"), type_)) == f"CAST(col AS {expected})" + @pytest.mark.parametrize( + ("base", "expected"), + [ + (types.String, "VARCHAR"), + (types.LargeBinary, "VARBINARY"), + (types.Float, "REAL"), + (types.Double, "DOUBLE"), + (types.DateTime, "TIMESTAMP(6)"), + ], + ) + def test_cast_renders_subclass_with_own_visit_name_as_base(self, base, expected): + type_ = type("Custom", (base,), {"__visit_name__": "pyathena_custom"})() + assert self._compile_sql(cast(column("col"), type_)) == f"CAST(col AS {expected})" + assert self._compile_sql(cast(column("col"), types.ARRAY(type_))) == ( + f"CAST(col AS ARRAY({expected}))" + ) + @pytest.mark.parametrize( "type_", [AthenaStruct(), types.ARRAY(AthenaStruct()), AthenaMap(String, AthenaStruct())] ) From 684eccd4fa1d5e801e94c66cceed31d4b9d84b94 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 15:20:36 +0900 Subject: [PATCH 3/6] Drop inaccurate rationale from the type syntax docs Athena's DDL parser accepts INTEGER inside nested types (#886), so INT is not required there, and the empty STRUCT sentence only needs the fact. Co-Authored-By: Claude Opus 5.5 --- docs/sqlalchemy.md | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index 179502aaa..f0535fa54 100644 --- a/docs/sqlalchemy.md +++ b/docs/sqlalchemy.md @@ -754,7 +754,7 @@ Athena parses DDL statements such as `CREATE TABLE` with Hive type syntax, and q | `AthenaArray`, `ARRAY` | `ARRAY` | `ARRAY(item)` | Complex types apply the same syntax to their nested types. -An `AthenaStruct` without fields raises `CompileError` in both, because Athena has no empty STRUCT type. +An `AthenaStruct` without fields raises `CompileError` in both. ## Floating-point types @@ -1139,7 +1139,7 @@ An outer `TypeDecorator` retains its result processor as well as native ARRAY or Raw `text()` queries and direct DB API queries retain the cursor's existing conversion behavior described below; they do not receive this projection automatically. Compared with earlier releases, reflected ARRAY columns are no longer reported as `String`. -ARRAY DDL now renders integer elements as `INT` and row elements as `STRUCT<...>`, which Athena requires for nested DDL types. +ARRAY DDL now renders integer elements as `INT` and row elements as `STRUCT<...>`. Code that inspects reflected types or compares compiled SQL strings should account for these changes. #### Basic Usage From 857014657720868939b8132c6dd89fa62a585315 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 15:39:26 +0900 Subject: [PATCH 4/6] Match CAST types by class as before and keep unknown-type rejection Independent review found three CAST regressions from dispatching by visit name: - A subclass whose visit name names another handled type (for example oracle.DATE, a DateTime subclass named DATE) rendered by that name instead of TIMESTAMP(6). - A compilation rule registered for a TypeDecorator stopped applying, because the decorator was resolved before dispatch. - An ARRAY assignment whose value type resolves to NullType rendered CAST(... AS NULL) instead of raising, because only nested elements were checked. AthenaDMLTypeCompiler.process() now matches ARRAY, MAP, ROW, String, binary, Double, Float, and DateTime types by class after resolving variants and decorators, as the former visit_cast did, and dispatches the declared type otherwise. The former _complex_dml_type callers use process_element(), which rejects an unknown type at any level. This replaces the visit-name fallback added earlier, and a decorator's compilation rule now also applies inside ARRAY, MAP, and ROW casts. Co-Authored-By: Claude Opus 5.5 --- pyathena/sqlalchemy/array.py | 6 +- pyathena/sqlalchemy/compiler.py | 146 +++++++-------------- tests/pyathena/sqlalchemy/test_compiler.py | 40 +++++- 3 files changed, 82 insertions(+), 110 deletions(-) diff --git a/pyathena/sqlalchemy/array.py b/pyathena/sqlalchemy/array.py index 9fbba051c..60864db44 100644 --- a/pyathena/sqlalchemy/array.py +++ b/pyathena/sqlalchemy/array.py @@ -600,7 +600,7 @@ def process(self, expression, **kw): ): raise exc.CompileError("An ARRAY slice assignment requires a non-NULL array") rhs = compiler.process(value, **kw) - rhs_type = compiler._dml_type_compiler.process( + rhs_type = compiler._dml_type_compiler.process_element( expression.value_type, require_precision=True ) rhs = f"CAST({rhs} AS {rhs_type})" @@ -649,7 +649,7 @@ def _rebuild(self, array, array_type, path, rhs, **kw): array_type = self._type_inspector.array_type(array_type) if array_type is None: raise exc.CompileError("Partial ARRAY updates require an ARRAY column type") - array_sql_type = compiler._dml_type_compiler.process(array_type) + array_sql_type = compiler._dml_type_compiler.process_element(array_type) array = f"coalesce({array}, CAST(ARRAY[] AS {array_sql_type}))" bound = path[0] if isinstance(bound, Slice): @@ -658,7 +658,7 @@ def _rebuild(self, array, array_type, path, rhs, **kw): def _prefix_and_padding(self, array, start, array_type): prefix = f"slice({array}, 1, least({start} - 1, cardinality({array})))" - element_type = self.compiler._dml_type_compiler.process( + element_type = self.compiler._dml_type_compiler.process_element( _ArrayTypeInspector.item_type(array_type) ) padding = ( diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index a676faf6c..e459a8117 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -364,13 +364,16 @@ class AthenaDMLTypeCompiler(GenericTypeCompiler): ``AthenaStatementCompiler`` renders the types of CAST expressions with this compiler, while ``AthenaTypeCompiler`` renders the Hive types of - DDL statements. Character types render as VARCHAR, FLOAT and REAL as - REAL, binary types as VARBINARY, and DateTime types as ``TIMESTAMP(6)`` - or ``TIMESTAMP(precision)``. Complex types render as - ``ROW(name type)``, ``MAP(key, value)``, and ``ARRAY(item)``. + DDL statements. A type is resolved through its ``with_variant()`` type for this dialect - and the implementation of a TypeDecorator before it is rendered. + and the implementation of a TypeDecorator, and then matched by class: + ARRAY, ``AthenaMap``, and ``AthenaStruct`` render as ``ARRAY(item)``, + ``MAP(key, value)``, and ``ROW(name type)``; String types as VARCHAR; + binary types as VARBINARY; Double types as DOUBLE; other Float types as + REAL; and DateTime types as ``TIMESTAMP(6)`` or ``TIMESTAMP(precision)``. + Subclasses of these types render the same way whatever their visit name. + Other types render through their visit methods. Two keyword arguments of ``process()`` adjust the rendering: @@ -392,63 +395,44 @@ def _type_inspector(self) -> _ArrayTypeInspector: @override def process(self, type_: TypeEngine[Any], **kw: Any) -> str: - return self._type_inspector.dialect_type(type_)._compiler_dispatch(self, **kw) - - def _process_element(self, type_: TypeEngine[Any], **kw: Any) -> str: - """Render the element type of an ARRAY, MAP, or ROW. + resolved = self._type_inspector.dialect_type(type_) + if isinstance(resolved, types.ARRAY): + return self.visit_array(resolved, **kw) + if isinstance(resolved, AthenaMap): + return self.visit_map(resolved, **kw) + if isinstance(resolved, AthenaStruct): + return self.visit_struct(resolved, **kw) + if isinstance(resolved, types.String): + return "VARCHAR" + if isinstance(resolved, (types.LargeBinary, types.BINARY, types.VARBINARY)): + return "VARBINARY" + if isinstance(resolved, types.Double): + return "DOUBLE" + if isinstance(resolved, types.Float): + return "REAL" + if isinstance(resolved, (types.DateTime, AthenaTimestamp)): + return self.visit_TIMESTAMP(resolved, **kw) # type: ignore[arg-type] + # Dispatch the declared type so a compilation rule registered for a + # TypeDecorator still applies. + return super().process(type_, **kw) + + def process_element(self, type_: TypeEngine[Any], **kw: Any) -> str: + """Render a type that must be known, such as an ARRAY, MAP, or ROW element. Args: - type_: The element type. + type_: The type to render. **kw: Type-compiler keyword arguments. Returns: - The element type clause. + The type clause. Raises: - CompileError: If the element type is unknown. + CompileError: If the type is unknown. """ if isinstance(self._type_inspector.dialect_type(type_), types.NullType): raise exc.CompileError("Bound ARRAY values require an explicit element type") return self.process(type_, **kw) - @override - def visit_unsupported_compilation( # type: ignore[override] # base returns NoReturn - self, element: Any, err: Exception, **kw: Any - ) -> str: - """Render a type whose own visit name has no method as its nearest base type. - - A subclass of ``String`` with its own ``__visit_name__``, for example, - renders as VARCHAR. - - Args: - element: The type to render. - err: The error from the missing visit method. - **kw: Type-compiler keyword arguments. - - Returns: - The type clause of the nearest base type that this compiler renders. - - Raises: - UnsupportedCompilationError: If no base type has a visit method. - """ - for base in type(element).__mro__[1:]: - visit_name = base.__dict__.get("__visit_name__") - if isinstance(visit_name, str) and hasattr(self, f"visit_{visit_name}"): - return cast("str", getattr(self, f"visit_{visit_name}")(element, **kw)) - return super().visit_unsupported_compilation(element, err, **kw) - - @override - def visit_FLOAT(self, type_: types.Float[Any], **kw: Any) -> str: - return "REAL" - - @override - def visit_REAL(self, type_: types.REAL[Any], **kw: Any) -> str: - return "REAL" - - @override - def visit_DOUBLE_PRECISION(self, type_, **kw) -> str: - return "DOUBLE" - @override def visit_NUMERIC(self, type_: types.Numeric[Any], **kw: Any) -> str: return self.visit_DECIMAL(type_, **kw) # type: ignore[arg-type] @@ -495,54 +479,10 @@ def visit_TIMESTAMP(self, type_: types.TIMESTAMP, **kw: Any) -> str: return f"TIMESTAMP({type_.precision})" return "TIMESTAMP(6)" - @override - def visit_DATETIME(self, type_: types.DateTime, **kw: Any) -> str: - return self.visit_TIMESTAMP(type_, **kw) # type: ignore[arg-type] - @override def visit_TIME(self, type_: types.Time, **kw: Any) -> str: raise exc.CompileError(f"Data type `{type_}` is not supported") - @override - def visit_CHAR(self, type_: types.CHAR, **kw: Any) -> str: - return "VARCHAR" - - @override - def visit_NCHAR(self, type_: types.NCHAR, **kw: Any) -> str: - return "VARCHAR" - - @override - def visit_VARCHAR(self, type_: types.String, **kw: Any) -> str: - return "VARCHAR" - - @override - def visit_NVARCHAR(self, type_: types.NVARCHAR, **kw: Any) -> str: - return "VARCHAR" - - @override - def visit_TEXT(self, type_: types.Text, **kw: Any) -> str: - return "VARCHAR" - - @override - def visit_CLOB(self, type_: types.CLOB, **kw: Any) -> str: - return "VARCHAR" - - @override - def visit_NCLOB(self, type_: types.Text, **kw: Any) -> str: - return "VARCHAR" - - @override - def visit_BLOB(self, type_: types.LargeBinary, **kw: Any) -> str: - return "VARBINARY" - - @override - def visit_BINARY(self, type_: types.BINARY, **kw: Any) -> str: - return "VARBINARY" - - @override - def visit_VARBINARY(self, type_: types.VARBINARY, **kw: Any) -> str: - return "VARBINARY" - def visit_JSON(self, type_: types.JSON, **kw: Any) -> str: """Render a JSON type. @@ -559,7 +499,7 @@ def visit_JSON(self, type_: types.JSON, **kw: Any) -> str: def visit_null(self, type_, **kw): return "NULL" - def visit_struct(self, type_, **kw): + def visit_struct(self, type_: Any, **kw: Any) -> str: """Render a STRUCT type as ``ROW(name type, ...)``. Args: @@ -579,7 +519,7 @@ def visit_struct(self, type_, **kw): raise exc.CompileError("ROW requires at least one field") preparer = self.dialect.identifier_preparer fields = ", ".join( - f"{preparer.quote(name)} {self._process_element(field_type, **kw)}" + f"{preparer.quote(name)} {self.process_element(field_type, **kw)}" for name, field_type in type_.fields.items() ) return f"ROW({fields})" @@ -596,7 +536,7 @@ def visit_STRUCT(self, type_, **kw): """ return self.visit_struct(type_, **kw) - def visit_map(self, type_, **kw): + def visit_map(self, type_: Any, **kw: Any) -> str: """Render a MAP type as ``MAP(key, value)``. Args: @@ -612,8 +552,8 @@ def visit_map(self, type_, **kw): """ if not isinstance(type_, AthenaMap): raise exc.CompileError(f"Cannot render `{type_!r}` as MAP") - key_type_str = self._process_element(type_.key_type, **kw) - value_type_str = self._process_element(type_.value_type, **kw) + key_type_str = self.process_element(type_.key_type, **kw) + value_type_str = self.process_element(type_.value_type, **kw) return f"MAP({key_type_str}, {value_type_str})" def visit_MAP(self, type_, **kw): @@ -628,7 +568,7 @@ def visit_MAP(self, type_, **kw): """ return self.visit_map(type_, **kw) - def visit_array(self, type_, **kw): + def visit_array(self, type_: Any, **kw: Any) -> str: """Render an ARRAY type as ``ARRAY(item)``. Args: @@ -643,7 +583,7 @@ def visit_array(self, type_, **kw): """ if not isinstance(type_, types.ARRAY): raise exc.CompileError(f"Cannot render `{type_!r}` as ARRAY") - return f"ARRAY({self._process_element(_ArrayTypeInspector.item_type(type_), **kw)})" + return f"ARRAY({self.process_element(_ArrayTypeInspector.item_type(type_), **kw)})" def visit_ARRAY(self, type_, **kw): """Render an ARRAY type through ``visit_array``. @@ -865,7 +805,9 @@ def _array_slice_step(self, sql, step, array_type, **kw): if _ArrayTypeInspector.has_unknown_element(array_type): empty = f"slice({sql}, 1, 0)" else: - empty_type = self._dml_type_compiler.process(array_type, timestamp_precision=False) + empty_type = self._dml_type_compiler.process_element( + array_type, timestamp_precision=False + ) empty = f"CAST(ARRAY[] AS {empty_type})" return f"IF({step_sql} = 1, {sql}, slice({empty}, {failure}, 0))" diff --git a/tests/pyathena/sqlalchemy/test_compiler.py b/tests/pyathena/sqlalchemy/test_compiler.py index 0d0dbf25d..fe8045be0 100644 --- a/tests/pyathena/sqlalchemy/test_compiler.py +++ b/tests/pyathena/sqlalchemy/test_compiler.py @@ -32,6 +32,7 @@ union, ) from sqlalchemy.engine.url import make_url +from sqlalchemy.ext.compiler import compiles from sqlalchemy.sql import literal, literal_column, operators from sqlalchemy.sql.compiler import FROM_LINTING from sqlalchemy.sql.ddl import CreateTable @@ -695,13 +696,41 @@ def test_cast_renders_types_rejected_or_respelled_in_ddl(self, type_, expected): (types.DateTime, "TIMESTAMP(6)"), ], ) - def test_cast_renders_subclass_with_own_visit_name_as_base(self, base, expected): - type_ = type("Custom", (base,), {"__visit_name__": "pyathena_custom"})() + @pytest.mark.parametrize("visit_name", ["pyathena_custom", "DATE", "INTEGER", "JSON"]) + def test_cast_renders_subclass_by_base_class(self, base, expected, visit_name): + # oracle.DATE, for example, subclasses DateTime with the visit name DATE. + type_ = type("Custom", (base,), {"__visit_name__": visit_name})() assert self._compile_sql(cast(column("col"), type_)) == f"CAST(col AS {expected})" assert self._compile_sql(cast(column("col"), types.ARRAY(type_))) == ( f"CAST(col AS ARRAY({expected}))" ) + def test_cast_applies_compilation_rule_of_decorator(self): + class _Wide(types.TypeDecorator): + impl = types.Integer + cache_ok = True + + @compiles(_Wide, "awsathena") + def _compile_wide(type_, compiler, **kw): + return "BIGINT" + + assert self._compile_sql(cast(column("col"), _Wide())) == "CAST(col AS BIGINT)" + assert self._compile_sql(cast(column("col"), types.ARRAY(_Wide()))) == ( + "CAST(col AS ARRAY(BIGINT))" + ) + + def test_array_assignment_rejects_unknown_value_type(self): + items = Table( + "items", + MetaData(), + Column( + "a", + AthenaArray(types.NullType()).with_variant(AthenaArray(Integer), "awsathena"), + ), + ) + with pytest.raises(exc.CompileError, match="explicit element type"): + items.update().values({items.c.a[1]: 1}).compile(dialect=self.dialect) + @pytest.mark.parametrize( "type_", [AthenaStruct(), types.ARRAY(AthenaStruct()), AthenaMap(String, AthenaStruct())] ) @@ -1023,9 +1052,10 @@ def test_integer_subclass_and_decorator_columns_use_int(self): Column("subclassed", type("MyInteger", (Integer,), {})()), Column("text_value", types.CLOB), ) - assert "decorated INT" in ddl - assert "subclassed INT" in ddl - assert "text_value STRING" in ddl + assert "decorated INT,\n" in ddl + assert "subclassed INT,\n" in ddl + assert "text_value STRING\n" in ddl + assert "INTEGER" not in ddl def test_unsupported_type_inside_struct_column_still_raises(self): with pytest.raises(exc.CompileError, match="not supported"): From 084d773b1f384c3aae2526aec92ecb26edab492e Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 15:52:27 +0900 Subject: [PATCH 5/6] Render CAST element types from their resolved implementation Dispatching a nested TypeDecorator by its declared type let a registered compilation rule bypass the ARRAY decimal precision check, for example AthenaArray of a Numeric decorator whose rule returns DECIMAL. The former _complex_dml_type resolved every element before rendering it, so process_element() now does the same, and only a top-level CAST dispatches the declared type. Co-Authored-By: Claude Opus 5.5 --- pyathena/sqlalchemy/compiler.py | 15 +++++++++++---- tests/pyathena/sqlalchemy/test_compiler.py | 16 +++++++++++++++- 2 files changed, 26 insertions(+), 5 deletions(-) diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index e459a8117..499ee58ed 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -373,7 +373,10 @@ class AthenaDMLTypeCompiler(GenericTypeCompiler): binary types as VARBINARY; Double types as DOUBLE; other Float types as REAL; and DateTime types as ``TIMESTAMP(6)`` or ``TIMESTAMP(precision)``. Subclasses of these types render the same way whatever their visit name. - Other types render through their visit methods. + Other types render through their visit methods; ``process()`` dispatches + the declared type, so a compilation rule registered for a TypeDecorator + applies to it, while ``process_element()`` and the elements of ARRAY, + MAP, and ROW types dispatch the resolved type. Two keyword arguments of ``process()`` adjust the rendering: @@ -417,7 +420,10 @@ def process(self, type_: TypeEngine[Any], **kw: Any) -> str: return super().process(type_, **kw) def process_element(self, type_: TypeEngine[Any], **kw: Any) -> str: - """Render a type that must be known, such as an ARRAY, MAP, or ROW element. + """Render a resolved type that must be known, such as an ARRAY, MAP, or ROW element. + + Unlike ``process()``, a TypeDecorator is always rendered as its + implementation, so a compilation rule registered for it does not apply. Args: type_: The type to render. @@ -429,9 +435,10 @@ def process_element(self, type_: TypeEngine[Any], **kw: Any) -> str: Raises: CompileError: If the type is unknown. """ - if isinstance(self._type_inspector.dialect_type(type_), types.NullType): + resolved = self._type_inspector.dialect_type(type_) + if isinstance(resolved, types.NullType): raise exc.CompileError("Bound ARRAY values require an explicit element type") - return self.process(type_, **kw) + return self.process(resolved, **kw) @override def visit_NUMERIC(self, type_: types.Numeric[Any], **kw: Any) -> str: diff --git a/tests/pyathena/sqlalchemy/test_compiler.py b/tests/pyathena/sqlalchemy/test_compiler.py index fe8045be0..a6958a5c2 100644 --- a/tests/pyathena/sqlalchemy/test_compiler.py +++ b/tests/pyathena/sqlalchemy/test_compiler.py @@ -7,6 +7,7 @@ import warnings from datetime import date, datetime +from decimal import Decimal import pytest from sqlalchemy import ( @@ -715,10 +716,23 @@ def _compile_wide(type_, compiler, **kw): return "BIGINT" assert self._compile_sql(cast(column("col"), _Wide())) == "CAST(col AS BIGINT)" + # Element types render their resolved implementation. assert self._compile_sql(cast(column("col"), types.ARRAY(_Wide()))) == ( - "CAST(col AS ARRAY(BIGINT))" + "CAST(col AS ARRAY(INTEGER))" ) + def test_array_bind_rejects_decorated_numeric_without_precision(self): + class _Amount(types.TypeDecorator): + impl = types.Numeric + cache_ok = True + + @compiles(_Amount, "awsathena") + def _compile_amount(type_, compiler, **kw): + return "DECIMAL" + + with pytest.raises(exc.CompileError, match="explicit Numeric precision"): + self._compile_sql(select(literal([Decimal("1.23")], AthenaArray(_Amount())))) + def test_array_assignment_rejects_unknown_value_type(self): items = Table( "items", From 3c61716555cbfd83ca58da1b3f616e8f8b612d2a Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 3 Oct 2026 16:53:50 +0900 Subject: [PATCH 6/6] Check ARRAY decimal precision before dispatching the CAST type The former _complex_dml_type rejected any Numeric without precision by class, but the check had moved into visit_DECIMAL, so a Numeric subclass with its own visit name and compilation rule skipped it. Check the resolved type in process() before dispatch, as before. Also cover a with_variant() Integer column, whose variant DDL now applies instead of the former INT special case. Co-Authored-By: Claude Opus 5.5 --- pyathena/sqlalchemy/compiler.py | 18 +++++++++--------- tests/pyathena/sqlalchemy/test_compiler.py | 14 ++++++++++++++ 2 files changed, 23 insertions(+), 9 deletions(-) diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index 499ee58ed..8c2f81183 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -415,6 +415,15 @@ def process(self, type_: TypeEngine[Any], **kw: Any) -> str: return "REAL" if isinstance(resolved, (types.DateTime, AthenaTimestamp)): return self.visit_TIMESTAMP(resolved, **kw) # type: ignore[arg-type] + if ( + kw.get("require_precision") + and isinstance(resolved, types.Numeric) + and resolved.precision is None + ): + raise exc.CompileError( + "ARRAY decimal values require explicit Numeric precision; " + "specify precision and scale to avoid implicit rounding" + ) # Dispatch the declared type so a compilation rule registered for a # TypeDecorator still applies. return super().process(type_, **kw) @@ -444,15 +453,6 @@ def process_element(self, type_: TypeEngine[Any], **kw: Any) -> str: def visit_NUMERIC(self, type_: types.Numeric[Any], **kw: Any) -> str: return self.visit_DECIMAL(type_, **kw) # type: ignore[arg-type] - @override - def visit_DECIMAL(self, type_: types.DECIMAL[Any], **kw: Any) -> str: - if kw.get("require_precision") and type_.precision is None: - raise exc.CompileError( - "ARRAY decimal values require explicit Numeric precision; " - "specify precision and scale to avoid implicit rounding" - ) - return super().visit_DECIMAL(type_, **kw) - def visit_tinyint(self, type_, **kw): """Render a tinyint type. diff --git a/tests/pyathena/sqlalchemy/test_compiler.py b/tests/pyathena/sqlalchemy/test_compiler.py index a6958a5c2..0c3548a76 100644 --- a/tests/pyathena/sqlalchemy/test_compiler.py +++ b/tests/pyathena/sqlalchemy/test_compiler.py @@ -733,6 +733,18 @@ def _compile_amount(type_, compiler, **kw): with pytest.raises(exc.CompileError, match="explicit Numeric precision"): self._compile_sql(select(literal([Decimal("1.23")], AthenaArray(_Amount())))) + def test_array_bind_rejects_numeric_subclass_without_precision(self): + class _Money(types.Numeric): + __visit_name__ = "pyathena_money" + cache_ok = True + + @compiles(_Money, "awsathena") + def _compile_money(type_, compiler, **kw): + return "DECIMAL" + + with pytest.raises(exc.CompileError, match="explicit Numeric precision"): + self._compile_sql(select(literal([Decimal("1.23")], AthenaArray(_Money())))) + def test_array_assignment_rejects_unknown_value_type(self): items = Table( "items", @@ -1064,10 +1076,12 @@ def test_integer_subclass_and_decorator_columns_use_int(self): ddl = self._ddl( Column("decorated", decorated(Integer())), Column("subclassed", type("MyInteger", (Integer,), {})()), + Column("variant", Integer().with_variant(types.BigInteger(), "awsathena")), Column("text_value", types.CLOB), ) assert "decorated INT,\n" in ddl assert "subclassed INT,\n" in ddl + assert "variant BIGINT,\n" in ddl assert "text_value STRING\n" in ddl assert "INTEGER" not in ddl