• Home
  • Features
  • Pricing
  • Docs
  • Announcements
  • Sign In

agronholm / sqlacodegen / 21780194160

07 Feb 2026 12:35PM UTC coverage: 97.596% (+0.08%) from 97.519%
21780194160

Pull #455

github

web-flow
Merge 1ca663745 into e9bfd344c
Pull Request #455: Support enum arrays

54 of 54 new or added lines in 3 files covered. (100.0%)

15 existing lines in 2 files now uncovered.

1746 of 1789 relevant lines covered (97.6%)

4.88 hits per line

Source File
Press 'n' to go to next uncovered line, 'b' for previous

96.64
/src/sqlacodegen/generators.py
1
from __future__ import annotations
5✔
2

3
import inspect
5✔
4
import re
5✔
5
import sys
5✔
6
from abc import ABCMeta, abstractmethod
5✔
7
from collections import defaultdict
5✔
8
from collections.abc import Collection, Iterable, Sequence
5✔
9
from dataclasses import dataclass
5✔
10
from importlib import import_module
5✔
11
from inspect import Parameter
5✔
12
from itertools import count
5✔
13
from keyword import iskeyword
5✔
14
from pprint import pformat
5✔
15
from textwrap import indent
5✔
16
from typing import Any, ClassVar, Literal, cast
5✔
17

18
import inflect
5✔
19
import sqlalchemy
5✔
20
from sqlalchemy import (
5✔
21
    ARRAY,
22
    Boolean,
23
    CheckConstraint,
24
    Column,
25
    Computed,
26
    Constraint,
27
    DefaultClause,
28
    Enum,
29
    ForeignKey,
30
    ForeignKeyConstraint,
31
    Identity,
32
    Index,
33
    MetaData,
34
    PrimaryKeyConstraint,
35
    String,
36
    Table,
37
    Text,
38
    TypeDecorator,
39
    UniqueConstraint,
40
)
41
from sqlalchemy.dialects.postgresql import DOMAIN, JSON, JSONB
5✔
42
from sqlalchemy.engine import Connection, Engine
5✔
43
from sqlalchemy.exc import CompileError
5✔
44
from sqlalchemy.sql.elements import TextClause
5✔
45
from sqlalchemy.sql.type_api import UserDefinedType
5✔
46
from sqlalchemy.types import TypeEngine
5✔
47

48
from .models import (
5✔
49
    ColumnAttribute,
50
    JoinType,
51
    Model,
52
    ModelClass,
53
    RelationshipAttribute,
54
    RelationshipType,
55
)
56
from .utils import (
5✔
57
    decode_postgresql_sequence,
58
    get_column_names,
59
    get_common_fk_constraints,
60
    get_compiled_expression,
61
    get_constraint_sort_key,
62
    get_stdlib_module_names,
63
    qualified_table_name,
64
    render_callable,
65
    uses_default_name,
66
)
67

68
_re_boolean_check_constraint = re.compile(r"(?:.*?\.)?(.*?) IN \(0, 1\)")
5✔
69
_re_column_name = re.compile(r'(?:(["`]?).*\1\.)?(["`]?)(.*)\2')
5✔
70
_re_enum_check_constraint = re.compile(r"(?:.*?\.)?(.*?) IN \((.+)\)")
5✔
71
_re_enum_item = re.compile(r"'(.*?)(?<!\\)'")
5✔
72
_re_invalid_identifier = re.compile(r"(?u)\W")
5✔
73

74

75
@dataclass
5✔
76
class LiteralImport:
5✔
77
    pkgname: str
5✔
78
    name: str
5✔
79

80

81
@dataclass
5✔
82
class Base:
5✔
83
    """Representation of MetaData for Tables, respectively Base for classes"""
84

85
    literal_imports: list[LiteralImport]
5✔
86
    declarations: list[str]
5✔
87
    metadata_ref: str
5✔
88
    decorator: str | None = None
5✔
89
    table_metadata_declaration: str | None = None
5✔
90

91

92
class CodeGenerator(metaclass=ABCMeta):
5✔
93
    valid_options: ClassVar[set[str]] = set()
5✔
94

95
    def __init__(
5✔
96
        self, metadata: MetaData, bind: Connection | Engine, options: Sequence[str]
97
    ):
98
        self.metadata: MetaData = metadata
5✔
99
        self.bind: Connection | Engine = bind
5✔
100
        self.options: set[str] = set(options)
5✔
101

102
        # Validate options
103
        invalid_options = {opt for opt in options if opt not in self.valid_options}
5✔
104
        if invalid_options:
5✔
105
            raise ValueError("Unrecognized options: " + ", ".join(invalid_options))
×
106

107
    @property
5✔
108
    @abstractmethod
5✔
109
    def views_supported(self) -> bool:
5✔
110
        pass
×
111

112
    @abstractmethod
5✔
113
    def generate(self) -> str:
5✔
114
        """
115
        Generate the code for the given metadata.
116
        .. note:: May modify the metadata.
117
        """
118

119

120
@dataclass(eq=False)
5✔
121
class TablesGenerator(CodeGenerator):
5✔
122
    valid_options: ClassVar[set[str]] = {
5✔
123
        "noindexes",
124
        "noconstraints",
125
        "nocomments",
126
        "nonativeenums",
127
        "nosyntheticenums",
128
        "include_dialect_options",
129
        "keep_dialect_types",
130
    }
131
    stdlib_module_names: ClassVar[set[str]] = get_stdlib_module_names()
5✔
132

133
    def __init__(
5✔
134
        self,
135
        metadata: MetaData,
136
        bind: Connection | Engine,
137
        options: Sequence[str],
138
        *,
139
        indentation: str = "    ",
140
    ):
141
        super().__init__(metadata, bind, options)
5✔
142
        self.indentation: str = indentation
5✔
143
        self.imports: dict[str, set[str]] = defaultdict(set)
5✔
144
        self.module_imports: set[str] = set()
5✔
145

146
        # Render SchemaItem.info and dialect kwargs (Table/Column) into output
147
        self.include_dialect_options_and_info: bool = (
5✔
148
            "include_dialect_options" in self.options
149
        )
150
        # Keep dialect-specific types instead of adapting to generic SQLAlchemy types
151
        self.keep_dialect_types: bool = "keep_dialect_types" in self.options
5✔
152

153
        # Track Python enum classes: maps (table_name, column_name) -> enum_class_name
154
        self.enum_classes: dict[tuple[str, str], str] = {}
5✔
155
        # Track enum values: maps enum_class_name -> list of values
156
        self.enum_values: dict[str, list[str]] = {}
5✔
157

158
    @property
5✔
159
    def views_supported(self) -> bool:
5✔
160
        return True
×
161

162
    def generate_base(self) -> None:
5✔
163
        self.base = Base(
5✔
164
            literal_imports=[LiteralImport("sqlalchemy", "MetaData")],
165
            declarations=["metadata = MetaData()"],
166
            metadata_ref="metadata",
167
        )
168

169
    def generate(self) -> str:
5✔
170
        self.generate_base()
5✔
171

172
        sections: list[str] = []
5✔
173

174
        # Remove unwanted elements from the metadata
175
        for table in list(self.metadata.tables.values()):
5✔
176
            if self.should_ignore_table(table):
5✔
177
                self.metadata.remove(table)
×
178
                continue
×
179

180
            if "noindexes" in self.options:
5✔
181
                table.indexes.clear()
5✔
182

183
            if "noconstraints" in self.options:
5✔
184
                table.constraints.clear()
5✔
185

186
            if "nocomments" in self.options:
5✔
187
                table.comment = None
5✔
188

189
            for column in table.columns:
5✔
190
                if "nocomments" in self.options:
5✔
191
                    column.comment = None
5✔
192

193
        # Use information from column constraints to figure out the intended column
194
        # types
195
        for table in self.metadata.tables.values():
5✔
196
            self.fix_column_types(table)
5✔
197

198
        # Generate the models
199
        models: list[Model] = self.generate_models()
5✔
200

201
        # Render module level variables
202
        if variables := self.render_module_variables(models):
5✔
203
            sections.append(variables + "\n")
5✔
204

205
        # Render enum classes
206
        if enum_classes := self.render_enum_classes():
5✔
207
            sections.append(enum_classes + "\n")
5✔
208

209
        # Render models
210
        if rendered_models := self.render_models(models):
5✔
211
            sections.append(rendered_models)
5✔
212

213
        # Render collected imports
214
        groups = self.group_imports()
5✔
215
        if imports := "\n\n".join(
5✔
216
            "\n".join(line for line in group) for group in groups
217
        ):
218
            sections.insert(0, imports)
5✔
219

220
        return "\n\n".join(sections) + "\n"
5✔
221

222
    def collect_imports(self, models: Iterable[Model]) -> None:
5✔
223
        for literal_import in self.base.literal_imports:
5✔
224
            self.add_literal_import(literal_import.pkgname, literal_import.name)
5✔
225

226
        for model in models:
5✔
227
            self.collect_imports_for_model(model)
5✔
228

229
    def collect_imports_for_model(self, model: Model) -> None:
5✔
230
        if model.__class__ is Model:
5✔
231
            self.add_import(Table)
5✔
232

233
        for column in model.table.c:
5✔
234
            self.collect_imports_for_column(column)
5✔
235

236
        for constraint in model.table.constraints:
5✔
237
            self.collect_imports_for_constraint(constraint)
5✔
238

239
        for index in model.table.indexes:
5✔
240
            self.collect_imports_for_constraint(index)
5✔
241

242
    def collect_imports_for_column(self, column: Column[Any]) -> None:
5✔
243
        self.add_import(column.type)
5✔
244

245
        if isinstance(column.type, ARRAY):
5✔
246
            self.add_import(column.type.item_type.__class__)
5✔
247
        elif isinstance(column.type, (JSONB, JSON)):
5✔
248
            if (
5✔
249
                not isinstance(column.type.astext_type, Text)
250
                or column.type.astext_type.length is not None
251
            ):
252
                self.add_import(column.type.astext_type)
5✔
253
        elif isinstance(column.type, DOMAIN):
5✔
254
            self.add_import(column.type.data_type.__class__)
5✔
255

256
        if column.default:
5✔
257
            self.add_import(column.default)
5✔
258

259
        if column.server_default:
5✔
260
            if isinstance(column.server_default, (Computed, Identity)):
5✔
261
                self.add_import(column.server_default)
5✔
262
            elif isinstance(column.server_default, DefaultClause):
5✔
263
                self.add_literal_import("sqlalchemy", "text")
5✔
264

265
    def collect_imports_for_constraint(self, constraint: Constraint | Index) -> None:
