Skip to content

Commit 03464f2

Browse files
author
John Lyu
committed
Add multi-file model output
1 parent 42b3b39 commit 03464f2

6 files changed

Lines changed: 210 additions & 6 deletions

File tree

CHANGES.rst

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,8 @@ Version history
55

66
- Added autoincrement to primary key columns to prevent missing field errors.
77
(`#473 <https://github.com/agronholm/sqlacodegen/issues/473>`_; PR by @jtmonroe)
8+
- Added ``--output-directory`` for writing generated models into separate files.
9+
(`#88 <https://github.com/agronholm/sqlacodegen/issues/88>`_)
810

911
**4.0.3**
1012

README.rst

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -70,10 +70,15 @@ Examples::
7070
sqlacodegen postgresql:///some_local_db
7171
sqlacodegen --generator tables mysql+pymysql://user:password@localhost/dbname
7272
sqlacodegen --generator dataclasses sqlite:///database.db
73+
sqlacodegen sqlite:///database.db --output-directory models
7374
# --engine-arg values are parsed with ast.literal_eval
7475
sqlacodegen oracle+oracledb://user:pass@127.0.0.1:1521/XE --engine-arg thick_mode=True
7576
sqlacodegen oracle+oracledb://user:pass@127.0.0.1:1521/XE --engine-arg thick_mode=True --engine-arg connect_args='{"user": "user", "dsn": "..."}'
7677

78+
To write generated models into separate files, use ``--output-directory``. This creates
79+
the directory if needed and writes one Python file per generated model, plus shared
80+
support files such as ``__base.py`` and ``__init__.py``.
81+
7782
To see the list of generic options::
7883

7984
sqlacodegen --help

src/sqlacodegen/cli.py

Lines changed: 17 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55
import sys
66
from contextlib import ExitStack
77
from importlib.metadata import entry_points, version
8+
from pathlib import Path
89
from typing import Any, TextIO
910

1011
from sqlalchemy.engine import create_engine
@@ -88,7 +89,14 @@ def main() -> None:
8889
"(values are parsed with ast.literal_eval)"
8990
),
9091
)
91-
parser.add_argument("--outfile", help="file to write output to (default: stdout)")
92+
output_group = parser.add_mutually_exclusive_group()
93+
output_group.add_argument(
94+
"--outfile", help="file to write output to (default: stdout)"
95+
)
96+
output_group.add_argument(
97+
"--output-directory",
98+
help="directory to write generated models to, one file per model",
99+
)
92100
args = parser.parse_args()
93101

94102
if args.version:
@@ -133,6 +141,14 @@ def main() -> None:
133141
engine, schema, (generator.views_supported and not args.noviews), tables
134142
)
135143

144+
if args.output_directory:
145+
output_directory = Path(args.output_directory)
146+
output_directory.mkdir(parents=True, exist_ok=True)
147+
for name, contents in generator.generate(multi_file=True).items():
148+
(output_directory / f"{name}.py").write_text(contents, encoding="utf-8")
149+
150+
return
151+
136152
# Open the target file (if given)
137153
with ExitStack() as stack:
138154
outfile: TextIO

src/sqlacodegen/generators.py

Lines changed: 93 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,7 @@
1414
from keyword import iskeyword
1515
from pprint import pformat
1616
from textwrap import indent
17-
from typing import Any, ClassVar, Literal, cast
17+
from typing import Any, ClassVar, Literal, cast, overload
1818

