diff --git a/config/engine.conf/src/main/java/io/aklivity/zilla/config/engine/StoredConfig.java b/config/engine.conf/src/main/java/io/aklivity/zilla/config/engine/StoredConfig.java new file mode 100644 index 0000000000..9062f59f8e --- /dev/null +++ b/config/engine.conf/src/main/java/io/aklivity/zilla/config/engine/StoredConfig.java @@ -0,0 +1,39 @@ +/* + * Copyright 2021-2026 Aklivity Inc + * + * Licensed under the Aklivity Community License (the "License"); you may not use + * this file except in compliance with the License. You may obtain a copy of the + * License at + * + * https://www.aklivity.io/aklivity-community-license/ + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, WITHOUT + * WARRANTIES OF ANY KIND, either express or implied. See the License for the + * specific language governing permissions and limitations under the License. + */ +package io.aklivity.zilla.config.engine; + +import java.util.Map; +import java.util.function.Function; + +public final class StoredConfig extends NamedConfig +{ + StoredConfig( + String name, + Map extensions) + { + super(name, extensions); + } + + public static StoredConfigBuilder builder( + Function mapper) + { + return new StoredConfigBuilder<>(mapper); + } + + public static StoredConfigBuilder builder() + { + return new StoredConfigBuilder<>(StoredConfig.class::cast); + } +} diff --git a/config/engine.conf/src/main/java/io/aklivity/zilla/config/engine/StoredConfigBuilder.java b/config/engine.conf/src/main/java/io/aklivity/zilla/config/engine/StoredConfigBuilder.java new file mode 100644 index 0000000000..f61946eba7 --- /dev/null +++ b/config/engine.conf/src/main/java/io/aklivity/zilla/config/engine/StoredConfigBuilder.java @@ -0,0 +1,50 @@ +/* + * Copyright 2021-2026 Aklivity Inc + * + * Licensed under the Aklivity Community License (the "License"); you may not use + * this file except in compliance with the License. You may obtain a copy of the + * License at + * + * https://www.aklivity.io/aklivity-community-license/ + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, WITHOUT + * WARRANTIES OF ANY KIND, either express or implied. See the License for the + * specific language governing permissions and limitations under the License. + */ +package io.aklivity.zilla.config.engine; + +import java.util.function.Function; + +public final class StoredConfigBuilder extends ConfigBuilder.Extensible> +{ + private final Function mapper; + + private String name; + + StoredConfigBuilder( + Function mapper) + { + this.mapper = mapper; + } + + @Override + @SuppressWarnings("unchecked") + protected Class> thisType() + { + return (Class>) getClass(); + } + + public StoredConfigBuilder name( + String name) + { + this.name = name; + return this; + } + + @Override + public T build() + { + return mapper.apply(new StoredConfig(name, extensions())); + } +} diff --git a/examples/tcp.echo.embedding/etc/zilla.yaml b/examples/tcp.echo.embedding/etc/zilla.yaml index 9af38c04fb..10480cc785 100644 --- a/examples/tcp.echo.embedding/etc/zilla.yaml +++ b/examples/tcp.echo.embedding/etc/zilla.yaml @@ -3,6 +3,9 @@ name: example embeddings: moderator0: type: glove +stores: + cache0: + type: memory bindings: north_tcp_server: type: tcp @@ -22,6 +25,7 @@ bindings: - "You will never believe what happened next." - "I have a massive secret but I absolutely cannot tell anyone here." threshold: 0.94 + store: cache0 telemetry: exporters: stdout_logs_exporter: diff --git a/incubator/model-vector.conf/src/main/java/io/aklivity/zilla/config/model/vector/VectorModelConfig.java b/incubator/model-vector.conf/src/main/java/io/aklivity/zilla/config/model/vector/VectorModelConfig.java index 60b51f1ad2..f861b9d54f 100644 --- a/incubator/model-vector.conf/src/main/java/io/aklivity/zilla/config/model/vector/VectorModelConfig.java +++ b/incubator/model-vector.conf/src/main/java/io/aklivity/zilla/config/model/vector/VectorModelConfig.java @@ -23,24 +23,28 @@ import io.aklivity.zilla.config.engine.EmbeddedConfig; import io.aklivity.zilla.config.engine.ModelConfig; import io.aklivity.zilla.config.engine.NamedConfig; +import io.aklivity.zilla.config.engine.StoredConfig; public final class VectorModelConfig extends ModelConfig { public final EmbeddedConfig embedding; public final List reject; public final double threshold; + public final StoredConfig store; VectorModelConfig( EmbeddedConfig embedding, List reject, double threshold, + StoredConfig store, Map extensions, List refs) { - super("vector", null, null, extensions, withEmbedding(embedding, refs)); + super("vector", null, null, extensions, withRefs(embedding, store, refs)); this.embedding = embedding; this.reject = reject; this.threshold = threshold; + this.store = store; } public static VectorModelConfigBuilder builder( @@ -54,8 +58,9 @@ public static VectorModelConfigBuilder builder() return new VectorModelConfigBuilder<>(VectorModelConfig.class::cast); } - private static List withEmbedding( + private static List withRefs( EmbeddedConfig embedding, + StoredConfig store, List refs) { List all = new ArrayList<>(); @@ -63,6 +68,10 @@ private static List withEmbedding( { all.add(embedding); } + if (store != null) + { + all.add(store); + } if (refs != null) { all.addAll(refs); diff --git a/incubator/model-vector.conf/src/main/java/io/aklivity/zilla/config/model/vector/VectorModelConfigBuilder.java b/incubator/model-vector.conf/src/main/java/io/aklivity/zilla/config/model/vector/VectorModelConfigBuilder.java index 6b382446ef..e395018d7d 100644 --- a/incubator/model-vector.conf/src/main/java/io/aklivity/zilla/config/model/vector/VectorModelConfigBuilder.java +++ b/incubator/model-vector.conf/src/main/java/io/aklivity/zilla/config/model/vector/VectorModelConfigBuilder.java @@ -20,6 +20,7 @@ import io.aklivity.zilla.config.engine.ConfigBuilder; import io.aklivity.zilla.config.engine.EmbeddedConfig; +import io.aklivity.zilla.config.engine.StoredConfig; public class VectorModelConfigBuilder extends ConfigBuilder.Extensible> { @@ -28,6 +29,7 @@ public class VectorModelConfigBuilder extends ConfigBuilder.Extensible reject; private double threshold; + private StoredConfig store; VectorModelConfigBuilder( Function mapper) @@ -74,9 +76,16 @@ public VectorModelConfigBuilder threshold( return this; } + public VectorModelConfigBuilder store( + String store) + { + this.store = store != null ? StoredConfig.builder().name(store).build() : null; + return this; + } + @Override public T build() { - return mapper.apply(new VectorModelConfig(embedding, reject, threshold, extensions(), refs())); + return mapper.apply(new VectorModelConfig(embedding, reject, threshold, store, extensions(), refs())); } } diff --git a/incubator/model-vector.conf/src/main/java/io/aklivity/zilla/config/model/vector/internal/VectorModelConfigAdapter.java b/incubator/model-vector.conf/src/main/java/io/aklivity/zilla/config/model/vector/internal/VectorModelConfigAdapter.java index 04584c5b35..696f5b0a63 100644 --- a/incubator/model-vector.conf/src/main/java/io/aklivity/zilla/config/model/vector/internal/VectorModelConfigAdapter.java +++ b/incubator/model-vector.conf/src/main/java/io/aklivity/zilla/config/model/vector/internal/VectorModelConfigAdapter.java @@ -37,6 +37,7 @@ public final class VectorModelConfigAdapter extends ConfigAdapter.Extensible> extensions) @@ -66,6 +67,11 @@ public JsonValue adaptToJson( builder.add(THRESHOLD_NAME, model.threshold); + if (model.store != null) + { + builder.add(STORE_NAME, model.store.name); + } + injectExtensions(model, builder); return builder.build(); @@ -89,10 +95,15 @@ public ModelConfig adaptFromJson( ? object.getJsonNumber(THRESHOLD_NAME).doubleValue() : 0.0; + String store = object.containsKey(STORE_NAME) + ? object.getString(STORE_NAME) + : null; + VectorModelConfigBuilder builder = VectorModelConfig.builder() .embedding(embedding) .reject(reject) - .threshold(threshold); + .threshold(threshold) + .store(store); injectExtensions(object, builder); diff --git a/incubator/model-vector.conf/src/test/java/io/aklivity/zilla/config/model/vector/internal/VectorModelConfigAdapterTest.java b/incubator/model-vector.conf/src/test/java/io/aklivity/zilla/config/model/vector/internal/VectorModelConfigAdapterTest.java index 4f207bdbf7..879db9fc5b 100644 --- a/incubator/model-vector.conf/src/test/java/io/aklivity/zilla/config/model/vector/internal/VectorModelConfigAdapterTest.java +++ b/incubator/model-vector.conf/src/test/java/io/aklivity/zilla/config/model/vector/internal/VectorModelConfigAdapterTest.java @@ -56,7 +56,8 @@ public void shouldReadVectorModel() "reject phrase one", "reject phrase two" ], - "threshold": 0.85 + "threshold": 0.85, + "store": "cache0" }"""; // WHEN @@ -68,6 +69,30 @@ public void shouldReadVectorModel() assertThat(config.embedding.name, equalTo("moderator0")); assertThat(config.reject, contains("reject phrase one", "reject phrase two")); assertThat(config.threshold, equalTo(0.85)); + assertThat(config.store.name, equalTo("cache0")); + } + + @Test + public void shouldDefaultStoreToNullWhenAbsent() + { + // GIVEN -- the store property is required by the vector model's own JSON schema, but the + // adapter itself stays defensive about a config built or parsed without going through it + String json = """ + { + "model": "vector", + "embedding": "moderator0", + "reject": + [ + "reject phrase one" + ], + "threshold": 0.85 + }"""; + + // WHEN + VectorModelConfig config = jsonb.fromJson(json, VectorModelConfig.class); + + // THEN + assertThat(config.store, nullValue()); } @Test @@ -83,13 +108,15 @@ public void shouldWriteVectorModel() "\"reject phrase one\"," + "\"reject phrase two\"" + "]," + - "\"threshold\":0.85" + + "\"threshold\":0.85," + + "\"store\":\"cache0\"" + "}"; VectorModelConfig config = VectorModelConfig.builder() .embedding("moderator0") .reject("reject phrase one") .reject("reject phrase two") .threshold(0.85) + .store("cache0") .build(); // WHEN diff --git a/incubator/model-vector.spec/src/main/scripts/io/aklivity/zilla/specs/model/vector/config/model.yaml b/incubator/model-vector.spec/src/main/scripts/io/aklivity/zilla/specs/model/vector/config/model.yaml index 85af243cad..3d6459f300 100644 --- a/incubator/model-vector.spec/src/main/scripts/io/aklivity/zilla/specs/model/vector/config/model.yaml +++ b/incubator/model-vector.spec/src/main/scripts/io/aklivity/zilla/specs/model/vector/config/model.yaml @@ -26,4 +26,5 @@ bindings: reject: - "You will never believe what happened next." threshold: 0.85 + store: cache0 exit: test diff --git a/incubator/model-vector.spec/src/main/scripts/io/aklivity/zilla/specs/model/vector/config/reject.yaml b/incubator/model-vector.spec/src/main/scripts/io/aklivity/zilla/specs/model/vector/config/reject.yaml index 64a9c04b1a..8b9048344c 100644 --- a/incubator/model-vector.spec/src/main/scripts/io/aklivity/zilla/specs/model/vector/config/reject.yaml +++ b/incubator/model-vector.spec/src/main/scripts/io/aklivity/zilla/specs/model/vector/config/reject.yaml @@ -18,6 +18,9 @@ name: test embeddings: embedding0: type: test +stores: + cache0: + type: test bindings: net0: type: test @@ -29,4 +32,5 @@ bindings: reject: - "reject this message" threshold: 0.99 + store: cache0 exit: app0 diff --git a/incubator/model-vector.spec/src/main/scripts/io/aklivity/zilla/specs/model/vector/schema/vector.schema.patch.json b/incubator/model-vector.spec/src/main/scripts/io/aklivity/zilla/specs/model/vector/schema/vector.schema.patch.json index abadcd5a81..85f023a384 100644 --- a/incubator/model-vector.spec/src/main/scripts/io/aklivity/zilla/specs/model/vector/schema/vector.schema.patch.json +++ b/incubator/model-vector.spec/src/main/scripts/io/aklivity/zilla/specs/model/vector/schema/vector.schema.patch.json @@ -48,13 +48,19 @@ "type": "number", "minimum": 0, "maximum": 1 + }, + "store": + { + "title": "Store", + "type": "string" } }, "required": [ "embedding", "reject", - "threshold" + "threshold", + "store" ] } } diff --git a/incubator/model-vector/src/main/java/io/aklivity/zilla/runtime/model/vector/internal/VectorModelHandlerImpl.java b/incubator/model-vector/src/main/java/io/aklivity/zilla/runtime/model/vector/internal/VectorModelHandlerImpl.java index ac7585a126..9e5760a458 100644 --- a/incubator/model-vector/src/main/java/io/aklivity/zilla/runtime/model/vector/internal/VectorModelHandlerImpl.java +++ b/incubator/model-vector/src/main/java/io/aklivity/zilla/runtime/model/vector/internal/VectorModelHandlerImpl.java @@ -14,69 +14,79 @@ */ package io.aklivity.zilla.runtime.model.vector.internal; +import java.nio.charset.StandardCharsets; +import java.security.MessageDigest; +import java.security.NoSuchAlgorithmException; +import java.time.Duration; import java.util.LinkedList; import java.util.List; import io.aklivity.zilla.config.model.vector.VectorModelConfig; import io.aklivity.zilla.runtime.common.vector.Vectors; import io.aklivity.zilla.runtime.engine.EngineContext; +import io.aklivity.zilla.runtime.engine.concurrent.Signaler; import io.aklivity.zilla.runtime.engine.embedding.EmbeddingHandler; import io.aklivity.zilla.runtime.engine.model.ModelCache; import io.aklivity.zilla.runtime.engine.model.ModelEnvelope; import io.aklivity.zilla.runtime.engine.model.ModelHandler; import io.aklivity.zilla.runtime.engine.model.ModelPipeline; import io.aklivity.zilla.runtime.engine.model.ModelTransform; +import io.aklivity.zilla.runtime.engine.store.StoreHandler; // Per-worker factory for the vector model, resolving the named embedding once and embedding the // configured reject phrases once, then vending a fresh per-stream VectorModelPipeline that reuses // this handler's resolved embedding and reject vectors on every supplyDecoder/supplyEncoder call. +// +// Every namespace binding is replicated to and attached independently on every EngineWorker, so +// without deduplication every worker (and, with a distributed store, every replica) would embed +// the same reject phrases independently -- a redundant N-way burst against whatever embedding +// provider is configured, fired all at once at attach time. The required store instead lets only +// the worker that wins a short-lived lock do the real embed call; every other worker polls the +// cache with capped backoff until the winner's result appears, including after the lock owner +// fails without writing one (the lock simply expires and the next poll wins it instead). final class VectorModelHandlerImpl implements ModelHandler { private static final Runnable NOOP = () -> { }; + private static final Duration LOCK_TTL = Duration.ofSeconds(30); + private static final long INITIAL_RETRY_DELAY_MILLIS = 100L; + private static final int CACHE_RETRY_SIGNAL_ID = 1; + private static final String NULL_VECTOR_TOKEN = "-"; + private static final String VECTOR_DELIMITER = ";"; + private static final String COMPONENT_DELIMITER = ","; + private final EmbeddingHandler handler; + private final List reject; private final float[][] rejectVectors; private final double threshold; private final List pending; + private final StoreHandler store; + private final Signaler signaler; + private final String cacheKey; + private final String lockKey; private boolean ready; private int rejectVectorsReceived; + private String lockToken; + private long retryDelayMillis = INITIAL_RETRY_DELAY_MILLIS; VectorModelHandlerImpl( EngineContext context, VectorModelConfig config) { this.handler = context.supplyEmbedding(config.embedding.id); + this.reject = config.reject; this.threshold = config.threshold; this.pending = new LinkedList<>(); this.rejectVectors = new float[config.reject.size()][]; + this.store = context.supplyStore(config.store.id); + this.signaler = context.signaler(); + this.cacheKey = "model.vector.reject." + digest(config.reject); + this.lockKey = cacheKey + ".lock"; - handler.embed(0L, 0L, 0L, config.reject, new EmbeddingHandler.CompletionCallback() - { - @Override - public void completed( - long contextId, - float[][] results) - { - for (int i = 0; i < results.length; i++) - { - onRejectVectorEmbedded(i, results[i]); - } - } - - @Override - public void failed( - long contextId, - Throwable ex) - { - for (int i = 0; i < rejectVectors.length; i++) - { - onRejectVectorEmbedded(i, null); - } - } - }); + store.get(cacheKey, this::onCacheGet); } @Override @@ -156,6 +166,135 @@ boolean matches( return matched; } + private void embedRejectPhrases() + { + handler.embed(0L, 0L, 0L, reject, new EmbeddingHandler.CompletionCallback() + { + @Override + public void completed( + long contextId, + float[][] results) + { + onEmbedComplete(results); + } + + @Override + public void failed( + long contextId, + Throwable ex) + { + onEmbedFailed(); + } + }); + } + + private void onEmbedComplete( + float[][] results) + { + store.put(cacheKey, encode(results), null, ignored -> unlock()); + + for (int i = 0; i < results.length; i++) + { + onRejectVectorEmbedded(i, results[i]); + } + } + + private void onEmbedFailed() + { + // Never cache a failure -- an unlucky transient error would otherwise permanently poison + // every worker's (and, distributed, every replica's) result. Release the lock instead so + // whichever worker polls next re-attempts the real embed call for itself. + unlock(); + + for (int i = 0; i < rejectVectors.length; i++) + { + onRejectVectorEmbedded(i, null); + } + } + + private void onCacheGet( + String key, + String value) + { + float[][] cached = value != null ? decode(value) : null; + if (cached != null) + { + settle(cached); + } + else + { + store.lock(lockKey, LOCK_TTL, this::onLockAcquire); + } + } + + private void onLockAcquire( + String key, + String token) + { + if (token != null) + { + this.lockToken = token; + embedRejectPhrases(); + } + else + { + scheduleCacheRetry(); + } + } + + private void scheduleCacheRetry() + { + signaler.signalAt(System.currentTimeMillis() + retryDelayMillis, CACHE_RETRY_SIGNAL_ID, this::onCacheRetry); + } + + private void onCacheRetry( + int signalId) + { + if (!ready) + { + store.get(cacheKey, this::onCacheRetryGet); + } + } + + private void onCacheRetryGet( + String key, + String value) + { + float[][] cached = value != null ? decode(value) : null; + if (cached != null) + { + settle(cached); + } + else + { + retryDelayMillis = Math.min(retryDelayMillis * 2L, LOCK_TTL.toMillis()); + store.lock(lockKey, LOCK_TTL, this::onLockAcquire); + } + } + + private void unlock() + { + if (lockToken != null) + { + store.unlock(lockKey, lockToken, this::onUnlockComplete); + lockToken = null; + } + } + + private void onUnlockComplete( + String token) + { + } + + private void settle( + float[][] vectors) + { + for (int i = 0; i < vectors.length; i++) + { + onRejectVectorEmbedded(i, vectors[i]); + } + } + private void onRejectVectorEmbedded( int index, float[] vector) @@ -171,4 +310,91 @@ private void onRejectVectorEmbedded( drain.forEach(Runnable::run); } } + + private static String encode( + float[][] vectors) + { + StringBuilder encoded = new StringBuilder(); + for (int i = 0; i < vectors.length; i++) + { + if (i > 0) + { + encoded.append(VECTOR_DELIMITER); + } + + float[] vector = vectors[i]; + if (vector == null) + { + encoded.append(NULL_VECTOR_TOKEN); + } + else + { + for (int j = 0; j < vector.length; j++) + { + if (j > 0) + { + encoded.append(COMPONENT_DELIMITER); + } + encoded.append(vector[j]); + } + } + } + return encoded.toString(); + } + + private float[][] decode( + String encoded) + { + String[] parts = encoded.split(VECTOR_DELIMITER, -1); + float[][] vectors = parts.length == reject.size() ? new float[parts.length][] : null; + + if (vectors != null) + { + for (int i = 0; i < parts.length; i++) + { + String part = parts[i]; + vectors[i] = NULL_VECTOR_TOKEN.equals(part) ? null : asVector(part); + } + } + + return vectors; + } + + private static float[] asVector( + String part) + { + String[] components = part.split(COMPONENT_DELIMITER, -1); + float[] vector = new float[components.length]; + for (int i = 0; i < components.length; i++) + { + vector[i] = Float.parseFloat(components[i]); + } + return vector; + } + + private static String digest( + List reject) + { + StringBuilder canonical = new StringBuilder(); + for (String phrase : reject) + { + canonical.append(phrase.length()).append(':').append(phrase); + } + + try + { + MessageDigest sha256 = MessageDigest.getInstance("SHA-256"); + byte[] hash = sha256.digest(canonical.toString().getBytes(StandardCharsets.UTF_8)); + StringBuilder hex = new StringBuilder(hash.length * 2); + for (byte b : hash) + { + hex.append(String.format("%02x", b)); + } + return hex.toString(); + } + catch (NoSuchAlgorithmException ex) + { + throw new IllegalStateException(ex); + } + } } diff --git a/incubator/model-vector/src/test/java/io/aklivity/zilla/runtime/model/vector/internal/VectorModelContextTest.java b/incubator/model-vector/src/test/java/io/aklivity/zilla/runtime/model/vector/internal/VectorModelContextTest.java index 6478e73731..c8ba2dd2a2 100644 --- a/incubator/model-vector/src/test/java/io/aklivity/zilla/runtime/model/vector/internal/VectorModelContextTest.java +++ b/incubator/model-vector/src/test/java/io/aklivity/zilla/runtime/model/vector/internal/VectorModelContextTest.java @@ -28,6 +28,7 @@ import io.aklivity.zilla.runtime.engine.embedding.EmbeddingHandler; import io.aklivity.zilla.runtime.engine.model.ModelContext; import io.aklivity.zilla.runtime.engine.model.ModelHandler; +import io.aklivity.zilla.runtime.engine.store.StoreHandler; public class VectorModelContextTest { @@ -36,12 +37,14 @@ public void shouldSupplyHandler() { EngineContext engine = mock(EngineContext.class); when(engine.supplyEmbedding(anyLong())).thenReturn(mock(EmbeddingHandler.class)); + when(engine.supplyStore(anyLong())).thenReturn(mock(StoreHandler.class)); ModelContext context = new VectorModelContext(engine); ModelConfig config = VectorModelConfig.builder() .embedding("moderator0") .reject("reject phrase") .threshold(0.85) + .store("cache0") .build(); ModelHandler handler = context.supplyHandler(config); diff --git a/incubator/model-vector/src/test/java/io/aklivity/zilla/runtime/model/vector/internal/VectorModelHandlerImplTest.java b/incubator/model-vector/src/test/java/io/aklivity/zilla/runtime/model/vector/internal/VectorModelHandlerImplTest.java new file mode 100644 index 0000000000..43775bf070 --- /dev/null +++ b/incubator/model-vector/src/test/java/io/aklivity/zilla/runtime/model/vector/internal/VectorModelHandlerImplTest.java @@ -0,0 +1,341 @@ +/* + * Copyright 2021-2026 Aklivity Inc + * + * Licensed under the Aklivity Community License (the "License"); you may not use + * this file except in compliance with the License. You may obtain a copy of the + * License at + * + * https://www.aklivity.io/aklivity-community-license/ + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, WITHOUT + * WARRANTIES OF ANY KIND, either express or implied. See the License for the + * specific language governing permissions and limitations under the License. + */ +package io.aklivity.zilla.runtime.model.vector.internal; + +import static org.hamcrest.MatcherAssert.assertThat; +import static org.hamcrest.Matchers.equalTo; +import static org.mockito.ArgumentMatchers.anyLong; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +import java.nio.charset.StandardCharsets; +import java.time.Instant; +import java.util.ArrayDeque; +import java.util.List; +import java.util.Queue; +import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.ConcurrentMap; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.function.IntConsumer; + +import org.junit.Before; +import org.junit.Test; + +import io.aklivity.zilla.config.model.vector.VectorModelConfig; +import io.aklivity.zilla.runtime.common.agrona.buffer.DirectBufferEx; +import io.aklivity.zilla.runtime.common.agrona.buffer.UnsafeBufferEx; +import io.aklivity.zilla.runtime.engine.EngineContext; +import io.aklivity.zilla.runtime.engine.concurrent.Signaler; +import io.aklivity.zilla.runtime.engine.embedding.EmbeddingHandler; +import io.aklivity.zilla.runtime.engine.model.ModelPipelineResult; +import io.aklivity.zilla.runtime.engine.model.ModelStatus; +import io.aklivity.zilla.runtime.engine.store.StoreHandler; +import io.aklivity.zilla.runtime.engine.test.internal.embedding.TestEmbeddingHandler; +import io.aklivity.zilla.runtime.engine.test.internal.store.TestStoreHandler; + +// Simulates two per-worker VectorModelHandlerImpl instances sharing one store (as store-memory +// shares its backing map across every EngineWorker in one process) to prove only one of them ever +// calls EmbeddingHandler.embed() for the same reject phrases, and both still resolve correctly. +// Each simulated worker gets its own task queue (its embedding handler and store handler both +// defer completions onto it) so the test can drain one worker at a time, deterministically, while +// the two share only the underlying store maps -- exactly the part being deduplicated. +@SuppressWarnings({ "rawtypes", "unchecked" }) +public class VectorModelHandlerImplTest +{ + private static final int FLAGS_FIN = 0x01; + private static final Runnable NOOP = () -> + { + }; + + // Raw: TestStoreHandler's own lock-entry/watcher value types are package-private, so these + // fields can only be typed generically enough to still be passed to its constructor unchanged. + private final ConcurrentMap entries = new ConcurrentHashMap<>(); + private final ConcurrentMap listeners = new ConcurrentHashMap(); + private final ConcurrentMap locks = new ConcurrentHashMap(); + private final FakeSignaler signaler = new FakeSignaler(); + private final AtomicInteger embedCalls = new AtomicInteger(); + + private VectorModelConfig config; + + @Before + public void init() + { + config = VectorModelConfig.builder() + .embedding("moderator0") + .reject("reject this message") + .threshold(0.99) + .store("cache0") + .build(); + } + + @Test + public void shouldEmbedOnceAcrossTwoWorkersSharingAStore() + { + // GIVEN -- first worker attaches: cache miss, wins the lock, starts (but hasn't finished) a + // real, counted embed call + Queue tasks1 = new ArrayDeque<>(); + VectorModelHandlerImpl first = new VectorModelHandlerImpl(newWorkerContext(tasks1), config); + drainOne(tasks1); + drainOne(tasks1); + assertThat(embedCalls.get(), equalTo(1)); + + // WHEN -- second worker attaches before the first worker's embed call has completed: + // cache still empty, loses the lock race, and schedules a retry instead of embedding + Queue tasks2 = new ArrayDeque<>(); + VectorModelHandlerImpl second = new VectorModelHandlerImpl(newWorkerContext(tasks2), config); + drainOne(tasks2); + drainOne(tasks2); + + // THEN -- second lost the lock race and scheduled a retry instead of embedding + assertThat(embedCalls.get(), equalTo(1)); + assertThat(signaler.pending(), equalTo(1)); + + // WHEN -- the winner's embed call completes and caches the result + drainOne(tasks1); + + // WHEN -- the loser's scheduled retry fires and picks up the cached result + signaler.fireNext(); + drain(tasks2); + + // THEN -- exactly one real reject-phrase embed call total, ever, across both workers + // (isRejected below triggers its own, unrelated, per-message embed calls, so this is the + // last point at which embedCalls only reflects reject-phrase warm-up work) + assertThat(embedCalls.get(), equalTo(1)); + + // THEN -- both workers resolve messages correctly regardless of which one actually embedded + assertThat(isRejected(first, tasks1, "reject this message"), equalTo(true)); + assertThat(isRejected(first, tasks1, "an unrelated message"), equalTo(false)); + assertThat(isRejected(second, tasks2, "reject this message"), equalTo(true)); + assertThat(isRejected(second, tasks2, "an unrelated message"), equalTo(false)); + } + + @Test + public void shouldNotCacheAFailureAndShouldReleaseTheLockForTheNextAttempt() + { + // GIVEN -- embed() itself fails for the one worker holding the lock + Queue tasks1 = new ArrayDeque<>(); + EngineContext context = mock(EngineContext.class); + EmbeddingHandler failing = new EmbeddingHandler() + { + @Override + public void embed( + long traceId, + long bindingId, + long contextId, + List texts, + CompletionCallback completion) + { + embedCalls.incrementAndGet(); + tasks1.add(() -> completion.failed(contextId, new RuntimeException("boom"))); + } + }; + when(context.supplyEmbedding(anyLong())).thenReturn(failing); + when(context.supplyStore(anyLong())).thenReturn(newStoreHandler(tasks1)); + when(context.signaler()).thenReturn(signaler); + + VectorModelHandlerImpl handler = new VectorModelHandlerImpl(context, config); + drain(tasks1); + + // THEN -- the failure was never cached, and the lock was released rather than left held + assertThat(embedCalls.get(), equalTo(1)); + assertThat(entries.isEmpty(), equalTo(true)); + assertThat(locks.isEmpty(), equalTo(true)); + + // WHEN -- a second worker attaches after the failure, using a real embedding backend + Queue tasks2 = new ArrayDeque<>(); + VectorModelHandlerImpl retried = new VectorModelHandlerImpl(newWorkerContext(tasks2), config); + drain(tasks2); + + // THEN -- exactly one more real reject-phrase embed call, by whichever worker asks next + // (isRejected below triggers its own, unrelated, per-message embed calls, so this is the + // last point at which embedCalls only reflects reject-phrase warm-up work) + assertThat(embedCalls.get(), equalTo(2)); + + // THEN -- the first worker's own local (failed) vectors never reject anything, but the + // second worker's successful, cached result does + assertThat(isRejected(handler, tasks1, "reject this message"), equalTo(false)); + assertThat(isRejected(retried, tasks2, "reject this message"), equalTo(true)); + } + + private EngineContext newWorkerContext( + Queue tasks) + { + EngineContext context = mock(EngineContext.class); + EmbeddingHandler delegate = new TestEmbeddingHandler(tasks::add); + EmbeddingHandler counting = new EmbeddingHandler() + { + @Override + public void embed( + long traceId, + long bindingId, + long contextId, + List texts, + CompletionCallback completion) + { + embedCalls.incrementAndGet(); + delegate.embed(traceId, bindingId, contextId, texts, completion); + } + }; + when(context.supplyEmbedding(anyLong())).thenReturn(counting); + when(context.supplyStore(anyLong())).thenReturn(newStoreHandler(tasks)); + when(context.signaler()).thenReturn(signaler); + return context; + } + + private StoreHandler newStoreHandler( + Queue tasks) + { + return new TestStoreHandler(null, tasks::add, entries, listeners, locks); + } + + private static void drain( + Queue tasks) + { + while (!tasks.isEmpty()) + { + tasks.poll().run(); + } + } + + private static void drainOne( + Queue tasks) + { + tasks.poll().run(); + } + + private boolean isRejected( + VectorModelHandlerImpl handler, + Queue tasks, + String text) + { + VectorModelPipeline pipeline = new VectorModelPipeline(handler, NOOP); + byte[] bytes = text.getBytes(StandardCharsets.UTF_8); + UnsafeBufferEx src = new UnsafeBufferEx(bytes); + UnsafeBufferEx dst = new UnsafeBufferEx(new byte[128]); + + pipeline.transform(0L, 0L, 0L, FLAGS_FIN, src, 0, src.capacity(), dst, 0, dst.capacity()); + drain(tasks); + ModelPipelineResult result = pipeline.transform(0L, 0L, 0L, 0x00, src, 0, 0, dst, 0, dst.capacity()); + + return result.status() == ModelStatus.REJECTED; + } + + private static final class FakeSignaler implements Signaler + { + private final Queue scheduled = new ArrayDeque<>(); + + int pending() + { + return scheduled.size(); + } + + void fireNext() + { + scheduled.poll().accept(0); + } + + @Override + public long signalAt( + long timeMillis, + int signalId, + IntConsumer handler) + { + scheduled.add(handler); + return 1L; + } + + @Override + public long signalAt( + Instant time, + int signalId, + IntConsumer handler) + { + throw new UnsupportedOperationException(); + } + + @Override + public void signalNow( + long originId, + long routedId, + long streamId, + long traceId, + int signalId, + int contextId) + { + throw new UnsupportedOperationException(); + } + + @Override + public void signalNow( + long originId, + long routedId, + long streamId, + long traceId, + int signalId, + int contextId, + DirectBufferEx buffer, + int offset, + int length) + { + throw new UnsupportedOperationException(); + } + + @Override + public long signalAt( + long timeMillis, + long originId, + long routedId, + long streamId, + long traceId, + int signalId, + int contextId) + { + throw new UnsupportedOperationException(); + } + + @Override + public long signalAt( + Instant time, + long originId, + long routedId, + long streamId, + long traceId, + int signalId, + int contextId) + { + throw new UnsupportedOperationException(); + } + + @Override + public long signalTask( + Runnable task, + long originId, + long routedId, + long streamId, + long traceId, + int signalId, + int contextId) + { + throw new UnsupportedOperationException(); + } + + @Override + public boolean cancel( + long cancelId) + { + throw new UnsupportedOperationException(); + } + } +} diff --git a/incubator/model-vector/src/test/java/io/aklivity/zilla/runtime/model/vector/internal/VectorModelPipelineTest.java b/incubator/model-vector/src/test/java/io/aklivity/zilla/runtime/model/vector/internal/VectorModelPipelineTest.java index 45feea65cd..7db946e870 100644 --- a/incubator/model-vector/src/test/java/io/aklivity/zilla/runtime/model/vector/internal/VectorModelPipelineTest.java +++ b/incubator/model-vector/src/test/java/io/aklivity/zilla/runtime/model/vector/internal/VectorModelPipelineTest.java @@ -23,6 +23,8 @@ import java.nio.charset.StandardCharsets; import java.util.ArrayDeque; import java.util.Queue; +import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.ConcurrentMap; import org.junit.Before; import org.junit.Test; @@ -30,11 +32,15 @@ import io.aklivity.zilla.config.model.vector.VectorModelConfig; import io.aklivity.zilla.runtime.common.agrona.buffer.UnsafeBufferEx; import io.aklivity.zilla.runtime.engine.EngineContext; +import io.aklivity.zilla.runtime.engine.concurrent.Signaler; import io.aklivity.zilla.runtime.engine.embedding.EmbeddingHandler; import io.aklivity.zilla.runtime.engine.model.ModelPipelineResult; import io.aklivity.zilla.runtime.engine.model.ModelStatus; +import io.aklivity.zilla.runtime.engine.store.StoreHandler; import io.aklivity.zilla.runtime.engine.test.internal.embedding.TestEmbeddingHandler; +import io.aklivity.zilla.runtime.engine.test.internal.store.TestStoreHandler; +@SuppressWarnings({ "rawtypes", "unchecked" }) public class VectorModelPipelineTest { private static final int FLAGS_FIN = 0x01; @@ -50,10 +56,19 @@ public void init() EmbeddingHandler embedding = new TestEmbeddingHandler(tasks::add); when(context.supplyEmbedding(anyLong())).thenReturn(embedding); + // Raw: TestStoreHandler's own lock-entry/watcher value types are package-private. + ConcurrentMap entries = new ConcurrentHashMap<>(); + ConcurrentMap listeners = new ConcurrentHashMap(); + ConcurrentMap locks = new ConcurrentHashMap(); + StoreHandler store = new TestStoreHandler(null, tasks::add, entries, listeners, locks); + when(context.supplyStore(anyLong())).thenReturn(store); + when(context.signaler()).thenReturn(mock(Signaler.class)); + VectorModelConfig config = VectorModelConfig.builder() .embedding("moderator0") .reject("reject this message") .threshold(0.99) + .store("cache0") .build(); handler = new VectorModelHandlerImpl(context, config);