diff --git a/redline/api/src/main/java/run/endive/redline/experimental/api/internal/InterruptWatchdog.java b/redline/api/src/main/java/run/endive/redline/experimental/api/internal/InterruptWatchdog.java new file mode 100644 index 000000000..330c54e87 --- /dev/null +++ b/redline/api/src/main/java/run/endive/redline/experimental/api/internal/InterruptWatchdog.java @@ -0,0 +1,180 @@ +package run.endive.redline.experimental.api.internal; + +import java.util.Iterator; +import java.util.Set; +import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.locks.LockSupport; +import java.util.logging.Level; +import java.util.logging.Logger; + +/** Raises {@link CtxBuffer#INTERRUPT_FLAG} for watched calls whose thread is interrupted. */ +public final class InterruptWatchdog { + + private static final Logger LOG = Logger.getLogger(InterruptWatchdog.class.getName()); + + private static final long POLL_INTERVAL_NANOS = + millisProperty("endive.redline.interruptPollMillis", 100); + + // how long the poller waits without calls before it exits + private static final long IDLE_EXIT_NANOS = + millisProperty("endive.redline.interruptIdleMillis", 60_000); + + // STATE holds the RUNNING and IDLE flags of the poller plus CALL per watched call + private static final int RUNNING = 1; + + private static final int IDLE = 2; + + private static final int CALL = 4; + + private static final AtomicInteger STATE = new AtomicInteger(); + + private static final Set ACTIVE = ConcurrentHashMap.newKeySet(); + + private static volatile Thread poller; + + private InterruptWatchdog() {} + + /** Raises the interrupt flag in a machine's context buffer. */ + @FunctionalInterface + public interface InterruptSink { + void requestInterrupt(); + } + + /** Watches the current thread until the returned handle is passed to {@link #exit}. */ + public static Registration enter(InterruptSink sink) { + var registration = new Registration(Thread.currentThread(), sink); + ACTIVE.add(registration); + int state = STATE.getAndAdd(CALL); + try { + if ((state & RUNNING) == 0) { + ensurePoller(); + } else if ((state & IDLE) != 0 && (STATE.getAndUpdate(s -> s & ~IDLE) & IDLE) != 0) { + LockSupport.unpark(poller); + } + } catch (RuntimeException | Error e) { + exit(registration); + throw e; + } + return registration; + } + + /** Stops watching; once this returns the poller can no longer raise the flag. */ + public static void exit(Registration registration) { + registration.deactivate(); + ACTIVE.remove(registration); + STATE.getAndAdd(-CALL); + } + + private static long millisProperty(String name, long defaultMillis) { + return TimeUnit.MILLISECONDS.toNanos(Math.max(1, Long.getLong(name, defaultMillis))); + } + + @SuppressWarnings("ThreadPriorityCheck") // not the priority of whichever caller starts it + private static void ensurePoller() { + if ((STATE.getAndUpdate(s -> s | RUNNING) & RUNNING) != 0) { + return; + } + try { + // no thread locals or class loader from the caller either + var thread = + new Thread( + null, + InterruptWatchdog::pollLoop, + "endive-redline-interrupt", + 0, + false); + thread.setDaemon(true); + thread.setContextClassLoader(null); + thread.setPriority(Thread.NORM_PRIORITY); + poller = thread; + thread.start(); + } catch (RuntimeException | Error e) { + STATE.getAndAdd(-RUNNING); + throw e; + } + } + + private static void pollLoop() { + boolean retired = false; + try { + while (true) { + // an interrupt status would make every park return at once + Thread.interrupted(); + if (STATE.get() < CALL && idleUntilExit()) { + retired = true; + return; + } + pollAll(); + LockSupport.parkNanos(POLL_INTERVAL_NANOS); + } + } finally { + if (!retired && STATE.updateAndGet(s -> s & ~(RUNNING | IDLE)) >= CALL) { + // died on an error: hand the watched calls to a new poller + ensurePoller(); + } + } + } + + // parks until a call clears IDLE or the idle time passes; true if the poller should exit + private static boolean idleUntilExit() { + if (!STATE.compareAndSet(RUNNING, RUNNING | IDLE)) { + return false; + } + long deadline = System.nanoTime() + IDLE_EXIT_NANOS; + long left; + while ((STATE.get() & IDLE) != 0 && (left = deadline - System.nanoTime()) > 0) { + LockSupport.parkNanos(left); + Thread.interrupted(); + } + if (STATE.compareAndSet(RUNNING | IDLE, 0)) { + return true; + } + STATE.getAndUpdate(s -> s & ~IDLE); + return false; + } + + private static void pollAll() { + for (Iterator it = ACTIVE.iterator(); it.hasNext(); ) { + var registration = it.next(); + try { + registration.poll(); + } catch (RuntimeException | Error e) { + // a broken sink: stop watching that call rather than every call + it.remove(); + LOG.log(Level.WARNING, "Stopped watching a call whose interrupt flag failed", e); + } + } + } + + /** One watched call. */ + public static final class Registration { + + private final Thread caller; + private final InterruptSink sink; + private boolean active = true; + + private Registration(Thread caller, InterruptSink sink) { + this.caller = caller; + this.sink = sink; + } + + // the lock pairs with deactivate(), so the flag is never raised after exit() + private void poll() { + if (caller.isInterrupted()) { + synchronized (this) { + if (active) { + sink.requestInterrupt(); + } + } + } + } + + private void deactivate() { + synchronized (this) { + active = false; + } + } + } +} diff --git a/redline/runner-jffi-tests/pom.xml b/redline/runner-jffi-tests/pom.xml index e84f52570..6e7ce7dce 100644 --- a/redline/runner-jffi-tests/pom.xml +++ b/redline/runner-jffi-tests/pom.xml @@ -78,6 +78,16 @@ + + org.apache.maven.plugins + maven-surefire-plugin + + + + 200 + + + run.endive test-gen-plugin diff --git a/redline/runner-jffi-tests/src/test/java/run/endive/redline/experimental/runner/jffi/internal/InterruptFlagTest.java b/redline/runner-jffi-tests/src/test/java/run/endive/redline/experimental/runner/jffi/internal/InterruptFlagTest.java index 55c8fdf3c..00f2e9af2 100644 --- a/redline/runner-jffi-tests/src/test/java/run/endive/redline/experimental/runner/jffi/internal/InterruptFlagTest.java +++ b/redline/runner-jffi-tests/src/test/java/run/endive/redline/experimental/runner/jffi/internal/InterruptFlagTest.java @@ -1,24 +1,20 @@ package run.endive.redline.experimental.runner.jffi.internal; +import static java.util.concurrent.TimeUnit.MILLISECONDS; import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertFalse; +import java.util.List; import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.Test; import run.endive.corpus.CorpusResources; -import run.endive.redline.experimental.api.internal.RedlineTarget; -import run.endive.redline.experimental.compiler.internal.NativeCompiler; -import run.endive.redline.experimental.runner.jffi.JffiNativeMachineFactory; import run.endive.runtime.HostFunction; import run.endive.runtime.ImportValues; +import run.endive.testing.NativeInstanceBuilder; import run.endive.wasm.Parser; import run.endive.wasm.types.FunctionType; -/** - * The watchdog raises the interrupt flag from another thread, so it can land after - * the call it was meant to stop has passed its last check. The flag must not then - * sit in the context and stop a later call that nobody interrupted. - */ +/** An interrupt the host handled during a call must not stop a later call. */ public class InterruptFlagTest { @AfterEach @@ -28,45 +24,80 @@ public void clearInterruptStatus() { } @Test - public void aFlagRaisedMidCallDoesNotStopTheNextCall() { + public void anInterruptHandledByTheHostDoesNotStopTheNextCall() { var module = Parser.parse(CorpusResources.getResource("compiled/interrupt-midcall.wat.wasm")); - - var machineRef = new JffiNativeMachine[1]; var imports = ImportValues.builder() .addFunction( new HostFunction( "host", "raiseFlag", - FunctionType.of(java.util.List.of(), java.util.List.of()), + FunctionType.of(List.of(), List.of()), (inst, args) -> { - machineRef[0].requestInterrupt(); + interruptAndHandle(); return null; })) .build(); try (var instance = - JffiNativeMachineFactory.builder(module) - .withImportValues(imports) - .withCompilerFunction( - m -> - NativeCompiler.compileAll( - RedlineTarget.detectHost().orElseThrow().triple(), - m)) - .build()) { - machineRef[0] = (JffiNativeMachine) instance.getMachine(); - - // Returns normally: the entry check ran before the flag was raised. + NativeInstanceBuilder.builder(module).withImportValues(imports).build()) { + // nothing polls after the host function, so this returns normally instance.export("callHost").apply(); assertEquals( 42, (int) instance.export("answer").apply()[0], - "a flag left over from the previous call must not stop this one"); - assertFalse( - Thread.currentThread().isInterrupted(), - "no interrupt happened, so the caller must not be left interrupted"); + "an interrupt handled during the previous call must not stop this one"); + assertFalse(Thread.currentThread().isInterrupted()); } } + + @Test + public void anInterruptHandledInANestedCallDoesNotStopTheOuterCall() { + var module = + Parser.parse(CorpusResources.getResource("compiled/reentrant-interrupt.wat.wasm")); + var noParams = FunctionType.of(List.of(), List.of()); + var imports = + ImportValues.builder() + .addFunction( + new HostFunction( + "host", + "reenter", + noParams, + (inst, args) -> { + inst.export("raise").apply(); + return null; + })) + .addFunction( + new HostFunction( + "host", + "raiseFlag", + noParams, + (inst, args) -> { + interruptAndHandle(); + return null; + })) + .addFunction( + new HostFunction("host", "tick", noParams, (inst, args) -> null)) + .build(); + + try (var instance = + NativeInstanceBuilder.builder(module).withImportValues(imports).build()) { + assertEquals( + 1000, + (int) instance.export("run").apply()[0], + "an interrupt handled in the nested call must not stop the outer one"); + } + } + + // interrupts this thread, gives the watchdog time to see it, then handles it as a host would + private static void interruptAndHandle() { + Thread.currentThread().interrupt(); + long end = System.nanoTime() + MILLISECONDS.toNanos(500); + while (System.nanoTime() < end) { + Thread.onSpinWait(); + } + Thread.interrupted(); + } } diff --git a/redline/runner-jffi-tests/src/test/java/run/endive/redline/experimental/runner/jffi/internal/InterruptionTest.java b/redline/runner-jffi-tests/src/test/java/run/endive/redline/experimental/runner/jffi/internal/InterruptionTest.java index fff74104e..9d1b62e17 100644 --- a/redline/runner-jffi-tests/src/test/java/run/endive/redline/experimental/runner/jffi/internal/InterruptionTest.java +++ b/redline/runner-jffi-tests/src/test/java/run/endive/redline/experimental/runner/jffi/internal/InterruptionTest.java @@ -1,63 +1,194 @@ package run.endive.redline.experimental.runner.jffi.internal; import static java.util.concurrent.TimeUnit.SECONDS; -import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; import static org.junit.jupiter.api.Assertions.assertThrows; import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.junit.jupiter.api.Assertions.fail; +import java.lang.management.ManagementFactory; +import java.util.List; +import java.util.concurrent.CountDownLatch; import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicReference; +import java.util.function.Consumer; import org.junit.jupiter.api.Test; import run.endive.corpus.CorpusResources; -import run.endive.redline.experimental.api.internal.RedlineTarget; -import run.endive.redline.experimental.compiler.internal.NativeCompiler; -import run.endive.redline.experimental.runner.jffi.JffiNativeMachineFactory; +import run.endive.runtime.HostFunction; +import run.endive.runtime.ImportValues; import run.endive.runtime.Instance; +import run.endive.runtime.WasmInterruptedException; +import run.endive.testing.NativeInstanceBuilder; import run.endive.wasm.Parser; -import run.endive.wasm.WasmEngineException; +import run.endive.wasm.types.FunctionType; public class InterruptionTest { + // well below the reentrant stack guard + private static final int DEPTH = 20; + + private static final int ROUNDS = 20; + @Test public void shouldInterruptLoopViaThread() throws InterruptedException { - try (var instance = buildInstance("compiled/infinite-loop.c.wasm")) { - var function = instance.export("run"); - assertThreadInterruption(function::apply); - } + var instance = + buildInstance("compiled/infinite-loop.c.wasm", ImportValues.builder().build()); + var function = instance.export("run"); + assertThreadInterruption(function::apply, null, instance); } @Test public void shouldInterruptCallViaThread() throws InterruptedException { - try (var instance = buildInstance("compiled/power.c.wasm")) { - var function = instance.export("run"); - assertThreadInterruption(() -> function.apply(100)); + var instance = buildInstance("compiled/power.c.wasm", ImportValues.builder().build()); + var function = instance.export("run"); + assertThreadInterruption(() -> function.apply(100), null, instance); + } + + @Test + public void shouldInterruptNestedLoopViaThread() throws InterruptedException { + var running = new CountDownLatch(1); + var instance = + buildInstance( + "compiled/reentrant-interrupt.wat.wasm", + reentrantImports(inst -> inst.export("spin").apply(), running)); + var function = instance.export("run"); + assertThreadInterruption(function::apply, running, instance); + } + + @Test + public void shouldInterruptLoopInAnotherMachineViaThread() throws InterruptedException { + var running = new CountDownLatch(1); + var inner = + buildInstance( + "compiled/reentrant-interrupt.wat.wasm", + reentrantImports(inst -> {}, running)); + var outer = + buildInstance( + "compiled/reentrant-interrupt.wat.wasm", + reentrantImports(inst -> inner.export("spin").apply(), running)); + var function = outer.export("run"); + assertThreadInterruption(function::apply, running, outer, inner); + } + + @Test + public void shouldInterruptAfterTheWatchdogWentIdle() throws InterruptedException { + var running = new CountDownLatch(1); + var instance = + buildInstance( + "compiled/reentrant-interrupt.wat.wasm", + reentrantImports(inst -> inst.export("spin").apply(), running)); + instance.export("raise").apply(); + // longer than the watchdog's idle time in this module's test configuration + Thread.sleep(1000); + var function = instance.export("run"); + assertThreadInterruption(function::apply, running, instance); + } + + @Test + public void callingStartsNoThreads() { + var threads = ManagementFactory.getThreadMXBean(); + int[] depth = {0}; + var imports = + ImportValues.builder() + .addFunction( + new HostFunction( + "host", + "reenter", + FunctionType.of(List.of(), List.of()), + (inst, args) -> { + if (depth[0]++ < DEPTH) { + inst.export("recurse").apply(); + } + depth[0] = 0; + return null; + })) + .build(); + try (var instance = buildInstance("compiled/reentrant-recursion.wat.wasm", imports)) { + var recurse = instance.export("recurse"); + // starts the shared watchdog + recurse.apply(); + + long before = threads.getTotalStartedThreadCount(); + for (int i = 0; i < ROUNDS; i++) { + recurse.apply(); + } + long started = threads.getTotalStartedThreadCount() - before; + + // loose bound: JIT compiler threads may start meanwhile + assertTrue(started < ROUNDS, started + " threads started"); } } - private static void assertThreadInterruption(Runnable function) throws InterruptedException { + // interrupts once `running` is released, or after 100ms without one + private static void assertThreadInterruption( + Runnable function, CountDownLatch running, Instance... instances) + throws InterruptedException { AtomicBoolean interrupted = new AtomicBoolean(); + var failure = new AtomicReference(); Thread thread = new Thread( () -> { - var e = assertThrows(WasmEngineException.class, function::run); - assertEquals("interrupted", e.getMessage()); + assertThrows(WasmInterruptedException.class, function::run); interrupted.set(true); }); thread.setDaemon(true); + thread.setUncaughtExceptionHandler( + (t, e) -> { + failure.set(e); + if (running != null) { + running.countDown(); + } + }); thread.start(); - Thread.sleep(100); + boolean started = true; + if (running == null) { + Thread.sleep(100); + } else { + started = running.await(10, SECONDS); + } thread.interrupt(); SECONDS.timedJoin(thread, 10); + if (failure.get() != null) { + fail("the call failed", failure.get()); + } + assertTrue(started, "the call never reached its loop"); + // a call still running uses the instances' memory, so they are leaked, not closed + assertFalse(thread.isAlive(), "the call was not interrupted"); + for (var instance : instances) { + instance.close(); + } assertTrue(interrupted.get()); } - private static Instance buildInstance(String resource) { - var module = Parser.parse(CorpusResources.getResource(resource)); - return JffiNativeMachineFactory.builder(module) - .withCompilerFunction( - m -> - NativeCompiler.compileAll( - RedlineTarget.detectHost().orElseThrow().triple(), m)) + private static ImportValues reentrantImports( + Consumer reenter, CountDownLatch running) { + var noParams = FunctionType.of(List.of(), List.of()); + return ImportValues.builder() + .addFunction( + new HostFunction( + "host", + "reenter", + noParams, + (inst, args) -> { + reenter.accept(inst); + return null; + })) + .addFunction(new HostFunction("host", "raiseFlag", noParams, (inst, args) -> null)) + .addFunction( + new HostFunction( + "host", + "tick", + noParams, + (inst, args) -> { + running.countDown(); + return null; + })) .build(); } + + private static Instance buildInstance(String resource, ImportValues imports) { + var module = Parser.parse(CorpusResources.getResource(resource)); + return NativeInstanceBuilder.builder(module).withImportValues(imports).build(); + } } diff --git a/redline/runner-jffi/src/main/java/run/endive/redline/experimental/runner/jffi/internal/JffiNativeMachine.java b/redline/runner-jffi/src/main/java/run/endive/redline/experimental/runner/jffi/internal/JffiNativeMachine.java index 4adfb8d2b..d5d570744 100644 --- a/redline/runner-jffi/src/main/java/run/endive/redline/experimental/runner/jffi/internal/JffiNativeMachine.java +++ b/redline/runner-jffi/src/main/java/run/endive/redline/experimental/runner/jffi/internal/JffiNativeMachine.java @@ -17,12 +17,14 @@ import java.util.List; import java.util.Map; import run.endive.redline.experimental.api.internal.CtxBuffer; +import run.endive.redline.experimental.api.internal.InterruptWatchdog; import run.endive.redline.experimental.api.internal.RedlineTarget; import run.endive.redline.experimental.api.internal.TypeMapUtils; import run.endive.redline.experimental.bridge.internal.CraneliftBridge; import run.endive.runtime.Instance; import run.endive.runtime.Machine; import run.endive.runtime.TrapException; +import run.endive.runtime.WasmInterruptedException; import run.endive.runtime.WasmRuntimeException; import run.endive.wasm.WasmEngineException; import run.endive.wasm.types.FunctionType; @@ -85,8 +87,7 @@ public final class JffiNativeMachine implements Machine { } private final Instance instance; - private final CallContext[] entryTrampolineCallCtxs; // entry trampoline CallContext per func - private final long[] entryTrampolineAddrs; // entry trampoline native addr per func + private final Function[] entryTrampolines; // entry trampoline per func private final FunctionType[] funcTypes; // wasm FunctionType per func private final long codeRegionAddr; private final int codeRegionOsPages; @@ -113,6 +114,7 @@ public final class JffiNativeMachine implements Machine { private JffiNativeMemory nativeMemory; private volatile Throwable pendingException; private int callDepth; + private final InterruptWatchdog.InterruptSink interruptFlag = this::raiseInterruptFlag; // Keep closure handles alive to prevent GC private final Closure.Handle trampolineHandle; @@ -137,8 +139,7 @@ public JffiNativeMachine( .FUNCTION) .count(); int totalFuncs = numImports + module.codeSection().functionBodyCount(); - this.entryTrampolineCallCtxs = new CallContext[totalFuncs]; - this.entryTrampolineAddrs = new long[totalFuncs]; + this.entryTrampolines = new Function[totalFuncs]; this.funcTypes = new FunctionType[totalFuncs]; this.importHandles = new Closure.Handle[numImports]; @@ -373,9 +374,10 @@ public JffiNativeMachine( for (int i = 0; i < compiledCode.length; i++) { if (compiledCode[i] != null) { int funcId = numImports + i; - entryTrampolineAddrs[funcId] = entryTrampolinePtrs.get(funcTypesByBody[i]); - entryTrampolineCallCtxs[funcId] = - createEntryTrampolineCallContext(funcTypesByBody[i]); + entryTrampolines[funcId] = + new Function( + entryTrampolinePtrs.get(funcTypesByBody[i]), + createEntryTrampolineCallContext(funcTypesByBody[i])); } } } @@ -931,7 +933,7 @@ private static WasmEngineException trapException(int trapCode) { return new TrapException("unaligned atomic"); } if (trapCode == CtxBuffer.TRAP_INTERRUPTED) { - return new TrapException("interrupted"); + return new WasmInterruptedException("Thread interrupted"); } return new WasmEngineException("trap: unknown code " + trapCode); } @@ -939,8 +941,7 @@ private static WasmEngineException trapException(int trapCode) { // --- Native function invocation --- private long invokeViaEntryTrampoline( - CallContext trampolineCallCtx, - long trampolineAddr, + Function trampoline, FunctionType funcType, long funcAddr, long memBase, @@ -948,6 +949,8 @@ private long invokeViaEntryTrampoline( long[] wasmArgs) { // nativeArgCount = funcPtr + memBase + ctxPtr + wasm params int nativeArgCount = 3 + wasmArgs.length; + CallContext trampolineCallCtx = trampoline.getCallContext(); + long trampolineAddr = trampoline.getFunctionAddress(); switch (nativeArgCount) { case 3: @@ -978,25 +981,17 @@ private long invokeViaEntryTrampoline( default: // >6 args: use HeapInvocationBuffer return invokeViaBufferWithTrampoline( - trampolineCallCtx, - trampolineAddr, - funcType, - funcAddr, - memBase, - ctxPtr, - wasmArgs); + trampoline, funcType, funcAddr, memBase, ctxPtr, wasmArgs); } } private long invokeViaBufferWithTrampoline( - CallContext trampolineCallCtx, - long trampolineAddr, + Function func, FunctionType funcType, long funcAddr, long memBase, long ctxPtr, long[] wasmArgs) { - var func = new Function(trampolineAddr, trampolineCallCtx); var buffer = new HeapInvocationBuffer(func); buffer.putAddress(funcAddr); // funcPtr (first arg to entry trampoline) buffer.putAddress(memBase); @@ -1042,8 +1037,6 @@ public long[] call(int funcId, long[] args) throws WasmEngineException { var funcType = funcTypes[funcId]; long funcAddr = MEM.getLong(funcTableAddr + (long) funcId * 8); - long trampolineAddr = entryTrampolineAddrs[funcId]; - var trampolineCallCtx = entryTrampolineCallCtxs[funcId]; try { boolean outermostCall = callDepth++ == 0; @@ -1076,45 +1069,28 @@ public long[] call(int funcId, long[] args) throws WasmEngineException { MEM.putInt(ctxBufferAddr + CtxBuffer.MEMORY_PAGES, mem.pages()); } - if (Thread.interrupted()) { - throw new TrapException("interrupted"); + if (Thread.currentThread().isInterrupted()) { + throw new WasmInterruptedException("Thread interrupted"); } - Thread caller = Thread.currentThread(); - Thread watchdog = - new Thread( - () -> { - // Keeps raising rather than returning after the - // first: a nested call clears the flag when it - // finishes, and the outer call still needs it. - while (!Thread.currentThread().isInterrupted()) { - if (caller.isInterrupted()) { - requestInterrupt(); - } - try { - Thread.sleep(1); - } catch (InterruptedException e) { - return; - } - } - }); - watchdog.setDaemon(true); - watchdog.start(); + // nested calls run inside the outermost call's watch + InterruptWatchdog.Registration watchdog = + outermostCall ? InterruptWatchdog.enter(interruptFlag) : null; long result; try { result = invokeViaEntryTrampoline( - trampolineCallCtx, - trampolineAddr, + entryTrampolines[funcId], funcType, funcAddr, cachedMemBase, ctxBufferAddr, args); } finally { - watchdog.interrupt(); - // The flag only ever means "stop this call". Left set it would - // trap the next one on a thread nobody interrupted. + if (watchdog != null) { + // exit first, so the poller cannot raise the flag after the clear + InterruptWatchdog.exit(watchdog); + } clearInterrupt(); } @@ -1134,7 +1110,6 @@ public long[] call(int funcId, long[] args) throws WasmEngineException { if (trapCode != 0) { MEM.putInt(ctxBufferAddr + CtxBuffer.TRAP_CODE, 0); if (trapCode == CtxBuffer.TRAP_INTERRUPTED) { - CHECKED_MEM.putLong(ctxBufferAddr + CtxBuffer.INTERRUPT_FLAG, 0L); Thread.currentThread().interrupt(); } throw trapException(trapCode); @@ -1165,16 +1140,15 @@ public long[] call(int funcId, long[] args) throws WasmEngineException { // Prevent the JIT from considering this machine unreachable during // the native call, which would let GC collect and close() free // native memory while code is executing. - // (ctxBuffer, funcTypesArray, code region) while code is executing. Reference.reachabilityFence(this); } } - public void requestInterrupt() { + private void raiseInterruptFlag() { CHECKED_MEM.putLong(ctxBufferAddr + CtxBuffer.INTERRUPT_FLAG, 1L); } - public void clearInterrupt() { + private void clearInterrupt() { CHECKED_MEM.putLong(ctxBufferAddr + CtxBuffer.INTERRUPT_FLAG, 0L); } diff --git a/redline/runner-tests/pom.xml b/redline/runner-tests/pom.xml index 95a6b06be..399f1ddac 100644 --- a/redline/runner-tests/pom.xml +++ b/redline/runner-tests/pom.xml @@ -68,6 +68,16 @@ + + org.apache.maven.plugins + maven-surefire-plugin + + + + 200 + + + run.endive test-gen-plugin diff --git a/redline/runner-tests/src/test/java/run/endive/redline/experimental/runner/internal/InterruptFlagTest.java b/redline/runner-tests/src/test/java/run/endive/redline/experimental/runner/internal/InterruptFlagTest.java index d81115e90..a77f9112b 100644 --- a/redline/runner-tests/src/test/java/run/endive/redline/experimental/runner/internal/InterruptFlagTest.java +++ b/redline/runner-tests/src/test/java/run/endive/redline/experimental/runner/internal/InterruptFlagTest.java @@ -1,24 +1,20 @@ package run.endive.redline.experimental.runner.internal; +import static java.util.concurrent.TimeUnit.MILLISECONDS; import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertFalse; +import java.util.List; import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.Test; import run.endive.corpus.CorpusResources; -import run.endive.redline.experimental.api.internal.RedlineTarget; -import run.endive.redline.experimental.compiler.internal.NativeCompiler; -import run.endive.redline.experimental.runner.NativeMachineFactory; import run.endive.runtime.HostFunction; import run.endive.runtime.ImportValues; +import run.endive.testing.NativeInstanceBuilder; import run.endive.wasm.Parser; import run.endive.wasm.types.FunctionType; -/** - * The watchdog raises the interrupt flag from another thread, so it can land after - * the call it was meant to stop has passed its last check. The flag must not then - * sit in the context and stop a later call that nobody interrupted. - */ +/** An interrupt the host handled during a call must not stop a later call. */ public class InterruptFlagTest { @AfterEach @@ -28,45 +24,80 @@ public void clearInterruptStatus() { } @Test - public void aFlagRaisedMidCallDoesNotStopTheNextCall() { + public void anInterruptHandledByTheHostDoesNotStopTheNextCall() { var module = Parser.parse(CorpusResources.getResource("compiled/interrupt-midcall.wat.wasm")); - - var machineRef = new NativeMachine[1]; var imports = ImportValues.builder() .addFunction( new HostFunction( "host", "raiseFlag", - FunctionType.of(java.util.List.of(), java.util.List.of()), + FunctionType.of(List.of(), List.of()), (inst, args) -> { - machineRef[0].requestInterrupt(); + interruptAndHandle(); return null; })) .build(); try (var instance = - NativeMachineFactory.builder(module) - .withImportValues(imports) - .withCompilerFunction( - m -> - NativeCompiler.compileAll( - RedlineTarget.detectHost().orElseThrow().triple(), - m)) - .build()) { - machineRef[0] = (NativeMachine) instance.getMachine(); - - // Returns normally: the entry check ran before the flag was raised. + NativeInstanceBuilder.builder(module).withImportValues(imports).build()) { + // nothing polls after the host function, so this returns normally instance.export("callHost").apply(); assertEquals( 42, (int) instance.export("answer").apply()[0], - "a flag left over from the previous call must not stop this one"); - assertFalse( - Thread.currentThread().isInterrupted(), - "no interrupt happened, so the caller must not be left interrupted"); + "an interrupt handled during the previous call must not stop this one"); + assertFalse(Thread.currentThread().isInterrupted()); } } + + @Test + public void anInterruptHandledInANestedCallDoesNotStopTheOuterCall() { + var module = + Parser.parse(CorpusResources.getResource("compiled/reentrant-interrupt.wat.wasm")); + var noParams = FunctionType.of(List.of(), List.of()); + var imports = + ImportValues.builder() + .addFunction( + new HostFunction( + "host", + "reenter", + noParams, + (inst, args) -> { + inst.export("raise").apply(); + return null; + })) + .addFunction( + new HostFunction( + "host", + "raiseFlag", + noParams, + (inst, args) -> { + interruptAndHandle(); + return null; + })) + .addFunction( + new HostFunction("host", "tick", noParams, (inst, args) -> null)) + .build(); + + try (var instance = + NativeInstanceBuilder.builder(module).withImportValues(imports).build()) { + assertEquals( + 1000, + (int) instance.export("run").apply()[0], + "an interrupt handled in the nested call must not stop the outer one"); + } + } + + // interrupts this thread, gives the watchdog time to see it, then handles it as a host would + private static void interruptAndHandle() { + Thread.currentThread().interrupt(); + long end = System.nanoTime() + MILLISECONDS.toNanos(500); + while (System.nanoTime() < end) { + Thread.onSpinWait(); + } + Thread.interrupted(); + } } diff --git a/redline/runner-tests/src/test/java/run/endive/redline/experimental/runner/internal/InterruptionTest.java b/redline/runner-tests/src/test/java/run/endive/redline/experimental/runner/internal/InterruptionTest.java index 755f10d0b..d6ef31798 100644 --- a/redline/runner-tests/src/test/java/run/endive/redline/experimental/runner/internal/InterruptionTest.java +++ b/redline/runner-tests/src/test/java/run/endive/redline/experimental/runner/internal/InterruptionTest.java @@ -1,63 +1,194 @@ package run.endive.redline.experimental.runner.internal; import static java.util.concurrent.TimeUnit.SECONDS; -import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; import static org.junit.jupiter.api.Assertions.assertThrows; import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.junit.jupiter.api.Assertions.fail; +import java.lang.management.ManagementFactory; +import java.util.List; +import java.util.concurrent.CountDownLatch; import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicReference; +import java.util.function.Consumer; import org.junit.jupiter.api.Test; import run.endive.corpus.CorpusResources; -import run.endive.redline.experimental.api.internal.RedlineTarget; -import run.endive.redline.experimental.compiler.internal.NativeCompiler; -import run.endive.redline.experimental.runner.NativeMachineFactory; +import run.endive.runtime.HostFunction; +import run.endive.runtime.ImportValues; import run.endive.runtime.Instance; +import run.endive.runtime.WasmInterruptedException; +import run.endive.testing.NativeInstanceBuilder; import run.endive.wasm.Parser; -import run.endive.wasm.WasmEngineException; +import run.endive.wasm.types.FunctionType; public class InterruptionTest { + // well below the reentrant stack guard + private static final int DEPTH = 20; + + private static final int ROUNDS = 20; + @Test public void shouldInterruptLoopViaThread() throws InterruptedException { - try (var instance = buildInstance("compiled/infinite-loop.c.wasm")) { - var function = instance.export("run"); - assertThreadInterruption(function::apply); - } + var instance = + buildInstance("compiled/infinite-loop.c.wasm", ImportValues.builder().build()); + var function = instance.export("run"); + assertThreadInterruption(function::apply, null, instance); } @Test public void shouldInterruptCallViaThread() throws InterruptedException { - try (var instance = buildInstance("compiled/power.c.wasm")) { - var function = instance.export("run"); - assertThreadInterruption(() -> function.apply(100)); + var instance = buildInstance("compiled/power.c.wasm", ImportValues.builder().build()); + var function = instance.export("run"); + assertThreadInterruption(() -> function.apply(100), null, instance); + } + + @Test + public void shouldInterruptNestedLoopViaThread() throws InterruptedException { + var running = new CountDownLatch(1); + var instance = + buildInstance( + "compiled/reentrant-interrupt.wat.wasm", + reentrantImports(inst -> inst.export("spin").apply(), running)); + var function = instance.export("run"); + assertThreadInterruption(function::apply, running, instance); + } + + @Test + public void shouldInterruptLoopInAnotherMachineViaThread() throws InterruptedException { + var running = new CountDownLatch(1); + var inner = + buildInstance( + "compiled/reentrant-interrupt.wat.wasm", + reentrantImports(inst -> {}, running)); + var outer = + buildInstance( + "compiled/reentrant-interrupt.wat.wasm", + reentrantImports(inst -> inner.export("spin").apply(), running)); + var function = outer.export("run"); + assertThreadInterruption(function::apply, running, outer, inner); + } + + @Test + public void shouldInterruptAfterTheWatchdogWentIdle() throws InterruptedException { + var running = new CountDownLatch(1); + var instance = + buildInstance( + "compiled/reentrant-interrupt.wat.wasm", + reentrantImports(inst -> inst.export("spin").apply(), running)); + instance.export("raise").apply(); + // longer than the watchdog's idle time in this module's test configuration + Thread.sleep(1000); + var function = instance.export("run"); + assertThreadInterruption(function::apply, running, instance); + } + + @Test + public void callingStartsNoThreads() { + var threads = ManagementFactory.getThreadMXBean(); + int[] depth = {0}; + var imports = + ImportValues.builder() + .addFunction( + new HostFunction( + "host", + "reenter", + FunctionType.of(List.of(), List.of()), + (inst, args) -> { + if (depth[0]++ < DEPTH) { + inst.export("recurse").apply(); + } + depth[0] = 0; + return null; + })) + .build(); + try (var instance = buildInstance("compiled/reentrant-recursion.wat.wasm", imports)) { + var recurse = instance.export("recurse"); + // starts the shared watchdog + recurse.apply(); + + long before = threads.getTotalStartedThreadCount(); + for (int i = 0; i < ROUNDS; i++) { + recurse.apply(); + } + long started = threads.getTotalStartedThreadCount() - before; + + // loose bound: JIT compiler threads may start meanwhile + assertTrue(started < ROUNDS, started + " threads started"); } } - private static void assertThreadInterruption(Runnable function) throws InterruptedException { + // interrupts once `running` is released, or after 100ms without one + private static void assertThreadInterruption( + Runnable function, CountDownLatch running, Instance... instances) + throws InterruptedException { AtomicBoolean interrupted = new AtomicBoolean(); + var failure = new AtomicReference(); Thread thread = new Thread( () -> { - var e = assertThrows(WasmEngineException.class, function::run); - assertEquals("interrupted", e.getMessage()); + assertThrows(WasmInterruptedException.class, function::run); interrupted.set(true); }); thread.setDaemon(true); + thread.setUncaughtExceptionHandler( + (t, e) -> { + failure.set(e); + if (running != null) { + running.countDown(); + } + }); thread.start(); - Thread.sleep(100); + boolean started = true; + if (running == null) { + Thread.sleep(100); + } else { + started = running.await(10, SECONDS); + } thread.interrupt(); SECONDS.timedJoin(thread, 10); + if (failure.get() != null) { + fail("the call failed", failure.get()); + } + assertTrue(started, "the call never reached its loop"); + // a call still running uses the instances' memory, so they are leaked, not closed + assertFalse(thread.isAlive(), "the call was not interrupted"); + for (var instance : instances) { + instance.close(); + } assertTrue(interrupted.get()); } - private static Instance buildInstance(String resource) { - var module = Parser.parse(CorpusResources.getResource(resource)); - return NativeMachineFactory.builder(module) - .withCompilerFunction( - m -> - NativeCompiler.compileAll( - RedlineTarget.detectHost().orElseThrow().triple(), m)) + private static ImportValues reentrantImports( + Consumer reenter, CountDownLatch running) { + var noParams = FunctionType.of(List.of(), List.of()); + return ImportValues.builder() + .addFunction( + new HostFunction( + "host", + "reenter", + noParams, + (inst, args) -> { + reenter.accept(inst); + return null; + })) + .addFunction(new HostFunction("host", "raiseFlag", noParams, (inst, args) -> null)) + .addFunction( + new HostFunction( + "host", + "tick", + noParams, + (inst, args) -> { + running.countDown(); + return null; + })) .build(); } + + private static Instance buildInstance(String resource, ImportValues imports) { + var module = Parser.parse(CorpusResources.getResource(resource)); + return NativeInstanceBuilder.builder(module).withImportValues(imports).build(); + } } diff --git a/redline/runner/src/main/java/run/endive/redline/experimental/runner/internal/NativeMachine.java b/redline/runner/src/main/java/run/endive/redline/experimental/runner/internal/NativeMachine.java index b9da6235b..7051207fe 100644 --- a/redline/runner/src/main/java/run/endive/redline/experimental/runner/internal/NativeMachine.java +++ b/redline/runner/src/main/java/run/endive/redline/experimental/runner/internal/NativeMachine.java @@ -14,12 +14,14 @@ import java.util.Map; import java.util.function.Function; import run.endive.redline.experimental.api.internal.CtxBuffer; +import run.endive.redline.experimental.api.internal.InterruptWatchdog; import run.endive.redline.experimental.api.internal.RedlineTarget; import run.endive.redline.experimental.api.internal.TypeMapUtils; import run.endive.redline.experimental.bridge.internal.CraneliftBridge; import run.endive.runtime.Instance; import run.endive.runtime.Machine; import run.endive.runtime.TrapException; +import run.endive.runtime.WasmInterruptedException; import run.endive.runtime.WasmRuntimeException; import run.endive.wasm.WasmEngineException; import run.endive.wasm.types.FunctionType; @@ -102,6 +104,7 @@ public final class NativeMachine implements Machine { private NativeMemory nativeMemory; private volatile Throwable pendingException; private int callDepth; + private final InterruptWatchdog.InterruptSink interruptFlag = this::raiseInterruptFlag; private boolean ownsMemory; private boolean closed; @@ -972,7 +975,7 @@ private static WasmEngineException trapException(int trapCode) { case CtxBuffer.TRAP_INDIRECT_CALL_TYPE_MISMATCH -> new TrapException("indirect call type mismatch"); case CtxBuffer.TRAP_UNALIGNED_ATOMIC -> new TrapException("unaligned atomic"); - case CtxBuffer.TRAP_INTERRUPTED -> new TrapException("interrupted"); + case CtxBuffer.TRAP_INTERRUPTED -> new WasmInterruptedException("Thread interrupted"); default -> new WasmEngineException("trap: unknown code " + trapCode); }; } @@ -1079,37 +1082,21 @@ public long[] call(int funcId, long[] args) throws WasmEngineException { ctxBuffer.set(ValueLayout.JAVA_INT, CtxBuffer.MEMORY_PAGES, mem.pages()); } - if (Thread.interrupted()) { - throw new TrapException("interrupted"); + if (Thread.currentThread().isInterrupted()) { + throw new WasmInterruptedException("Thread interrupted"); } - Thread caller = Thread.currentThread(); - Thread watchdog = - new Thread( - () -> { - // Keeps raising rather than returning after the - // first: a nested call clears the flag when it - // finishes, and the outer call still needs it. - while (!Thread.currentThread().isInterrupted()) { - if (caller.isInterrupted()) { - requestInterrupt(); - } - try { - Thread.sleep(1); - } catch (InterruptedException e) { - return; - } - } - }); - watchdog.setDaemon(true); - watchdog.start(); + // nested calls run inside the outermost call's watch + InterruptWatchdog.Registration watchdog = + outermostCall ? InterruptWatchdog.enter(interruptFlag) : null; long result; try { result = (long) handle.invokeExact(cachedMemBase, ctxBuffer, args); } finally { - watchdog.interrupt(); - // The flag only ever means "stop this call". Left set it would - // trap the next one on a thread nobody interrupted. + if (watchdog != null) { + // exit first, so the poller cannot raise the flag after the clear + InterruptWatchdog.exit(watchdog); + } clearInterrupt(); } @@ -1129,7 +1116,6 @@ public long[] call(int funcId, long[] args) throws WasmEngineException { if (trapCode != 0) { ctxBuffer.set(ValueLayout.JAVA_INT, CtxBuffer.TRAP_CODE, 0); if (trapCode == CtxBuffer.TRAP_INTERRUPTED) { - ctxBuffer.set(ValueLayout.JAVA_LONG, CtxBuffer.INTERRUPT_FLAG, 0L); Thread.currentThread().interrupt(); } throw trapException(trapCode); @@ -1164,11 +1150,11 @@ public long[] call(int funcId, long[] args) throws WasmEngineException { } } - public void requestInterrupt() { + private void raiseInterruptFlag() { ctxBuffer.set(ValueLayout.JAVA_LONG, CtxBuffer.INTERRUPT_FLAG, 1L); } - public void clearInterrupt() { + private void clearInterrupt() { ctxBuffer.set(ValueLayout.JAVA_LONG, CtxBuffer.INTERRUPT_FLAG, 0L); } diff --git a/wasm-corpus/src/main/resources/compiled/reentrant-interrupt.wat.wasm b/wasm-corpus/src/main/resources/compiled/reentrant-interrupt.wat.wasm new file mode 100644 index 000000000..623e23197 Binary files /dev/null and b/wasm-corpus/src/main/resources/compiled/reentrant-interrupt.wat.wasm differ diff --git a/wasm-corpus/src/main/resources/wat/reentrant-interrupt.wat b/wasm-corpus/src/main/resources/wat/reentrant-interrupt.wat new file mode 100644 index 000000000..ef7ff7fb7 --- /dev/null +++ b/wasm-corpus/src/main/resources/wat/reentrant-interrupt.wat @@ -0,0 +1,22 @@ +;; "run" re-enters through the host, then loops so the interrupt flag is polled. +(module + (import "host" "reenter" (func $reenter)) + (import "host" "raiseFlag" (func $raiseFlag)) + (import "host" "tick" (func $tick)) + + (func (export "run") (result i32) + (local $i i32) + (call $reenter) + (loop $l + (local.set $i (i32.add (local.get $i) (i32.const 1))) + (br_if $l (i32.lt_u (local.get $i) (i32.const 1000)))) + (local.get $i)) + + (func (export "raise") + (call $raiseFlag)) + + (func (export "spin") + (loop $l + (call $tick) + (br $l))) +)