1919
import inflect
2020
import sqlalchemy
@@ -110,8 +110,16 @@ def __init__(
110110
def views_supported(self) -> bool:
111111
pass
112112

113+
@overload
113114
@abstractmethod
114-
def generate(self) -> str:
115+
def generate(self, multi_file: Literal[False] = False) -> str: ...
116+
117+
@overload
118+
@abstractmethod
119+
def generate(self, multi_file: Literal[True]) -> dict[str, str]: ...
120+
121+
@abstractmethod
122+
def generate(self, multi_file: bool = False) -> str | dict[str, str]:
115123
"""
116124
Generate the code for the given metadata.
117125
.. note:: May modify the metadata.
@@ -167,10 +175,14 @@ def generate_base(self) -> None:
167175
metadata_ref="metadata",
168176
)
169177

170-
def generate(self) -> str:
171-
self.generate_base()
178+
@overload
179+
def generate(self, multi_file: Literal[False] = False) -> str: ...
172180

173-
sections: list[str] = []
181+
@overload
182+
def generate(self, multi_file: Literal[True]) -> dict[str, str]: ...
183+
184+
def generate(self, multi_file: bool = False) -> str | dict[str, str]:
185+
self.generate_base()
174186

175187
# Remove unwanted elements from the metadata
176188
for table in list(self.metadata.tables.values()):
@@ -199,6 +211,14 @@ def generate(self) -> str:
199211
# Generate the models
200212
models: list[Model] = self.generate_models()
201213

214+
if multi_file:
215+
return self.render_multi_file(models)
216+
else:
217+
return self.render_all(models)
218+
219+
def render_all(self, models: list[Model]) -> str:
220+
sections: list[str] = []
221+
202222
# Render module level variables
203223
if variables := self.render_module_variables(models):
204224
sections.append(variables + "\n")
@@ -220,6 +240,74 @@ def generate(self) -> str:
220240

221241
return "\n\n".join(sections) + "\n"
222242

243+
def render_multi_file(self, models: list[Model]) -> dict[str, str]:
244+
rendered_files = {
245+
"__base": self.render_base_model(),
246+
"__init__": "",
247+
}
248+
base = self.base
249+
imports = self.imports
250+
module_imports = self.module_imports
251+
252+
try:
253+
for model in models:
254+
self.imports = defaultdict(set)
255+
self.module_imports = set()
256+
self.base = self.get_base_model_import(base, model)
257+
self.collect_imports([model])
258+
rendered_files[model.name] = self.render_all([model])
259+
finally:
260+
self.base = base
261+
self.imports = imports
262+
self.module_imports = module_imports
263+
264+
return rendered_files
265+
266+
def render_base_model(self) -> str:
267+
imports = self.imports
268+
module_imports = self.module_imports
269+
try:
270+
self.imports = defaultdict(set)
271+
self.module_imports = set()
272+
for literal_import in self.base.literal_imports:
273+
self.add_literal_import(literal_import.pkgname, literal_import.name)
274+
275+
declarations = list(self.base.declarations)
276+
if self.base.table_metadata_declaration is not None:
277+
declarations.append(self.base.table_metadata_declaration)
278+
279+
sections = []
280+
groups = self.group_imports()
281+
if rendered_imports := "\n\n".join(
282+
"\n".join(line for line in group) for group in groups
283+
):
284+
sections.append(rendered_imports)
285+
286+
if declarations:
287+
sections.append("\n".join(declarations))
288+
289+
return "\n\n".join(sections) + "\n"
290+
finally:
291+
self.imports = imports
292+
self.module_imports = module_imports
293+
294+
def get_base_model_import(self, base: Base, model: Model) -> Base:
295+
literal_imports = []
296+
if base.metadata_ref == "metadata":
297+
literal_imports.append(LiteralImport(".__base", "metadata"))
298+
elif base.declarations and isinstance(model, ModelClass):
299+
literal_imports.append(
300+
LiteralImport(".__base", base.metadata_ref.partition(".")[0])
301+
)
302+
303+
return Base(
304+
literal_imports=literal_imports,
305+
declarations=[],
306+
metadata_ref=base.metadata_ref,
307+
decorator=base.decorator,
308+
table_metadata_declaration=None,
309+
)
310+
223311
def collect_imports(self, models: Iterable[Model]) -> None:
224312
for literal_import in self.base.literal_imports:
225313
self.add_literal_import(literal_import.pkgname, literal_import.name)

tests/test_cli.py

Lines changed: 45 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -87,6 +87,51 @@ class Foo(Base):
8787
)
8888

8989

90+
def test_cli_declarative_output_directory(db_path: Path, tmp_path: Path) -> None:
91+
output_path = tmp_path / "models"
92+
subprocess.run(
93+
[
94+
"sqlacodegen",
95+
f"sqlite:///{db_path}",
96+
"--generator",
97+
"declarative",
98+
"--output-directory",
99+
str(output_path),
100+
],
101+
check=True,
102+
)
103+
104+
assert sorted(path.name for path in output_path.iterdir()) == [
105+
"Foo.py",
106+
"__base.py",
107+
"__init__.py",
108+
]
109+
assert (
110+
(output_path / "__base.py").read_text()
111+
== """\
112+
from sqlalchemy.orm import DeclarativeBase
113+
114+
class Base(DeclarativeBase):
115+
pass
116+
"""
117+
)
118+
assert (output_path / "__init__.py").read_text() == ""
119+
assert (
120+
(output_path / "Foo.py").read_text()
121+
== """\
122+
from .__base import Base
123+
from sqlalchemy import Integer, Text
124+
from sqlalchemy.orm import Mapped, mapped_column
125+
126+
class Foo(Base):
127+
__tablename__ = 'foo'
128+
129+
id: Mapped[int] = mapped_column(Integer, primary_key=True)
130+
name: Mapped[str] = mapped_column(Text, nullable=False)
131+
"""
132+
)
133+
134+
90135
def test_cli_dataclass(db_path: Path, tmp_path: Path) -> None:
91136
output_path = tmp_path / "outfile"
92137
subprocess.run(

tests/test_generator_declarative.py

Lines changed: 48 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -77,6 +77,54 @@ class SimpleItems(Base):
7777
)
7878

7979

80+
def test_multi_file_generation(metadata: MetaData, engine: Engine) -> None:
81+
generator = DeclarativeGenerator(metadata, engine, [])
82+
Table(
83+
"foo",
84+
metadata,
85+
Column("id", INTEGER, primary_key=True),
86+
Column("name", Text, nullable=False),
87+
)
88+
Table(
89+
"bar",
90+
metadata,
91+
Column("id", INTEGER, primary_key=True),
92+
Column("enabled", INTEGER, nullable=False),
93+
)
94+
95+
assert generator.generate(multi_file=True) == {
96+
"__base": """\
97+
from sqlalchemy.orm import DeclarativeBase
98+
99+
class Base(DeclarativeBase):
100+
pass
101+
""",
102+
"__init__": "",
103+
"Foo": """\
104+
from .__base import Base
105+
from sqlalchemy import Integer, Text
106+
from sqlalchemy.orm import Mapped, mapped_column
107+
108+
class Foo(Base):
109+
__tablename__ = 'foo'
110+
111+
id: Mapped[int] = mapped_column(Integer, primary_key=True)
112+
name: Mapped[str] = mapped_column(Text, nullable=False)
113+
""",
114+
"Bar": """\
115+
from .__base import Base
116+
from sqlalchemy import Integer
117+
from sqlalchemy.orm import Mapped, mapped_column
118+
119+
class Bar(Base):
120+
__tablename__ = 'bar'
121+
122+
id: Mapped[int] = mapped_column(Integer, primary_key=True)
123+
enabled: Mapped[int] = mapped_column(Integer, nullable=False)
124+
""",
125+
}
126+
127+
80128
def test_index_with_kwargs(generator: CodeGenerator) -> None:
81129
simple_items = Table(
82130
"simple_items",

0 commit comments

Comments
 (0)