1414from keyword import iskeyword
1515from pprint import pformat
1616from textwrap import indent
17- from typing import Any , ClassVar , Literal , cast
17+ from typing import Any , ClassVar , Literal , cast , overload
1818
1919import inflect
2020import 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 )
0 commit comments