Skip to content

Commit 0537325

Browse files
refactor: Introduce PythonCodeGeneratorContext to manage per-namespace generation state and extract common constants.
1 parent 57ebabb commit 0537325

5 files changed

Lines changed: 135 additions & 79 deletions

File tree

src/main/java/com/regnosys/rosetta/generator/python/PythonCodeGenerator.java

Lines changed: 48 additions & 62 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,6 @@
11
package 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
@@ -13,6 +12,8 @@
1312
import com.regnosys.rosetta.generator.python.object.PythonModelObjectGenerator;
1413
import com.regnosys.rosetta.generator.python.util.PythonCodeGeneratorUtil;
1514
import com.regnosys.rosetta.generator.python.util.PythonCodeWriter;
15+
import static com.regnosys.rosetta.generator.python.util.PythonCodeGeneratorConstants.*;
16+
1617
import com.regnosys.rosetta.rosetta.RosettaEnumeration;
1718
import com.regnosys.rosetta.rosetta.RosettaModel;
1819
import com.regnosys.rosetta.rosetta.simple.Data;
@@ -25,7 +26,7 @@
2526

2627
import org.jgrapht.Graph;
2728
import org.jgrapht.graph.DefaultEdge;
28-
import org.jgrapht.graph.DirectedAcyclicGraph;
29+
2930
import org.jgrapht.traverse.TopologicalOrderIterator;
3031

3132
import 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
}
Lines changed: 48 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,48 @@
1+
package com.regnosys.rosetta.generator.python;
2+
3+
import java.util.ArrayList;
4+
import java.util.HashMap;
5+
import java.util.HashSet;
6+
import java.util.List;
7+
import java.util.Map;
8+
import java.util.Set;
9+
10+
import org.jgrapht.Graph;
11+
import org.jgrapht.graph.DefaultEdge;
12+
import org.jgrapht.graph.DirectedAcyclicGraph;
13+
14+
public class PythonCodeGeneratorContext {
15+
private List<String> subfolders = null;
16+
private Map<String, CharSequence> objects = null; // Python code for types by nameSpace, by type name
17+
private Graph<String, DefaultEdge> dependencyDAG = null;
18+
private Set<String> enumImports = null;
19+
20+
public PythonCodeGeneratorContext() {
21+
this.subfolders = new ArrayList<>();
22+
this.objects = new HashMap<>();
23+
this.dependencyDAG = new DirectedAcyclicGraph<>(DefaultEdge.class);
24+
this.enumImports = new HashSet<>();
25+
}
26+
27+
public List<String> getSubfolders() {
28+
return subfolders;
29+
}
30+
31+
public Map<String, CharSequence> getObjects() {
32+
return objects;
33+
}
34+
35+
public Graph<String, DefaultEdge> getDependencyDAG() {
36+
return dependencyDAG;
37+
}
38+
39+
public Set<String> getEnumImports() {
40+
return enumImports;
41+
}
42+
43+
public void addSubfolder(String subfolder) {
44+
if (!subfolders.contains(subfolder)) {
45+
subfolders.add(subfolder);
46+
}
47+
}
48+
}

src/main/java/com/regnosys/rosetta/generator/python/object/PythonModelObjectGenerator.java

Lines changed: 13 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -10,9 +10,7 @@
1010
import jakarta.inject.Inject;
1111
import org.jgrapht.Graph;
1212
import org.jgrapht.graph.DefaultEdge;
13-
import org.jgrapht.graph.DirectedAcyclicGraph;
1413
import org.jgrapht.graph.GraphCycleProhibitedException;
15-
import org.jgrapht.traverse.TopologicalOrderIterator;
1614