5✔
266
        if isinstance(constraint, Index):
5✔
267
            if len(constraint.columns) > 1 or not uses_default_name(constraint):
5✔
268
                self.add_literal_import("sqlalchemy", "Index")
5✔
269
        elif isinstance(constraint, PrimaryKeyConstraint):
5✔
270
            if not uses_default_name(constraint):
5✔
271
                self.add_literal_import("sqlalchemy", "PrimaryKeyConstraint")
5✔
272
        elif isinstance(constraint, UniqueConstraint):
5✔
273
            if len(constraint.columns) > 1 or not uses_default_name(constraint):
5✔
274
                self.add_literal_import("sqlalchemy", "UniqueConstraint")
5✔
275
        elif isinstance(constraint, ForeignKeyConstraint):
5✔
276
            if len(constraint.columns) > 1 or not uses_default_name(constraint):
5✔
277
                self.add_literal_import("sqlalchemy", "ForeignKeyConstraint")
5✔
278
            else:
279
                self.add_import(ForeignKey)
5✔
280
        else:
281
            self.add_import(constraint)
5✔
282

283
    def add_import(self, obj: Any) -> None:
5✔
284
        # Don't store builtin imports
285
        if getattr(obj, "__module__", "builtins") == "builtins":
5✔
286
            return
×
287

288
        type_ = type(obj) if not isinstance(obj, type) else obj
5✔
289
        pkgname = type_.__module__
5✔
290

291
        # The column types have already been adapted towards generic types if possible,
292
        # so if this is still a vendor specific type (e.g., MySQL INTEGER) be sure to
293
        # use that rather than the generic sqlalchemy type as it might have different
294
        # constructor parameters.
295
        if pkgname.startswith("sqlalchemy.dialects."):
5✔
296
            dialect_pkgname = ".".join(pkgname.split(".")[0:3])
5✔
297
            dialect_pkg = import_module(dialect_pkgname)
5✔
298

299
            if type_.__name__ in dialect_pkg.__all__:
5✔
300
                pkgname = dialect_pkgname
5✔
301
        elif type_ is getattr(sqlalchemy, type_.__name__, None):
5✔
302
            pkgname = "sqlalchemy"
5✔
303
        else:
304
            pkgname = type_.__module__
5✔
305

306
        self.add_literal_import(pkgname, type_.__name__)
5✔
307

308
    def add_literal_import(self, pkgname: str, name: str) -> None:
5✔
309
        names = self.imports.setdefault(pkgname, set())
5✔
310
        names.add(name)
5✔
311

312
    def remove_literal_import(self, pkgname: str, name: str) -> None:
5✔
313
        names = self.imports.setdefault(pkgname, set())
5✔
314
        if name in names:
5✔
315
            names.remove(name)
×
316

317
    def add_module_import(self, pgkname: str) -> None:
5✔
318
        self.module_imports.add(pgkname)
5✔
319

320
    def group_imports(self) -> list[list[str]]:
5✔
321
        future_imports: list[str] = []
5✔
322
        stdlib_imports: list[str] = []
5✔
323
        thirdparty_imports: list[str] = []
5✔
324

325
        def get_collection(package: str) -> list[str]:
5✔
326
            collection = thirdparty_imports
5✔
327
            if package == "__future__":
5✔
328
                collection = future_imports
×
329
            elif package in self.stdlib_module_names:
5✔
330
                collection = stdlib_imports
5✔
331
            elif package in sys.modules:
5✔
332
                if "site-packages" not in (sys.modules[package].__file__ or ""):
5✔
333
                    collection = stdlib_imports
5✔
334
            return collection
5✔
335

336
        for package in sorted(self.imports):
5✔
337
            imports = ", ".join(sorted(self.imports[package]))
5✔
338

339
            collection = get_collection(package)
5✔
340
            collection.append(f"from {package} import {imports}")
5✔
341

342
        for module in sorted(self.module_imports):
5✔
343
            collection = get_collection(module)
5✔
344
            collection.append(f"import {module}")
5✔
345

346
        return [
5✔
347
            group
348
            for group in (future_imports, stdlib_imports, thirdparty_imports)
349
            if group
350
        ]
351

352
    def generate_models(self) -> list[Model]:
5✔
353
        models = [Model(table) for table in self.metadata.sorted_tables]
5✔
354

355
        # Collect the imports
356
        self.collect_imports(models)
5✔
357

358
        # Generate names for models
359
        global_names = {
5✔
360
            name for namespace in self.imports.values() for name in namespace
361
        }
362
        for model in models:
5✔
363
            self.generate_model_name(model, global_names)
5✔
364
            global_names.add(model.name)
5✔
365

366
        return models
5✔
367

368
    def generate_model_name(self, model: Model, global_names: set[str]) -> None:
5✔
369
        preferred_name = f"t_{model.table.name}"
5✔
370
        model.name = self.find_free_name(preferred_name, global_names)
5✔
371

372
    def render_module_variables(self, models: list[Model]) -> str:
5✔
373
        declarations = self.base.declarations
5✔
374

375
        if any(not isinstance(model, ModelClass) for model in models):
5✔
376
            if self.base.table_metadata_declaration is not None:
5✔
377
                declarations.append(self.base.table_metadata_declaration)
×
378

379
        return "\n".join(declarations)
5✔
380

381
    def render_models(self, models: list[Model]) -> str:
5✔
382
        rendered: list[str] = []
5✔
383
        for model in models:
5✔
384
            rendered_table = self.render_table(model.table)
5✔
385
            rendered.append(f"{model.name} = {rendered_table}")
5✔
386

387
        return "\n\n".join(rendered)
5✔
388

389
    def render_table(self, table: Table) -> str:
5✔
390
        args: list[str] = [f"{table.name!r}, {self.base.metadata_ref}"]
5✔
391
        kwargs: dict[str, object] = {}
5✔
392
        for column in table.columns:
5✔
393
            # Cast is required because of a bug in the SQLAlchemy stubs regarding
394
            # Table.columns
395
            args.append(self.render_column(column, True, is_table=True))
5✔
396

397
        for constraint in sorted(table.constraints, key=get_constraint_sort_key):
5✔
398
            if uses_default_name(constraint):
5✔
399
                if isinstance(constraint, PrimaryKeyConstraint):
5✔
400
                    continue
5✔
401
                elif isinstance(constraint, (ForeignKeyConstraint, UniqueConstraint)):
5✔
402
                    if len(constraint.columns) == 1:
5✔
403
                        continue
5✔
404

405
            args.append(self.render_constraint(constraint))
5✔
406

407
        for index in sorted(table.indexes, key=lambda i: cast(str, i.name)):
5✔
408
            # One-column indexes should be rendered as index=True on columns
409
            if len(index.columns) > 1 or not uses_default_name(index):
5✔
410
                args.append(self.render_index(index))
5✔
411

412
        if table.schema:
5✔
413
            kwargs["schema"] = repr(table.schema)
5✔
414

415
        table_comment = getattr(table, "comment", None)
5✔
416
        if table_comment:
5✔
417
            kwargs["comment"] = repr(table.comment)
5✔
418

419
        # add info + dialect kwargs for callable context (opt-in)
420
        if self.include_dialect_options_and_info:
5✔
421
            self._add_dialect_kwargs_and_info(table, kwargs, values_for_dict=False)
5✔
422

423
        return render_callable("Table", *args, kwargs=kwargs, indentation="    ")
5✔
424

425
    def render_index(self, index: Index) -> str:
5✔
426
        extra_args = [repr(col.name) for col in index.columns]
5✔
427
        kwargs = {
5✔
428
            key: repr(value) if isinstance(value, str) else value
429
            for key, value in sorted(index.kwargs.items(), key=lambda item: item[0])
430
        }
431
        if index.unique:
5✔
432
            kwargs["unique"] = True
5✔
433

434
        return render_callable("Index", repr(index.name), *extra_args, kwargs=kwargs)
5✔
435

436
    # TODO find better solution for is_table
437
    def render_column(
5✔
438
        self, column: Column[Any], show_name: bool, is_table: bool = False
439
    ) -> str:
440
        args = []
5✔
441
        kwargs: dict[str, Any] = {}
5✔
442
        kwarg = []
5✔
443
        is_part_of_composite_pk = (
5✔
444
            column.primary_key and len(column.table.primary_key) > 1
445
        )
446
        dedicated_fks = [
5✔
447
            c
448
            for c in column.foreign_keys
449
            if c.constraint
450
            and len(c.constraint.columns) == 1
451
            and uses_default_name(c.constraint)
452
        ]
453
        is_unique = any(
5✔
454
            isinstance(c, UniqueConstraint)
455
            and set(c.columns) == {column}
456
            and uses_default_name(c)
457
            for c in column.table.constraints
458
        )
459
        is_unique = is_unique or any(
5✔
460
            i.unique and set(i.columns) == {column} and uses_default_name(i)
461
            for i in column.table.indexes
462
        )
463
        is_primary = (
5✔
464
            any(
465
                isinstance(c, PrimaryKeyConstraint)
466
                and column.name in c.columns
467
                and uses_default_name(c)
468
                for c in column.table.constraints
469
            )
470
            or column.primary_key
471
        )
472
        has_index = any(
5✔
473
            set(i.columns) == {column} and uses_default_name(i)
474
            for i in column.table.indexes
475
        )
476

477
        if show_name:
5✔
478
            args.append(repr(column.name))
5✔
479

480
        # Render the column type if there are no foreign keys on it or any of them
481
        # points back to itself
482
        if not dedicated_fks or any(fk.column is column for fk in dedicated_fks):
5✔
483
            args.append(self.render_column_type(column))
5✔
484

485
        for fk in dedicated_fks:
5✔
486
            args.append(self.render_constraint(fk))
5✔
487

488
        if column.default:
5✔
489
            args.append(repr(column.default))
5✔
490

491
        if column.key != column.name:
5✔
492
            kwargs["key"] = column.key
×
493
        if is_primary:
5✔
494
            kwargs["primary_key"] = True
5✔
495
        if not column.nullable and not column.primary_key:
5✔
496
            kwargs["nullable"] = False
5✔
497
        if column.nullable and is_part_of_composite_pk:
5✔
498
            kwargs["nullable"] = True
5✔
499

500
        if is_unique:
5✔
501
            column.unique = True
5✔
502
            kwargs["unique"] = True
5✔
503
        if has_index:
5✔
504
            column.index = True
5✔
505
            kwarg.append("index")
5✔
506
            kwargs["index"] = True
5✔
507

508
        if isinstance(column.server_default, DefaultClause):
