Update Elasticsearch Store: Use KNN insettad of script_score

- knn instead of script_score, removed initialization
 - only using normalized similarities, adjusted unit test
 - making l2norm's distances consistent with others
 - update dependency version and docs
 - upate autoconfigure ITs
This commit is contained in:
Laura Trotta
2024-04-12 17:42:13 +02:00
committed by Christian Tzolov
parent c9dd336ea3
commit 78f3797f8d
11 changed files with 184 additions and 203 deletions

View File

@@ -164,6 +164,7 @@
<sap.hanadb.version>2.20.11</sap.hanadb.version>
<oracle.version>23.4.0.24.05</oracle.version>
<postgresql.version>42.7.2</postgresql.version>
<elasticsearch-java.version>8.13.3</elasticsearch-java.version>
<milvus.version>2.3.4</milvus.version>
<pinecone.version>0.8.0</pinecone.version>
<fastjson.version>2.0.46</fastjson.version>

View File

@@ -134,11 +134,18 @@ Properties starting with the `spring.ai.vectorstore.elasticsearch.*` prefix are
|`spring.ai.vectorstore.elasticsearch.index-name` | The name of the index to store the vectors. | spring-ai-document-index
|`spring.ai.vectorstore.elasticsearch.dimensions` | The number of dimensions in the vector. | 1536
|`spring.ai.vectorstore.elasticsearch.dense-vector-indexing` | Whether to use dense vector indexing. | true
|`spring.ai.vectorstore.elasticsearch.similarity` | The similarity function to use. | `cosine`
|`spring.ai.vectorstore.elasticsearch.initialize-schema`| whether to initialize the required schema | `false`
|===
The following similarity functions are available:
* cosine
* l2_norm
* dot_product
More details about each in the https://www.elastic.co/guide/en/elasticsearch/reference/master/dense-vector.html#dense-vector-params[Elasticsearch Documentation] on dense vectors.
== Metadata Filtering
You can leverage the generic, portable xref:api/vectordbs.adoc#metadata-filters[metadata filters] with Elasticsearch as well.
@@ -214,10 +221,11 @@ Read the link:https://www.elastic.co/guide/en/elasticsearch/client/java-api-clie
----
@Bean
public RestClient restClient() {
RestClientBuilder builder = RestClient.builder(new HttpHost("<host>", 9200, "http"));
Header[] defaultHeaders = new Header[] { new BasicHeader("Authorization", "Basic <encoded username and password>") };
builder.setDefaultHeaders(defaultHeaders);
return builder.build();
RestClient.builder(new HttpHost("<host>", 9200, "http"))
.setDefaultHeaders(new Header[]{
new BasicHeader("Authorization", "Basic <encoded username and password>")
})
.build();
}
----

View File

@@ -289,6 +289,7 @@
<optional>true</optional>
</dependency>
<!-- Elasticsearch Vector Store-->
<dependency>
<groupId>org.springframework.ai</groupId>
<artifactId>spring-ai-elasticsearch-store</artifactId>
@@ -326,6 +327,14 @@
<optional>true</optional>
</dependency>
<!-- Elastic Search vector store -->
<dependency>
<groupId>co.elastic.clients</groupId>
<artifactId>elasticsearch-java</artifactId>
<version>${elasticsearch-java.version}</version>
<optional>true</optional>
</dependency>
<!-- test dependencies -->
<dependency>

View File

