diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index 02ae2b5a..71b41c90 100644 --- a/docs/sqlalchemy.md +++ b/docs/sqlalchemy.md @@ -825,12 +825,17 @@ This generates the following SQL structure: ```sql CREATE TABLE users ( - id INTEGER, - profile ROW(name STRING, age INTEGER, email STRING), - settings ROW(theme STRING, notifications ROW(email STRING, push STRING)) + id INT, + profile STRUCT, + settings STRUCT> ) ``` +`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`. + #### Querying STRUCT data PyAthena automatically converts STRUCT data between different formats: @@ -970,13 +975,16 @@ This generates the following SQL structure: ```sql CREATE TABLE products ( - id INTEGER, + id INT, attributes MAP, - metrics MAP, - categories MAP + metrics MAP, + categories MAP ) ``` +`CREATE TABLE` renders integer MAP keys and values as `INT`. +`CAST` still spells those integers as `INTEGER`. + #### Querying MAP data PyAthena automatically converts MAP data between different formats: diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index 600b6519..4211643b 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -90,6 +90,9 @@ class AthenaTypeCompiler(GenericTypeCompiler): - 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(...)``. + See Also: AWS Athena Data Types: https://docs.aws.amazon.com/athena/latest/ug/data-types.html @@ -121,7 +124,7 @@ def visit_TINYINT(self, type_: types.Integer, **kw: Any) -> str: return "TINYINT" def visit_INTEGER(self, type_: types.Integer, **kw: Any) -> str: - return "INT" if kw.get("_athena_array_ddl") else "INTEGER" + return "INT" if kw.get("_athena_hive_ddl") else "INTEGER" def visit_SMALLINT(self, type_: types.SmallInteger, **kw: Any) -> str: return "SMALLINT" @@ -199,31 +202,51 @@ 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. + + ``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. + + Args: + kw: Type-compiler keyword arguments. When Hive spelling applies, + ``_athena_hive_ddl`` is set so nested types keep it. + + Returns: + True when the type should use Hive DDL syntax. + """ + if kw.get("_athena_hive_ddl") or isinstance(kw.get("type_expression"), Column): + kw["_athena_hive_ddl"] = True + return True + return False + def visit_struct(self, type_, **kw): - if isinstance(type_, AthenaStruct): - if type_.fields: - field_specs = [] - for field_name, field_type in type_.fields.items(): - field_type_str = self.process(field_type, **kw) - preparer = ( - AthenaDDLIdentifierPreparer(self.dialect) - if kw.get("_athena_array_ddl") - else self.dialect.identifier_preparer - ) - name = preparer.quote(field_name) - separator = ":" if kw.get("_athena_array_ddl") else " " - field_specs.append(f"{name}{separator}{field_type_str}") - if kw.get("_athena_array_ddl"): - return f"STRUCT<{', '.join(field_specs)}>" - return f"ROW({', '.join(field_specs)})" + # Empty structs keep the existing ROW() rendering in every context. + if not isinstance(type_, AthenaStruct) or not type_.fields: return "ROW()" - return "ROW()" + hive_ddl = self._enable_hive_column_ddl(kw) + preparer = ( + AthenaDDLIdentifierPreparer(self.dialect) + if hive_ddl + else self.dialect.identifier_preparer + ) + 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)})" def visit_STRUCT(self, type_, **kw): return self.visit_struct(type_, **kw) def visit_map(self, type_, **kw): 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}>" @@ -234,7 +257,7 @@ def visit_MAP(self, type_, **kw): def visit_array(self, type_, **kw): if isinstance(type_, types.ARRAY): - kw["_athena_array_ddl"] = True + kw["_athena_hive_ddl"] = True item_type_str = self.process(_ArrayTypeInspector.item_type(type_), **kw) return f"ARRAY<{item_type_str}>" return "ARRAY" @@ -1158,6 +1181,7 @@ def get_column_specification(self, column: Column[Any], **kwargs) -> str: # 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.process(column.type, type_expression=column) text = [f"{self.preparer.format_column(column)} {type_}"] if column.comment: diff --git a/tests/pyathena/sqlalchemy/test_base.py b/tests/pyathena/sqlalchemy/test_base.py index 69e91352..bf010643 100644 --- a/tests/pyathena/sqlalchemy/test_base.py +++ b/tests/pyathena/sqlalchemy/test_base.py @@ -3496,8 +3496,8 @@ def test_create_table_with_map_types(self, engine): # Verify MAP types are correctly compiled assert "attributes MAP" in ddl_string - assert "metrics MAP" in ddl_string - assert "complex_map MAP" in ddl_string + assert "metrics MAP" in ddl_string + assert "complex_map MAP>" in ddl_string assert "nested_map MAP>" in ddl_string def test_create_table_with_struct_types(self, engine): @@ -3539,12 +3539,12 @@ def test_create_table_with_struct_types(self, engine): ddl_string = str(create_ddl) # Verify STRUCT types are correctly compiled - assert "user_info ROW(name STRING, age INTEGER, email STRING)" in ddl_string + assert "user_info STRUCT" in ddl_string assert ( - "nested_struct ROW(personal ROW(first_name STRING, last_name STRING), " - "preferences MAP)" in ddl_string + "nested_struct STRUCT, " + "preferences:MAP>" in ddl_string ) - assert "struct_with_array ROW(tags ARRAY, scores ARRAY)" in ddl_string + assert "struct_with_array STRUCT, scores:ARRAY>" in ddl_string def test_create_table_with_complex_nested_types(self, engine): """Test DDL compilation for complex nested combinations of ARRAY, MAP, and STRUCT.""" @@ -3582,6 +3582,57 @@ def test_create_table_with_complex_nested_types(self, engine): ) assert expected_type in ddl_string + def test_external_parquet_struct_columns_round_trip(self, engine): + """Create a Parquet table of top-level and MAP-nested STRUCTs and read the fields back.""" + _, conn = engine + table_name = "test_external_parquet_struct_columns" + location = f"{ENV.s3_staging_dir}{ENV.schema}/{table_name}/" + table = Table( + table_name, + MetaData(schema=ENV.schema), + Column( + "profile", + AthenaStruct( + ("name", types.String), + ("age", types.Integer), + ( + "address", + AthenaStruct(("city", types.String), ("zip", types.Integer)), + ), + ), + ), + Column( + "labels", + AthenaMap( + types.String, + AthenaStruct(("value", types.String), ("count", types.Integer)), + ), + ), + awsathena_location=location, + awsathena_file_format="PARQUET", + ) + ddl = str(CreateTable(table).compile(dialect=conn.dialect)) + assert "profile STRUCT>" in ddl + assert "labels MAP>" in ddl + table.create(bind=conn) + conn.execute( + text( + f"INSERT INTO {ENV.schema}.{table_name} VALUES (" + "CAST(ROW('Ada', 36, ROW('London', 12345)) AS " + "ROW(name VARCHAR, age INTEGER, address ROW(city VARCHAR, zip INTEGER))), " + "MAP(ARRAY['home'], ARRAY[CAST(ROW('Lovelace', 2) AS " + "ROW(value VARCHAR, count INTEGER))]))" + ) + ) + row = conn.execute( + text( + "SELECT profile.name, profile.age, profile.address.city, " + "profile.address.zip, labels['home'].value, labels['home'].count " + f"FROM {ENV.schema}.{table_name}" + ) + ).one() + assert tuple(row) == ("Ada", 36, "London", 12345, "Lovelace", 2) + def test_sqlalchemy_execute_with_execution_options_callback(self, engine): """Test callback functionality through SQLAlchemy execution_options.""" engine, conn = engine diff --git a/tests/pyathena/sqlalchemy/test_compiler.py b/tests/pyathena/sqlalchemy/test_compiler.py index 9315ca95..c25e886e 100644 --- a/tests/pyathena/sqlalchemy/test_compiler.py +++ b/tests/pyathena/sqlalchemy/test_compiler.py @@ -127,6 +127,31 @@ def test_visit_struct_single_field(self): result = compiler.visit_struct(struct_type) assert result == "ROW(name STRING)" or result == "ROW(name VARCHAR)" + def test_visit_struct_and_map_without_column_context_stay_row(self): + dialect = AthenaDialect() + compiler = AthenaTypeCompiler(dialect) + 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>" + ) + def test_visit_map_default(self): dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) @@ -606,6 +631,43 @@ def test_cast_resolves_variants_and_decorators(self, type_, expected): for target in (type_, decorated(type_), decorated(decorated(type_))): assert self._compile_sql(cast(column("col"), target)) == f"CAST(col AS {expected})" + @pytest.mark.parametrize( + ("type_", "expected"), + [ + ( + AthenaStruct(("name", String), ("age", Integer)), + "ROW(name VARCHAR, age INTEGER)", + ), + ( + AthenaStruct( + ("personal", AthenaStruct(("name", String), ("age", Integer))), + ("scores", types.ARRAY(Integer)), + ("attrs", AthenaMap(String, String)), + ), + "ROW(personal ROW(name VARCHAR, age INTEGER), scores ARRAY(INTEGER), " + "attrs MAP(VARCHAR, VARCHAR))", + ), + ( + AthenaMap(String, AthenaStruct(("value", String), ("count", Integer))), + "MAP(VARCHAR, ROW(value VARCHAR, count INTEGER))", + ), + ( + AthenaMap(String, AthenaMap(Integer, AthenaStruct(("n", Integer)))), + "MAP(VARCHAR, MAP(INTEGER, ROW(n INTEGER)))", + ), + ( + types.ARRAY(AthenaMap(String, AthenaStruct(("n", Integer)))), + "ARRAY(MAP(VARCHAR, ROW(n INTEGER)))", + ), + ( + AthenaStruct(("tags", AthenaMap(String, types.ARRAY(Integer)))), + "ROW(tags MAP(VARCHAR, ARRAY(INTEGER)))", + ), + ], + ) + def test_complex_cast_keeps_dml_syntax(self, type_, expected): + assert self._compile_sql(cast(column("col"), type_)) == f"CAST(col AS {expected})" + def test_timestamp_precision_applies_to_compared_values(self): col = column("col", AthenaTimestamp(precision=3)) value = datetime(2012, 10, 15, 12, 57, 18, 789999) @@ -682,7 +744,10 @@ def test_like_pattern_is_not_truncated(self, pattern): class TestAthenaDDLCompiler: - """Compile-only (no AWS) tests for the DDL compiler's S3 Tables support. + """Compile-only (no AWS) tests for the DDL compiler. + + Covers column type rendering in CREATE TABLE, where STRUCT and nested MAP + types use Hive syntax, and S3 Tables support. S3 Tables are queried by setting the connection ``catalog_name`` to ``s3tablescatalog/`` and using the namespace as the table @@ -700,6 +765,16 @@ def _s3tables_dialect(self, **connect_opts): } return dialect + def _ddl(self, *columns): + table = Table( + "events", + MetaData(schema="analytics"), + *columns, + awsathena_location="s3://bucket/events/", + awsathena_file_format="PARQUET", + ) + return str(CreateTable(table).compile(dialect=AthenaDialect())) + def test_create_table_s3tables_catalog_omits_location(self): table = Table( "tbl", @@ -811,3 +886,114 @@ def test_create_connect_args_stores_connect_options_for_subclass_dialects(self): ) ddl = str(CreateTable(table).compile(dialect=dialect)) assert "LOCATION" not in ddl + + def test_create_table_renders_hive_struct_syntax(self): + ddl = self._ddl( + Column("id", Integer), + Column("scores", AthenaArray(Integer)), + Column( + "profile", + AthenaStruct( + ("name", String), + ("age", Integer), + ("ratio", Float(24)), + ("address", AthenaStruct(("city", String), ("zip", Integer))), + ), + ), + Column( + "labels", + AthenaMap(String, AthenaStruct(("value", String), ("count", Integer))), + ), + Column( + "nested_maps", + AthenaMap(String, AthenaMap(Integer, AthenaStruct(("n", Integer)))), + ), + Column( + "mixed", + AthenaStruct( + ("tags", AthenaArray(String)), + ("attrs", AthenaMap(String, Integer)), + ("scores", AthenaArray(Integer)), + ), + ), + Column( + "deep", + AthenaArray( + AthenaMap(String, AthenaStruct(("value", String), ("flag", types.Boolean))) + ), + ), + ) + assert "id INT" in ddl + assert "scores ARRAY" in ddl + assert ( + "profile STRUCT>" + ) in ddl + assert "labels MAP>" in ddl + assert "nested_maps MAP>>" in ddl + assert ( + "mixed STRUCT, attrs:MAP, scores:ARRAY>" + ) in ddl + assert "deep ARRAY>>" in ddl + assert "ROW(" not in ddl + assert "INTEGER" not in ddl + assert "REAL" not in ddl + assert "FLOAT(" not in ddl + + def test_struct_field_quoting_follows_ddl_preparer(self): + struct_type = AthenaStruct( + ("date", String), + ("select", Integer), + ('na"me', String), + ("a`b", String), + ("first name", String), + ("_hidden", Integer), + ) + ddl = self._ddl(Column("payload", struct_type)) + assert ( + 'payload STRUCT<`date`:STRING, `select`:INT, `na"me`:STRING, ' + "`a``b`:STRING, `first name`:STRING, `_hidden`:INT>" + ) in ddl + cast_sql = str(cast(column("payload"), struct_type).compile(dialect=AthenaDialect())) + assert cast_sql == ( + 'CAST(payload AS ROW(date VARCHAR, "select" INTEGER, "na""me" VARCHAR, ' + '"a`b" VARCHAR, "first name" VARCHAR, _hidden INTEGER))' + ) + + def test_empty_struct_column_stays_row(self): + ddl = self._ddl( + Column("empty", AthenaStruct()), + Column("filled", AthenaStruct(("n", Integer))), + ) + assert "empty ROW()" in ddl + assert "filled STRUCT" in ddl + assert "STRUCT<>" not in ddl + + def test_unsupported_type_inside_struct_column_still_raises(self): + with pytest.raises(exc.CompileError, match="not supported"): + self._ddl(Column("payload", AthenaStruct(("when", types.Time)))) + + def test_scalar_and_float_column_ddl_is_unchanged(self): + ddl = self._ddl( + Column("id", Integer), + Column("label", String), + Column("flag", types.Boolean), + Column("ratio", Float), + Column("ratio_prec", Float(24)), + Column("real_value", types.REAL), + Column("float_value", types.FLOAT), + Column("wide", types.Double), + Column("amount", Numeric(10, 2)), + ) + assert "id INT" in ddl + assert "label STRING" in ddl + assert "flag BOOLEAN" in ddl + assert "ratio FLOAT" in ddl + assert "ratio_prec FLOAT" in ddl + assert "real_value FLOAT" in ddl + assert "float_value FLOAT" in ddl + assert "wide DOUBLE" in ddl + assert "amount DECIMAL(10, 2)" in ddl + assert "INTEGER" not in ddl + assert "REAL" not in ddl + assert "FLOAT(" not in ddl