5✔
509
            kwargs["server_default"] = render_callable(
5✔
510
                "text", repr(cast(TextClause, column.server_default.arg).text)
511
            )
512
        elif isinstance(column.server_default, Computed):
5✔
513
            expression = str(column.server_default.sqltext)
5✔
514

515
            computed_kwargs = {}
5✔
516
            if column.server_default.persisted is not None:
5✔
517
                computed_kwargs["persisted"] = column.server_default.persisted
5✔
518

519
            args.append(
5✔
520
                render_callable("Computed", repr(expression), kwargs=computed_kwargs)
521
            )
522
        elif isinstance(column.server_default, Identity):
5✔
523
            args.append(repr(column.server_default))
5✔
524
        elif column.server_default:
5✔
525
            kwargs["server_default"] = repr(column.server_default)
×
526

527
        comment = getattr(column, "comment", None)
5✔
528
        if comment:
5✔
529
            kwargs["comment"] = repr(comment)
5✔
530

531
        # add column info + dialect kwargs for callable context (opt-in)
532
        if self.include_dialect_options_and_info:
5✔
533
            self._add_dialect_kwargs_and_info(column, kwargs, values_for_dict=False)
5✔
534

535
        return self.render_column_callable(is_table, *args, **kwargs)
5✔
536

537
    def render_column_callable(self, is_table: bool, *args: Any, **kwargs: Any) -> str:
5✔
538
        if is_table:
5✔
539
            self.add_import(Column)
5✔
540
            return render_callable("Column", *args, kwargs=kwargs)
5✔
541
        else:
542
            return render_callable("mapped_column", *args, kwargs=kwargs)
5✔
543

544
    def render_column_type(self, column: Column[Any]) -> str:
5✔
545
        column_type = column.type
5✔
546
        # Check if this is an enum column with a Python enum class
547
        if isinstance(column_type, Enum) and column is not None:
5✔
548
            if enum_class_name := self.enum_classes.get(
5✔
549
                (column.table.name, column.name)
550
            ):
551
                # Import SQLAlchemy Enum (will be handled in collect_imports)
552
                self.add_import(Enum)
5✔
553
                return f"Enum({enum_class_name}, values_callable=lambda cls: [member.value for member in cls])"
5✔
554

555
        args = []
5✔
556
        kwargs: dict[str, Any] = {}
5✔
557

558
        # Check if this is an ARRAY column with an Enum item type mapped to a Python enum class
559
        if isinstance(column_type, ARRAY) and isinstance(column_type.item_type, Enum):
5✔
560
            if enum_class_name := self.enum_classes.get(
5✔
561
                (column.table.name, column.name)
562
            ):
563
                self.add_import(ARRAY)
5✔
564
                self.add_import(Enum)
5✔
565
                rendered_enum = f"Enum({enum_class_name}, values_callable=lambda cls: [member.value for member in cls])"
5✔
566
                if column_type.dimensions is not None:
5✔
567
                    kwargs["dimensions"] = repr(column_type.dimensions)
5✔
568

569
                return render_callable("ARRAY", rendered_enum, kwargs=kwargs)
5✔
570

571
        sig = inspect.signature(column_type.__class__.__init__)
5✔
572
        defaults = {param.name: param.default for param in sig.parameters.values()}
5✔
573
        missing = object()
5✔
574
        use_kwargs = False
5✔
575
        for param in list(sig.parameters.values())[1:]:
5✔
576
            # Remove annoyances like _warn_on_bytestring
577
            if param.name.startswith("_"):
5✔
578
                continue
5✔
579
            elif param.kind in (Parameter.VAR_POSITIONAL, Parameter.VAR_KEYWORD):
5✔
580
                use_kwargs = True
5✔
581
                continue
5✔
582

583
            value = getattr(column_type, param.name, missing)
5✔
584

585
            if isinstance(value, (JSONB, JSON)):
5✔
586
                # Remove astext_type if it's the default
587
                if (
5✔
588
                    isinstance(value.astext_type, Text)
589
                    and value.astext_type.length is None
590
                ):
591
                    value.astext_type = None  # type: ignore[assignment]
5✔
592
                else:
593
                    self.add_import(Text)
5✔
594

595
            default = defaults.get(param.name, missing)
5✔
596
            if isinstance(value, TextClause):
5✔
597
                self.add_literal_import("sqlalchemy", "text")
5✔
598
                rendered_value = render_callable("text", repr(value.text))
5✔
599
            else:
600
                rendered_value = repr(value)
5✔
601

602
            if value is missing or value == default:
5✔
603
                use_kwargs = True
5✔
604
            elif use_kwargs:
5✔
605
                kwargs[param.name] = rendered_value
5✔
606
            else:
607
                args.append(rendered_value)
5✔
608

609
        vararg = next(
5✔
610
            (
611
                param.name
612
                for param in sig.parameters.values()
613
                if param.kind is Parameter.VAR_POSITIONAL
614
            ),
615
            None,
616
        )
617
        if vararg and hasattr(column_type, vararg):
5✔
618
            varargs_repr = [repr(arg) for arg in getattr(column_type, vararg)]
5✔
619
            args.extend(varargs_repr)
5✔
620

621
        # These arguments cannot be autodetected from the Enum initializer
622
        if isinstance(column_type, Enum):
5✔
623
            for colname in "name", "schema":
5✔
624
                if (value := getattr(column_type, colname)) is not None:
5✔
625
                    kwargs[colname] = repr(value)
5✔
626

627
        if isinstance(column_type, (JSONB, JSON)):
5✔
628
            # Remove astext_type if it's the default
629
            if (
5✔
630
                isinstance(column_type.astext_type, Text)
631
                and column_type.astext_type.length is None
632
            ):
633
                del kwargs["astext_type"]
5✔
634

635
        if args or kwargs:
5✔
636
            return render_callable(column_type.__class__.__name__, *args, kwargs=kwargs)
5✔
637
        else:
638
            return column_type.__class__.__name__
5✔
639

640
    def render_constraint(self, constraint: Constraint | ForeignKey) -> str:
5✔
641
        def add_fk_options(*opts: Any) -> None:
5✔
642
            args.extend(repr(opt) for opt in opts)
5✔
643
            for attr in "ondelete", "onupdate", "deferrable", "initially", "match":
5✔
644
                value = getattr(constraint, attr, None)
5✔
645
                if value:
5✔
646
                    kwargs[attr] = repr(value)
5✔
647

648
        args: list[str] = []
5✔
649
        kwargs: dict[str, Any] = {}
5✔
650
        if isinstance(constraint, ForeignKey):
5✔
651
            remote_column = (
5✔
652
                f"{constraint.column.table.fullname}.{constraint.column.name}"
653
            )
654
            add_fk_options(remote_column)
5✔
655
        elif isinstance(constraint, ForeignKeyConstraint):
5✔
656
            local_columns = get_column_names(constraint)
5✔
657
            remote_columns = [
5✔
658
                f"{fk.column.table.fullname}.{fk.column.name}"
659
                for fk in constraint.elements
660
            ]
661
            add_fk_options(local_columns, remote_columns)
5✔
662
        elif isinstance(constraint, CheckConstraint):
5✔
663
            args.append(repr(get_compiled_expression(constraint.sqltext, self.bind)))
5✔
664
        elif isinstance(constraint, (UniqueConstraint, PrimaryKeyConstraint)):
5✔
665
            args.extend(repr(col.name) for col in constraint.columns)
5✔
666
        else:
667
            raise TypeError(
×
668
                f"Cannot render constraint of type {constraint.__class__.__name__}"
669
            )
670

671
        if isinstance(constraint, Constraint) and not uses_default_name(constraint):
5✔
672
            kwargs["name"] = repr(constraint.name)
5✔
673

674
        return render_callable(constraint.__class__.__name__, *args, kwargs=kwargs)
5✔
675

676
    def _add_dialect_kwargs_and_info(
5✔
677
        self, obj: Any, target_kwargs: dict[str, object], *, values_for_dict: bool
678
    ) -> None:
679
        """
680
        Merge SchemaItem-like object's .info and .dialect_kwargs into target_kwargs.
681
        - values_for_dict=True: keep raw values so pretty-printer emits repr() (for __table_args__ dict)
682
        - values_for_dict=False: set values to repr() strings (for callable kwargs)
683
        """
684
        info_dict = getattr(obj, "info", None)
5✔
685
        if info_dict:
5✔
686
            target_kwargs["info"] = info_dict if values_for_dict else repr(info_dict)
5✔
687

688
        dialect_keys: list[str]
689
        try:
5✔
690
            dialect_keys = sorted(getattr(obj, "dialect_kwargs"))
5✔
691
        except Exception:
×
692
            return
×
693

694
        dialect_kwargs = getattr(obj, "dialect_kwargs", {})
5✔
695
        for key in dialect_keys:
5✔
696
            try:
5✔
697
                value = dialect_kwargs[key]
5✔
698
            except Exception:
×
699
                continue
×
700

701
            # Render values:
702
            # - callable context (values_for_dict=False): produce a string expression.
703
            #   primitives use repr(value); custom objects stringify then repr().
704
            # - dict context (values_for_dict=True): pass raw primitives / str;
705
            #   custom objects become str(value) so pformat quotes them.
706
            if values_for_dict:
5✔
707
                if isinstance(value, type(None) | bool | int | float):
5✔
708
                    target_kwargs[key] = value
×
709
                elif isinstance(value, str | dict | list):
5✔
710
                    target_kwargs[key] = value
5✔
711
                else:
712
                    target_kwargs[key] = str(value)
5✔
713
            else:
714
                if isinstance(
5✔
715
                    value, type(None) | bool | int | float | str | dict | list
716
                ):
717
                    target_kwargs[key] = repr(value)
5✔
718
                else:
719
                    target_kwargs[key] = repr(str(value))
5✔
720

721
    def should_ignore_table(self, table: Table) -> bool:
5✔
722
        # Support for Alembic and sqlalchemy-migrate -- never expose the schema version
723
        # tables
724
        return table.name in ("alembic_version", "migrate_version")
5✔
725

726
    def find_free_name(
5✔
727
        self, name: str, global_names: set[str], local_names: Collection[str] = ()
728
    ) -> str:
729
        """
730
        Generate an attribute name that does not clash with other local or global names.
731
        """
732
        name = name.strip()
5✔
733
        assert name, "Identifier cannot be empty"
5✔
734
        name = _re_invalid_identifier.sub("_", name)
5✔
735
        if name[0].isdigit():
5✔
736
            name = "_" + name
5✔
737
        elif iskeyword(name) or name == "metadata":
5✔
738
            name += "_"
5✔
739

740
        original = name
