From 0bf3c30c055a8b4632b8370b86722b977dab1cf5 Mon Sep 17 00:00:00 2001 From: DDT <1786035110@qq.com> Date: Thu, 17 Sep 2026 01:13:28 +0800 Subject: [PATCH 1/2] fix: eliminate TOCTOU race in MemoryRuntimeManager.acquire() --- .../server/runtime/MemoryRuntimeManager.java | 44 +++++++++++------- .../runtime/MemoryRuntimeManagerTest.java | 46 +++++++++++++++++++ 2 files changed, 74 insertions(+), 16 deletions(-) diff --git a/memind-server/src/main/java/com/openmemind/ai/memory/server/runtime/MemoryRuntimeManager.java b/memind-server/src/main/java/com/openmemind/ai/memory/server/runtime/MemoryRuntimeManager.java index 1db1f11c..8128f0b4 100644 --- a/memind-server/src/main/java/com/openmemind/ai/memory/server/runtime/MemoryRuntimeManager.java +++ b/memind-server/src/main/java/com/openmemind/ai/memory/server/runtime/MemoryRuntimeManager.java @@ -16,10 +16,12 @@ import com.openmemind.ai.memory.core.Memory; import com.openmemind.ai.memory.core.builder.MemoryBuildOptions; import java.util.concurrent.atomic.AtomicReference; +import java.util.concurrent.locks.ReentrantLock; public final class MemoryRuntimeManager { private final AtomicReference current; + private final ReentrantLock lock = new ReentrantLock(); public MemoryRuntimeManager() { this.current = new AtomicReference<>(); @@ -30,27 +32,31 @@ public MemoryRuntimeManager(RuntimeHandle initial) { } public RuntimeLease acquire() { - while (true) { + lock.lock(); + try { RuntimeHandle handle = requireCurrentHandle(); - if (handle.draining().get()) { - continue; - } handle.inFlightRequests().incrementAndGet(); - if (handle == current.get() && !handle.draining().get()) { - return new RuntimeLease(handle, () -> release(handle)); - } - release(handle); + return new RuntimeLease(handle, () -> release(handle)); + } + finally { + lock.unlock(); } } public void swap(Memory memory, MemoryBuildOptions options, long version) { - RuntimeHandle next = new RuntimeHandle(memory, options, version); - RuntimeHandle previous = current.getAndSet(next); - if (previous == null) { - return; + lock.lock(); + try { + RuntimeHandle next = new RuntimeHandle(memory, options, version); + RuntimeHandle previous = current.getAndSet(next); + if (previous == null) { + return; + } + previous.draining().set(true); + tryClose(previous); + } + finally { + lock.unlock(); } - previous.draining().set(true); - tryClose(previous); } public long currentVersion() { @@ -62,8 +68,14 @@ public RuntimeHandle currentHandle() { } private void release(RuntimeHandle handle) { - handle.inFlightRequests().decrementAndGet(); - tryClose(handle); + lock.lock(); + try { + handle.inFlightRequests().decrementAndGet(); + tryClose(handle); + } + finally { + lock.unlock(); + } } private RuntimeHandle requireCurrentHandle() { diff --git a/memind-server/src/test/java/com/openmemind/ai/memory/server/runtime/MemoryRuntimeManagerTest.java b/memind-server/src/test/java/com/openmemind/ai/memory/server/runtime/MemoryRuntimeManagerTest.java index 63d63da2..bf4efa1a 100644 --- a/memind-server/src/test/java/com/openmemind/ai/memory/server/runtime/MemoryRuntimeManagerTest.java +++ b/memind-server/src/test/java/com/openmemind/ai/memory/server/runtime/MemoryRuntimeManagerTest.java @@ -18,6 +18,12 @@ import com.openmemind.ai.memory.core.Memory; import com.openmemind.ai.memory.core.builder.MemoryBuildOptions; import com.openmemind.ai.memory.server.support.TestMemory; +import java.util.ArrayList; +import java.util.List; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.Future; import java.util.concurrent.atomic.AtomicInteger; import org.junit.jupiter.api.Test; @@ -72,6 +78,46 @@ void swapCanInstallFirstRuntimeWhenManagerStartsEmpty() { } } + @Test + void concurrentAcquireDuringSwapClosesOldRuntimeExactlyOnce() throws Exception { + AtomicInteger closeCount = new AtomicInteger(); + Memory oldMemory = closeTrackingMemory(closeCount); + MemoryRuntimeManager manager = + new MemoryRuntimeManager( + new RuntimeHandle(oldMemory, MemoryBuildOptions.defaults(), 1)); + + int threads = 4; + int rounds = 200; + ExecutorService executor = Executors.newFixedThreadPool(threads); + try { + CountDownLatch start = new CountDownLatch(1); + List> futures = new ArrayList<>(); + for (int i = 0; i < threads; i++) { + futures.add(executor.submit(() -> { + start.await(); + for (int j = 0; j < rounds; j++) { + try (RuntimeLease ignored = manager.acquire()) { + // hold and release a lease while the runtime is being swapped + } + } + return null; + })); + } + start.countDown(); + for (int version = 2; version <= 20; version++) { + manager.swap(new TestMemory(), MemoryBuildOptions.defaults(), version); + } + for (Future future : futures) { + future.get(); + } + } + finally { + executor.shutdownNow(); + } + + assertThat(closeCount).hasValue(1); + } + private static Memory closeTrackingMemory(AtomicInteger closeCount) { return new TestMemory(closeCount::incrementAndGet); } From cb27a7a7262b3299745809a5a6e2d91130e5723b Mon Sep 17 00:00:00 2001 From: DDT <1786035110@qq.com> Date: Thu, 17 Sep 2026 01:27:31 +0800 Subject: [PATCH 2/2] fix: perform runtime close outside the manager lock --- .../server/runtime/MemoryRuntimeManager.java | 34 +++++++--- .../runtime/MemoryRuntimeManagerTest.java | 64 +++++++++---------- 2 files changed, 57 insertions(+), 41 deletions(-) diff --git a/memind-server/src/main/java/com/openmemind/ai/memory/server/runtime/MemoryRuntimeManager.java b/memind-server/src/main/java/com/openmemind/ai/memory/server/runtime/MemoryRuntimeManager.java index 8128f0b4..14050219 100644 --- a/memind-server/src/main/java/com/openmemind/ai/memory/server/runtime/MemoryRuntimeManager.java +++ b/memind-server/src/main/java/com/openmemind/ai/memory/server/runtime/MemoryRuntimeManager.java @@ -44,6 +44,7 @@ public RuntimeLease acquire() { } public void swap(Memory memory, MemoryBuildOptions options, long version) { + Memory toClose; lock.lock(); try { RuntimeHandle next = new RuntimeHandle(memory, options, version); @@ -52,11 +53,12 @@ public void swap(Memory memory, MemoryBuildOptions options, long version) { return; } previous.draining().set(true); - tryClose(previous); + toClose = shouldClose(previous); } finally { lock.unlock(); } + closeMemory(toClose); } public long currentVersion() { @@ -68,14 +70,16 @@ public RuntimeHandle currentHandle() { } private void release(RuntimeHandle handle) { + Memory toClose; lock.lock(); try { handle.inFlightRequests().decrementAndGet(); - tryClose(handle); + toClose = shouldClose(handle); } finally { lock.unlock(); } + closeMemory(toClose); } private RuntimeHandle requireCurrentHandle() { @@ -88,13 +92,27 @@ private RuntimeHandle requireCurrentHandle() { return handle; } - private void tryClose(RuntimeHandle handle) { + /** + * Returns the memory to close when the given handle is drained and has no + * in-flight requests, otherwise {@code null}. Must be called while holding the + * lock so that at most one caller observes the transition to (draining, zero + * in-flight) and therefore closes the memory exactly once. + */ + private Memory shouldClose(RuntimeHandle handle) { if (handle.draining().get() && handle.inFlightRequests().get() == 0) { - try { - handle.memory().close(); - } catch (Exception e) { - throw new IllegalStateException("Failed to close drained memory runtime", e); - } + return handle.memory(); + } + return null; + } + + private void closeMemory(Memory memory) { + if (memory == null) { + return; + } + try { + memory.close(); + } catch (Exception e) { + throw new IllegalStateException("Failed to close drained memory runtime", e); } } } diff --git a/memind-server/src/test/java/com/openmemind/ai/memory/server/runtime/MemoryRuntimeManagerTest.java b/memind-server/src/test/java/com/openmemind/ai/memory/server/runtime/MemoryRuntimeManagerTest.java index bf4efa1a..8eeda783 100644 --- a/memind-server/src/test/java/com/openmemind/ai/memory/server/runtime/MemoryRuntimeManagerTest.java +++ b/memind-server/src/test/java/com/openmemind/ai/memory/server/runtime/MemoryRuntimeManagerTest.java @@ -18,12 +18,11 @@ import com.openmemind.ai.memory.core.Memory; import com.openmemind.ai.memory.core.builder.MemoryBuildOptions; import com.openmemind.ai.memory.server.support.TestMemory; -import java.util.ArrayList; -import java.util.List; import java.util.concurrent.CountDownLatch; import java.util.concurrent.ExecutorService; import java.util.concurrent.Executors; import java.util.concurrent.Future; +import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicInteger; import org.junit.jupiter.api.Test; @@ -79,43 +78,42 @@ void swapCanInstallFirstRuntimeWhenManagerStartsEmpty() { } @Test - void concurrentAcquireDuringSwapClosesOldRuntimeExactlyOnce() throws Exception { - AtomicInteger closeCount = new AtomicInteger(); - Memory oldMemory = closeTrackingMemory(closeCount); - MemoryRuntimeManager manager = - new MemoryRuntimeManager( - new RuntimeHandle(oldMemory, MemoryBuildOptions.defaults(), 1)); - - int threads = 4; - int rounds = 200; - ExecutorService executor = Executors.newFixedThreadPool(threads); - try { - CountDownLatch start = new CountDownLatch(1); - List> futures = new ArrayList<>(); - for (int i = 0; i < threads; i++) { - futures.add(executor.submit(() -> { - start.await(); - for (int j = 0; j < rounds; j++) { - try (RuntimeLease ignored = manager.acquire()) { - // hold and release a lease while the runtime is being swapped - } - } - return null; - })); + void closeOfOldRuntimeDoesNotBlockAcquireOfNewRuntime() throws Exception { + CountDownLatch closeEntered = new CountDownLatch(1); + CountDownLatch allowClose = new CountDownLatch(1); + Memory blockingMemory = new TestMemory(() -> { + closeEntered.countDown(); + try { + allowClose.await(); } - start.countDown(); - for (int version = 2; version <= 20; version++) { - manager.swap(new TestMemory(), MemoryBuildOptions.defaults(), version); - } - for (Future future : futures) { - future.get(); + catch (InterruptedException e) { + Thread.currentThread().interrupt(); } + }); + MemoryRuntimeManager manager = new MemoryRuntimeManager( + new RuntimeHandle(blockingMemory, MemoryBuildOptions.defaults(), 1)); + + Thread swapper = new Thread( + () -> manager.swap(new TestMemory(), MemoryBuildOptions.defaults(), 2)); + swapper.start(); + assertThat(closeEntered.await(5, TimeUnit.SECONDS)).isTrue(); + + ExecutorService executor = Executors.newSingleThreadExecutor(); + try { + Future version = executor.submit(() -> { + try (RuntimeLease lease = manager.acquire()) { + return lease.handle().version(); + } + }); + // The old runtime is still mid-close; acquiring the new runtime must + // not block on that close (close happens outside the manager lock). + assertThat(version.get(5, TimeUnit.SECONDS)).isEqualTo(2L); } finally { + allowClose.countDown(); executor.shutdownNow(); + swapper.join(5000); } - - assertThat(closeCount).hasValue(1); } private static Memory closeTrackingMemory(AtomicInteger closeCount) {