diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index f0535fa54..6e285f19c 100644 --- a/docs/sqlalchemy.md +++ b/docs/sqlalchemy.md @@ -755,6 +755,7 @@ Athena parses DDL statements such as `CREATE TABLE` with Hive type syntax, and q Complex types apply the same syntax to their nested types. An `AthenaStruct` without fields raises `CompileError` in both. +`NullType`, including an unrecognized type that reflection reports as `NullType`, raises `CompileError` in DDL at any depth. ## Floating-point types @@ -846,6 +847,10 @@ CREATE TABLE users ( `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). +Reflected STRUCT and ROW columns use `AthenaStruct` with their field names and types. +A field type that the dialect does not recognize is reflected as `NullType` with a warning. +Selecting such a column returns the value from the cursor, as described under Data format support below; SQLAlchemy does not convert it to the reflected field types. + #### Querying STRUCT data PyAthena automatically converts STRUCT data between different formats: @@ -991,6 +996,10 @@ CREATE TABLE products ( `CREATE TABLE` renders integer MAP keys and values as `INT`. `CAST` still spells those integers as `INTEGER`. +Reflected MAP columns use `AthenaMap` with their key and value types. +A key or value type that the dialect does not recognize is reflected as `NullType` with a warning. +Selecting such a column returns the value from the cursor, as described under Data format support below; SQLAlchemy does not convert it to the reflected key and value types. + #### Querying MAP data PyAthena automatically converts MAP data between different formats: diff --git a/pyathena/sqlalchemy/base.py b/pyathena/sqlalchemy/base.py index d33331e6b..d620fce6e 100644 --- a/pyathena/sqlalchemy/base.py +++ b/pyathena/sqlalchemy/base.py @@ -695,6 +695,34 @@ def get_columns(self, connection: Connection, table_name: str, schema: str | Non return self._get_columns(connection, table_name, schema=schema, **kw) def _get_column_type(self, type_: str, _nested: bool = False): + """Map an Athena column type string to a SQLAlchemy type. + + Accepts both the Hive (``struct``, ``map``) and the + Trino (``row(a integer)``, ``map(integer, integer)``) spellings, and + parses the element, key, value, and field types of ARRAY, MAP, and + STRUCT/ROW types. + + Args: + type_: The column type reported by Athena. + _nested: Whether ``type_`` is nested in another type. A nested MAP + or STRUCT/ROW that cannot be parsed raises, so that the + enclosing type is reported as unrecognized. + + Returns: + The SQLAlchemy type. A type name that is not recognized, such as + ``foo`` in ``struct``, becomes ``NullType`` in place with a + warning. A type that cannot be parsed, such as ``map`` or + ``varchar(x)``, makes its innermost enclosing ARRAY ``NullType`` + with a warning; without an enclosing ARRAY, a top-level MAP or + STRUCT/ROW becomes ``NullType`` instead. + + Raises: + ValueError: If a type cannot be parsed and neither an enclosing + ARRAY nor a top-level MAP or STRUCT/ROW handles it, for example + a top-level ``varchar(x)`` or a nested ``map``. + TypeError: In the same case, for a DECIMAL type with more + arguments than SQLAlchemy's ``DECIMAL`` accepts. + """ type_ = type_.strip() match = self._pattern_column_type.match(type_) if match: @@ -710,6 +738,12 @@ def _get_column_type(self, type_: str, _nested: bool = False): except (TypeError, ValueError): util.warn(f"Did not recognize type '{type_}'") return types.NullType() + if not _nested and name in ("map", "row", "struct") and length: + try: + return self._get_column_type(type_, _nested=True) + except (TypeError, ValueError): + util.warn(f"Did not recognize type '{type_}'") + return types.NullType() if _nested and name == "map" and length: key, value = _split_type_arguments(length) return AthenaMap( diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index 8c2f81183..189855074 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -90,8 +90,8 @@ class AthenaTypeCompiler(GenericTypeCompiler): 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``. + ``ARRAY``. TIME, JSON, a STRUCT without fields, and ``NullType`` + have no Athena DDL type and raise ``CompileError``. See Also: AWS Athena Data Types: @@ -238,10 +238,6 @@ def visit_unicode(self, type_, **kw): def visit_unicode_text(self, type_, **kw): return "STRING" - @override - def visit_null(self, type_, **kw): - return "NULL" - def visit_tinyint(self, type_, **kw): """Render a tinyint type through ``visit_TINYINT``. diff --git a/tests/pyathena/sqlalchemy/test_base.py b/tests/pyathena/sqlalchemy/test_base.py index 8af87da47..9e0536bc7 100644 --- a/tests/pyathena/sqlalchemy/test_base.py +++ b/tests/pyathena/sqlalchemy/test_base.py @@ -183,6 +183,7 @@ def open_cursor(cursor_class, converter=None): assert [column["name"] for column in columns] == ["id", "payload", "label", "dt"] assert isinstance(columns[0]["type"], types.INTEGER) assert isinstance(columns[1]["type"], AthenaStruct) + assert list(columns[1]["type"].fields) == ["a", "b"] assert type(columns[2]["type"]) is types.String assert type(columns[3]["type"]) is types.String assert [column["comment"] for column in columns] == ["identifier", None, None, None] @@ -196,6 +197,62 @@ def open_cursor(cursor_class, converter=None): assert "WHERE table_schema = 'my_schema' AND table_name = 'o''neil'" in operation assert kwargs == {"result_reuse_enable": False} + @pytest.mark.parametrize( + ("map_type", "struct_type"), + [ + # Metadata API (Hive) spellings. + ("map", "struct>>"), + # information_schema (Trino) spellings. + ("map(integer, varchar)", 'row(a integer, "b c" array(row(x integer)))'), + ], + ) + def test_top_level_map_and_struct_reflect_their_types(self, map_type, struct_type): + dialect = AthenaDialect() + map_ = dialect._get_column_type(map_type) + struct = dialect._get_column_type(struct_type) + + assert isinstance(map_, AthenaMap) + assert isinstance(map_.key_type, types.INTEGER) + assert isinstance(map_.value_type, (types.String, types.VARCHAR)) + assert isinstance(struct, AthenaStruct) + assert list(struct.fields) == ["a", "b c"] + assert isinstance(struct.fields["a"], types.INTEGER) + assert isinstance(struct.fields["b c"], AthenaArray) + assert list(struct.fields["b c"].item_type.fields) == ["x"] + + table = Table( + "t", + MetaData(), + Column("m", map_), + Column("s", struct), + awsathena_location="s3://bucket/path/", + ) + ddl = str(CreateTable(table).compile(dialect=dialect)) + assert "\tm MAP,\n" in ddl + assert "\ts STRUCT>>\n" in ddl + + def test_unrecognized_field_type_reflects_null_type_and_blocks_ddl(self): + dialect = AthenaDialect() + with pytest.warns(sqlalchemy.exc.SAWarning, match="Did not recognize type"): + struct = dialect._get_column_type("struct") + assert isinstance(struct.fields["a"], types.INTEGER) + assert isinstance(struct.fields["b"], types.NullType) + table = Table( + "t", MetaData(), Column("payload", struct), awsathena_location="s3://bucket/path/" + ) + with pytest.raises( + sqlalchemy.exc.CompileError, + match=r"column 'payload'.*Can't generate DDL for NullType", + ): + CreateTable(table).compile(dialect=dialect) + + @pytest.mark.parametrize( + "type_", ["map", "struct", "row(a)", "struct>", "map>"] + ) + def test_unrecognized_top_level_map_or_struct_reflects_null_type(self, type_): + with pytest.warns(sqlalchemy.exc.SAWarning, match="Did not recognize type"): + assert isinstance(AthenaDialect()._get_column_type(type_), types.NullType) + def test_empty_metadata_comment_is_no_comment(self): # Glue can carry an empty comment, so the metadata path must agree with # the information_schema path rather than reflecting it as a comment. @@ -1499,10 +1556,15 @@ def test_reflect_select(self, engine): assert isinstance(one_row_complex.c.col_binary.type, types.BINARY) assert isinstance(one_row_complex.c.col_array.type, AthenaArray) assert isinstance(one_row_complex.c.col_array.type.item_type, types.INTEGER) - assert isinstance(one_row_complex.c.col_map.type, types.String) - # With struct support, col_struct should now be recognized as AthenaStruct - + assert isinstance(one_row_complex.c.col_map.type, AthenaMap) + assert isinstance(one_row_complex.c.col_map.type.key_type, types.INTEGER) + assert isinstance(one_row_complex.c.col_map.type.value_type, types.INTEGER) assert isinstance(one_row_complex.c.col_struct.type, AthenaStruct) + assert list(one_row_complex.c.col_struct.type.fields) == ["a", "b"] + assert isinstance(one_row_complex.c.col_struct.type.fields["a"], types.INTEGER) + ddl = str(CreateTable(one_row_complex).compile(dialect=engine.dialect)) + assert "\tcol_map MAP,\n" in ddl + assert "\tcol_struct STRUCT,\n" in ddl assert isinstance( one_row_complex.c.col_decimal.type, types.DECIMAL, @@ -1552,11 +1614,13 @@ def test_get_column_type(self, engine): assert isinstance(dialect._get_column_type("date"), types.DATE) assert isinstance(dialect._get_column_type("binary"), types.BINARY) assert isinstance(dialect._get_column_type("array"), AthenaArray) - assert isinstance(dialect._get_column_type("map"), types.String) - # With struct support, struct types should be recognized as AthenaStruct - - assert isinstance(dialect._get_column_type("struct"), AthenaStruct) - assert isinstance(dialect._get_column_type("row"), AthenaStruct) + assert isinstance(dialect._get_column_type("map"), AthenaMap) + struct = dialect._get_column_type("struct") + assert isinstance(struct, AthenaStruct) + assert list(struct.fields) == ["a", "b"] + row = dialect._get_column_type("row") + assert isinstance(row, AthenaStruct) + assert list(row.fields) == ["name", "age"] decimal_with_args = dialect._get_column_type("decimal(10,1)") assert isinstance(decimal_with_args, types.DECIMAL) assert decimal_with_args.precision == 10 diff --git a/tests/pyathena/sqlalchemy/test_compiler.py b/tests/pyathena/sqlalchemy/test_compiler.py index 0c3548a76..b5a841ba4 100644 --- a/tests/pyathena/sqlalchemy/test_compiler.py +++ b/tests/pyathena/sqlalchemy/test_compiler.py @@ -218,6 +218,25 @@ def test_ddl_and_cast_types(self, type_, ddl, cast_type): 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.NullType(), + AthenaArray(types.NullType()), + AthenaMap(Integer, types.NullType()), + AthenaStruct(("a", Integer), ("b", types.NullType())), + ], + ) + def test_null_type_is_rejected_in_ddl(self, type_): + dialect = AthenaDialect() + with pytest.raises(exc.CompileError, match="Can't generate DDL for NullType"): + dialect.type_compiler_instance.process(type_) + + def test_null_type_cast_is_unchanged(self): + assert str(cast(column("x"), types.NullType()).compile(dialect=AthenaDialect())) == ( + "CAST(x AS NULL)" + ) + @pytest.mark.parametrize( "type_", [types.JSON(), AthenaArray(types.JSON), AthenaMap(String, types.JSON)] )