5✔
741
        for i in count():
5✔
742
            if name not in global_names and name not in local_names:
5✔
743
                break
5✔
744

745
            name = original + (str(i) if i else "_")
5✔
746

747
        return name
5✔
748

749
    def _enum_name_to_class_name(self, enum_name: str) -> str:
5✔
750
        """Convert a database enum name to a Python class name (PascalCase)."""
751
        return "".join(part.capitalize() for part in enum_name.split("_") if part)
5✔
752

753
    def _create_enum_class(
5✔
754
        self, table_name: str, column_name: str, values: list[str]
755
    ) -> str:
756
        """
757
        Create a Python enum class name and register it.
758

759
        Returns the enum class name to use in generated code.
760
        """
761
        # Generate enum class name from table and column names
762
        # Convert to PascalCase: user_status -> UserStatus
763
        base_name = "".join(
5✔
764
            part.capitalize()
765
            for part in table_name.split("_") + column_name.split("_")
766
            if part
767
        )
768

769
        # Ensure uniqueness
770
        enum_class_name = base_name
5✔
771
        for counter in count(1):
5✔
772
            if enum_class_name not in self.enum_values:
5✔
773
                break
5✔
774

775
            # Check if it's the same enum (same values)
776
            if self.enum_values[enum_class_name] == values:
5✔
777
                # Reuse existing enum class
778
                return enum_class_name
5✔
779

780
            enum_class_name = f"{base_name}{counter}"
5✔
781

782
        # Register the new enum class
783
        self.enum_values[enum_class_name] = values
5✔
784
        return enum_class_name
5✔
785

786
    def render_enum_classes(self) -> str:
5✔
787
        """Render Python enum class definitions."""
788
        if not self.enum_values:
5✔
789
            return ""
5✔
790

791
        self.add_module_import("enum")
5✔
792

793
        enum_defs = []
5✔
794
        for enum_class_name, values in sorted(self.enum_values.items()):
5✔
795
            # Create enum members with valid Python identifiers
796
            members = []
5✔
797
            for value in values:
5✔
798
                # Unescape SQL escape sequences (e.g., \' -> ')
799
                # The value from the CHECK constraint has SQL escaping
800
                unescaped_value = value.replace("\\'", "'").replace("\\\\", "\\")
5✔
801

802
                # Create a valid identifier from the enum value
803
                member_name = _re_invalid_identifier.sub("_", unescaped_value).upper()
5✔
804
                if not member_name:
5✔
805
                    member_name = "EMPTY"
×
806
                elif member_name[0].isdigit():
5✔
807
                    member_name = "_" + member_name
×
808
                elif iskeyword(member_name):
5✔
809
                    member_name += "_"
×
810
                #
811
                # # Re-escape for Python string literal
812
                # python_escaped = unescaped_value.replace("\\", "\\\\").replace(
813
                #     "'", "\\'"
814
                # )
815
                members.append(f"    {member_name} = {unescaped_value!r}")
5✔
816

817
            enum_def = f"class {enum_class_name}(str, enum.Enum):\n" + "\n".join(
5✔
818
                members
819
            )
820
            enum_defs.append(enum_def)
5✔
821

822
        return "\n\n\n".join(enum_defs)
5✔
823

824
    def fix_column_types(self, table: Table) -> None:
5✔
825
        """Adjust the reflected column types."""
826
        # Detect check constraints for boolean and enum columns
827
        for constraint in table.constraints.copy():
5✔
828
            if isinstance(constraint, CheckConstraint):
5✔
829
                sqltext = get_compiled_expression(constraint.sqltext, self.bind)
5✔
830

831
                # Turn any integer-like column with a CheckConstraint like
832
                # "column IN (0, 1)" into a Boolean
833
                if match := _re_boolean_check_constraint.match(sqltext):
5✔
834
                    if colname_match := _re_column_name.match(match.group(1)):
5✔
835
                        colname = colname_match.group(3)
5✔
836
                        table.constraints.remove(constraint)
5✔
837
                        table.c[colname].type = Boolean()
5✔
838
                        continue
5✔
839

840
                # Turn VARCHAR columns with CHECK constraints like "column IN ('a', 'b')"
841
                # into synthetic Enum types with Python enum classes
842
                if (
5✔
843
                    "nosyntheticenums" not in self.options
844
                    and (match := _re_enum_check_constraint.match(sqltext))
845
                    and (colname_match := _re_column_name.match(match.group(1)))
846
                ):
847
                    colname = colname_match.group(3)
5✔
848
                    items = match.group(2)
5✔
849
                    if isinstance(table.c[colname].type, String) and not isinstance(
5✔
850
                        table.c[colname].type, Enum
851
                    ):
852
                        options = _re_enum_item.findall(items)
5✔
853
                        # Create Python enum class
854
                        enum_class_name = self._create_enum_class(
5✔
855
                            table.name, colname, options
856
                        )
857
                        self.enum_classes[(table.name, colname)] = enum_class_name
5✔
858
                        # Convert to Enum type but KEEP the constraint
859
                        table.c[colname].type = Enum(*options, native_enum=False)
5✔
860
                        continue
5✔
861

862
        for column in table.c:
5✔
863
            # Handle native database Enum types (e.g., PostgreSQL ENUM)
864
            if (
5✔
865
                "nonativeenums" not in self.options
866
                and isinstance(column.type, Enum)
867
                and column.type.enums
868
            ):
869
                if column.type.name:
5✔
870
                    # Named enum - create shared enum class if not already created
871
                    if (table.name, column.name) not in self.enum_classes:
5✔
872
                        # Check if we've already created an enum for this name
873
                        existing_class = None
5✔
874
                        for (t, c), cls in self.enum_classes.items():
5✔
875
                            if cls == self._enum_name_to_class_name(column.type.name):
5✔
876
                                existing_class = cls
5✔
877
                                break
5✔
878

879
                        if existing_class:
5✔
880
                            enum_class_name = existing_class
5✔
881
                        else:
882
                            # Create new enum class from the enum's name
883
                            enum_class_name = self._enum_name_to_class_name(
5✔
884
                                column.type.name
885
                            )
886
                            # Register the enum values if not already registered
887
                            if enum_class_name not in self.enum_values:
5✔
888
                                self.enum_values[enum_class_name] = list(
5✔
889
                                    column.type.enums
890
                                )
891

892
                        self.enum_classes[(table.name, column.name)] = enum_class_name
5✔
893
                else:
894
                    # Unnamed enum - create enum class per column
895
                    if (table.name, column.name) not in self.enum_classes:
5✔
896
                        enum_class_name = self._create_enum_class(
5✔
897
                            table.name, column.name, list(column.type.enums)
898
                        )
899
                        self.enum_classes[(table.name, column.name)] = enum_class_name
5✔
900

901
            # Handle ARRAY columns with Enum item types (e.g., PostgreSQL ARRAY(ENUM))
902
            elif (
5✔
903
                "nonativeenums" not in self.options
904
                and isinstance(column.type, ARRAY)
905
                and isinstance(column.type.item_type, Enum)
906
                and column.type.item_type.enums
907
            ):
908
                enum_type = column.type.item_type
5✔
909
                if enum_type.name:
5✔
910
                    # Named enum inside ARRAY - create shared enum class
911
                    if (table.name, column.name) not in self.enum_classes:
5✔
912
                        existing_class = None
5✔
913
                        for (t, c), cls in self.enum_classes.items():
5✔
914
                            if cls == self._enum_name_to_class_name(enum_type.name):
5✔
915
                                existing_class = cls
5✔
916
                                break
5✔
917

918
                        if existing_class:
5✔
919
                            enum_class_name = existing_class
5✔
920
                        else:
921
                            enum_class_name = self._enum_name_to_class_name(
5✔
922
                                enum_type.name
923
                            )
924
                            if enum_class_name not in self.enum_values:
5✔
925
                                self.enum_values[enum_class_name] = list(
5✔
926
                                    enum_type.enums
927
                                )
928

929
                        self.enum_classes[(table.name, column.name)] = enum_class_name
5✔
930
                else:
931
                    # Unnamed enum inside ARRAY - create enum class per column
932
                    if (table.name, column.name) not in self.enum_classes:
5✔
933
                        enum_class_name = self._create_enum_class(
5✔
934
                            table.name, column.name, list(enum_type.enums)
935
                        )
936
                        self.enum_classes[(table.name, column.name)] = enum_class_name
5✔
937

938
            if not self.keep_dialect_types:
5✔
939
                try:
5✔
940
                    column.type = self.get_adapted_type(column.type)
5✔
941
                except CompileError:
5✔
942
                    continue
5✔
943

944
            # PostgreSQL specific fix: detect sequences from server_default
945
            if column.server_default and self.bind.dialect.name == "postgresql":
5✔
946
                if isinstance(column.server_default, DefaultClause) and isinstance(
5✔
947
                    column.server_default.arg, TextClause
948
                ):
949
                    schema, seqname = decode_postgresql_sequence(
5✔
950
                        column.server_default.arg
951
                    )
952
                    if seqname:
5✔
953
                        # Add an explicit sequence
954
                        if seqname != f"{column.table.name}_{column.name}_seq":
5✔
955
                            column.default = sqlalchemy.Sequence(seqname, schema=schema)
5✔
956

957
                        column.server_default = None
5✔
958

959
    def get_adapted_type(self, coltype: Any) -> Any:
5✔
960
        compiled_type = coltype.compile(self.bind.engine.dialect)
5✔
961
        for supercls in coltype.__class__.__mro__:
5✔
962
            if not supercls.__name__.startswith("_") and hasattr(
5✔
963
                supercls, "__visit_name__"
964
            ):
965
                # Don't try to adapt UserDefinedType as it's not a proper column type
966
                if supercls is UserDefinedType or issubclass(supercls, TypeDecorator):
5✔
967
                    return coltype
5✔
968

969
                # Hack to fix adaptation of the Enum class which is broken since
970
                # SQLAlchemy 1.2
971
                kw = {}
5✔
972
                if supercls is Enum:
5✔
973
                    kw["name"] = coltype.name
5✔
974
                    if coltype.schema:
5✔
975
                        kw["schema"] = coltype.schema
5✔
976

977
                # Hack to fix Postgres DOMAIN type adaptation, broken as of SQLAlchemy 2.0.42
978
                # For additional information - https://github.com/agronholm/sqlacodegen/issues/416#issuecomment-3417480599
979
                if supercls is DOMAIN:
5✔
980
                    if coltype.default:
5✔
UNCOV
981
                        kw["default"] = coltype.default
×
982
                    if coltype.constraint_name is not None:
