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..f0535fa54 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. + ## 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 @@ -1122,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 @@ -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..60864db44 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_element( + 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_element(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_element( + _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..8c2f81183 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,21 @@ class AthenaTypeCompiler(GenericTypeCompiler): - """Type compiler for Amazon Athena SQL 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 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(...)``. + """Type compiler for Amazon Athena DDL types. + + 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``. + + 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``. See Also: AWS Athena Data Types: @@ -145,7 +140,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 +168,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 +215,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 +258,278 @@ 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) + + def visit_struct(self, type_, **kw): + """Render a STRUCT type as ``STRUCT``. + + Args: + type_: The type to render. + **kw: Type-compiler keyword arguments. + + Returns: + The STRUCT type clause. - ``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. + Raises: + CompileError: If the type is not an ``AthenaStruct`` or has no fields. + """ + 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 through ``visit_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. """ - if kw.get("_athena_hive_ddl") or isinstance(kw.get("type_expression"), Column): - kw["_athena_hive_ddl"] = True - return True - return False + return self.visit_struct(type_, **kw) - def visit_struct(self, type_, **kw): - """Render a STRUCT type. + 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. + + A type is resolved through its ``with_variant()`` type for this dialect + 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; ``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: + + - ``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: + 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] + 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) + + def process_element(self, type_: TypeEngine[Any], **kw: Any) -> str: + """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. + **kw: Type-compiler keyword arguments. + + Returns: + The type clause. + + Raises: + CompileError: If the type is unknown. + """ + 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(resolved, **kw) + + @override + def visit_NUMERIC(self, type_: types.Numeric[Any], **kw: Any) -> str: + return self.visit_DECIMAL(type_, **kw) # type: ignore[arg-type] + + 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" - 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()``. + 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" + + @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_TIME(self, type_: types.Time, **kw: Any) -> str: + raise exc.CompileError(f"Data type `{type_}` is not supported") + + 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_: Any, **kw: Any) -> str: + """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 +539,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``. + def visit_map(self, type_: Any, **kw: Any) -> str: + """Render a MAP type as ``MAP(key, value)``. Args: type_: The type to render. @@ -339,13 +552,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``. @@ -359,11 +575,8 @@ 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``. + def visit_array(self, type_: Any, **kw: Any) -> str: + """Render an ARRAY type as ``ARRAY(item)``. Args: type_: The type to render. @@ -371,12 +584,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 +629,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 +809,13 @@ 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_element( + 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 +1054,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. - """ - 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. + CompileError: For a 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 +1077,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 +1471,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..0c3548a76 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 ( @@ -32,6 +33,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 @@ -88,68 +90,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 +149,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 +192,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 +675,95 @@ 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( + ("base", "expected"), + [ + (types.String, "VARCHAR"), + (types.LargeBinary, "VARBINARY"), + (types.Float, "REAL"), + (types.Double, "DOUBLE"), + (types.DateTime, "TIMESTAMP(6)"), + ], + ) + @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)" + # Element types render their resolved implementation. + assert self._compile_sql(cast(column("col"), types.ARRAY(_Wide()))) == ( + "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_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", + 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())] + ) + 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 +1056,34 @@ 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("variant", Integer().with_variant(types.BigInteger(), "awsathena")), + Column("text_value", types.CLOB), ) - assert "empty ROW()" in ddl - assert "filled STRUCT" in ddl - assert "STRUCT<>" not in ddl + 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 def test_unsupported_type_inside_struct_column_still_raises(self): with pytest.raises(exc.CompileError, match="not supported"):