Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions CHANGES.rst
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,11 @@ Version history

- Added autoincrement to primary key columns to prevent missing field errors.
(`#473 <https://github.com/agronholm/sqlacodegen/issues/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 <https://github.com/agronholm/sqlacodegen/issues/441>`_)

**4.0.3**

Expand Down
6 changes: 6 additions & 0 deletions src/sqlacodegen/generators.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Comment thread
agronholm marked this conversation as resolved.
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(
Expand Down
33 changes: 32 additions & 1 deletion tests/test_generator_tables.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()

Expand All @@ -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(
Expand Down