5✔
983
                        kw["constraint_name"] = coltype.constraint_name
5✔
984
                    if coltype.not_null:
5✔
UNCOV
985
                        kw["not_null"] = coltype.not_null
×
986
                    if coltype.check is not None:
5✔
987
                        kw["check"] = coltype.check
5✔
988
                    if coltype.create_type:
5✔
989
                        kw["create_type"] = coltype.create_type
5✔
990

991
                try:
5✔
992
                    new_coltype = coltype.adapt(supercls)
5✔
993
                except TypeError:
5✔
994
                    # If the adaptation fails, don't try again
995
                    break
5✔
996

997
                for key, value in kw.items():
5✔
998
                    setattr(new_coltype, key, value)
5✔
999

1000
                if isinstance(coltype, ARRAY):
5✔
1001
                    new_coltype.item_type = self.get_adapted_type(new_coltype.item_type)
5✔
1002

1003
                try:
5✔
1004
                    # If the adapted column type does not render the same as the
1005
                    # original, don't substitute it
1006
                    if new_coltype.compile(self.bind.engine.dialect) != compiled_type:
5✔
1007
                        break
5✔
1008
                except CompileError:
5✔
1009
                    # If the adapted column type can't be compiled, don't substitute it
1010
                    break
5✔
1011

1012
                # Stop on the first valid non-uppercase column type class
1013
                coltype = new_coltype
5✔
1014
                if supercls.__name__ != supercls.__name__.upper():
5✔
1015
                    break
5✔
1016

1017
        return coltype
5✔
1018

1019

1020
class DeclarativeGenerator(TablesGenerator):
5✔
1021
    valid_options: ClassVar[set[str]] = TablesGenerator.valid_options | {
5✔
1022
        "use_inflect",
1023
        "nojoined",
1024
        "nobidi",
1025
        "noidsuffix",
1026
        "nofknames",
1027
    }
1028

1029
    def __init__(
5✔
1030
        self,
1031
        metadata: MetaData,
1032
        bind: Connection | Engine,
1033
        options: Sequence[str],
1034
        *,
1035
        indentation: str = "    ",
1036
        base_class_name: str = "Base",
1037
    ):
1038
        super().__init__(metadata, bind, options, indentation=indentation)
5✔
1039
        self.base_class_name: str = base_class_name
5✔
1040
        self.inflect_engine = inflect.engine()
5✔
1041

1042
    def generate_base(self) -> None:
5✔
1043
        self.base = Base(
5✔
1044
            literal_imports=[LiteralImport("sqlalchemy.orm", "DeclarativeBase")],
1045
            declarations=[
1046
                f"class {self.base_class_name}(DeclarativeBase):",
1047
                f"{self.indentation}pass",
1048
            ],
1049
            metadata_ref=f"{self.base_class_name}.metadata",
1050
        )
1051

1052
    def collect_imports(self, models: Iterable[Model]) -> None:
5✔
1053
        super().collect_imports(models)
5✔
1054
        if any(isinstance(model, ModelClass) for model in models):
5✔
1055
            self.add_literal_import("sqlalchemy.orm", "Mapped")
5✔
1056
            self.add_literal_import("sqlalchemy.orm", "mapped_column")
5✔
1057

1058
    def collect_imports_for_model(self, model: Model) -> None:
5✔
1059
        super().collect_imports_for_model(model)
5✔
1060
        if isinstance(model, ModelClass):
5✔
1061
            if model.relationships:
5✔
1062
                self.add_literal_import("sqlalchemy.orm", "relationship")
5✔
1063

1064
    def generate_models(self) -> list[Model]:
5✔
1065
        models_by_table_name: dict[str, Model] = {}
5✔
1066

1067
        # Pick association tables from the metadata into their own set, don't process
1068
        # them normally
1069
        links: defaultdict[str, list[Model]] = defaultdict(lambda: [])
5✔
1070
        for table in self.metadata.sorted_tables:
5✔
1071
            qualified_name = qualified_table_name(table)
5✔
1072

1073
            # Link tables have exactly two foreign key constraints and all columns are
1074
            # involved in them
1075
            fk_constraints = sorted(
5✔
1076
                table.foreign_key_constraints, key=get_constraint_sort_key
1077
            )
1078
            if len(fk_constraints) == 2 and all(
5✔
1079
                col.foreign_keys for col in table.columns
1080
            ):
1081
                model = models_by_table_name[qualified_name] = Model(table)
5✔
1082
                tablename = fk_constraints[0].elements[0].column.table.name
5✔
1083
                links[tablename].append(model)
5✔
1084
                continue
5✔
1085

1086
            # Only form model classes for tables that have a primary key and are not
1087
            # association tables
1088
            if not table.primary_key:
5✔
1089
                models_by_table_name[qualified_name] = Model(table)
5✔
1090
            else:
1091
                model = ModelClass(table)
5✔
1092
                models_by_table_name[qualified_name] = model
5✔
1093

1094
                # Fill in the columns
1095
                for column in table.c:
5✔
1096
                    column_attr = ColumnAttribute(model, column)
5✔
1097
                    model.columns.append(column_attr)
5✔
1098

1099
        # Add relationships
1100
        for model in models_by_table_name.values():
5✔
1101
            if isinstance(model, ModelClass):
5✔
1102
                self.generate_relationships(
5✔
1103
                    model, models_by_table_name, links[model.table.name]
1104
                )
1105

1106
        # Nest inherited classes in their superclasses to ensure proper ordering
1107
        if "nojoined" not in self.options:
5✔
1108
            for model in list(models_by_table_name.values()):
5✔
1109
                if not isinstance(model, ModelClass):
5✔
1110
                    continue
5✔
1111

1112
                pk_column_names = {col.name for col in model.table.primary_key.columns}
5✔
1113
                for constraint in model.table.foreign_key_constraints:
5✔
1114
                    if set(get_column_names(constraint)) == pk_column_names:
5✔
1115
                        target = models_by_table_name[
5✔
1116
                            qualified_table_name(constraint.elements[0].column.table)
1117
                        ]
1118
                        if isinstance(target, ModelClass):
5✔
1119
                            model.parent_class = target
5✔
1120
                            target.children.append(model)
5✔
1121

1122
        # Change base if we only have tables
1123
        if not any(
5✔
1124
            isinstance(model, ModelClass) for model in models_by_table_name.values()
1125
        ):
1126
            super().generate_base()
5✔
1127

1128
        # Collect the imports
1129
        self.collect_imports(models_by_table_name.values())
5✔
1130

1131
        # Rename models and their attributes that conflict with imports or other
1132
        # attributes
1133
        global_names = {
5✔
1134
            name for namespace in self.imports.values() for name in namespace
1135
        }
1136
        for model in models_by_table_name.values():
5✔
1137
            self.generate_model_name(model, global_names)
5✔
1138
            global_names.add(model.name)
5✔
1139

1140
        return list(models_by_table_name.values())
5✔
1141

1142
    def generate_relationships(
5✔
1143
        self,
1144
        source: ModelClass,
1145
        models_by_table_name: dict[str, Model],
1146
        association_tables: list[Model],
1147
    ) -> list[RelationshipAttribute]:
1148
        relationships: list[RelationshipAttribute] = []
5✔
1149
        reverse_relationship: RelationshipAttribute | None
1150

1151
        # Add many-to-one (and one-to-many) relationships
1152
        pk_column_names = {col.name for col in source.table.primary_key.columns}
5✔
1153
        for constraint in sorted(
5✔
1154
            source.table.foreign_key_constraints, key=get_constraint_sort_key
1155
        ):
1156
            target = models_by_table_name[
5✔
1157
                qualified_table_name(constraint.elements[0].column.table)
1158
            ]
1159
            if isinstance(target, ModelClass):
5✔
1160
                if "nojoined" not in self.options:
5✔
1161
                    if set(get_column_names(constraint)) == pk_column_names:
5✔
1162
                        parent = models_by_table_name[
5✔
1163
                            qualified_table_name(constraint.elements[0].column.table)
1164
                        ]
1165
                        if isinstance(parent, ModelClass):
5✔
1166
                            source.parent_class = parent
5✔
1167
                            parent.children.append(source)
5✔
1168
                            continue
5✔
1169

1170
                # Add uselist=False to One-to-One relationships
1171
                column_names = get_column_names(constraint)
5✔
1172
                if any(
5✔
1173
                    isinstance(c, (PrimaryKeyConstraint, UniqueConstraint))
1174
                    and {col.name for col in c.columns} == set(column_names)
1175
                    for c in constraint.table.constraints
1176
                ):
1177
                    r_type = RelationshipType.ONE_TO_ONE
5✔
1178
                else:
1179
                    r_type = RelationshipType.MANY_TO_ONE
5✔
1180

1181
                relationship = RelationshipAttribute(r_type, source, target, constraint)
5✔
1182
                source.relationships.append(relationship)
5✔
1183

1184
                # For self referential relationships, remote_side needs to be set
1185
                if source is target:
5✔
1186
                    relationship.remote_side = [
5✔
1187
                        source.get_column_attribute(col.name)
1188
                        for col in constraint.referred_table.primary_key
1189
                    ]
1190

1191
                # If the two tables share more than one foreign key constraint,
1192
                # SQLAlchemy needs an explicit primaryjoin to figure out which column(s)
1193
                # it needs
1194
                common_fk_constraints = get_common_fk_constraints(
5✔
1195
                    source.table, target.table
1196
                )
1197
                if len(common_fk_constraints) > 1:
5✔
1198
                    relationship.foreign_keys = [
5✔
1199
                        source.get_column_attribute(key)
1200
                        for key in constraint.column_keys
1201
                    ]
1202

1203
                # Generate the opposite end of the relationship in the target class
1204
                if "nobidi" not in self.options:
5✔
1205
                    if r_type is RelationshipType.MANY_TO_ONE:
5✔
1206
                        r_type = RelationshipType.ONE_TO_MANY
5✔
1207

1208
                    reverse_relationship = RelationshipAttribute(
5✔
1209
                        r_type,
1210
                        target,
1211
                        source,
1212
                        constraint,
1213
                        foreign_keys=relationship.foreign_keys,
1214
                        backref=relationship,
1215
                    )
1216
                    relationship.backref = reverse_relationship
5✔
1217
                    target.relationships.append(reverse_relationship)
5✔
1218

1219
                    # For self referential relationships, remote_side needs to be set
1220
                    if source is target:
5✔
1221
                        reverse_relationship.remote_side = [
5✔
1222
                            source.get_column_attribute(colname)
1223
                            for colname in constraint.column_keys
1224
                        ]
1225

1226
        # Add many-to-many relationships
