diff --git a/CHANGES.rst b/CHANGES.rst index 75be5d1..eaa1ff6 100644 --- a/CHANGES.rst +++ b/CHANGES.rst @@ -5,6 +5,11 @@ Version history - Added autoincrement to primary key columns to prevent missing field errors. (`#473 `_; PR by @jtmonroe) +- Preserve dialect-specific ``ARRAY`` types (e.g. ``postgresql.ARRAY``) instead + of adapting them to the generic ``sqlalchemy.ARRAY``. The generic type does + not implement operators like ``.contains()``, so adapting silently broke + PostgreSQL array queries on generated models. + (`#441 `_) **4.0.3** diff --git a/src/sqlacodegen/generators.py b/src/sqlacodegen/generators.py index 0dfba36..a10a21e 100644 --- a/src/sqlacodegen/generators.py +++ b/src/sqlacodegen/generators.py @@ -1036,6 +1036,12 @@ def fix_enum_column(col_name: str, enum_type: Enum) -> None: column.server_default = None def get_adapted_type(self, coltype: Any) -> Any: + # Keep dialect-specific ARRAY subclasses; the generic sqlalchemy.ARRAY + # is missing operators like .contains() (GH-441). + if isinstance(coltype, ARRAY) and type(coltype) is not ARRAY: + coltype.item_type = self.get_adapted_type(coltype.item_type) + return coltype + compiled_type = coltype.compile(self.bind.engine.dialect) for supercls in coltype.__class__.__mro__: if not supercls.__name__.startswith("_") and hasattr( diff --git a/tests/test_generator_tables.py b/tests/test_generator_tables.py index 8633e3b..e467dc2 100644 --- a/tests/test_generator_tables.py +++ b/tests/test_generator_tables.py @@ -130,7 +130,8 @@ def test_arrays(generator: CodeGenerator) -> None: validate_code( generator.generate(), """\ - from sqlalchemy import ARRAY, Column, Double, Integer, MetaData, Table + from sqlalchemy import Column, Double, Integer, MetaData, Table + from sqlalchemy.dialects.postgresql import ARRAY metadata = MetaData() @@ -144,6 +145,36 @@ def test_arrays(generator: CodeGenerator) -> None: ) +@pytest.mark.parametrize("engine", ["postgresql"], indirect=["engine"]) +def test_array_preserves_dialect_for_runtime_operators( + generator: CodeGenerator, +) -> None: + """Regression test for GH-441.""" + Table( + "simple_items", + generator.metadata, + Column("id", postgresql.TEXT, primary_key=True), + Column("tags", postgresql.ARRAY(postgresql.TEXT)), + ) + + validate_code( + generator.generate(), + """\ + from sqlalchemy import Column, MetaData, Table, Text + from sqlalchemy.dialects.postgresql import ARRAY + + metadata = MetaData() + + + t_simple_items = Table( + 'simple_items', metadata, + Column('id', Text, primary_key=True), + Column('tags', ARRAY(Text())) + ) + """, + ) + + @pytest.mark.parametrize("engine", ["postgresql"], indirect=["engine"]) def test_jsonb(generator: CodeGenerator) -> None: Table(