11package com .regnosys .rosetta .generator .python ;
2- // TODO: collect imports as a set rather than an array
3-
42// TODO: re-engineer type generation to use an object that has the features carried throughout the generation (imports, etc.)
3+
54// TODO: function support
65// TODO: review and consolidate unit tests
76// TODO: review migrating choice alias processor to PythonModelObjectGenerator
1312import com .regnosys .rosetta .generator .python .object .PythonModelObjectGenerator ;
1413import com .regnosys .rosetta .generator .python .util .PythonCodeGeneratorUtil ;
1514import com .regnosys .rosetta .generator .python .util .PythonCodeWriter ;
15+ import static com .regnosys .rosetta .generator .python .util .PythonCodeGeneratorConstants .*;
16+
1617import com .regnosys .rosetta .rosetta .RosettaEnumeration ;
1718import com .regnosys .rosetta .rosetta .RosettaModel ;
1819import com .regnosys .rosetta .rosetta .simple .Data ;
2526
2627import org .jgrapht .Graph ;
2728import org .jgrapht .graph .DefaultEdge ;
28- import org . jgrapht . graph . DirectedAcyclicGraph ;
29+
2930import org .jgrapht .traverse .TopologicalOrderIterator ;
3031
3132import java .util .*;
@@ -95,29 +96,35 @@ public class PythonCodeGenerator extends AbstractExternalGenerator {
9596
9697 private static final Logger LOGGER = LoggerFactory .getLogger (PythonCodeGenerator .class );
9798
98- private List <String > subfolders = null ;
99- private Map <String , Map <String , CharSequence >> objects = null ; // Python code for types by nameSpace, by type name
100- private Graph <String , DefaultEdge > dependencyDAG = null ;
101- private Set <String > enumImports = null ;
99+ private Map <String , PythonCodeGeneratorContext > contexts = null ;
102100
103101 public PythonCodeGenerator () {
104- super ("python" );
102+ super (PYTHON );
105103 }
106104
107105 @ Override
108106 public Map <String , ? extends CharSequence > beforeAllGenerate (ResourceSet set ,
109107 Collection <? extends RosettaModel > models , String version ) {
110- subfolders = new ArrayList <>();
111- objects = new HashMap <>();
112- dependencyDAG = new DirectedAcyclicGraph <>(DefaultEdge .class );
113- enumImports = new HashSet <>();
114- pojoGenerator .beforeAllGenerate (dependencyDAG , enumImports );
108+
109+ contexts = new HashMap <>();
115110 return Collections .emptyMap ();
116111 }
117112
118113 @ Override
119114 public Map <String , ? extends CharSequence > generate (Resource resource , RosettaModel model , String version ) {
120- String cleanVersion = cleanVersion (version );
115+ if (model == null ) {
116+ throw new IllegalArgumentException ("Model is null" );
117+ }
118+ LOGGER .debug ("Processing module: {}" , model .getName ());
119+
120+ String nameSpace = PythonCodeGeneratorUtil .getNamespace (model );
121+ PythonCodeGeneratorContext context = contexts .get (nameSpace );
122+ if (context == null ) {
123+ context = new PythonCodeGeneratorContext ();
124+ contexts .put (nameSpace , context );
125+ }
126+
127+ String cleanVersion = PythonCodeGeneratorUtil .cleanVersion (version );
121128
122129 Map <String , CharSequence > result = new HashMap <>();
123130
@@ -132,21 +139,15 @@ public PythonCodeGenerator() {
132139 .map (Function .class ::cast ).collect (Collectors .toList ());
133140
134141 if (!rosettaClasses .isEmpty () || !rosettaEnums .isEmpty () || !rosettaFunctions .isEmpty ()) {
135- addSubfolder (model .getName ());
142+ context . addSubfolder (model .getName ());
136143 if (!rosettaFunctions .isEmpty ()) {
137- addSubfolder (model .getName () + ".functions" );
144+ context . addSubfolder (model .getName () + ".functions" );
138145 }
139146 }
140147
141- LOGGER .debug ("Processing module: {}" , model .getName ());
142-
143- String nameSpace = PythonCodeGeneratorUtil .getNamespace (model );
144- Map <String , CharSequence > currentObject = objects .get (nameSpace );
145- if (currentObject == null ) {
146- currentObject = new HashMap <String , CharSequence >();
147- objects .put (nameSpace , currentObject );
148- }
149- currentObject .putAll (pojoGenerator .generate (rosettaClasses , cleanVersion ));
148+ Map <String , CharSequence > currentObject = context .getObjects ();
149+ currentObject .putAll (pojoGenerator .generate (rosettaClasses , cleanVersion , context .getDependencyDAG (),
150+ context .getEnumImports ()));
150151 result .putAll (enumGenerator .generate (rosettaEnums , cleanVersion ));
151152 result .putAll (funcGenerator .generate (rosettaFunctions , cleanVersion ));
152153
@@ -158,26 +159,30 @@ public PythonCodeGenerator() {
158159 ResourceSet set ,
159160 Collection <? extends RosettaModel > models ,
160161 String version ) {
161- String cleanVersion = cleanVersion (version );
162162 Map <String , CharSequence > result = new HashMap <>();
163-
164- List <String > workspaces = getWorkspaces (subfolders );
165- result .putAll (generateWorkspaces (workspaces , cleanVersion ));
166- result .putAll (generateInits (subfolders ));
167-
168- for (String nameSpace : objects .keySet ()) {
169- Map <String , CharSequence > currentObject = objects .get (nameSpace );
170- if (currentObject != null && !currentObject .isEmpty ()) {
171- result .put ("pyproject.toml" , PythonCodeGeneratorUtil .createPYProjectTomlFile (nameSpace , cleanVersion ));
172- result .putAll (processDAG (nameSpace , currentObject ));
173- }
163+ String cleanVersion = PythonCodeGeneratorUtil .cleanVersion (version );
164+ for (String nameSpace : contexts .keySet ()) {
165+ PythonCodeGeneratorContext context = contexts .get (nameSpace );
166+ List <String > subfolders = context .getSubfolders ();
167+ result .putAll (generateWorkspaces (getWorkspaces (subfolders ), cleanVersion ));
168+ result .putAll (generateInits (subfolders ));
169+ result .putAll (processDAG (nameSpace , context , cleanVersion ));
174170 }
175171 return result ;
176172 }
177173
178- private Map <String , CharSequence > processDAG (String nameSpace , Map <String , CharSequence > nameSpaceObjects ) {
174+ private Map <String , CharSequence > processDAG (String nameSpace , PythonCodeGeneratorContext context ,
175+ String cleanVersion ) {
176+ if (nameSpace == null || context == null || cleanVersion == null ) {
177+ throw new IllegalArgumentException ("Invalid arguments" );
178+ }
179179 Map <String , CharSequence > result = new HashMap <>();
180- if (dependencyDAG != null ) {
180+ Map <String , CharSequence > nameSpaceObjects = context .getObjects ();
181+ Graph <String , DefaultEdge > dependencyDAG = context .getDependencyDAG ();
182+ Set <String > enumImports = context .getEnumImports ();
183+
184+ if (nameSpaceObjects != null && !nameSpaceObjects .isEmpty () && dependencyDAG != null && enumImports != null ) {
185+ result .put (PYPROJECT_TOML , PythonCodeGeneratorUtil .createPYProjectTomlFile (nameSpace , cleanVersion ));
181186 PythonCodeWriter bundleWriter = new PythonCodeWriter ();
182187 TopologicalOrderIterator <String , DefaultEdge > topologicalOrderIterator = new TopologicalOrderIterator <>(
183188 dependencyDAG );
@@ -205,7 +210,7 @@ private Map<String, CharSequence> processDAG(String nameSpace, Map<String, CharS
205210
206211 // create the stub
207212 String [] parsedName = name .split ("\\ ." );
208- String stubFileName = "src/" + String .join ("/" , parsedName ) + ".py" ;
213+ String stubFileName = SRC + String .join ("/" , parsedName ) + ".py" ;
209214
210215 PythonCodeWriter stubWriter = new PythonCodeWriter ();
211216 stubWriter .appendLine ("# pylint: disable=unused-import" );
@@ -224,25 +229,11 @@ private Map<String, CharSequence> processDAG(String nameSpace, Map<String, CharS
224229 }
225230 bundleWriter .newLine ();
226231 bundleWriter .appendLine ("# EOF" );
227- result .put ("src/" + nameSpace + "/_bundle.py" , bundleWriter .toString ());
232+ result .put (SRC + nameSpace + "/_bundle.py" , bundleWriter .toString ());
228233 }
229234 return result ;
230235 }
231236
232- private String cleanVersion (String version ) {
233- if (version == null || version .equals ("${project.version}" )) {
234- return "0.0.0" ;
235- }
236-
237- String [] versionParts = version .split ("\\ ." );
238- if (versionParts .length > 2 ) {
239- String thirdPart = versionParts [2 ].replaceAll ("[^\\ d]" , "" );
240- return versionParts [0 ] + "." + versionParts [1 ] + "." + thirdPart ;
241- }
242-
243- return "0.0.0" ;
244- }
245-
246237 private List <String > getWorkspaces (List <String > subfolders ) {
247238 return subfolders .stream ().map (subfolder -> subfolder .split ("\\ ." )[0 ]).distinct ().collect (Collectors .toList ());
248239 }
@@ -251,7 +242,7 @@ private Map<String, String> generateWorkspaces(List<String> workspaces, String v
251242 Map <String , String > result = new HashMap <>();
252243
253244 for (String workspace : workspaces ) {
254- result .put (PythonCodeGeneratorUtil .toPyFileName (workspace , "__init__" ),
245+ result .put (PythonCodeGeneratorUtil .toPyFileName (workspace , INIT ),
255246 PythonCodeGeneratorUtil .createTopLevelInitFile (version ));
256247 result .put (PythonCodeGeneratorUtil .toPyFileName (workspace , "version" ),
257248 PythonCodeGeneratorUtil .createVersionFile (version ));
@@ -268,16 +259,11 @@ private Map<String, String> generateInits(List<String> subfolders) {
268259 String [] parts = subfolder .split ("\\ ." );
269260 for (int i = 1 ; i < parts .length ; i ++) {
270261 String key = String .join ("." , Arrays .copyOfRange (parts , 0 , i + 1 ));
271- result .putIfAbsent (PythonCodeGeneratorUtil .toPyFileName (key , "__init__" ), " " );
262+ result .putIfAbsent (PythonCodeGeneratorUtil .toPyFileName (key , INIT ), " " );
272263 }
273264 }
274265
275266 return result ;
276267 }
277268
278- private void addSubfolder (String subfolder ) {
279- if (!subfolders .contains (subfolder )) {
280- subfolders .add (subfolder );
281- }
282- }
283269}
0 commit comments