1227
        for association_table in association_tables:
5✔
1228
            fk_constraints = sorted(
5✔
1229
                association_table.table.foreign_key_constraints,
1230
                key=get_constraint_sort_key,
1231
            )
1232
            target = models_by_table_name[
5✔
1233
                qualified_table_name(fk_constraints[1].elements[0].column.table)
1234
            ]
1235
            if isinstance(target, ModelClass):
5✔
1236
                relationship = RelationshipAttribute(
5✔
1237
                    RelationshipType.MANY_TO_MANY,
1238
                    source,
1239
                    target,
1240
                    fk_constraints[1],
1241
                    association_table,
1242
                )
1243
                source.relationships.append(relationship)
5✔
1244

1245
                # Generate the opposite end of the relationship in the target class
1246
                reverse_relationship = None
5✔
1247
                if "nobidi" not in self.options:
5✔
1248
                    reverse_relationship = RelationshipAttribute(
5✔
1249
                        RelationshipType.MANY_TO_MANY,
1250
                        target,
1251
                        source,
1252
                        fk_constraints[0],
1253
                        association_table,
1254
                        relationship,
1255
                    )
1256
                    relationship.backref = reverse_relationship
5✔
1257
                    target.relationships.append(reverse_relationship)
5✔
1258

1259
                # Add a primary/secondary join for self-referential many-to-many
1260
                # relationships
1261
                if source is target:
5✔
1262
                    both_relationships = [relationship]
5✔
1263
                    reverse_flags = [False, True]
5✔
1264
                    if reverse_relationship:
5✔
1265
                        both_relationships.append(reverse_relationship)
5✔
1266

1267
                    for relationship, reverse in zip(both_relationships, reverse_flags):
5✔
1268
                        if (
5✔
1269
                            not relationship.association_table
1270
                            or not relationship.constraint
1271
                        ):
UNCOV
1272
                            continue
×
1273

1274
                        constraints = sorted(
5✔
1275
                            relationship.constraint.table.foreign_key_constraints,
1276
                            key=get_constraint_sort_key,
1277
                            reverse=reverse,
1278
                        )
1279
                        pri_pairs = zip(
5✔
1280
                            get_column_names(constraints[0]), constraints[0].elements
1281
                        )
1282
                        sec_pairs = zip(
5✔
1283
                            get_column_names(constraints[1]), constraints[1].elements
1284
                        )
1285
                        relationship.primaryjoin = [
5✔
1286
                            (
1287
                                relationship.source,
1288
                                elem.column.name,
1289
                                relationship.association_table,
1290
                                col,
1291
                            )
1292
                            for col, elem in pri_pairs
1293
                        ]
1294
                        relationship.secondaryjoin = [
5✔
1295
                            (
1296
                                relationship.target,
1297
                                elem.column.name,
1298
                                relationship.association_table,
1299
                                col,
1300
                            )
1301
                            for col, elem in sec_pairs
1302
                        ]
1303

1304
        return relationships
5✔
1305

1306
    def generate_model_name(self, model: Model, global_names: set[str]) -> None:
5✔
1307
        if isinstance(model, ModelClass):
5✔
1308
            preferred_name = _re_invalid_identifier.sub("_", model.table.name)
5✔
1309
            preferred_name = "".join(
5✔
1310
                part[:1].upper() + part[1:] for part in preferred_name.split("_")
1311
            )
1312
            if "use_inflect" in self.options:
5✔
1313
                singular_name = self.inflect_engine.singular_noun(preferred_name)
5✔
1314
                if singular_name:
5✔
1315
                    preferred_name = singular_name
5✔
1316

1317
            model.name = self.find_free_name(preferred_name, global_names)
5✔
1318

1319
            # Fill in the names for column attributes
1320
            local_names: set[str] = set()
5✔
1321
            for column_attr in model.columns:
5✔
1322
                self.generate_column_attr_name(column_attr, global_names, local_names)
5✔
1323
                local_names.add(column_attr.name)
5✔
1324

1325
            # Fill in the names for relationship attributes
1326
            for relationship in model.relationships:
5✔
1327
                self.generate_relationship_name(relationship, global_names, local_names)
5✔
1328
                local_names.add(relationship.name)
5✔
1329
        else:
1330
            super().generate_model_name(model, global_names)
5✔
1331

1332
    def generate_column_attr_name(
5✔
1333
        self,
1334
        column_attr: ColumnAttribute,
1335
        global_names: set[str],
1336
        local_names: set[str],
1337
    ) -> None:
1338
        column_attr.name = self.find_free_name(
5✔
1339
            column_attr.column.name, global_names, local_names
1340
        )
1341

1342
    def generate_relationship_name(
5✔
1343
        self,
1344
        relationship: RelationshipAttribute,
1345
        global_names: set[str],
1346
        local_names: set[str],
1347
    ) -> None:
1348
        def strip_id_suffix(name: str) -> str:
5✔
1349
            # Strip _id only if at the end or followed by underscore (e.g., "course_id" -> "course", "course_id_1" -> "course_1")
1350
            # But don't strip from "parent_id1" (where id is followed by a digit without underscore)
1351
            return re.sub(r"_id(?=_|$)", "", name)
5✔
1352

1353
        def get_m2m_qualified_name(default_name: str) -> str:
5✔
1354
            """Generate qualified name for many-to-many relationship when multiple junction tables exist."""
1355
            # Check if there are multiple M2M relationships to the same target
1356
            target_m2m_relationships = [
5✔
1357
                r
1358
                for r in relationship.source.relationships
1359
                if r.target is relationship.target
1360
                and r.type == RelationshipType.MANY_TO_MANY
1361
            ]
1362

1363
            # Only use junction-based naming when there are multiple M2M to same target
1364
            if len(target_m2m_relationships) > 1:
5✔
1365
                if relationship.source is relationship.target:
5✔
1366
                    # Self-referential: use FK column name from junction table
1367
                    # (e.g., "parent_id" -> "parent", "child_id" -> "child")
1368
                    if relationship.constraint:
5✔
1369
                        column_names = [c.name for c in relationship.constraint.columns]
5✔
1370
                        if len(column_names) == 1:
5✔
1371
                            fk_qualifier = strip_id_suffix(column_names[0])
5✔
1372
                        else:
UNCOV
1373
                            fk_qualifier = "_".join(
×
1374
                                strip_id_suffix(col_name) for col_name in column_names
1375
                            )
1376
                        return fk_qualifier
5✔
1377
                elif relationship.association_table:
5✔
1378
                    # Normal: use junction table name as qualifier
1379
                    junction_name = relationship.association_table.table.name
5✔
1380
                    fk_qualifier = strip_id_suffix(junction_name)
5✔
1381
                    return f"{relationship.target.table.name}_{fk_qualifier}"
5✔
1382
            else:
1383
                # Single M2M: use simple name from junction table FK column
1384
                # (e.g., "right_id" -> "right" instead of "right_table")
1385
                if relationship.constraint and "noidsuffix" not in self.options:
5✔
1386
                    column_names = [c.name for c in relationship.constraint.columns]
5✔
1387
                    if len(column_names) == 1:
5✔
1388
                        stripped_name = strip_id_suffix(column_names[0])
5✔
1389
                        if stripped_name != column_names[0]:
5✔
1390
                            return stripped_name
5✔
1391

1392
            return default_name
5✔
1393

1394
        def get_fk_qualified_name(constraint: ForeignKeyConstraint) -> str:
5✔
1395
            """Generate qualified name for one-to-many/one-to-one relationship using FK column names."""
1396
            column_names = [c.name for c in constraint.columns]
5✔
1397

1398
            if len(column_names) == 1:
5✔
1399
                # Single column FK: strip _id suffix if present
1400
                fk_qualifier = strip_id_suffix(column_names[0])
5✔
1401
            else:
1402
                # Multi-column FK: concatenate all column names (strip _id from each)
1403
                fk_qualifier = "_".join(
5✔
1404
                    strip_id_suffix(col_name) for col_name in column_names
1405
                )
1406

1407
            # For self-referential relationships, don't prepend the table name
1408
            if relationship.source is relationship.target:
5✔
UNCOV
1409
                return fk_qualifier
×
1410
            else:
1411
                return f"{relationship.target.table.name}_{fk_qualifier}"
5✔
1412

1413
        def resolve_preferred_name() -> str:
5✔
1414
            resolved_name = relationship.target.table.name
5✔
1415

1416
            # For reverse relationships with multiple FKs to the same table, use the FK
1417
            # column name to create a more descriptive relationship name
1418
            # For M2M relationships with multiple junction tables, use the junction table name
1419
            use_fk_based_naming = "nofknames" not in self.options and (
5✔
1420
                (
1421
                    relationship.constraint
1422
                    and relationship.type
1423
                    in (RelationshipType.ONE_TO_MANY, RelationshipType.ONE_TO_ONE)
1424
                    and relationship.foreign_keys
1425
                )
1426
                or (
1427
                    relationship.type == RelationshipType.MANY_TO_MANY
1428
                    and relationship.association_table
1429
                )
1430
            )
1431

1432
            if use_fk_based_naming:
5✔
1433
                if relationship.type == RelationshipType.MANY_TO_MANY:
5✔
1434
                    resolved_name = get_m2m_qualified_name(resolved_name)
5✔
1435
                elif relationship.constraint:
5✔
1436
                    resolved_name = get_fk_qualified_name(relationship.constraint)
5✔
1437

1438
            # If there's a constraint with a single column that contains "_id", use the
1439
            # stripped version as the relationship name
1440
            elif relationship.constraint and "noidsuffix" not in self.options:
5✔
1441
                is_source = relationship.source.table is relationship.constraint.table
5✔
1442
                if is_source or relationship.type not in (
5✔
1443
                    RelationshipType.ONE_TO_ONE,
1444
                    RelationshipType.ONE_TO_MANY,
1445
                ):
1446
                    column_names = [c.name for c in relationship.constraint.columns]
5✔
1447
                    if len(column_names) == 1:
5✔
1448
                        stripped_name = strip_id_suffix(column_names[0])
5✔
1449
                        # Only use the stripped name if it actually changed (had _id in it)
1450
                        if stripped_name != column_names[0]:
5✔
1451
                            resolved_name = stripped_name
5✔
1452
                    else:
1453
                        # For composite FKs, check if there are multiple FKs to the same target
1454
                        target_relationships = [
5✔
1455
                            r
1456
                            for r in relationship.source.relationships
1457
                            if r.target is relationship.target
1458
                            and r.type == relationship.type
1459
                        ]
1460
                        if len(target_relationships) > 1:
5✔
1461
                            # Multiple FKs to same table - use concatenated column names
