diff --git a/src/bun.js/bindings/webcore/Worker.cpp b/src/bun.js/bindings/webcore/Worker.cpp index 8db7b462b33a..181a5bc68494 100644 --- a/src/bun.js/bindings/webcore/Worker.cpp +++ b/src/bun.js/bindings/webcore/Worker.cpp @@ -133,12 +133,23 @@ extern "C" void WebWorker__setRef( void Worker::setKeepAlive(bool keepAlive) { + Locker locker { m_implLock }; + if (!impl_) + return; WebWorker__setRef(impl_, keepAlive); } +void Worker::clearZigImpl() +{ + Locker locker { m_implLock }; + impl_ = nullptr; +} + bool Worker::updatePtr() { + Locker locker { m_implLock }; if (!WebWorker__updatePtr(impl_, this)) { + impl_ = nullptr; m_onlineClosingFlags = ClosingFlag; m_terminationFlags.fetch_or(TerminatedFlag); return false; @@ -216,7 +227,10 @@ ExceptionOr> Worker::create(ScriptExecutionContext& context, const S return Exception { TypeError, errorMessage.toWTFString(BunString::ZeroCopy) }; } - worker->impl_ = impl; + { + Locker locker { worker->m_implLock }; + worker->impl_ = impl; + } worker->m_workerCreationTime = MonotonicTime::now(); return worker; @@ -263,6 +277,9 @@ void Worker::terminate() { // m_contextProxy.terminateWorkerGlobalScope(); m_terminationFlags.fetch_or(TerminateRequestedFlag); + Locker locker { m_implLock }; + if (!impl_) + return; WebWorker__notifyNeedTermination(impl_); } @@ -467,6 +484,10 @@ void Worker::forEachWorker(const FunctionclearZigImpl(); worker->dispatchExit(exitCode); // no longer referenced by Zig worker->deref(); diff --git a/src/bun.js/bindings/webcore/Worker.h b/src/bun.js/bindings/webcore/Worker.h index bbc73053ddfe..fc0d2055bfe9 100644 --- a/src/bun.js/bindings/webcore/Worker.h +++ b/src/bun.js/bindings/webcore/Worker.h @@ -76,6 +76,7 @@ class Worker final : public ThreadSafeRefCounted, public EventTargetWith void dispatchEvent(Event&); void dispatchCloseEvent(Event&); void setKeepAlive(bool); + void clearZigImpl(); void postTaskToWorkerGlobalScope(Function&&); @@ -119,7 +120,8 @@ class Worker final : public ThreadSafeRefCounted, public EventTargetWith // Tracks TerminateRequestedFlag and TerminatedFlag std::atomic m_terminationFlags { 0 }; const ScriptExecutionContextIdentifier m_clientIdentifier; - void* impl_ { nullptr }; + Lock m_implLock; + void* impl_ WTF_GUARDED_BY_LOCK(m_implLock) { nullptr }; }; JSValue createNodeWorkerThreadsBinding(Zig::GlobalObject* globalObject); diff --git a/test/js/web/workers/worker-terminate-race.test.ts b/test/js/web/workers/worker-terminate-race.test.ts new file mode 100644 index 000000000000..48333ceaafea --- /dev/null +++ b/test/js/web/workers/worker-terminate-race.test.ts @@ -0,0 +1,58 @@ +import { expect, test } from "bun:test"; +import { bunEnv, bunExe } from "harness"; + +// The Zig WebWorker struct is freed on the worker thread once the worker +// exits. These tests hammer ref()/unref()/terminate() from the parent +// thread while the worker thread is tearing down, which used to read the +// freed struct (ASAN use-after-poison in WebWorker__setRef / +// WebWorker__notifyNeedTermination). + +async function run(src: string) { + await using proc = Bun.spawn({ + cmd: [bunExe(), "-e", src], + env: bunEnv, + stdout: "pipe", + stderr: "pipe", + }); + const [stdout, stderr, exitCode] = await Promise.all([proc.stdout.text(), proc.stderr.text(), proc.exited]); + expect(stdout).toBe(""); + if (exitCode !== 0) { + expect(stderr).toBe(""); + } + expect(exitCode).toBe(0); +} + +test.concurrent("Worker: ref/unref after terminate does not use-after-free", async () => { + await run(` + const w = new Worker("data:text/javascript,", {}); + w.terminate(); + for (let i = 0; i < 100000; i++) { + w.unref(); + w.ref(); + } + w.terminate(); + w.unref(); + `); +}); + +test.concurrent("Worker: ref/unref racing natural exit does not use-after-free", async () => { + await run(` + const w = new Worker("data:text/javascript,", {}); + const end = Date.now() + 2000; + while (Date.now() < end) { + w.unref(); + w.ref(); + } + w.unref(); + `); +}); + +test.concurrent("Worker: terminate racing natural exit does not use-after-free", async () => { + await run(` + const w = new Worker("data:text/javascript,", {}); + const end = Date.now() + 2000; + while (Date.now() < end) { + w.terminate(); + } + `); +});