From 95675a85f8d774b3560abcdbe4207c137b86d04d Mon Sep 17 00:00:00 2001 From: dafriz Date: Thu, 2 Jan 2025 21:38:58 +1100 Subject: [PATCH] Set ElasticSearch size to match requested topK used in KNN search --- .../ElasticsearchVectorStore.java | 19 +++++---- .../ElasticsearchVectorStoreIT.java | 40 +++++++++++++++++++ 2 files changed, 49 insertions(+), 10 deletions(-) diff --git a/vector-stores/spring-ai-elasticsearch-store/src/main/java/org/springframework/ai/vectorstore/elasticsearch/ElasticsearchVectorStore.java b/vector-stores/spring-ai-elasticsearch-store/src/main/java/org/springframework/ai/vectorstore/elasticsearch/ElasticsearchVectorStore.java index 445a7a2e1..f9b69f746 100644 --- a/vector-stores/spring-ai-elasticsearch-store/src/main/java/org/springframework/ai/vectorstore/elasticsearch/ElasticsearchVectorStore.java +++ b/vector-stores/spring-ai-elasticsearch-store/src/main/java/org/springframework/ai/vectorstore/elasticsearch/ElasticsearchVectorStore.java @@ -242,16 +242,15 @@ public class ElasticsearchVectorStore extends AbstractObservationVectorStore imp final float finalThreshold = threshold; float[] vectors = this.embeddingModel.embed(searchRequest.getQuery()); - SearchResponse res = this.elasticsearchClient.search( - sr -> sr.index(this.options.getIndexName()) - .knn(knn -> knn.queryVector(EmbeddingUtils.toList(vectors)) - .similarity(finalThreshold) - .k((long) searchRequest.getTopK()) - .field("embedding") - .numCandidates((long) (1.5 * searchRequest.getTopK())) - .filter(fl -> fl.queryString( - qs -> qs.query(getElasticsearchQueryString(searchRequest.getFilterExpression()))))), - Document.class); + SearchResponse res = this.elasticsearchClient.search(sr -> sr.index(this.options.getIndexName()) + .knn(knn -> knn.queryVector(EmbeddingUtils.toList(vectors)) + .similarity(finalThreshold) + .k((long) searchRequest.getTopK()) + .field("embedding") + .numCandidates((long) (1.5 * searchRequest.getTopK())) + .filter(fl -> fl + .queryString(qs -> qs.query(getElasticsearchQueryString(searchRequest.getFilterExpression()))))) + .size(searchRequest.getTopK()), Document.class); return res.hits().hits().stream().map(this::toDocument).collect(Collectors.toList()); } diff --git a/vector-stores/spring-ai-elasticsearch-store/src/test/java/org/springframework/ai/vectorstore/elasticsearch/ElasticsearchVectorStoreIT.java b/vector-stores/spring-ai-elasticsearch-store/src/test/java/org/springframework/ai/vectorstore/elasticsearch/ElasticsearchVectorStoreIT.java index d5589b724..e35b98c16 100644 --- a/vector-stores/spring-ai-elasticsearch-store/src/test/java/org/springframework/ai/vectorstore/elasticsearch/ElasticsearchVectorStoreIT.java +++ b/vector-stores/spring-ai-elasticsearch-store/src/test/java/org/springframework/ai/vectorstore/elasticsearch/ElasticsearchVectorStoreIT.java @@ -20,6 +20,7 @@ import java.io.IOException; import java.nio.charset.StandardCharsets; import java.time.Duration; import java.time.ZonedDateTime; +import java.util.ArrayList; import java.util.Date; import java.util.List; import java.util.Map; @@ -400,6 +401,45 @@ class ElasticsearchVectorStoreIT { }); } + @Test + public void overDefaultSizeTest() { + + var overDefaultSize = 12; + + getContextRunner().run(context -> { + + ElasticsearchVectorStore vectorStore = context.getBean("vectorStore_cosine", + ElasticsearchVectorStore.class); + + var testDocs = new ArrayList(); + for (int i = 0; i < overDefaultSize; i++) { + testDocs.add(new Document(String.valueOf(i), "Great Depression " + i, Map.of())); + } + vectorStore.add(testDocs); + + Awaitility.await() + .until(() -> vectorStore.similaritySearch( + SearchRequest.builder().query("Great Depression").topK(1).similarityThresholdAll().build()), + hasSize(1)); + + List results = vectorStore.similaritySearch(SearchRequest.builder() + .query("Great Depression") + .topK(overDefaultSize) + .similarityThresholdAll() + .build()); + + assertThat(results).hasSize(overDefaultSize); + + // Remove all documents from the store + vectorStore.delete(testDocs.stream().map(Document::getId).toList()); + + Awaitility.await() + .until(() -> vectorStore.similaritySearch( + SearchRequest.builder().query("Great Depression").topK(1).similarityThresholdAll().build()), + hasSize(0)); + }); + } + @SpringBootConfiguration @EnableAutoConfiguration(exclude = { DataSourceAutoConfiguration.class }) public static class TestApplication {