1462
                            resolved_name = "_".join(
5✔
1463
                                strip_id_suffix(col_name) for col_name in column_names
1464
                            )
1465

1466
            if "use_inflect" in self.options:
5✔
1467
                inflected_name: str | Literal[False]
1468
                if relationship.type in (
5✔
1469
                    RelationshipType.ONE_TO_MANY,
1470
                    RelationshipType.MANY_TO_MANY,
1471
                ):
1472
                    if not self.inflect_engine.singular_noun(resolved_name):
5✔
1473
                        resolved_name = self.inflect_engine.plural_noun(resolved_name)
5✔
1474
                else:
1475
                    inflected_name = self.inflect_engine.singular_noun(resolved_name)
5✔
1476
                    if inflected_name:
5✔
1477
                        resolved_name = inflected_name
5✔
1478

1479
            return resolved_name
5✔
1480

1481
        if (
5✔
1482
            relationship.type
1483
            in (RelationshipType.ONE_TO_MANY, RelationshipType.ONE_TO_ONE)
1484
            and relationship.source is relationship.target
1485
            and relationship.backref
1486
            and relationship.backref.name
1487
        ):
1488
            preferred_name = relationship.backref.name + "_reverse"
5✔
1489
        else:
1490
            preferred_name = resolve_preferred_name()
5✔
1491

1492
        relationship.name = self.find_free_name(
5✔
1493
            preferred_name, global_names, local_names
1494
        )
1495

1496
    def render_models(self, models: list[Model]) -> str:
5✔
1497
        rendered: list[str] = []
5✔
1498
        for model in models:
5✔
1499
            if isinstance(model, ModelClass):
5✔
1500
                rendered.append(self.render_class(model))
5✔
1501
            else:
1502
                rendered.append(f"{model.name} = {self.render_table(model.table)}")
5✔
1503

1504
        return "\n\n\n".join(rendered)
5✔
1505

1506
    def render_class(self, model: ModelClass) -> str:
5✔
1507
        sections: list[str] = []
5✔
1508

1509
        # Render class variables / special declarations
1510
        class_vars: str = self.render_class_variables(model)
5✔
1511
        if class_vars:
5✔
1512
            sections.append(class_vars)
5✔
1513

1514
        # Render column attributes
1515
        rendered_column_attributes: list[str] = []
5✔
1516
        for nullable in (False, True):
5✔
1517
            for column_attr in model.columns:
5✔
1518
                if column_attr.column.nullable is nullable:
5✔
1519
                    rendered_column_attributes.append(
5✔
1520
                        self.render_column_attribute(column_attr)
1521
                    )
1522

1523
        if rendered_column_attributes:
5✔
1524
            sections.append("\n".join(rendered_column_attributes))
5✔
1525

1526
        # Render relationship attributes
1527
        rendered_relationship_attributes: list[str] = [
5✔
1528
            self.render_relationship(relationship)
1529
            for relationship in model.relationships
1530
        ]
1531

1532
        if rendered_relationship_attributes:
5✔
1533
            sections.append("\n".join(rendered_relationship_attributes))
5✔
1534

1535
        declaration = self.render_class_declaration(model)
5✔
1536
        rendered_sections = "\n\n".join(
5✔
1537
            indent(section, self.indentation) for section in sections
1538
        )
1539
        return f"{declaration}\n{rendered_sections}"
5✔
1540

1541
    def render_class_declaration(self, model: ModelClass) -> str:
5✔
1542
        parent_class_name = (
5✔
1543
            model.parent_class.name if model.parent_class else self.base_class_name
1544
        )
1545
        return f"class {model.name}({parent_class_name}):"
5✔
1546

1547
    def render_class_variables(self, model: ModelClass) -> str:
5✔
1548
        variables = [f"__tablename__ = {model.table.name!r}"]
5✔
1549

1550
        # Render constraints and indexes as __table_args__
1551
        table_args = self.render_table_args(model.table)
5✔
1552
        if table_args:
5✔
1553
            variables.append(f"__table_args__ = {table_args}")
5✔
1554

1555
        return "\n".join(variables)
5✔
1556

1557
    def render_table_args(self, table: Table) -> str:
5✔
1558
        args: list[str] = []
5✔
1559
        kwargs: dict[str, object] = {}
5✔
1560

1561
        # Render constraints
1562
        for constraint in sorted(table.constraints, key=get_constraint_sort_key):
5✔
1563
            if uses_default_name(constraint):
5✔
1564
                if isinstance(constraint, PrimaryKeyConstraint):
5✔
1565
                    continue
5✔
1566
                if (
5✔
1567
                    isinstance(constraint, (ForeignKeyConstraint, UniqueConstraint))
1568
                    and len(constraint.columns) == 1
1569
                ):
1570
                    continue
5✔
1571

1572
            args.append(self.render_constraint(constraint))
5✔
1573

1574
        # Render indexes
1575
        for index in sorted(table.indexes, key=lambda i: cast(str, i.name)):
5✔
1576
            if len(index.columns) > 1 or not uses_default_name(index):
5✔
1577
                args.append(self.render_index(index))
5✔
1578

1579
        if table.schema:
5✔
1580
            kwargs["schema"] = table.schema
5✔
1581

1582
        if table.comment:
5✔
1583
            kwargs["comment"] = table.comment
5✔
1584

1585
        # add info + dialect kwargs for dict context (__table_args__) (opt-in)
1586
        if self.include_dialect_options_and_info:
5✔
1587
            self._add_dialect_kwargs_and_info(table, kwargs, values_for_dict=True)
5✔
1588

1589
        if kwargs:
5✔
1590
            formatted_kwargs = pformat(kwargs)
5✔
1591
            if not args:
5✔
1592
                return formatted_kwargs
5✔
1593
            else:
1594
                args.append(formatted_kwargs)
5✔
1595

1596
        if args:
5✔
1597
            rendered_args = f",\n{self.indentation}".join(args)
5✔
1598
            if len(args) == 1:
5✔
1599
                rendered_args += ","
5✔
1600

1601
            return f"(\n{self.indentation}{rendered_args}\n)"
5✔
1602
        else:
1603
            return ""
5✔
1604

1605
    def render_column_python_type(self, column: Column[Any]) -> str:
5✔
1606
        def get_type_qualifiers() -> tuple[str, TypeEngine[Any], str]:
5✔
1607
            column_type = column.type
5✔
1608
            pre: list[str] = []
5✔
1609
            post_size = 0
5✔
1610
            if column.nullable:
5✔
1611
                self.add_literal_import("typing", "Optional")
5✔
1612
                pre.append("Optional[")
5✔
1613
                post_size += 1
5✔
1614

1615
            if isinstance(column_type, ARRAY):
5✔
1616
                dim = getattr(column_type, "dimensions", None) or 1
5✔
1617
                pre.extend("list[" for _ in range(dim))
5✔
1618
                post_size += dim
5✔
1619

1620
                column_type = column_type.item_type
5✔
1621

1622
            return "".join(pre), column_type, "]" * post_size
5✔
1623

1624
        def render_python_type(column_type: TypeEngine[Any]) -> str:
5✔
1625
            # Check if this is an enum column with a Python enum class
1626
            if isinstance(column_type, Enum):
5✔
1627
                table_name = column.table.name
5✔
1628
                column_name = column.name
5✔
1629
                if (table_name, column_name) in self.enum_classes:
5✔
1630
                    enum_class_name = self.enum_classes[(table_name, column_name)]
5✔
1631
                    return enum_class_name
5✔
1632

1633
            if isinstance(column_type, DOMAIN):
5✔
1634
                column_type = column_type.data_type
5✔
1635

1636
            try:
5✔
1637
                python_type = column_type.python_type
5✔
1638
                python_type_module = python_type.__module__
5✔
1639
                python_type_name = python_type.__name__
5✔
1640
            except NotImplementedError:
5✔
1641
                self.add_literal_import("typing", "Any")
5✔
1642
                return "Any"
5✔
1643

1644
            if python_type_module == "builtins":
5✔
1645
                return python_type_name
5✔
1646

1647
            self.add_module_import(python_type_module)
5✔
1648
            return f"{python_type_module}.{python_type_name}"
5✔
1649

1650
        pre, col_type, post = get_type_qualifiers()
5✔
1651
        column_python_type = f"{pre}{render_python_type(col_type)}{post}"
5✔
1652
        return column_python_type
5✔
1653

1654
    def render_column_attribute(self, column_attr: ColumnAttribute) -> str:
5✔
1655
        column = column_attr.column
5✔
1656
        rendered_column = self.render_column(column, column_attr.name != column.name)
5✔
1657
        rendered_column_python_type = self.render_column_python_type(column)
5✔
1658

1659
        return f"{column_attr.name}: Mapped[{rendered_column_python_type}] = {rendered_column}"
5✔
1660

1661
    def render_relationship(self, relationship: RelationshipAttribute) -> str:
5✔
1662
        def render_column_attrs(column_attrs: list[ColumnAttribute]) -> str:
5✔
1663
            rendered = []
5✔
1664
            for attr in column_attrs:
5✔
1665
                if attr.model is relationship.source:
5✔
1666
                    rendered.append(attr.name)
5✔
1667
                else:
UNCOV
1668
                    rendered.append(repr(f"{attr.model.name}.{attr.name}"))
×
1669

1670
            return "[" + ", ".join(rendered) + "]"
5✔
1671

1672
        def render_foreign_keys(column_attrs: list[ColumnAttribute]) -> str:
5✔
1673
            rendered = []
5✔
1674
            render_as_string = False
5✔
1675
            # Assume that column_attrs are all in relationship.source or none
1676
            for attr in column_attrs:
5✔
1677
                if attr.model is relationship.source:
5✔
1678
                    rendered.append(attr.name)
5✔
1679
                else:
1680
                    rendered.append(f"{attr.model.name}.{attr.name}")
5✔
1681
                    render_as_string = True
5✔
1682

1683
            if render_as_string:
5✔
1684
                return "'[" + ", ".join(rendered) + "]'"
5✔
1685
            else:
1686
                return "[" + ", ".join(rendered) + "]"
5✔
1687

1688
        def render_join(terms: list[JoinType]) -> str:
5✔
1689
            rendered_joins = []
5✔
1690
            for source, source_col, target, target_col in terms:
5✔
1691
                rendered = f"lambda: {source.name}.{source_col} == {target.name}."
5✔
1692
                if target.__class__ is Model:
5✔
1693
                    rendered += "c."
5✔
1694

1695
                rendered += str(target_col)
5✔
1696
                rendered_joins.append(rendered)
5✔
1697

