Skip to content

Commit ff7eadc

Browse files
authored
Deduplicate Lua callback adapters (#1290)
* Deduplicate Lua callback adapters * Preserve renamed Lua callback targets
1 parent 5ebbbdc commit ff7eadc

3 files changed

Lines changed: 140 additions & 37 deletions

File tree

de.peeeq.wurstscript/src/main/java/de/peeeq/wurstscript/translation/lua/translation/ExprTranslation.java

Lines changed: 1 addition & 36 deletions
Original file line numberDiff line numberDiff line change
@@ -53,42 +53,7 @@ public static LuaExpr translate(ImDealloc e, LuaTranslator tr) {
5353
}
5454

5555
public static LuaExpr translate(ImFuncRef e, LuaTranslator tr) {
56-
// return LuaAst.LuaExprFuncRef(tr.luaFunc.getFor(e.getFunc()));
57-
// alternative: use xpcall to get stacktraces (did not work)
58-
boolean returnsValue = !(e.getFunc().getReturnType() instanceof ImVoid);
59-
LuaVariable dots = LuaAst.LuaVariable("...", LuaAst.LuaNoExpr());
60-
LuaStatements callbackBody = LuaAst.LuaStatements();
61-
if (returnsValue) {
62-
LuaVariable tempRes = LuaAst.LuaVariable("tempRes", LuaAst.LuaExprNull());
63-
callbackBody.add(tempRes);
64-
callbackBody.add(LuaAst.LuaExprFunctionCallByName("xpcall",
65-
LuaAst.LuaExprlist(
66-
LuaAst.LuaExprFunctionAbstraction(
67-
LuaAst.LuaParams(dots.copy()),
68-
LuaAst.LuaStatements(
69-
LuaAst.LuaAssignment(LuaAst.LuaExprVarAccess(tempRes),
70-
LuaAst.LuaExprFunctionCall(tr.luaFunc.getFor(e.getFunc()), LuaAst.LuaExprlist(LuaAst.LuaExprVarAccess(dots.copy())))))
71-
),
72-
LuaAst.LuaLiteral("function(err) if err == \"" + WURST_ABORT_THREAD_SENTINEL + "\" then return end BJDebugMsg(\"lua callback error: \" .. tostring(err)) xpcall(function() " + callErrorFunc(tr, "tostring(err)", "in lua callback error handler") + " end, function(err2) if err2 == \"" + WURST_ABORT_THREAD_SENTINEL + "\" then return end BJDebugMsg(\"error reporting error: \" .. tostring(err2)) BJDebugMsg(\"while reporting: \" .. tostring(err)) end) end"),
73-
LuaAst.LuaExprVarAccess(dots.copy())
74-
)
75-
));
76-
callbackBody.add(LuaAst.LuaReturn(LuaAst.LuaExprVarAccess(tempRes)));
77-
} else {
78-
callbackBody.add(LuaAst.LuaExprFunctionCallByName("xpcall",
79-
LuaAst.LuaExprlist(
80-
LuaAst.LuaExprFunctionAbstraction(
81-
LuaAst.LuaParams(dots.copy()),
82-
LuaAst.LuaStatements(
83-
LuaAst.LuaExprFunctionCall(tr.luaFunc.getFor(e.getFunc()), LuaAst.LuaExprlist(LuaAst.LuaExprVarAccess(dots.copy())))
84-
)
85-
),
86-
LuaAst.LuaLiteral("function(err) if err == \"" + WURST_ABORT_THREAD_SENTINEL + "\" then return end BJDebugMsg(\"lua callback error: \" .. tostring(err)) xpcall(function() " + callErrorFunc(tr, "tostring(err)", "in lua callback error handler") + " end, function(err2) if err2 == \"" + WURST_ABORT_THREAD_SENTINEL + "\" then return end BJDebugMsg(\"error reporting error: \" .. tostring(err2)) BJDebugMsg(\"while reporting: \" .. tostring(err)) end) end"),
87-
LuaAst.LuaExprVarAccess(dots.copy())
88-
)
89-
));
90-
}
91-
return LuaAst.LuaExprFunctionAbstraction(LuaAst.LuaParams(dots), callbackBody);
56+
return LuaAst.LuaExprFuncRef(tr.callbackAdapterFor(e.getFunc()));
9257
}
9358

9459
static String callErrorFunc(LuaTranslator tr, String msg) {

de.peeeq.wurstscript/src/main/java/de/peeeq/wurstscript/translation/lua/translation/LuaTranslator.java

Lines changed: 64 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -137,6 +137,8 @@ private ImProg getProg() {
137137

138138
List<ExprTranslation.TupleFunc> tupleEqualsFuncs = new ArrayList<>();
139139
List<ExprTranslation.TupleFunc> tupleCopyFuncs = new ArrayList<>();
140+
private final Map<ImFunction, LuaFunction> callbackAdapters = new IdentityHashMap<>();
141+
private LuaFunction callbackErrorHandler;
140142

141143
// Array-default infrastructure (metatables/helper functions) shared across
142144
// every array of a given entry type, instead of allocated per array
@@ -385,6 +387,68 @@ public LuaCompilationUnit translate() {
385387
return luaModel;
386388
}
387389

390+
/**
391+
* Function references need an xpcall boundary, but the boundary is a property of the referenced
392+
* function rather than of each expression which names it. Emit one reusable adapter per target
393+
* so evaluating a function reference performs no closure allocation.
394+
*/
395+
LuaFunction callbackAdapterFor(ImFunction target) {
396+
LuaFunction existing = callbackAdapters.get(target);
397+
if (existing != null) {
398+
return existing;
399+
}
400+
401+
LuaFunction targetLua = luaFunc.getFor(target);
402+
LuaVariable dots = LuaAst.LuaVariable("...", LuaAst.LuaNoExpr());
403+
LuaFunction adapter = LuaAst.LuaFunction(
404+
uniqueName("__wurst_callback_" + targetLua.getName()),
405+
LuaAst.LuaParams(dots), LuaAst.LuaStatements());
406+
callbackAdapters.put(target, adapter);
407+
408+
LuaFunction errorHandler = callbackErrorHandler();
409+
LuaExprFunctionCallByName xpcall = LuaAst.LuaExprFunctionCallByName("xpcall",
410+
LuaAst.LuaExprlist(
411+
LuaAst.LuaExprFuncRef(targetLua),
412+
LuaAst.LuaExprFuncRef(errorHandler),
413+
LuaAst.LuaExprVarAccess(dots.copy())));
414+
if (target.getReturnType() instanceof ImVoid) {
415+
adapter.getBody().add(xpcall);
416+
} else {
417+
// Keep exactly the first callback result. Returning select(2, xpcall(...)) directly
418+
// could leak additional Lua return values into a surrounding argument list.
419+
LuaVariable ignored = LuaAst.LuaVariable("_", LuaAst.LuaNoExpr());
420+
LuaVariable result = LuaAst.LuaVariable("result", LuaAst.LuaNoExpr());
421+
adapter.getBody().add(ignored);
422+
adapter.getBody().add(result);
423+
adapter.getBody().add(LuaAst.LuaAssignment(
424+
LuaAst.LuaLiteral("_, result"), xpcall));
425+
adapter.getBody().add(LuaAst.LuaReturn(LuaAst.LuaExprVarAccess(result)));
426+
}
427+
luaModel.add(adapter);
428+
return adapter;
429+
}
430+
431+
private LuaFunction callbackErrorHandler() {
432+
if (callbackErrorHandler != null) {
433+
return callbackErrorHandler;
434+
}
435+
LuaVariable err = LuaAst.LuaVariable("err", LuaAst.LuaNoExpr());
436+
callbackErrorHandler = LuaAst.LuaFunction(uniqueName("__wurst_callback_error"),
437+
LuaAst.LuaParams(err), LuaAst.LuaStatements());
438+
callbackErrorHandler.getBody().add(LuaAst.LuaLiteral(
439+
"if err == \"" + ExprTranslation.WURST_ABORT_THREAD_SENTINEL + "\" then return end"));
440+
callbackErrorHandler.getBody().add(LuaAst.LuaLiteral(
441+
"BJDebugMsg(\"lua callback error: \" .. tostring(err))"));
442+
callbackErrorHandler.getBody().add(LuaAst.LuaLiteral(
443+
"xpcall(function() " + ExprTranslation.callErrorFunc(this, "tostring(err)",
444+
"in lua callback error handler")
445+
+ " end, function(err2) if err2 == \"" + ExprTranslation.WURST_ABORT_THREAD_SENTINEL
446+
+ "\" then return end BJDebugMsg(\"error reporting error: \" .. tostring(err2))"
447+
+ " BJDebugMsg(\"while reporting: \" .. tostring(err)) end)"));
448+
luaModel.add(callbackErrorHandler);
449+
return callbackErrorHandler;
450+
}
451+
388452
/**
389453
* Rejects calls/references to functions that an earlier optimizer pass
390454
* detached from the IM program. Without this invariant the Lua printer

de.peeeq.wurstscript/src/test/java/tests/wurstscript/tests/LuaTranslationTests.java

Lines changed: 75 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -128,6 +128,15 @@ private List<String> uniqueMatches(String output, String regex, int group) {
128128
return result;
129129
}
130130

131+
private int countMatches(String output, String regex) {
132+
Matcher matcher = Pattern.compile(regex).matcher(output);
133+
int result = 0;
134+
while (matcher.find()) {
135+
result++;
136+
}
137+
return result;
138+
}
139+
131140
private String singleMatch(String output, String regex, int group) {
132141
Matcher matcher = Pattern.compile(regex).matcher(output);
133142
assertTrue("Expected pattern to occur: " + regex, matcher.find());
@@ -173,6 +182,11 @@ private String compileLuaWithRunArgs(String testName, boolean withStdLib, String
173182

174183
private String compileLuaWithCUs(String testName, boolean withStdLib, List<CU> extraCUs, String... lines) {
175184
RunArgs runArgs = new RunArgs().with("-lua", "-inline", "-localOptimizations", "-stacktraces");
185+
return compileLuaWithCUs(testName, withStdLib, extraCUs, runArgs, lines);
186+
}
187+
188+
private String compileLuaWithCUs(String testName, boolean withStdLib, List<CU> extraCUs,
189+
RunArgs runArgs, String... lines) {
176190
WurstGui gui = new WurstGuiCliImpl();
177191
WurstCompilerJassImpl compiler = new WurstCompilerJassImpl(null, gui, null, runArgs);
178192
List<CU> inputs = new ArrayList<>();
@@ -2475,12 +2489,72 @@ public void luaFunctionRefWrapperForwardsVarargs() throws IOException {
24752489
" ForForce(f, () -> skip)"
24762490
);
24772491
String compiled = Files.toString(new File("test-output/lua/LuaTranslationTests_luaFunctionRefWrapperForwardsVarargs.lua"), Charsets.UTF_8);
2478-
assertTrue(compiled.contains("xpcall(function (...)"));
2492+
assertContainsRegex(compiled, "function\\s+__wurst_callback_[A-Za-z0-9_]+\\(\\.\\.\\.\\)");
2493+
assertFalse(compiled.contains("xpcall(function (...)"));
2494+
assertContainsRegex(compiled,
2495+
"xpcall\\([A-Za-z0-9_]+, __wurst_callback_error[A-Za-z0-9_]*, \\.\\.\\.\\)");
24792496
assertTrue(compiled.contains(", ...)"));
24802497
assertFalse(compiled.contains("local temp = ..."));
24812498
assertFalse(compiled.contains("ForForce(f, function (...) \n\t\t\tlocal tempRes"));
24822499
}
24832500

2501+
@Test
2502+
public void luaFunctionRefsReuseOneAdapterAndPreserveSingleReturn() {
2503+
String compiled = compileLuaWithCUs(
2504+
"LuaTranslationTests_luaFunctionRefsReuseOneAdapterAndPreserveSingleReturn",
2505+
false,
2506+
Collections.emptyList(),
2507+
new RunArgs().with("-lua", "-inline", "-localOptimizations"),
2508+
"type boolexpr extends handle",
2509+
"package Test",
2510+
"@extern native Condition(code callback) returns boolexpr",
2511+
"function predicate() returns boolean",
2512+
" return true",
2513+
"init",
2514+
" let first = Condition(function predicate)",
2515+
" let second = Condition(function predicate)"
2516+
);
2517+
2518+
List<String> adapters = uniqueMatches(compiled,
2519+
"function\\s+(__wurst_callback_predicate[A-Za-z0-9_]*)\\(\\.\\.\\.\\)", 1);
2520+
assertEquals("one adapter must serve every reference to the same function:\n" + compiled,
2521+
1, adapters.size());
2522+
String adapter = adapters.get(0);
2523+
assertEquals("both Condition calls must reference the cached adapter", 2,
2524+
countMatches(compiled, "Condition\\(" + Pattern.quote(adapter) + "\\)"));
2525+
String adapterBody = getFunctionBody(compiled, adapter);
2526+
assertTrue(adapterBody.contains("_, result = xpcall(predicate,"));
2527+
assertTrue(adapterBody.contains("return result"));
2528+
assertFalse("callback sites must not allocate anonymous wrappers", compiled.contains("Condition(function ("));
2529+
}
2530+
2531+
@Test
2532+
public void luaFunctionRefAdapterTracksLateClassFunctionRename() {
2533+
String compiled = compileLuaWithCUs(
2534+
"LuaTranslationTests_luaFunctionRefAdapterTracksLateClassFunctionRename",
2535+
false,
2536+
Collections.emptyList(),
2537+
new RunArgs().with("-lua"),
2538+
"package Test",
2539+
"@extern native consume(code callback)",
2540+
"@extern native CallbackOwner_staticCallback()",
2541+
"class CallbackOwner",
2542+
" function start()",
2543+
" consume(function staticCallback)",
2544+
" private static function staticCallback()",
2545+
" consume(function staticCallback)",
2546+
"init",
2547+
" CallbackOwner_staticCallback()",
2548+
" new CallbackOwner().start()"
2549+
);
2550+
2551+
String callbackName = singleMatch(compiled,
2552+
"function\\s+(CallbackOwner_[A-Za-z0-9_]*staticCallback[A-Za-z0-9_]*)\\(\\)", 1);
2553+
assertTrue("adapter must track the class callback's final name:\n" + compiled,
2554+
compiled.contains("xpcall(" + callbackName + ","));
2555+
assertFalse(compiled.contains("xpcall(staticCallback,"));
2556+
}
2557+
24842558
@Test
24852559
public void luaFunctionRefStacktraceHandlerUsesWurstStackPosition() throws IOException {
24862560
CU errorHandling = new CU("ErrorHandling.wurst", String.join("\n",

0 commit comments

Comments
 (0)