From 12e3c7be1decb524fb38f076c3592968eea2e407 Mon Sep 17 00:00:00 2001 From: Christian Tzolov Date: Thu, 28 Sep 2023 06:25:52 +0200 Subject: [PATCH] Stabilize the vector store ITs - Ensures that the vector store similarity search by threshold tests use dynamically computed threshold that is between the top 2 results from the ordered search. This ensures that the threshold value is not affected by changes in the embedding API results. - Minor code style improvments. --- .../embedding/AzureOpenAiEmbeddingClient.java | 8 +--- .../embedding/OpenAiEmbeddingClient.java | 2 +- .../ai/vectorstore/MilvusVectorStore.java | 1 - .../ai/vectorstore/MilvusVectorStoreIT.java | 4 +- .../ai/vectorstore/Neo4jVectorStore.java | 2 + .../ai/vectorstore/Neo4jVectorStoreIT.java | 42 ++++++++++++------- .../ai/vectorstore/PgVectorStore.java | 3 +- .../ai/vectorstore/PgVectorStoreIT.java | 13 +++--- 8 files changed, 40 insertions(+), 35 deletions(-) diff --git a/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/embedding/AzureOpenAiEmbeddingClient.java b/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/embedding/AzureOpenAiEmbeddingClient.java index f936c3c7b..79622d2bb 100644 --- a/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/embedding/AzureOpenAiEmbeddingClient.java +++ b/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/embedding/AzureOpenAiEmbeddingClient.java @@ -60,11 +60,7 @@ public class AzureOpenAiEmbeddingClient implements EmbeddingClient { } private List extractEmbeddingsList(Embeddings embeddings) { - return embeddings.getData() - .stream() - .map(EmbeddingItem::getEmbedding) - .flatMap(List::stream) - .collect(Collectors.toList()); + return embeddings.getData().stream().map(EmbeddingItem::getEmbedding).flatMap(List::stream).toList(); } @Override @@ -72,7 +68,7 @@ public class AzureOpenAiEmbeddingClient implements EmbeddingClient { logger.debug("Retrieving embeddings"); Embeddings embeddings = this.azureOpenAiClient.getEmbeddings(this.model, new EmbeddingsOptions(texts)); logger.debug("Embeddings retrieved"); - return embeddings.getData().stream().map(emb -> emb.getEmbedding()).collect(Collectors.toList()); + return embeddings.getData().stream().map(emb -> emb.getEmbedding()).toList(); } @Override diff --git a/spring-ai-openai/src/main/java/org/springframework/ai/openai/embedding/OpenAiEmbeddingClient.java b/spring-ai-openai/src/main/java/org/springframework/ai/openai/embedding/OpenAiEmbeddingClient.java index a7ac0652c..7410ede44 100644 --- a/spring-ai-openai/src/main/java/org/springframework/ai/openai/embedding/OpenAiEmbeddingClient.java +++ b/spring-ai-openai/src/main/java/org/springframework/ai/openai/embedding/OpenAiEmbeddingClient.java @@ -61,7 +61,7 @@ public class OpenAiEmbeddingClient implements EmbeddingClient { public List> embed(List texts) { EmbeddingResponse embeddingResponse = embedForResponse(texts); - return embeddingResponse.getData().stream().map(emb -> emb.getEmbedding()).collect(Collectors.toList()); + return embeddingResponse.getData().stream().map(emb -> emb.getEmbedding()).toList(); } @Override diff --git a/vector-stores/spring-ai-milvus-store/src/main/java/org/springframework/ai/vectorstore/MilvusVectorStore.java b/vector-stores/spring-ai-milvus-store/src/main/java/org/springframework/ai/vectorstore/MilvusVectorStore.java index d44760c02..f6466346b 100644 --- a/vector-stores/spring-ai-milvus-store/src/main/java/org/springframework/ai/vectorstore/MilvusVectorStore.java +++ b/vector-stores/spring-ai-milvus-store/src/main/java/org/springframework/ai/vectorstore/MilvusVectorStore.java @@ -448,7 +448,6 @@ public class MilvusVectorStore implements VectorStore, SmartLifecycle { CreateCollectionParam createCollectionReq = CreateCollectionParam.newBuilder() .withDatabaseName(this.config.databaseName) .withCollectionName(this.config.collectionName) - // .withDatabaseName(this.collectionName) .withDescription("Spring AI Vector Store") .withConsistencyLevel(ConsistencyLevelEnum.STRONG) .withShardsNum(2) diff --git a/vector-stores/spring-ai-milvus-store/src/test/java/org/springframework/ai/vectorstore/MilvusVectorStoreIT.java b/vector-stores/spring-ai-milvus-store/src/test/java/org/springframework/ai/vectorstore/MilvusVectorStoreIT.java index aadbd47dd..763b296fb 100644 --- a/vector-stores/spring-ai-milvus-store/src/test/java/org/springframework/ai/vectorstore/MilvusVectorStoreIT.java +++ b/vector-stores/spring-ai-milvus-store/src/test/java/org/springframework/ai/vectorstore/MilvusVectorStoreIT.java @@ -198,7 +198,9 @@ public class MilvusVectorStoreIT { assertThat(distances).hasSize(3); - List results = vectorStore.similaritySearch("Great", 5, (1 - (distances.get(0) + 0.001))); + float threshold = (distances.get(0) + distances.get(1)) / 2; + + List results = vectorStore.similaritySearch("Great", 5, (1 - threshold)); assertThat(results).hasSize(1); Document resultDoc = results.get(0); diff --git a/vector-stores/spring-ai-neo4j-store/src/main/java/org/springframework/ai/vectorstore/Neo4jVectorStore.java b/vector-stores/spring-ai-neo4j-store/src/main/java/org/springframework/ai/vectorstore/Neo4jVectorStore.java index b82359d5a..e1e30a414 100644 --- a/vector-stores/spring-ai-neo4j-store/src/main/java/org/springframework/ai/vectorstore/Neo4jVectorStore.java +++ b/vector-stores/spring-ai-neo4j-store/src/main/java/org/springframework/ai/vectorstore/Neo4jVectorStore.java @@ -345,7 +345,9 @@ public class Neo4jVectorStore implements VectorStore, InitializingBean { private static Document recordToDocument(org.neo4j.driver.Record neoRecord) { var node = neoRecord.get("node").asNode(); + var score = neoRecord.get("score").asFloat(); var metaData = new HashMap(); + metaData.put("distance", 1 - score); node.keys().forEach(key -> { if (key.startsWith("metadata.")) { metaData.put(key.substring(key.indexOf(".") + 1), node.get(key).asObject()); diff --git a/vector-stores/spring-ai-neo4j-store/src/test/java/org/springframework/ai/vectorstore/Neo4jVectorStoreIT.java b/vector-stores/spring-ai-neo4j-store/src/test/java/org/springframework/ai/vectorstore/Neo4jVectorStoreIT.java index adecbe7fe..eddb6b64b 100644 --- a/vector-stores/spring-ai-neo4j-store/src/test/java/org/springframework/ai/vectorstore/Neo4jVectorStoreIT.java +++ b/vector-stores/spring-ai-neo4j-store/src/test/java/org/springframework/ai/vectorstore/Neo4jVectorStoreIT.java @@ -1,10 +1,19 @@ package org.springframework.ai.vectorstore; +import java.util.Collections; +import java.util.List; +import java.util.UUID; + import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import org.neo4j.driver.AuthTokens; import org.neo4j.driver.Driver; import org.neo4j.driver.GraphDatabase; +import org.testcontainers.containers.Neo4jContainer; +import org.testcontainers.junit.jupiter.Container; +import org.testcontainers.junit.jupiter.Testcontainers; +import org.testcontainers.utility.DockerImageName; + import org.springframework.ai.autoconfigure.openai.OpenAiAutoConfiguration; import org.springframework.ai.document.Document; import org.springframework.ai.embedding.EmbeddingClient; @@ -14,15 +23,6 @@ import org.springframework.boot.autoconfigure.EnableAutoConfiguration; import org.springframework.boot.autoconfigure.jdbc.DataSourceAutoConfiguration; import org.springframework.boot.test.context.runner.ApplicationContextRunner; import org.springframework.context.annotation.Bean; -import org.testcontainers.containers.Neo4jContainer; -import org.testcontainers.junit.jupiter.Container; -import org.testcontainers.junit.jupiter.Testcontainers; -import org.testcontainers.utility.DockerImageName; - -import java.util.Collections; -import java.util.List; -import java.util.UUID; -import java.util.stream.Collectors; import static org.assertj.core.api.Assertions.assertThat; @@ -72,10 +72,11 @@ class Neo4jVectorStoreIT { assertThat(resultDoc.getId()).isEqualTo(this.documents.get(2).getId()); assertThat(resultDoc.getText()).isEqualTo( "Great Depression Great Depression Great Depression Great Depression Great Depression Great Depression"); - assertThat(resultDoc.getMetadata()).isEqualTo(Collections.singletonMap("meta2", "meta2")); + assertThat(resultDoc.getMetadata()).containsKey("meta2"); + assertThat(resultDoc.getMetadata()).containsKey("distance"); // Remove all documents from the store - vectorStore.delete(this.documents.stream().map(Document::getId).collect(Collectors.toList())); + vectorStore.delete(this.documents.stream().map(Document::getId).toList()); List results2 = vectorStore.similaritySearch("Great", 1); assertThat(results2).isEmpty(); @@ -100,7 +101,8 @@ class Neo4jVectorStoreIT { Document resultDoc = results.get(0); assertThat(resultDoc.getId()).isEqualTo(document.getId()); assertThat(resultDoc.getText()).isEqualTo("Spring AI rocks!!"); - assertThat(resultDoc.getMetadata()).isEqualTo(Collections.singletonMap("meta1", "meta1")); + assertThat(resultDoc.getMetadata()).containsKey("meta1"); + assertThat(resultDoc.getMetadata()).containsKey("distance"); Document sameIdDocument = new Document(document.getId(), "The World is Big and Salvation Lurks Around the Corner", @@ -114,7 +116,8 @@ class Neo4jVectorStoreIT { resultDoc = results.get(0); assertThat(resultDoc.getId()).isEqualTo(document.getId()); assertThat(resultDoc.getText()).isEqualTo("The World is Big and Salvation Lurks Around the Corner"); - assertThat(resultDoc.getMetadata()).isEqualTo(Collections.singletonMap("meta2", "meta2")); + assertThat(resultDoc.getMetadata()).containsKey("meta2"); + assertThat(resultDoc.getMetadata()).containsKey("distance"); }); } @@ -128,16 +131,23 @@ class Neo4jVectorStoreIT { vectorStore.add(this.documents); - assertThat(vectorStore.similaritySearch("Great", 5, 0)).hasSize(3); + List fullResult = vectorStore.similaritySearch("Great", 5, 0); - List results = vectorStore.similaritySearch("Great", 5, 0.89); + List distances = fullResult.stream().map(doc -> (Float) doc.getMetadata().get("distance")).toList(); + + assertThat(distances).hasSize(3); + + float threshold = (distances.get(0) + distances.get(1)) / 2; + + List results = vectorStore.similaritySearch("Great", 5, 1 - threshold); assertThat(results).hasSize(1); Document resultDoc = results.get(0); assertThat(resultDoc.getId()).isEqualTo(this.documents.get(2).getId()); assertThat(resultDoc.getText()).isEqualTo( "Great Depression Great Depression Great Depression Great Depression Great Depression Great Depression"); - assertThat(resultDoc.getMetadata()).isEqualTo(Collections.singletonMap("meta2", "meta2")); + assertThat(resultDoc.getMetadata()).containsKey("meta2"); + assertThat(resultDoc.getMetadata()).containsKey("distance"); }); } diff --git a/vector-stores/spring-ai-pgvector-store/src/main/java/org/springframework/ai/vectorstore/PgVectorStore.java b/vector-stores/spring-ai-pgvector-store/src/main/java/org/springframework/ai/vectorstore/PgVectorStore.java index 9a6f50d81..6972babde 100644 --- a/vector-stores/spring-ai-pgvector-store/src/main/java/org/springframework/ai/vectorstore/PgVectorStore.java +++ b/vector-stores/spring-ai-pgvector-store/src/main/java/org/springframework/ai/vectorstore/PgVectorStore.java @@ -23,7 +23,6 @@ import java.util.Map; import java.util.Optional; import java.util.UUID; import java.util.concurrent.atomic.AtomicBoolean; -import java.util.stream.Collectors; import java.util.stream.IntStream; import com.fasterxml.jackson.core.JsonProcessingException; @@ -158,7 +157,7 @@ public class PgVectorStore implements VectorStore, SmartLifecycle { List doubleEmbedding = IntStream.range(0, floatArray.length) .mapToDouble(i -> floatArray[i]) .boxed() - .collect(Collectors.toList()); + .toList(); return doubleEmbedding; } diff --git a/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorStoreIT.java b/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorStoreIT.java index 30708277f..aab80b211 100644 --- a/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorStoreIT.java +++ b/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorStoreIT.java @@ -20,7 +20,6 @@ import java.util.Collections; import java.util.Iterator; import java.util.List; import java.util.UUID; -import java.util.stream.Collectors; import javax.sql.DataSource; @@ -102,7 +101,7 @@ public class PgVectorStoreIT { assertThat(resultDoc.getMetadata()).containsKeys("meta2", "distance"); // Remove all documents from the store - vectorStore.delete(documents.stream().map(doc -> doc.getId()).collect(Collectors.toList())); + vectorStore.delete(documents.stream().map(doc -> doc.getId()).toList()); List results2 = vectorStore.similaritySearch("Great", 1); assertThat(results2).hasSize(0); @@ -165,7 +164,7 @@ public class PgVectorStoreIT { List distances = fullResult.stream() .map(doc -> (Float) doc.getMetadata().get("distance")) - .collect(Collectors.toList()); + .toList(); assertThat(fullResult).hasSize(3); @@ -173,9 +172,9 @@ public class PgVectorStoreIT { fullResult.stream().forEach(doc -> System.out.println(doc.getMetadata().get("distance"))); - List embeddingDistance = ((PgVectorStore) vectorStore).embeddingDistance("Great"); + float threshold = (distances.get(0) + distances.get(1)) / 2; - List results = vectorStore.similaritySearch("Great", 5, (1 - (distances.get(0) + 0.01))); + List results = vectorStore.similaritySearch("Great", 5, (1 - threshold)); assertThat(results).hasSize(1); Document resultDoc = results.get(0); @@ -189,9 +188,7 @@ public class PgVectorStoreIT { private static boolean isSortedByDistance(List docs) { - List distances = docs.stream() - .map(doc -> (Float) doc.getMetadata().get("distance")) - .collect(Collectors.toList()); + List distances = docs.stream().map(doc -> (Float) doc.getMetadata().get("distance")).toList(); if (CollectionUtils.isEmpty(distances) || distances.size() == 1) { return true;