1698
            if len(rendered_joins) > 1:
5✔
UNCOV
1699
                rendered = ", ".join(rendered_joins)
×
UNCOV
1700
                return f"and_({rendered})"
×
1701
            else:
1702
                return rendered_joins[0]
5✔
1703

1704
        # Render keyword arguments
1705
        kwargs: dict[str, Any] = {}
5✔
1706
        if relationship.type is RelationshipType.ONE_TO_ONE and relationship.constraint:
5✔
1707
            if relationship.constraint.referred_table is relationship.source.table:
5✔
1708
                kwargs["uselist"] = False
5✔
1709

1710
        # Add the "secondary" keyword for many-to-many relationships
1711
        if relationship.association_table:
5✔
1712
            table_ref = relationship.association_table.table.name
5✔
1713
            if relationship.association_table.schema:
5✔
1714
                table_ref = f"{relationship.association_table.schema}.{table_ref}"
5✔
1715

1716
            kwargs["secondary"] = repr(table_ref)
5✔
1717

1718
        if relationship.remote_side:
5✔
1719
            kwargs["remote_side"] = render_column_attrs(relationship.remote_side)
5✔
1720

1721
        if relationship.foreign_keys:
5✔
1722
            kwargs["foreign_keys"] = render_foreign_keys(relationship.foreign_keys)
5✔
1723

1724
        if relationship.primaryjoin:
5✔
1725
            kwargs["primaryjoin"] = render_join(relationship.primaryjoin)
5✔
1726

1727
        if relationship.secondaryjoin:
5✔
1728
            kwargs["secondaryjoin"] = render_join(relationship.secondaryjoin)
5✔
1729

1730
        if relationship.backref:
5✔
1731
            kwargs["back_populates"] = repr(relationship.backref.name)
5✔
1732

1733
        rendered_relationship = render_callable(
5✔
1734
            "relationship", repr(relationship.target.name), kwargs=kwargs
1735
        )
1736

1737
        relationship_type: str
1738
        if relationship.type == RelationshipType.ONE_TO_MANY:
5✔
1739
            relationship_type = f"list['{relationship.target.name}']"
5✔
1740
        elif relationship.type in (
5✔
1741
            RelationshipType.ONE_TO_ONE,
1742
            RelationshipType.MANY_TO_ONE,
1743
        ):
1744
            relationship_type = f"'{relationship.target.name}'"
5✔
1745
            if relationship.constraint and any(
5✔
1746
                col.nullable for col in relationship.constraint.columns
1747
            ):
1748
                self.add_literal_import("typing", "Optional")
5✔
1749
                relationship_type = f"Optional[{relationship_type}]"
5✔
1750
        elif relationship.type == RelationshipType.MANY_TO_MANY:
5✔
1751
            relationship_type = f"list['{relationship.target.name}']"
5✔
1752
        else:
UNCOV
1753
            self.add_literal_import("typing", "Any")
×
UNCOV
1754
            relationship_type = "Any"
×
1755

1756
        return (
5✔
1757
            f"{relationship.name}: Mapped[{relationship_type}] "
1758
            f"= {rendered_relationship}"
1759
        )
1760

1761

1762
class DataclassGenerator(DeclarativeGenerator):
5✔
1763
    def __init__(
5✔
1764
        self,
1765
        metadata: MetaData,
1766
        bind: Connection | Engine,
1767
        options: Sequence[str],
1768
        *,
1769
        indentation: str = "    ",
1770
        base_class_name: str = "Base",
1771
        quote_annotations: bool = False,
1772
        metadata_key: str = "sa",
1773
    ):
1774
        super().__init__(
5✔
1775
            metadata,
1776
            bind,
1777
            options,
1778
            indentation=indentation,
1779
            base_class_name=base_class_name,
1780
        )
1781
        self.metadata_key: str = metadata_key
5✔
1782
        self.quote_annotations: bool = quote_annotations
5✔
1783

1784
    def generate_base(self) -> None:
5✔
1785
        self.base = Base(
5✔
1786
            literal_imports=[
1787
                LiteralImport("sqlalchemy.orm", "DeclarativeBase"),
1788
                LiteralImport("sqlalchemy.orm", "MappedAsDataclass"),
1789
            ],
1790
            declarations=[
1791
                (f"class {self.base_class_name}(MappedAsDataclass, DeclarativeBase):"),
1792
                f"{self.indentation}pass",
1793
            ],
1794
            metadata_ref=f"{self.base_class_name}.metadata",
1795
        )
1796

1797

1798
class SQLModelGenerator(DeclarativeGenerator):
5✔
1799
    def __init__(
5✔
1800
        self,
1801
        metadata: MetaData,
1802
        bind: Connection | Engine,
1803
        options: Sequence[str],
1804
        *,
1805
        indentation: str = "    ",
1806
        base_class_name: str = "SQLModel",
1807
    ):
1808
        super().__init__(
5✔
1809
            metadata,
1810
            bind,
1811
            options,
1812
            indentation=indentation,
1813
            base_class_name=base_class_name,
1814
        )
1815

1816
    @property
5✔
1817
    def views_supported(self) -> bool:
5✔
UNCOV
1818
        return False
×
1819

1820
    def render_column_callable(self, is_table: bool, *args: Any, **kwargs: Any) -> str:
5✔
1821
        self.add_import(Column)
5✔
1822
        return render_callable("Column", *args, kwargs=kwargs)
5✔
1823

1824
    def generate_base(self) -> None:
5✔
1825
        self.base = Base(
5✔
1826
            literal_imports=[],
1827
            declarations=[],
1828
            metadata_ref="",
1829
        )
1830

1831
    def collect_imports(self, models: Iterable[Model]) -> None:
5✔
1832
        super(DeclarativeGenerator, self).collect_imports(models)
5✔
1833
        if any(isinstance(model, ModelClass) for model in models):
5✔
1834
            self.remove_literal_import("sqlalchemy", "MetaData")
5✔
1835
            self.add_literal_import("sqlmodel", "SQLModel")
5✔
1836
            self.add_literal_import("sqlmodel", "Field")
5✔
1837

1838
    def collect_imports_for_model(self, model: Model) -> None:
5✔
1839
        super(DeclarativeGenerator, self).collect_imports_for_model(model)
5✔
1840
        if isinstance(model, ModelClass):
5✔
1841
            for column_attr in model.columns:
5✔
1842
                if column_attr.column.nullable:
5✔
1843
                    self.add_literal_import("typing", "Optional")
5✔
1844
                    break
5✔
1845

1846
            if model.relationships:
5✔
1847
                self.add_literal_import("sqlmodel", "Relationship")
5✔
1848

1849
    def render_module_variables(self, models: list[Model]) -> str:
5✔
1850
        declarations: list[str] = []
5✔
1851
        if any(not isinstance(model, ModelClass) for model in models):
5✔
UNCOV
1852
            if self.base.table_metadata_declaration is not None:
×
UNCOV
1853
                declarations.append(self.base.table_metadata_declaration)
×
1854

1855
        return "\n".join(declarations)
5✔
1856

1857
    def render_class_declaration(self, model: ModelClass) -> str:
5✔
1858
        if model.parent_class:
5✔
UNCOV
1859
            parent = model.parent_class.name
×
1860
        else:
1861
            parent = self.base_class_name
5✔
1862

1863
        superclass_part = f"({parent}, table=True)"
5✔
1864
        return f"class {model.name}{superclass_part}:"
5✔
1865

1866
    def render_class_variables(self, model: ModelClass) -> str:
5✔
1867
        variables = []
5✔
1868

1869
        if model.table.name != model.name.lower():
5✔
1870
            variables.append(f"__tablename__ = {model.table.name!r}")
5✔
1871

1872
        # Render constraints and indexes as __table_args__
1873
        table_args = self.render_table_args(model.table)
5✔
1874
        if table_args:
5✔
1875
            variables.append(f"__table_args__ = {table_args}")
5✔
1876

1877
        return "\n".join(variables)
5✔
1878

1879
    def render_column_attribute(self, column_attr: ColumnAttribute) -> str:
5✔
1880
        column = column_attr.column
5✔
1881
        rendered_column = self.render_column(column, True)
5✔
1882
        rendered_column_python_type = self.render_column_python_type(column)
5✔
1883

1884
        kwargs: dict[str, Any] = {}
5✔
1885
        if column.nullable:
5✔
1886
            kwargs["default"] = None
5✔
1887
        kwargs["sa_column"] = f"{rendered_column}"
5✔
1888

1889
        rendered_field = render_callable("Field", kwargs=kwargs)
5✔
1890

1891
        return f"{column_attr.name}: {rendered_column_python_type} = {rendered_field}"
5✔
1892

1893
    def render_relationship(self, relationship: RelationshipAttribute) -> str:
5✔
1894
        rendered = super().render_relationship(relationship).partition(" = ")[2]
5✔
1895
        args = self.render_relationship_args(rendered)
5✔
1896
        kwargs: dict[str, Any] = {}
5✔
1897
        annotation = repr(relationship.target.name)
5✔
1898

1899
        if relationship.type in (
5✔
1900
            RelationshipType.ONE_TO_MANY,
1901
            RelationshipType.MANY_TO_MANY,
1902
        ):
1903
            annotation = f"list[{annotation}]"
5✔
1904
        else:
1905
            self.add_literal_import("typing", "Optional")
5✔
1906
            annotation = f"Optional[{annotation}]"
5✔
1907

1908
        rendered_field = render_callable("Relationship", *args, kwargs=kwargs)
5✔
1909
        return f"{relationship.name}: {annotation} = {rendered_field}"
5✔
1910

1911
    def render_relationship_args(self, arguments: str) -> list[str]:
5✔
1912
        argument_list = arguments.split(",")
5✔
1913
        # delete ')' and ' ' from args
1914
        argument_list[-1] = argument_list[-1][:-1]
5✔
1915
        argument_list = [argument[1:] for argument in argument_list]
5✔
1916

1917
        rendered_args: list[str] = []
5✔
1918
        for arg in argument_list:
5✔
1919
            if "back_populates" in arg:
5✔
1920
                rendered_args.append(arg)
5✔
1921
            if "uselist=False" in arg:
5✔
1922
                rendered_args.append("sa_relationship_kwargs={'uselist': False}")
5✔
1923

1924
        return rendered_args
5✔
STATUS · Troubleshooting · Open an Issue · Sales · Support · CAREERS · ENTERPRISE · START FREE TRIAL · SCHEDULE DEMO
ANNOUNCEMENTS · TWITTER · TOS & SLA · Supported CI Services · What's a CI service? · Automated Testing

© 2026 Coveralls, Inc