Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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<RuntimeHandle> current;
private final ReentrantLock lock = new ReentrantLock();

public MemoryRuntimeManager() {
this.current = new AtomicReference<>();
Expand All @@ -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() {
Expand All @@ -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);
}
Comment on lines +74 to +78
finally {
lock.unlock();
}
closeMemory(toClose);
}

private RuntimeHandle requireCurrentHandle() {
Expand All @@ -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);
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -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;

Expand Down Expand Up @@ -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<Long> 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);
}
Expand Down