1715
import java.util.*;
1816
import java.util.stream.Collectors;
@@ -34,24 +32,15 @@ public class PythonModelObjectGenerator {
3432
@Inject
3533
private PythonChoiceAliasProcessor pythonChoiceAliasProcessor;
3634

37-
private Graph<String, DefaultEdge> dependencyDAG = null;
38-
private Set<String> enumImports = null;
39-
40-
public void beforeAllGenerate(Graph<String, DefaultEdge> dependencyDAGIn, Set<String> enumImportsIn) {
41-
dependencyDAG = dependencyDAGIn;
42-
enumImports = enumImportsIn;
43-
}
44-
4535
/**
4636
* Generate Python from the collection of Rosetta classes (of type Data).
47-
* Note: this function updates the dependency graph used by afterAllGenerate to
48-
* create the bundle
4937
*
5038
* @param rosettaClasses the collection of Rosetta Classes for this model
5139
* @param version the version for this collection of classes
5240
* @return a Map of all the generated Python indexed by the class name
5341
*/
54-
public Map<String, String> generate(Iterable<Data> rosettaClasses, String version) {
42+
public Map<String, String> generate(Iterable<Data> rosettaClasses, String version,
43+
Graph<String, DefaultEdge> dependencyDAG, Set<String> enumImports) {
5544
if (dependencyDAG == null) {
5645
throw new RuntimeException("Dependency DAG not initialized");
5746
}
@@ -65,7 +54,7 @@ public Map<String, String> generate(Iterable<Data> rosettaClasses, String versio
6554
String nameSpace = PythonCodeGeneratorUtil.getNamespace(model);
6655

6756
// Generate Python for the class
68-
String pythonClass = generateClass(rosettaClass, nameSpace, version);
57+
String pythonClass = generateClass(rosettaClass, nameSpace, version, enumImports);
6958

7059
// construct the class name using "." as a delimiter
7160
String className = model.getName() + "." + rosettaClass.getName();
@@ -76,13 +65,14 @@ public Map<String, String> generate(Iterable<Data> rosettaClasses, String versio
7665
Data superClass = rosettaClass.getSuperType();
7766
RosettaModel superModel = (RosettaModel) superClass.eContainer();
7867
String superClassName = superModel.getName() + "." + superClass.getName();
79-
addDependency(className, superClassName);
68+
69+
addDependency(dependencyDAG, className, superClassName);
8070
}
8171
}
8272
return result;
8373
}
8474

85-
private void addDependency(String className, String dependencyName) {
75+
private void addDependency(Graph<String, DefaultEdge> dependencyDAG, String className, String dependencyName) {
8676
dependencyDAG.addVertex(dependencyName);
8777
if (!className.equals(dependencyName)) {
8878
try {
@@ -93,11 +83,17 @@ private void addDependency(String className, String dependencyName) {
9383
}
9484
}
9585

96-
private String generateClass(Data rosettaClass, String nameSpace, String version) {
86+
private String generateClass(Data rosettaClass, String nameSpace, String version, Set<String> enumImports) {
87+
if (rosettaClass == null) {
88+
throw new RuntimeException("Rosetta class not initialized");
89+
}
9790
if (rosettaClass.getSuperType() != null && rosettaClass.getSuperType().getName() == null) {
9891
throw new RuntimeException(
9992
"The class superType for " + rosettaClass.getName() + " exists but its name is null");
10093
}
94+
if (enumImports == null) {
95+
throw new RuntimeException("Enum imports not initialized");
96+
}
10197

10298
Set<String> enumImportsFound = pythonAttributeProcessor.getImportsFromAttributes(rosettaClass);
10399
enumImports.addAll(enumImportsFound);
Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,12 @@
1+
package com.regnosys.rosetta.generator.python.util;
2+
3+
public final class PythonCodeGeneratorConstants {
4+
private PythonCodeGeneratorConstants() {
5+
// Restricted constructor
6+
}
7+
8+
public static final String PYTHON = "python";
9+
public static final String SRC = "src/";
10+
public static final String PYPROJECT_TOML = "pyproject.toml";
11+
public static final String INIT = "__init__";
12+
}

src/main/java/com/regnosys/rosetta/generator/python/util/PythonCodeGeneratorUtil.java

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -150,4 +150,18 @@ public static String createPYProjectTomlFile(String namespace, String version) {
150150
[tool.setuptools.packages.find]
151151
where = ["src"]""".formatted(namespace, version).stripIndent();
152152
}
153+
154+
public static String cleanVersion(String version) {
155+
if (version == null || version.equals("${project.version}")) {
156+
return "0.0.0";
157+
}
158+
159+
String[] versionParts = version.split("\\.");
160+
if (versionParts.length > 2) {
161+
String thirdPart = versionParts[2].replaceAll("[^\\d]", "");
162+
return versionParts[0] + "." + versionParts[1] + "." + thirdPart;
163+
}
164+
165+
return "0.0.0";
166+
}
153167
}

0 commit comments

Comments
 (0)