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..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 @@ -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,33 @@ 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; + Memory toClose; + lock.lock(); + try { + RuntimeHandle next = new RuntimeHandle(memory, options, version); + RuntimeHandle previous = current.getAndSet(next); + if (previous == null) { + return; + } + previous.draining().set(true); + toClose = shouldClose(previous); + } + finally { + lock.unlock(); } - previous.draining().set(true); - tryClose(previous); + closeMemory(toClose); } public long currentVersion() { @@ -62,8 +70,16 @@ public RuntimeHandle currentHandle() { } private void release(RuntimeHandle handle) { - handle.inFlightRequests().decrementAndGet(); - tryClose(handle); + Memory toClose; + lock.lock(); + try { + handle.inFlightRequests().decrementAndGet(); + toClose = shouldClose(handle); + } + finally { + lock.unlock(); + } + closeMemory(toClose); } private RuntimeHandle requireCurrentHandle() { @@ -76,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 63d63da2..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,6 +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.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; @@ -72,6 +77,45 @@ void swapCanInstallFirstRuntimeWhenManagerStartsEmpty() { } } + @Test + void closeOfOldRuntimeDoesNotBlockAcquireOfNewRuntime() throws Exception { + CountDownLatch closeEntered = new CountDownLatch(1); + CountDownLatch allowClose = new CountDownLatch(1); + Memory blockingMemory = new TestMemory(() -> { + closeEntered.countDown(); + try { + allowClose.await(); + } + 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); + } + } + private static Memory closeTrackingMemory(AtomicInteger closeCount) { return new TestMemory(closeCount::incrementAndGet); }