@@ -52,10 +52,7 @@ class ElasticsearchVectorStoreAutoConfiguration {
if (properties.getDimensions() != null) {
elasticsearchVectorStoreOptions.setDimensions(properties.getDimensions());
}
if (properties.isDenseVectorIndexing() != null) {
elasticsearchVectorStoreOptions.setDenseVectorIndexing(properties.isDenseVectorIndexing());
}
if (StringUtils.hasText(properties.getSimilarity())) {
if (properties.getSimilarity() != null) {
elasticsearchVectorStoreOptions.setSimilarity(properties.getSimilarity());
}

View File

@@ -16,6 +16,7 @@
package org.springframework.ai.autoconfigure.vectorstore.elasticsearch;
import org.springframework.ai.autoconfigure.CommonVectorStoreProperties;
import org.springframework.ai.vectorstore.SimilarityFunction;
import org.springframework.boot.context.properties.ConfigurationProperties;
/**
@@ -37,15 +38,10 @@ public class ElasticsearchVectorStoreProperties extends CommonVectorStorePropert
*/
private Integer dimensions;
/**
* Whether to use dense vector indexing.
*/
private Boolean denseVectorIndexing;
/**
* The similarity function to use.
*/
private String similarity;
private SimilarityFunction similarity;
public String getIndexName() {
return this.indexName;
@@ -63,19 +59,11 @@ public class ElasticsearchVectorStoreProperties extends CommonVectorStorePropert
this.dimensions = dimensions;
}
public Boolean isDenseVectorIndexing() {
return denseVectorIndexing;
}
public void setDenseVectorIndexing(Boolean denseVectorIndexing) {
this.denseVectorIndexing = denseVectorIndexing;
}
public String getSimilarity() {
public SimilarityFunction getSimilarity() {
return similarity;
}
public void setSimilarity(String similarity) {
public void setSimilarity(SimilarityFunction similarity) {
this.similarity = similarity;
}

View File

@@ -18,13 +18,12 @@ package org.springframework.ai.autoconfigure.vectorstore.elasticsearch;
import org.awaitility.Awaitility;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.junit.jupiter.params.ParameterizedTest;
import org.junit.jupiter.params.provider.ValueSource;
import org.springframework.ai.autoconfigure.openai.OpenAiAutoConfiguration;
import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration;
import org.springframework.ai.document.Document;
import org.springframework.ai.vectorstore.ElasticsearchVectorStore;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.SimilarityFunction;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.elasticsearch.ElasticsearchRestClientAutoConfiguration;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
@@ -51,8 +50,6 @@ class ElasticsearchVectorStoreAutoConfigurationIT {
"docker.elastic.co/elasticsearch/elasticsearch:8.12.2")
.withEnv("xpack.security.enabled", "false");
private static final String DEFAULT = "default cosine similarity";
private List<Document> documents = List.of(
new Document("1", getText("classpath:/test/data/spring.ai.txt"), Map.of("meta1", "meta1")),
new Document("2", getText("classpath:/test/data/time.shelter.txt"), Map.of()),
@@ -65,21 +62,14 @@ class ElasticsearchVectorStoreAutoConfigurationIT {
.withPropertyValues("spring.elasticsearch.uris=" + elasticsearchContainer.getHttpHostAddress(),
"spring.ai.openai.api-key=" + System.getenv("OPENAI_API_KEY"));
@ParameterizedTest(name = "{0} : {displayName} ")
@ValueSource(strings = { DEFAULT, """
double value = dotProduct(params.query_vector, 'embedding');
return sigmoid(1, Math.E, -value);
""", "1 / (1 + l1norm(params.query_vector, 'embedding'))",
"1 / (1 + l2norm(params.query_vector, 'embedding'))" })
public void addAndSearchTest(String similarityFunction) {
// No parametrized test based on similarity function,
// by default the bean will be created using cosine.
@Test
public void addAndSearchTest() {
this.contextRunner.run(context -> {
ElasticsearchVectorStore vectorStore = context.getBean(ElasticsearchVectorStore.class);
if (!DEFAULT.equals(similarityFunction)) {
vectorStore.withSimilarityFunction(similarityFunction);
}
vectorStore.add(documents);
Awaitility.await()
@@ -120,7 +110,7 @@ class ElasticsearchVectorStoreAutoConfigurationIT {
"spring.ai.vectorstore.elasticsearch.index-name=example",
"spring.ai.vectorstore.elasticsearch.dimensions=1024",
"spring.ai.vectorstore.elasticsearch.dense-vector-indexing=true",
"spring.ai.vectorstore.elasticsearch.similarity=dot_product")
"spring.ai.vectorstore.elasticsearch.similarity=cosine")
.run(context -> {
var properties = context.getBean(ElasticsearchVectorStoreProperties.class);
var elasticsearchVectorStore = context.getBean(ElasticsearchVectorStore.class);
@@ -128,8 +118,7 @@ class ElasticsearchVectorStoreAutoConfigurationIT {
assertThat(properties).isNotNull();
assertThat(properties.getIndexName()).isEqualTo("example");
assertThat(properties.getDimensions()).isEqualTo(1024);
assertThat(properties.isDenseVectorIndexing()).isTrue();
assertThat(properties.getSimilarity()).isEqualTo("dot_product");
assertThat(properties.getSimilarity()).isEqualTo(SimilarityFunction.cosine);
assertThat(elasticsearchVectorStore).isNotNull();
});

View File

@@ -35,6 +35,7 @@
<dependency>
<groupId>co.elastic.clients</groupId>
<artifactId>elasticsearch-java</artifactId>
<version>${elasticsearch-java.version}</version>
</dependency>
<!-- TESTING -->
@@ -45,7 +46,6 @@
<scope>test</scope>
</dependency>
<dependency>
<groupId>org.springframework.ai</groupId>
<artifactId>spring-ai-test</artifactId>

View File

@@ -16,17 +16,13 @@
package org.springframework.ai.vectorstore;
import co.elastic.clients.elasticsearch.ElasticsearchClient;
import co.elastic.clients.elasticsearch._types.mapping.DenseVectorProperty;
import co.elastic.clients.elasticsearch._types.mapping.Property;
import co.elastic.clients.elasticsearch._types.query_dsl.Query;
import co.elastic.clients.elasticsearch.core.BulkRequest;
import co.elastic.clients.elasticsearch.core.BulkResponse;
import co.elastic.clients.elasticsearch.core.SearchResponse;
import co.elastic.clients.elasticsearch.core.bulk.BulkResponseItem;
import co.elastic.clients.elasticsearch.core.search.Hit;
import co.elastic.clients.elasticsearch.indices.CreateIndexResponse;
import co.elastic.clients.json.JsonData;
import co.elastic.clients.json.jackson.JacksonJsonpMapper;
import co.elastic.clients.transport.endpoints.BooleanResponse;
import co.elastic.clients.transport.rest_client.RestClientTransport;
import com.fasterxml.jackson.databind.DeserializationFeature;
import com.fasterxml.jackson.databind.ObjectMapper;
@@ -46,16 +42,17 @@ import java.util.Objects;
import java.util.Optional;
import java.util.stream.Collectors;
import static java.lang.Math.sqrt;
import static org.springframework.ai.vectorstore.SimilarityFunction.l2_norm;
/**
* @author Jemin Huh
* @author Wei Jiang
* @author Laura Trotta
* @since 1.0.0
*/
public class ElasticsearchVectorStore implements VectorStore, InitializingBean {
// divided by 2 to get score in the range [0, 1]
public static final String COSINE_SIMILARITY_FUNCTION = "(cosineSimilarity(params.query_vector, 'embedding') + 1.0) / 2";
private static final Logger logger = LoggerFactory.getLogger(ElasticsearchVectorStore.class);
private final EmbeddingModel embeddingModel;
@@ -66,8 +63,6 @@ public class ElasticsearchVectorStore implements VectorStore, InitializingBean {
private final FilterExpressionConverter filterExpressionConverter;
private String similarityFunction;
private final boolean initializeSchema;
public ElasticsearchVectorStore(RestClient restClient, EmbeddingModel embeddingModel, boolean initializeSchema) {
@@ -84,30 +79,22 @@ public class ElasticsearchVectorStore implements VectorStore, InitializingBean {
this.embeddingModel = embeddingModel;
this.options = options;
this.filterExpressionConverter = new ElasticsearchAiSearchFilterExpressionConverter();
// the potential functions for vector fields at
// https://www.elastic.co/guide/en/elasticsearch/reference/current/query-dsl-script-score-query.html#vector-functions
this.similarityFunction = COSINE_SIMILARITY_FUNCTION;
}
public ElasticsearchVectorStore withSimilarityFunction(String similarityFunction) {
this.similarityFunction = similarityFunction;
return this;
}
@Override
public void add(List<Document> documents) {
BulkRequest.Builder builkRequestBuilder = new BulkRequest.Builder();
BulkRequest.Builder bulkRequestBuilder = new BulkRequest.Builder();
for (Document document : documents) {
if (Objects.isNull(document.getEmbedding()) || document.getEmbedding().isEmpty()) {
logger.debug("Calling EmbeddingModel for document id = " + document.getId());
document.setEmbedding(this.embeddingModel.embed(document));
}
builkRequestBuilder.operations(op -> op
bulkRequestBuilder.operations(op -> op
.index(idx -> idx.index(this.options.getIndexName()).id(document.getId()).document(document)));
}
BulkResponse bulkRequest = bulkRequest(builkRequestBuilder.build());
BulkResponse bulkRequest = bulkRequest(bulkRequestBuilder.build());
if (bulkRequest.errors()) {
List<BulkResponseItem> bulkResponseItems = bulkRequest.items();
@@ -121,10 +108,10 @@ public class ElasticsearchVectorStore implements VectorStore, InitializingBean {
@Override
public Optional<Boolean> delete(List<String> idList) {
BulkRequest.Builder builkRequestBuilder = new BulkRequest.Builder();
BulkRequest.Builder bulkRequestBuilder = new BulkRequest.Builder();
for (String id : idList)
builkRequestBuilder.operations(op -> op.delete(idx -> idx.index(this.options.getIndexName()).id(id)));
return Optional.of(bulkRequest(builkRequestBuilder.build()).errors());
bulkRequestBuilder.operations(op -> op.delete(idx -> idx.index(this.options.getIndexName()).id(id)));
return Optional.of(bulkRequest(bulkRequestBuilder.build()).errors());
}
private BulkResponse bulkRequest(BulkRequest bulkRequest) {
@@ -139,28 +126,34 @@ public class ElasticsearchVectorStore implements VectorStore, InitializingBean {
@Override
public List<Document> similaritySearch(SearchRequest searchRequest) {
Assert.notNull(searchRequest, "The search request must not be null.");
return similaritySearch(this.embeddingModel.embed(searchRequest.getQuery()), searchRequest.getTopK(),
Double.valueOf(searchRequest.getSimilarityThreshold()).floatValue(),
searchRequest.getFilterExpression());
}
try {
float threshold = (float) searchRequest.getSimilarityThreshold();
// reverting l2_norm distance to its original value
if (options.getSimilarity().equals(l2_norm)) {
threshold = 1 - threshold;
}
final float finalThreshold = threshold;
List<Float> vectors = this.embeddingModel.embed(searchRequest.getQuery())
.stream()
.map(Double::floatValue)
.toList();
public List<Document> similaritySearch(List<Double> embedding, int topK, double similarityThreshold,
Filter.Expression filterExpression) {
return similaritySearch(
new co.elastic.clients.elasticsearch.core.SearchRequest.Builder().index(options.getIndexName())
.query(getElasticsearchSimilarityQuery(embedding, filterExpression))
.size(topK)
.minScore(similarityThreshold)
.build());
}
SearchResponse<Document> res = elasticsearchClient.search(
sr -> sr.index(options.getIndexName())
.knn(knn -> knn.queryVector(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);
private Query getElasticsearchSimilarityQuery(List<Double> embedding, Filter.Expression filterExpression) {
return Query.of(queryBuilder -> queryBuilder.scriptScore(scriptScoreQueryBuilder -> scriptScoreQueryBuilder
.query(queryBuilder2 -> queryBuilder2.queryString(queryStringQuerybuilder -> queryStringQuerybuilder
.query(getElasticsearchQueryString(filterExpression))))
.script(scriptBuilder -> scriptBuilder
.inline(inlineScriptBuilder -> inlineScriptBuilder.source(this.similarityFunction)
.params("query_vector", JsonData.of(embedding))))));
return res.hits().hits().stream().map(this::toDocument).collect(Collectors.toList());
}
catch (IOException e) {
throw new RuntimeException(e);
}
}
private String getElasticsearchQueryString(Filter.Expression filterExpression) {
@@ -169,31 +162,31 @@ public class ElasticsearchVectorStore implements VectorStore, InitializingBean {
}
private List<Document> similaritySearch(co.elastic.clients.elasticsearch.core.SearchRequest searchRequest) {
try {
return this.elasticsearchClient.search(searchRequest, Document.class)
.hits()
.hits()
.stream()
.map(this::toDocument)
.collect(Collectors.toList());
}
catch (IOException e) {
throw new RuntimeException(e);
}
}
private Document toDocument(Hit<Document> hit) {
Document document = hit.source();
document.getMetadata().put("distance", 1 - hit.score().floatValue());
document.getMetadata().put("distance", calculateDistance(hit.score().floatValue()));
return document;
}
private boolean indexExists() {
// more info on score/distance calculation
// https://www.elastic.co/guide/en/elasticsearch/reference/current/knn-search.html#knn-similarity-search
private float calculateDistance(Float score) {
switch (options.getSimilarity()) {
case l2_norm:
// the returned value of l2_norm is the opposite of the other functions
// (closest to zero means more accurate), so to make it consistent
// with the other functions the reverse is returned applying a "1-"
// to the standard transformation
return (float) (1 - (sqrt((1 / score) - 1)));
// cosine and dot_product
default:
return (2 * score) - 1;
}
}
public boolean indexExists() {
try {
BooleanResponse response = this.elasticsearchClient.indices()
.exists(existRequestBuilder -> existRequestBuilder.index(options.getIndexName()));
return response.value();
return this.elasticsearchClient.indices().exists(ex -> ex.index(options.getIndexName())).value();
}
catch (IOException e) {
throw new RuntimeException(e);
@@ -203,18 +196,9 @@ public class ElasticsearchVectorStore implements VectorStore, InitializingBean {
private CreateIndexResponse createIndexMapping() {
try {
return this.elasticsearchClient.indices()
.create(createIndexBuilder -> createIndexBuilder.index(options.getIndexName())
.mappings(typeMappingBuilder -> {
typeMappingBuilder.properties("embedding",
new Property.Builder()
.denseVector(new DenseVectorProperty.Builder().dims(options.getDimensions())
.similarity(options.getSimilarity())
.index(options.isDenseVectorIndexing())
.build())
.build());
return typeMappingBuilder;
}));
.create(cr -> cr.index(options.getIndexName())
.mappings(map -> map.properties("embedding", p -> p.denseVector(
dv -> dv.similarity(options.getSimilarity().toString()).dims(options.getDimensions())))));
}
catch (IOException e) {
throw new RuntimeException(e);
@@ -233,4 +217,4 @@ public class ElasticsearchVectorStore implements VectorStore, InitializingBean {
}
}
}
}

View File

@@ -34,15 +34,10 @@ public class ElasticsearchVectorStoreOptions {
*/
private int dimensions = 1536;
/**
* Whether to use dense vector indexing.
*/
private boolean denseVectorIndexing = true;
/**
* The similarity function to use.
*/
private String similarity = "cosine";
private SimilarityFunction similarity = SimilarityFunction.cosine;
public String getIndexName() {
return indexName;
@@ -60,19 +55,11 @@ public class ElasticsearchVectorStoreOptions {
this.dimensions = dims;
}
public boolean isDenseVectorIndexing() {
return denseVectorIndexing;
}
public void setDenseVectorIndexing(boolean denseVectorIndexing) {
this.denseVectorIndexing = denseVectorIndexing;
}
public String getSimilarity() {
public SimilarityFunction getSimilarity() {
return similarity;
}
public void setSimilarity(String similarity) {
public void setSimilarity(SimilarityFunction similarity) {
this.similarity = similarity;
}

View File

@@ -0,0 +1,15 @@
package org.springframework.ai.vectorstore;
/**
* https://www.elastic.co/guide/en/elasticsearch/reference/master/dense-vector.html
* max_inner_product is currently not supported because the distance value is not
* normalized and would not comply with the requirement of being between 0 and 1
*
* @author Laura Trotta
* @since 1.0.0
*/
public enum SimilarityFunction {
l2_norm, dot_product, cosine
}

View File

@@ -25,6 +25,12 @@ import java.util.Map;
import java.util.UUID;
import java.util.concurrent.TimeUnit;
import co.elastic.clients.elasticsearch.ElasticsearchClient;
import co.elastic.clients.elasticsearch.cat.indices.IndicesRecord;
import co.elastic.clients.json.jackson.JacksonJsonpMapper;
import co.elastic.clients.transport.rest_client.RestClientTransport;
import com.fasterxml.jackson.databind.DeserializationFeature;
import com.fasterxml.jackson.databind.ObjectMapper;
import org.apache.http.HttpHost;
import org.awaitility.Awaitility;
import org.elasticsearch.client.RestClient;
@@ -36,7 +42,6 @@ import org.junit.jupiter.params.provider.ValueSource;
import org.testcontainers.elasticsearch.ElasticsearchContainer;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import org.testcontainers.shaded.com.fasterxml.jackson.databind.ObjectMapper;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
@@ -59,14 +64,10 @@ class ElasticsearchVectorStoreIT {
@Container
private static final ElasticsearchContainer elasticsearchContainer = new ElasticsearchContainer(
"docker.elastic.co/elasticsearch/elasticsearch:8.12.2")
"docker.elastic.co/elasticsearch/elasticsearch:8.13.3")
.withEnv("xpack.security.enabled", "false");
private static final String DEFAULT = "default cosine similarity";
protected final ObjectMapper objectMapper = new ObjectMapper();
private List<Document> documents = List.of(
private final List<Document> documents = List.of(
new Document("1", getText("classpath:/test/data/spring.ai.txt"), Map.of("meta1", "meta1")),
new Document("2", getText("classpath:/test/data/time.shelter.txt"), Map.of()),
new Document("3", getText("classpath:/test/data/great.depression.txt"), Map.of("meta2", "meta2")));
@@ -95,35 +96,33 @@ class ElasticsearchVectorStoreIT {
@BeforeEach
void cleanDatabase() {
getContextRunner().run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
vectorStore.delete(List.of("_all"));
// deleting indices and data before following tests
ElasticsearchClient elasticsearchClient = context.getBean(ElasticsearchClient.class);
List indices = elasticsearchClient.cat().indices().valueBody().stream().map(IndicesRecord::index).toList();
if (!indices.isEmpty()) {
elasticsearchClient.indices().delete(del -> del.index(indices));
}
});
}
@ParameterizedTest(name = "{0} : {displayName} ")
@ValueSource(strings = { DEFAULT, """
double value = dotProduct(params.query_vector, 'embedding');
return sigmoid(1, Math.E, -value);
""", "1 / (1 + l1norm(params.query_vector, 'embedding'))",
"1 / (1 + l2norm(params.query_vector, 'embedding'))" })
@ValueSource(strings = { "cosine", "l2_norm", "dot_product" })
public void addAndSearchTest(String similarityFunction) {
getContextRunner().run(context -> {
ElasticsearchVectorStore vectorStore = context.getBean(ElasticsearchVectorStore.class);
if (!DEFAULT.equals(similarityFunction)) {
vectorStore.withSimilarityFunction(similarityFunction);
}
ElasticsearchVectorStore vectorStore = context.getBean("vectorStore_" + similarityFunction,
ElasticsearchVectorStore.class);
vectorStore.add(documents);
Awaitility.await()
.until(() -> vectorStore
.similaritySearch(SearchRequest.query("Great Depression").withTopK(1).withSimilarityThreshold(0)),
.similaritySearch(SearchRequest.query("Great Depression").withTopK(1).withSimilarityThresholdAll()),
hasSize(1));
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("Great Depression").withTopK(1).withSimilarityThreshold(0));
.similaritySearch(SearchRequest.query("Great Depression").withTopK(1).withSimilarityThresholdAll());
assertThat(results).hasSize(1);
Document resultDoc = results.get(0);
@@ -138,25 +137,18 @@ class ElasticsearchVectorStoreIT {
Awaitility.await()
.until(() -> vectorStore
.similaritySearch(SearchRequest.query("Great Depression").withTopK(1).withSimilarityThreshold(0)),
.similaritySearch(SearchRequest.query("Great Depression").withTopK(1).withSimilarityThresholdAll()),
hasSize(0));
});
}
@ParameterizedTest(name = "{0} : {displayName} ")
@ValueSource(strings = { DEFAULT, """
double value = dotProduct(params.query_vector, 'embedding');
return sigmoid(1, Math.E, -value);
""", "1 / (1 + l1norm(params.query_vector, 'embedding'))",
"1 / (1 + l2norm(params.query_vector, 'embedding'))" })
@ValueSource(strings = { "cosine", "l2_norm", "dot_product" })
public void searchWithFilters(String similarityFunction) {
getContextRunner().run(context -> {
ElasticsearchVectorStore vectorStore = context.getBean(ElasticsearchVectorStore.class);
if (!DEFAULT.equals(similarityFunction)) {
vectorStore.withSimilarityFunction(similarityFunction);
}
ElasticsearchVectorStore vectorStore = context.getBean("vectorStore_" + similarityFunction,
ElasticsearchVectorStore.class);
var bgDocument = new Document("1", "The World is Big and Salvation Lurks Around the Corner",
Map.of("country", "BG", "year", 2020, "activationDate", new Date(1000)));
@@ -168,7 +160,9 @@ class ElasticsearchVectorStoreIT {
vectorStore.add(List.of(bgDocument, nlDocument, bgDocument2));
Awaitility.await()
.until(() -> vectorStore.similaritySearch(SearchRequest.query("The World").withTopK(5)), hasSize(3));
.until(() -> vectorStore
.similaritySearch(SearchRequest.query("The World").withTopK(5).withSimilarityThresholdAll()),
hasSize(3));
List<Document> results = vectorStore.similaritySearch(SearchRequest.query("The World")
.withTopK(5)
@@ -246,18 +240,12 @@ class ElasticsearchVectorStoreIT {
}
@ParameterizedTest(name = "{0} : {displayName} ")
@ValueSource(strings = { DEFAULT, """
double value = dotProduct(params.query_vector, 'embedding');
return sigmoid(1, Math.E, -value);
""", "1 / (1 + l1norm(params.query_vector, 'embedding'))",
"1 / (1 + l2norm(params.query_vector, 'embedding'))" })
@ValueSource(strings = { "cosine", "l2_norm", "dot_product" })
public void documentUpdateTest(String similarityFunction) {
getContextRunner().run(context -> {
ElasticsearchVectorStore vectorStore = context.getBean(ElasticsearchVectorStore.class);
if (!DEFAULT.equals(similarityFunction)) {
vectorStore.withSimilarityFunction(similarityFunction);
}
ElasticsearchVectorStore vectorStore = context.getBean("vectorStore_" + similarityFunction,
ElasticsearchVectorStore.class);
Document document = new Document(UUID.randomUUID().toString(), "Spring AI rocks!!",
Map.of("meta1", "meta1"));
@@ -265,11 +253,11 @@ class ElasticsearchVectorStoreIT {
Awaitility.await()
.until(() -> vectorStore
.similaritySearch(SearchRequest.query("Spring").withSimilarityThreshold(0).withTopK(5)),
.similaritySearch(SearchRequest.query("Spring").withSimilarityThresholdAll().withTopK(5)),
hasSize(1));
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("Spring").withSimilarityThreshold(0).withTopK(5));
.similaritySearch(SearchRequest.query("Spring").withSimilarityThresholdAll().withTopK(5));
assertThat(results).hasSize(1);
Document resultDoc = results.get(0);
@@ -282,7 +270,7 @@ class ElasticsearchVectorStoreIT {
"The World is Big and Salvation Lurks Around the Corner", Map.of("meta2", "meta2"));
vectorStore.add(List.of(sameIdDocument));
SearchRequest fooBarSearchRequest = SearchRequest.query("FooBar").withTopK(5);
SearchRequest fooBarSearchRequest = SearchRequest.query("FooBar").withTopK(5).withSimilarityThresholdAll();
Awaitility.await()
.until(() -> vectorStore.similaritySearch(fooBarSearchRequest).get(0).getContent(),
@@ -306,24 +294,15 @@ class ElasticsearchVectorStoreIT {
}
@ParameterizedTest(name = "{0} : {displayName} ")
@ValueSource(strings = { DEFAULT, """
double value = dotProduct(params.query_vector, 'embedding');
return sigmoid(1, Math.E, -value);
""", "1 / (1 + l1norm(params.query_vector, 'embedding'))",
"1 / (1 + l2norm(params.query_vector, 'embedding'))" })
@ValueSource(strings = { "cosine", "l2_norm", "dot_product" })
public void searchThresholdTest(String similarityFunction) {
getContextRunner().run(context -> {
ElasticsearchVectorStore vectorStore = context.getBean(ElasticsearchVectorStore.class);
if (!DEFAULT.equals(similarityFunction)) {
vectorStore.withSimilarityFunction(similarityFunction);
}
ElasticsearchVectorStore vectorStore = context.getBean("vectorStore_" + similarityFunction,
ElasticsearchVectorStore.class);
vectorStore.add(documents);
SearchRequest query = SearchRequest.query("Great Depression")
.withTopK(50)
.withSimilarityThreshold(SearchRequest.SIMILARITY_THRESHOLD_ACCEPT_ALL);
SearchRequest query = SearchRequest.query("Great Depression").withTopK(50).withSimilarityThresholdAll();
Awaitility.await().until(() -> vectorStore.similaritySearch(query), hasSize(3));
@@ -333,10 +312,10 @@ class ElasticsearchVectorStoreIT {
assertThat(distances).hasSize(3);
float threshold = (distances.get(0) + distances.get(1)) / 2;
float thresholdResult = (distances.get(0) + distances.get(1)) / 2;
List<Document> results = vectorStore.similaritySearch(
SearchRequest.query("Great Depression").withTopK(50).withSimilarityThreshold(1 - threshold));
SearchRequest.query("Great Depression").withTopK(50).withSimilarityThreshold(thresholdResult));
assertThat(results).hasSize(1);
Document resultDoc = results.get(0);
@@ -349,9 +328,8 @@ class ElasticsearchVectorStoreIT {
vectorStore.delete(documents.stream().map(Document::getId).toList());
Awaitility.await()
.until(() -> vectorStore
.similaritySearch(SearchRequest.query("Great Depression").withTopK(50).withSimilarityThreshold(0)),
hasSize(0));
.until(() -> vectorStore.similaritySearch(
SearchRequest.query("Great Depression").withTopK(50).withSimilarityThresholdAll()), hasSize(0));
});
}
@@ -359,11 +337,25 @@ class ElasticsearchVectorStoreIT {
@EnableAutoConfiguration(exclude = { DataSourceAutoConfiguration.class })
public static class TestApplication {
@Bean
public ElasticsearchVectorStore vectorStore(EmbeddingModel embeddingModel) {
return new ElasticsearchVectorStore(
RestClient.builder(HttpHost.create(elasticsearchContainer.getHttpHostAddress())).build(),
embeddingModel, true);
@Bean("vectorStore_cosine")
public ElasticsearchVectorStore vectorStoreDefault(EmbeddingModel embeddingModel, RestClient restClient) {
return new ElasticsearchVectorStore(restClient, embeddingModel, true);
}
@Bean("vectorStore_l2_norm")
public ElasticsearchVectorStore vectorStoreL2(EmbeddingModel embeddingModel, RestClient restClient) {
ElasticsearchVectorStoreOptions options = new ElasticsearchVectorStoreOptions();
options.setIndexName("index_l2");
options.setSimilarity(SimilarityFunction.l2_norm);
return new ElasticsearchVectorStore(options, restClient, embeddingModel, true);
}
@Bean("vectorStore_dot_product")
public ElasticsearchVectorStore vectorStoreDotProduct(EmbeddingModel embeddingModel, RestClient restClient) {
ElasticsearchVectorStoreOptions options = new ElasticsearchVectorStoreOptions();
options.setIndexName("index_dot_product");
options.setSimilarity(SimilarityFunction.dot_product);
return new ElasticsearchVectorStore(options, restClient, embeddingModel, true);
}
@Bean
@@ -371,6 +363,17 @@ class ElasticsearchVectorStoreIT {
return new OpenAiEmbeddingModel(new OpenAiApi(System.getenv("OPENAI_API_KEY")));
}
@Bean
RestClient restClient() {
return RestClient.builder(HttpHost.create(elasticsearchContainer.getHttpHostAddress())).build();
}
@Bean
ElasticsearchClient elasticsearchClient(RestClient restClient) {
return new ElasticsearchClient(new RestClientTransport(restClient, new JacksonJsonpMapper(
new ObjectMapper().configure(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES, false))));
}
}
}