Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -411,6 +411,17 @@ private static ImType getFirstType(ImType t) {
private static void toTupleExpressions(ImStmts body, ImTranslator translator, ImFunction f) {
Replacer replacer = new Replacer();
body.accept(new Element.DefaultVisitor() {
@Override
public void visit(ImTupleSelection selection) {
ImExpr selectedStorage = selectTupleStorageComponent(selection, translator, f);
if (selectedStorage != null) {
replacer.replace(selection, selectedStorage);
selectedStorage.accept(this);
return;
}
super.visit(selection);
}

@Override
public void visit(ImNull n) {
// Expand null<⦅T1, T2, ...⦆> ==> <null<T1>, null<T2>, ...>
Expand Down Expand Up @@ -566,6 +577,86 @@ public void visit(ImMethodCall mc) {
});
}

/**
* Select tuple storage before expanding it. Expanding first turns one array/member read into
* reads of every scalar backing variable, which then have to be preserved through discard
* calls because those reads can fail. A source-level field read only needs the selected
* backing component.
*/
private static @org.eclipse.jdt.annotation.Nullable ImExpr selectTupleStorageComponent(
ImTupleSelection selection, ImTranslator translator, ImFunction f) {
List<Integer> componentPath = new ArrayList<>();
ImExpr storage = selection;
while (storage instanceof ImTupleSelection current) {
if (current.isUsedAsLValue()) {
return null;
}
componentPath.add(current.getTupleIndex());
storage = current.getTupleExpr();
}
Collections.reverse(componentPath);

ImExpr expanded;
ImStmts prelude = JassIm.ImStmts();
if (storage instanceof ImVarAccess access && access.attrTyp() instanceof ImTupleType) {
VarsForTupleResult selected = selectTupleComponent(
translator.getVarsForTuple(access.getVar()), componentPath);
if (selected == null) {
return null;
}
expanded = selected.<ImExpr>map(
parts -> JassIm.ImTupleExpr(parts.collect(Collectors.toCollection(JassIm::ImExprs))),
JassIm::ImVarAccess);
} else if (storage instanceof ImVarArrayAccess access
&& access.attrTyp() instanceof ImTupleType) {
ImExprs indexes = captureIndexesOnceIfNeeded(access.getIndexes(), prelude, f);
VarsForTupleResult selected = selectTupleComponent(
translator.getVarsForTuple(access.getVar()), componentPath);
if (selected == null) {
return null;
}
expanded = selected.<ImExpr>map(
parts -> JassIm.ImTupleExpr(parts.collect(Collectors.toCollection(JassIm::ImExprs))),
var -> JassIm.ImVarArrayAccess(access.getTrace(), var, indexes.copy()));
} else if (storage instanceof ImMemberAccess access
&& access.attrTyp() instanceof ImTupleType) {
boolean indexesAreEffectful = access.getIndexes().stream()
.anyMatch(SideEffectAnalyzer::quickcheckHasSideeffects);
ImExpr receiver = captureOnceIfNeeded(access.getReceiver(), "tupleReceiver", prelude,
f, indexesAreEffectful);
ImExprs indexes = captureIndexesOnceIfNeeded(access.getIndexes(), prelude, f);
VarsForTupleResult selected = selectTupleComponent(
translator.getVarsForTuple(access.getVar()), componentPath);
if (selected == null) {
return null;
}
expanded = selected.<ImExpr>map(
parts -> JassIm.ImTupleExpr(parts.collect(Collectors.toCollection(JassIm::ImExprs))),
var -> JassIm.ImMemberAccess(access.getTrace(), receiver.copy(),
access.getTypeArguments().copy(), var, indexes.copy()));
} else {
return null;
}

if (prelude.isEmpty()) {
return expanded;
}
return JassIm.ImStatementExpr(prelude, expanded);
}

private static @org.eclipse.jdt.annotation.Nullable VarsForTupleResult selectTupleComponent(
VarsForTupleResult tuple, List<Integer> componentPath) {
VarsForTupleResult selected = tuple;
for (int index : componentPath) {
if (!(selected instanceof ImTranslator.TupleResult tupleResult)
|| index < 0 || index >= tupleResult.getItems().size()) {
return null;
}
selected = tupleResult.getItems().get(index);
}
return selected;
}

private static ImExpr captureOnceIfNeeded(ImExpr expr, String name, ImStmts stmts, ImFunction f,
boolean forceCapture) {
if (!forceCapture && !SideEffectAnalyzer.quickcheckHasSideeffects(expr)) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,16 @@ private String compiledLua(String testName) throws IOException {
return Files.toString(new File("test-output/lua/LuaBackendAuditTests_" + testName + ".lua"), Charsets.UTF_8);
}

private String luaFunctionBody(String compiled, String functionName) {
java.util.regex.Matcher matcher = java.util.regex.Pattern.compile(
"function " + java.util.regex.Pattern.quote(functionName)
+ "\\([^)]*\\)\\s*\\R(.*?)\\Rend",
java.util.regex.Pattern.DOTALL)
.matcher(compiled);
assertTrue("expected generated Lua function " + functionName, matcher.find());
return matcher.group(1);
}

private String compileOptimizedLua(String testName, String... lines) {
RunArgs runArgs = new RunArgs().with("-lua", "-inline", "-localOptimizations");
return compileLuaWithRunArgs(testName, runArgs, false, lines);
Expand Down Expand Up @@ -209,6 +219,69 @@ public void tuplesAreScalarizedWithoutLuaAllocations() throws IOException {
assertFalse("tuple arrays must be split into scalar arrays", compiled.contains("__wurst_arrIndex("));
}

@Test
public void tupleFieldReadsOnlyLoadTheSelectedStorageComponent() throws IOException {
String[] source = {
"package Test",
"native testSuccess()",
"tuple vec3(real x, real y, real z)",
"tuple segment(vec3 start, vec3 finish)",
"vec3 array points",
"segment array segments",
"int indexCalls",
"@noinline function nextIndex() returns int",
" indexCalls++",
" return 2",
"@noinline function readAt(int index) returns real",
" return points[index].z",
"@noinline function readAtNext() returns real",
" return points[nextIndex()].z",
"@noinline function readNested(int index) returns real",
" return segments[index].finish.y",
"class Entity",
" vec3 pos",
"@noinline function readPos(Entity entity) returns real",
" return entity.pos.z",
"init",
" points[2] = vec3(1., 2., 3.)",
" segments[2] = segment(vec3(4., 5., 6.), vec3(7., 8., 9.))",
" let entity = new Entity()",
" entity.pos = vec3(4., 5., 6.)",
" if readAt(2) == 3. and readAtNext() == 3. and indexCalls == 1",
" and readNested(2) == 8. and readPos(entity) == 6.",
" testSuccess()"
};
test().testLua(true).executeProg().lines(source);

String compiled = compileOptimizedLua(
"tupleFieldReadsOnlyLoadTheSelectedStorageComponentOptimized", source);
String plainArrayRead = luaFunctionBody(compiled, "readAt");
String effectfulArrayRead = luaFunctionBody(compiled, "readAtNext");
String nestedArrayRead = luaFunctionBody(compiled, "readNested");
String memberRead = luaFunctionBody(compiled, "readPos");

assertTrue(plainArrayRead.contains("points_z["));
assertFalse(plainArrayRead.contains("points_x["));
assertFalse(plainArrayRead.contains("points_y["));
assertTrue(effectfulArrayRead.contains("points_z["));
assertFalse(effectfulArrayRead.contains("points_x["));
assertFalse(effectfulArrayRead.contains("points_y["));
assertEquals("a selected tuple-array field must evaluate its index exactly once",
1, effectfulArrayRead.split("nextIndex\\(", -1).length - 1);
assertTrue(nestedArrayRead.contains("segments_finish_y["));
assertFalse(nestedArrayRead.contains("segments_start_"));
assertFalse(nestedArrayRead.contains("segments_finish_x["));
assertFalse(nestedArrayRead.contains("segments_finish_z["));
assertTrue(memberRead.contains("Entity_pos_z_storage["));
assertFalse(memberRead.contains("Entity_pos_x_storage["));
assertFalse(memberRead.contains("Entity_pos_y_storage["));
assertFalse("selected tuple storage reads must not need discard helpers",
plainArrayRead.contains("__wurst_tuple_discard_")
|| effectfulArrayRead.contains("__wurst_tuple_discard_")
|| nestedArrayRead.contains("__wurst_tuple_discard_")
|| memberRead.contains("__wurst_tuple_discard_"));
}

@Test
public void tupleMemberArrayAssignmentCapturesIndex() throws IOException {
test().testLua(true).luaOnly(false).executeProg().lines(
Expand Down
Loading