From 27674e047310d82fe5413872baec22dcfc7a0551 Mon Sep 17 00:00:00 2001 From: Matt Van Horn <455140+mvanhorn@users.noreply.github.com> Date: Mon, 28 Sep 2026 00:08:44 -0700 Subject: [PATCH 1/3] fix: render Hive STRUCT syntax in table column DDL Fixes #855 --- docs/sqlalchemy.md | 22 ++- pyathena/sqlalchemy/compiler.py | 59 +++++-- tests/pyathena/sqlalchemy/test_base.py | 88 +++++++++- tests/pyathena/sqlalchemy/test_compiler.py | 187 +++++++++++++++++++++ 4 files changed, 327 insertions(+), 29 deletions(-) diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index 02ae2b5a..8a14d5a6 100644 --- a/docs/sqlalchemy.md +++ b/docs/sqlalchemy.md @@ -825,12 +825,19 @@ 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`. +An empty `AthenaStruct()` column remains `ROW()`. +Code that compares compiled `CREATE TABLE` strings should expect `STRUCT<...>` and `INT` where earlier releases emitted `ROW(...)` and `INTEGER` for these column types. + #### Querying STRUCT data PyAthena automatically converts STRUCT data between different formats: @@ -970,13 +977,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..cafcf696 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,6 +124,7 @@ def visit_TINYINT(self, type_: types.Integer, **kw: Any) -> str: return "TINYINT" def visit_INTEGER(self, type_: types.Integer, **kw: Any) -> str: + # Hive DDL spells integers as INT inside ARRAY, STRUCT, and MAP columns. return "INT" if kw.get("_athena_array_ddl") else "INTEGER" def visit_SMALLINT(self, type_: types.SmallInteger, **kw: Any) -> str: @@ -199,31 +203,51 @@ def visit_tinyint(self, type_, **kw): def visit_enum(self, type_, **kw): return self.visit_string(type_, **kw) + def _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_array_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_array_ddl`` is set so nested types keep it. + + Returns: + True when the type should use Hive DDL syntax. + """ + if kw.get("_athena_array_ddl") or isinstance(kw.get("type_expression"), Column): + kw["_athena_array_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._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._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}>" @@ -1158,6 +1182,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 e086e7d6..86735f57 100644 --- a/tests/pyathena/sqlalchemy/test_base.py +++ b/tests/pyathena/sqlalchemy/test_base.py @@ -7,6 +7,7 @@ from types import SimpleNamespace from urllib.parse import quote_plus +import boto3 import numpy as np import pandas as pd import pytest @@ -45,6 +46,24 @@ ) +def _delete_s3_prefix(location: str) -> None: + """Delete objects stored under an external table location. + + Args: + location: The table's ``s3://bucket/prefix/`` location. + """ + bucket, _, prefix = location.removeprefix("s3://").partition("/") + if not bucket or not prefix or prefix == "/": + return + if not prefix.endswith("/"): + prefix = f"{prefix}/" + client = boto3.client("s3") + for page in client.get_paginator("list_objects_v2").paginate(Bucket=bucket, Prefix=prefix): + objects = [{"Key": item["Key"]} for item in page.get("Contents", [])] + if objects: + client.delete_objects(Bucket=bucket, Delete={"Objects": objects}) + + def unique_s3tables_table_name(base: str) -> str: """Return a unique S3 Tables table name. @@ -3468,8 +3487,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): @@ -3511,12 +3530,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.""" @@ -3554,6 +3573,63 @@ 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 + try: + 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) + finally: + try: + conn.execute(text(f"DROP TABLE IF EXISTS {ENV.schema}.{table_name}")) + finally: + _delete_s3_prefix(location) + 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..8e563945 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) @@ -681,6 +743,131 @@ def test_like_pattern_is_not_truncated(self, pattern): assert self._format_sql(stmt).endswith(f"WHERE x LIKE '{pattern}'") +class TestStructColumnDDL: + """CREATE TABLE renders STRUCT and nested MAP types with Hive syntax.""" + + 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_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 + + class TestAthenaDDLCompiler: """Compile-only (no AWS) tests for the DDL compiler's S3 Tables support. From af04b14116513e4d6ffbec54954f5abe86ae60ab Mon Sep 17 00:00:00 2001 From: Matt Van Horn <455140+mvanhorn@users.noreply.github.com> Date: Mon, 28 Sep 2026 12:03:38 -0700 Subject: [PATCH 2/3] refactor: rename the Hive DDL flag and drop per-test cleanup Rename the private _athena_array_ddl keyword to _athena_hive_ddl, since it now also covers STRUCT and MAP column DDL, and rename _hive_column_ddl to _enable_hive_column_ddl so the kw side effect is visible at call sites. Remove the S3 cleanup helper and try/finally from the live STRUCT round-trip test to follow the existing pattern, and drop the migration note and ROW() sentence from docs/sqlalchemy.md. --- docs/sqlalchemy.md | 2 - pyathena/sqlalchemy/compiler.py | 19 ++++----- tests/pyathena/sqlalchemy/test_base.py | 59 ++++++++------------------ 3 files changed, 26 insertions(+), 54 deletions(-) diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index 8a14d5a6..71b41c90 100644 --- a/docs/sqlalchemy.md +++ b/docs/sqlalchemy.md @@ -835,8 +835,6 @@ CREATE TABLE users ( 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`. -An empty `AthenaStruct()` column remains `ROW()`. -Code that compares compiled `CREATE TABLE` strings should expect `STRUCT<...>` and `INT` where earlier releases emitted `ROW(...)` and `INTEGER` for these column types. #### Querying STRUCT data diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index cafcf696..4211643b 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -124,8 +124,7 @@ def visit_TINYINT(self, type_: types.Integer, **kw: Any) -> str: return "TINYINT" def visit_INTEGER(self, type_: types.Integer, **kw: Any) -> str: - # Hive DDL spells integers as INT inside ARRAY, STRUCT, and MAP columns. - 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" @@ -203,23 +202,23 @@ def visit_tinyint(self, type_, **kw): def visit_enum(self, type_, **kw): return self.visit_string(type_, **kw) - def _hive_column_ddl(self, kw: dict[str, Any]) -> bool: + 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_array_ddl`` so nested fields use + 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_array_ddl`` is set so nested types keep it. + ``_athena_hive_ddl`` is set so nested types keep it. Returns: True when the type should use Hive DDL syntax. """ - if kw.get("_athena_array_ddl") or isinstance(kw.get("type_expression"), Column): - kw["_athena_array_ddl"] = True + if kw.get("_athena_hive_ddl") or isinstance(kw.get("type_expression"), Column): + kw["_athena_hive_ddl"] = True return True return False @@ -227,7 +226,7 @@ def visit_struct(self, type_, **kw): # Empty structs keep the existing ROW() rendering in every context. if not isinstance(type_, AthenaStruct) or not type_.fields: return "ROW()" - hive_ddl = self._hive_column_ddl(kw) + hive_ddl = self._enable_hive_column_ddl(kw) preparer = ( AthenaDDLIdentifierPreparer(self.dialect) if hive_ddl @@ -247,7 +246,7 @@ def visit_STRUCT(self, type_, **kw): def visit_map(self, type_, **kw): if isinstance(type_, AthenaMap): - self._hive_column_ddl(kw) + 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}>" @@ -258,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" diff --git a/tests/pyathena/sqlalchemy/test_base.py b/tests/pyathena/sqlalchemy/test_base.py index 86735f57..0d06ffb9 100644 --- a/tests/pyathena/sqlalchemy/test_base.py +++ b/tests/pyathena/sqlalchemy/test_base.py @@ -7,7 +7,6 @@ from types import SimpleNamespace from urllib.parse import quote_plus -import boto3 import numpy as np import pandas as pd import pytest @@ -46,24 +45,6 @@ ) -def _delete_s3_prefix(location: str) -> None: - """Delete objects stored under an external table location. - - Args: - location: The table's ``s3://bucket/prefix/`` location. - """ - bucket, _, prefix = location.removeprefix("s3://").partition("/") - if not bucket or not prefix or prefix == "/": - return - if not prefix.endswith("/"): - prefix = f"{prefix}/" - client = boto3.client("s3") - for page in client.get_paginator("list_objects_v2").paginate(Bucket=bucket, Prefix=prefix): - objects = [{"Key": item["Key"]} for item in page.get("Contents", [])] - if objects: - client.delete_objects(Bucket=bucket, Delete={"Objects": objects}) - - def unique_s3tables_table_name(base: str) -> str: """Return a unique S3 Tables table name. @@ -3605,30 +3586,24 @@ def test_external_parquet_struct_columns_round_trip(self, engine): ddl = str(CreateTable(table).compile(dialect=conn.dialect)) assert "profile STRUCT>" in ddl assert "labels MAP>" in ddl - try: - 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))]))" - ) + 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) - finally: - try: - conn.execute(text(f"DROP TABLE IF EXISTS {ENV.schema}.{table_name}")) - finally: - _delete_s3_prefix(location) + ) + 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.""" From 954cc9b4f9cb7f50fd4d620a20dc346240bc02ba Mon Sep 17 00:00:00 2001 From: Matt Van Horn <455140+mvanhorn@users.noreply.github.com> Date: Mon, 28 Sep 2026 22:31:50 -0700 Subject: [PATCH 3/3] test: fold struct column DDL tests into TestAthenaDDLCompiler Move the STRUCT/MAP column DDL tests and the _ddl helper into the existing TestAthenaDDLCompiler class, next to _s3tables_dialect, and generalize the class docstring to cover column type rendering as well as S3 Tables support. --- tests/pyathena/sqlalchemy/test_compiler.py | 251 ++++++++++----------- 1 file changed, 125 insertions(+), 126 deletions(-) diff --git a/tests/pyathena/sqlalchemy/test_compiler.py b/tests/pyathena/sqlalchemy/test_compiler.py index 8e563945..c25e886e 100644 --- a/tests/pyathena/sqlalchemy/test_compiler.py +++ b/tests/pyathena/sqlalchemy/test_compiler.py @@ -743,133 +743,11 @@ def test_like_pattern_is_not_truncated(self, pattern): assert self._format_sql(stmt).endswith(f"WHERE x LIKE '{pattern}'") -class TestStructColumnDDL: - """CREATE TABLE renders STRUCT and nested MAP types with Hive syntax.""" - - 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_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 - - 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 @@ -887,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", @@ -998,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