Add observability support to existing vector stores

Add observability support to:
 - Cassandra
 - Chroma
 - Elasticsearch
 - Milvus
 - Neo4j
 - OpenSearch
 - Qdrant
 - Redis
 - Typesense
 - Weaviate
 - Pinecone
 - Oracle
 - Gemifire
 - MongoDB
 - HanaDB

 Add autoconfiguration obsrvability for the above vector stores.
 Add integration tests for all vector stores.
This commit is contained in:
Christian Tzolov
2024-08-17 08:37:12 +02:00
committed by Mark Pollack
parent 3b7522b6c0
commit 93fa2bf45a
98 changed files with 4684 additions and 341 deletions

View File

@@ -23,7 +23,23 @@ public enum VectorStoreProvider {
// @formatter:off
PG_VECTOR("pg_vector"),
SIMPLE_VECTOR_STORE("simple_vector_store");
AZURE("azure"),
CASSANDRA("cassandra"),
CHROMA("chroma"),
ELASTICSEARCH("elasticsearch"),
MILVUS("milvus"),
NEO4J("neo4j"),
OPENSEARCH("opensearch"),
QDRANT("qdrant"),
REDIS("redis"),
TYPESENSE("typesense"),
WEAVIATE("weaviate"),
PINECONE("pinecone"),
ORACLE("oracle"),
MONGODB("mongodb"),
GEMFIRE("gemfire"),
HANA("hana"),
SIMPLE("simple");
// @formatter:on
private final String value;

View File

@@ -40,6 +40,7 @@ import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.observation.conventions.VectorStoreSimilarityMetric;
import org.springframework.ai.vectorstore.observation.AbstractObservationVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.core.io.Resource;
import com.fasterxml.jackson.core.JsonProcessingException;
@@ -47,6 +48,8 @@ import com.fasterxml.jackson.core.type.TypeReference;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.fasterxml.jackson.databind.ObjectWriter;
import io.micrometer.observation.ObservationRegistry;
/**
* SimpleVectorStore is a simple implementation of the VectorStore interface.
*
@@ -71,6 +74,14 @@ public class SimpleVectorStore extends AbstractObservationVectorStore {
protected EmbeddingModel embeddingModel;
public SimpleVectorStore(EmbeddingModel embeddingModel) {
this(embeddingModel, ObservationRegistry.NOOP, null);
}
public SimpleVectorStore(EmbeddingModel embeddingModel, ObservationRegistry observationRegistry,
VectorStoreObservationConvention customObservationConvention) {
super(observationRegistry, customObservationConvention);
Objects.requireNonNull(embeddingModel, "EmbeddingModel must not be null");
this.embeddingModel = embeddingModel;
}
@@ -265,7 +276,7 @@ public class SimpleVectorStore extends AbstractObservationVectorStore {
@Override
public VectorStoreObservationContext.Builder createObservationContextBuilder(String operationName) {
return VectorStoreObservationContext.builder(VectorStoreProvider.SIMPLE_VECTOR_STORE.value(), operationName)
return VectorStoreObservationContext.builder(VectorStoreProvider.SIMPLE.value(), operationName)
.withDimensions(this.embeddingModel.dimensions())
.withCollectionName("in-memory-map")
.withSimilarityMetric(VectorStoreSimilarityMetric.COSINE.value());

View File

@@ -38,18 +38,10 @@ public abstract class AbstractObservationVectorStore implements VectorStore {
@Nullable
private final VectorStoreObservationConvention customObservationConvention;
public AbstractObservationVectorStore() {
this(ObservationRegistry.NOOP, null);
}
public AbstractObservationVectorStore(ObservationRegistry observationRegistry) {
this(observationRegistry, null);
}
public AbstractObservationVectorStore(ObservationRegistry observationRegistry,
VectorStoreObservationConvention customSearchObservationConvention) {
VectorStoreObservationConvention customObservationConvention) {
this.observationRegistry = observationRegistry;
this.customObservationConvention = customSearchObservationConvention;
this.customObservationConvention = customObservationConvention;
}
@Override

View File

@@ -44,6 +44,9 @@ public class DefaultVectorStoreObservationConvention implements VectorStoreObser
private static final KeyValue TOP_K_NONE = KeyValue.of(HighCardinalityKeyNames.TOP_K, KeyValue.NONE_VALUE);
private static final KeyValue SIMILARITY_THRESHOLD_NONE = KeyValue.of(HighCardinalityKeyNames.SIMILARITY_THRESHOLD,
KeyValue.NONE_VALUE);
private static final KeyValue SIMILARITY_METRIC_NONE = KeyValue.of(HighCardinalityKeyNames.SIMILARITY_METRIC,
KeyValue.NONE_VALUE);
@@ -89,7 +92,7 @@ public class DefaultVectorStoreObservationConvention implements VectorStoreObser
public KeyValues getHighCardinalityKeyValues(VectorStoreObservationContext context) {
return KeyValues.of(query(context), metadataFilter(context), topK(context), dimensions(context),
similarityMetric(context), collectionName(context), namespace(context), fieldName(context),
indexName(context));
indexName(context), similarityThreshold(context));
}
protected KeyValue springAiKind() {
@@ -133,6 +136,14 @@ public class DefaultVectorStoreObservationConvention implements VectorStoreObser
return TOP_K_NONE;
}
protected KeyValue similarityThreshold(VectorStoreObservationContext context) {
if (context.getQueryRequest() != null && context.getQueryRequest().getSimilarityThreshold() >= 0) {
return KeyValue.of(HighCardinalityKeyNames.SIMILARITY_THRESHOLD,
"" + context.getQueryRequest().getSimilarityThreshold());
}
return SIMILARITY_THRESHOLD_NONE;
}
protected KeyValue similarityMetric(VectorStoreObservationContext context) {
if (StringUtils.hasText(context.getSimilarityMetric())) {
return KeyValue.of(HighCardinalityKeyNames.SIMILARITY_METRIC, context.getSimilarityMetric());

View File

@@ -27,7 +27,6 @@ import io.micrometer.observation.Observation;
* @author Christian Tzolov
* @since 1.0.0
*/
public class VectorStoreObservationContext extends Observation.Context {
public enum Operation {

View File

@@ -128,6 +128,17 @@ public enum VectorStoreObservationDocumentation implements ObservationDocumentat
return "db.vector.query.top_k";
}
},
/**
* Similarity threshold that accepts all search scores. A threshold value of 0.0
* means any similarity is accepted or disable the similarity threshold filtering.
* A threshold value of 1.0 means an exact match is required.
*/
SIMILARITY_THRESHOLD {
@Override
public String asString() {
return "db.vector.query.similarity_threshold";
}
},
/**
* The dimension of the vector.
*/

View File

@@ -519,6 +519,12 @@
<scope>test</scope>
</dependency>
<dependency>
<groupId>io.micrometer</groupId>
<artifactId>micrometer-observation-test</artifactId>
<scope>test</scope>
</dependency>
</dependencies>
</project>

View File

@@ -19,8 +19,14 @@ import com.azure.core.credential.AzureKeyCredential;
import com.azure.search.documents.indexes.SearchIndexClient;
import com.azure.search.documents.indexes.SearchIndexClientBuilder;
import io.micrometer.observation.ObservationRegistry;
import java.util.List;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.vectorstore.azure.AzureVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.ObjectProvider;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
@@ -48,9 +54,12 @@ public class AzureVectorStoreAutoConfiguration {
@Bean
@ConditionalOnMissingBean
public AzureVectorStore vectorStore(SearchIndexClient searchIndexClient, EmbeddingModel embeddingModel,
AzureVectorStoreProperties properties) {
AzureVectorStoreProperties properties, ObjectProvider<ObservationRegistry> observationRegistry,
ObjectProvider<VectorStoreObservationConvention> customObservationConvention) {
var vectorStore = new AzureVectorStore(searchIndexClient, embeddingModel, properties.isInitializeSchema());
var vectorStore = new AzureVectorStore(searchIndexClient, embeddingModel, properties.isInitializeSchema(),
List.of(), observationRegistry.getIfUnique(() -> ObservationRegistry.NOOP),
customObservationConvention.getIfAvailable(() -> null));
vectorStore.setIndexName(properties.getIndexName());

View File

@@ -20,9 +20,13 @@ import java.time.Duration;
import com.datastax.oss.driver.api.core.CqlSession;
import com.datastax.oss.driver.api.core.config.DefaultDriverOption;
import io.micrometer.observation.ObservationRegistry;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.vectorstore.CassandraVectorStore;
import org.springframework.ai.vectorstore.CassandraVectorStoreConfig;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.ObjectProvider;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.cassandra.CassandraAutoConfiguration;
import org.springframework.boot.autoconfigure.cassandra.DriverConfigLoaderBuilderCustomizer;
@@ -33,6 +37,7 @@ import org.springframework.context.annotation.Bean;
/**
* @author Mick Semb Wever
* @author Christian Tzolov
* @since 1.0.0
*/
@AutoConfiguration(after = CassandraAutoConfiguration.class)
@@ -43,7 +48,8 @@ public class CassandraVectorStoreAutoConfiguration {
@Bean
@ConditionalOnMissingBean
public CassandraVectorStore vectorStore(EmbeddingModel embeddingModel, CassandraVectorStoreProperties properties,
CqlSession cqlSession) {
CqlSession cqlSession, ObjectProvider<ObservationRegistry> observationRegistry,
ObjectProvider<VectorStoreObservationConvention> customObservationConvention) {
var builder = CassandraVectorStoreConfig.builder().withCqlSession(cqlSession);
@@ -61,7 +67,9 @@ public class CassandraVectorStoreAutoConfiguration {
builder = builder.returnEmbeddings();
}
return CassandraVectorStore.create(builder.build(), embeddingModel);
return new CassandraVectorStore(builder.build(), embeddingModel,
observationRegistry.getIfUnique(() -> ObservationRegistry.NOOP),
customObservationConvention.getIfAvailable(() -> null));
}
@Bean

View File

@@ -18,6 +18,8 @@ package org.springframework.ai.autoconfigure.vectorstore.chroma;
import org.springframework.ai.chroma.ChromaApi;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.vectorstore.ChromaVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.ObjectProvider;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
@@ -29,6 +31,8 @@ import org.springframework.web.client.RestClient;
import com.fasterxml.jackson.databind.ObjectMapper;
import io.micrometer.observation.ObservationRegistry;
/**
* @author Christian Tzolov
* @author Eddú Meléndez
@@ -72,9 +76,11 @@ public class ChromaVectorStoreAutoConfiguration {
@Bean
@ConditionalOnMissingBean
public ChromaVectorStore vectorStore(EmbeddingModel embeddingModel, ChromaApi chromaApi,
ChromaVectorStoreProperties storeProperties) {
ChromaVectorStoreProperties storeProperties, ObjectProvider<ObservationRegistry> observationRegistry,
ObjectProvider<VectorStoreObservationConvention> customObservationConvention) {
return new ChromaVectorStore(embeddingModel, chromaApi, storeProperties.getCollectionName(),
storeProperties.isInitializeSchema());
storeProperties.isInitializeSchema(), observationRegistry.getIfUnique(() -> ObservationRegistry.NOOP),
customObservationConvention.getIfAvailable(() -> null));
}
static class PropertiesChromaConnectionDetails implements ChromaConnectionDetails {

View File

@@ -20,6 +20,8 @@ import org.elasticsearch.client.RestClient;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.vectorstore.ElasticsearchVectorStore;
import org.springframework.ai.vectorstore.ElasticsearchVectorStoreOptions;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.ObjectProvider;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
@@ -28,10 +30,13 @@ import org.springframework.boot.context.properties.EnableConfigurationProperties
import org.springframework.context.annotation.Bean;
import org.springframework.util.StringUtils;
import io.micrometer.observation.ObservationRegistry;
/**
* @author Eddú Meléndez
* @author Wei Jiang
* @author Josh Long
* @author Christian Tzolov
* @since 1.0.0
*/
@@ -43,7 +48,8 @@ class ElasticsearchVectorStoreAutoConfiguration {
@Bean
@ConditionalOnMissingBean
ElasticsearchVectorStore vectorStore(ElasticsearchVectorStoreProperties properties, RestClient restClient,
EmbeddingModel embeddingModel) {
EmbeddingModel embeddingModel, ObjectProvider<ObservationRegistry> observationRegistry,
ObjectProvider<VectorStoreObservationConvention> customObservationConvention) {
ElasticsearchVectorStoreOptions elasticsearchVectorStoreOptions = new ElasticsearchVectorStoreOptions();
if (StringUtils.hasText(properties.getIndexName())) {
@@ -57,7 +63,8 @@ class ElasticsearchVectorStoreAutoConfiguration {
}
return new ElasticsearchVectorStore(elasticsearchVectorStoreOptions, restClient, embeddingModel,
properties.isInitializeSchema());
properties.isInitializeSchema(), observationRegistry.getIfUnique(() -> ObservationRegistry.NOOP),
customObservationConvention.getIfAvailable(() -> null));
}
}

View File

@@ -19,6 +19,8 @@ package org.springframework.ai.autoconfigure.vectorstore.gemfire;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.vectorstore.GemFireVectorStore;
import org.springframework.ai.vectorstore.GemFireVectorStoreConfig;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.ObjectProvider;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
@@ -26,8 +28,11 @@ import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty;
import org.springframework.boot.context.properties.EnableConfigurationProperties;
import org.springframework.context.annotation.Bean;
import io.micrometer.observation.ObservationRegistry;
/**
* @author Geet Rawat
* @author Christian Tzolov
*/
@AutoConfiguration
@ConditionalOnClass({ GemFireVectorStore.class, EmbeddingModel.class })
@@ -45,7 +50,8 @@ public class GemFireVectorStoreAutoConfiguration {
@Bean
@ConditionalOnMissingBean
public GemFireVectorStore gemfireVectorStore(EmbeddingModel embeddingModel, GemFireVectorStoreProperties properties,
GemFireConnectionDetails gemFireConnectionDetails) {
GemFireConnectionDetails gemFireConnectionDetails, ObjectProvider<ObservationRegistry> observationRegistry,
ObjectProvider<VectorStoreObservationConvention> customObservationConvention) {
var config = new GemFireVectorStoreConfig();
config.setHost(gemFireConnectionDetails.getHost())
@@ -57,7 +63,9 @@ public class GemFireVectorStoreAutoConfiguration {
.setVectorSimilarityFunction(properties.getVectorSimilarityFunction())
.setFields(properties.getFields())
.setSslEnabled(properties.isSslEnabled());
return new GemFireVectorStore(config, embeddingModel, properties.isInitializeSchema());
return new GemFireVectorStore(config, embeddingModel, properties.isInitializeSchema(),
observationRegistry.getIfUnique(() -> ObservationRegistry.NOOP),
customObservationConvention.getIfAvailable(() -> null));
}
private static class PropertiesGemFireConnectionDetails implements GemFireConnectionDetails {

View File

@@ -22,6 +22,8 @@ import org.springframework.ai.vectorstore.HanaCloudVectorStore;
import org.springframework.ai.vectorstore.HanaCloudVectorStoreConfig;
import org.springframework.ai.vectorstore.HanaVectorEntity;
import org.springframework.ai.vectorstore.HanaVectorRepository;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.ObjectProvider;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
@@ -29,8 +31,11 @@ import org.springframework.boot.autoconfigure.data.jpa.JpaRepositoriesAutoConfig
import org.springframework.boot.context.properties.EnableConfigurationProperties;
import org.springframework.context.annotation.Bean;
import io.micrometer.observation.ObservationRegistry;
/**
* @author Rahul Mittal
* @author Christian Tzolov
* @since 1.0.0
*/
@AutoConfiguration(after = { JpaRepositoriesAutoConfiguration.class })
@@ -41,13 +46,17 @@ public class HanaCloudVectorStoreAutoConfiguration {
@Bean
@ConditionalOnMissingBean
public HanaCloudVectorStore vectorStore(HanaVectorRepository<? extends HanaVectorEntity> repository,
EmbeddingModel embeddingModel, HanaCloudVectorStoreProperties properties) {
EmbeddingModel embeddingModel, HanaCloudVectorStoreProperties properties,
ObjectProvider<ObservationRegistry> observationRegistry,
ObjectProvider<VectorStoreObservationConvention> customObservationConvention) {
return new HanaCloudVectorStore(repository, embeddingModel,
HanaCloudVectorStoreConfig.builder()
.tableName(properties.getTableName())
.topK(properties.getTopK())
.build());
.build(),
observationRegistry.getIfUnique(() -> ObservationRegistry.NOOP),
customObservationConvention.getIfAvailable(() -> null));
}
}

View File

@@ -17,6 +17,7 @@ package org.springframework.ai.autoconfigure.vectorstore.milvus;
import java.util.concurrent.TimeUnit;
import io.micrometer.observation.ObservationRegistry;
import io.milvus.client.MilvusServiceClient;
import io.milvus.param.ConnectParam;
import io.milvus.param.IndexType;
@@ -25,6 +26,8 @@ import io.milvus.param.MetricType;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.vectorstore.MilvusVectorStore;
import org.springframework.ai.vectorstore.MilvusVectorStore.MilvusVectorStoreConfig;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.ObjectProvider;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
@@ -51,7 +54,8 @@ public class MilvusVectorStoreAutoConfiguration {
@Bean
@ConditionalOnMissingBean
public MilvusVectorStore vectorStore(MilvusServiceClient milvusClient, EmbeddingModel embeddingModel,
MilvusVectorStoreProperties properties) {
MilvusVectorStoreProperties properties, ObjectProvider<ObservationRegistry> observationRegistry,
ObjectProvider<VectorStoreObservationConvention> customObservationConvention) {
MilvusVectorStoreConfig config = MilvusVectorStoreConfig.builder()
.withCollectionName(properties.getCollectionName())
@@ -62,7 +66,9 @@ public class MilvusVectorStoreAutoConfiguration {
.withEmbeddingDimension(properties.getEmbeddingDimension())
.build();
return new MilvusVectorStore(milvusClient, embeddingModel, config, properties.isInitializeSchema());
return new MilvusVectorStore(milvusClient, embeddingModel, config, properties.isInitializeSchema(),
observationRegistry.getIfUnique(() -> ObservationRegistry.NOOP),
customObservationConvention.getIfAvailable(() -> null));
}
@Bean

View File

@@ -17,7 +17,9 @@ package org.springframework.ai.autoconfigure.vectorstore.mongo;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.vectorstore.MongoDBAtlasVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.BeansException;
import org.springframework.beans.factory.ObjectProvider;
import org.springframework.beans.factory.config.BeanPostProcessor;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
@@ -31,6 +33,8 @@ import org.springframework.data.mongodb.core.convert.MongoCustomConversions;
import org.springframework.util.MimeType;
import org.springframework.util.StringUtils;
import io.micrometer.observation.ObservationRegistry;
import java.util.Arrays;
/**
@@ -46,7 +50,8 @@ public class MongoDBAtlasVectorStoreAutoConfiguration {
@Bean
@ConditionalOnMissingBean
MongoDBAtlasVectorStore vectorStore(MongoTemplate mongoTemplate, EmbeddingModel embeddingModel,
MongoDBAtlasVectorStoreProperties properties) {
MongoDBAtlasVectorStoreProperties properties, ObjectProvider<ObservationRegistry> observationRegistry,
ObjectProvider<VectorStoreObservationConvention> customObservationConvention) {
var builder = MongoDBAtlasVectorStore.MongoDBVectorStoreConfig.builder();
@@ -61,7 +66,9 @@ public class MongoDBAtlasVectorStoreAutoConfiguration {
}
MongoDBAtlasVectorStore.MongoDBVectorStoreConfig config = builder.build();
return new MongoDBAtlasVectorStore(mongoTemplate, embeddingModel, config, properties.isInitializeSchema());
return new MongoDBAtlasVectorStore(mongoTemplate, embeddingModel, config, properties.isInitializeSchema(),
observationRegistry.getIfUnique(() -> ObservationRegistry.NOOP),
customObservationConvention.getIfAvailable(() -> null));
}
@Bean

View File

@@ -19,6 +19,8 @@ import org.neo4j.driver.Driver;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.vectorstore.Neo4jVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.ObjectProvider;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
@@ -26,9 +28,12 @@ import org.springframework.boot.autoconfigure.neo4j.Neo4jAutoConfiguration;
import org.springframework.boot.context.properties.EnableConfigurationProperties;
import org.springframework.context.annotation.Bean;
import io.micrometer.observation.ObservationRegistry;
/**
* @author Jingzhou Ou
* @author Josh Long
* @author Christian Tzolov
*/
@AutoConfiguration(after = Neo4jAutoConfiguration.class)
@ConditionalOnClass({ Neo4jVectorStore.class, EmbeddingModel.class, Driver.class })
@@ -38,7 +43,8 @@ public class Neo4jVectorStoreAutoConfiguration {
@Bean
@ConditionalOnMissingBean
public Neo4jVectorStore vectorStore(Driver driver, EmbeddingModel embeddingModel,
Neo4jVectorStoreProperties properties) {
Neo4jVectorStoreProperties properties, ObjectProvider<ObservationRegistry> observationRegistry,
ObjectProvider<VectorStoreObservationConvention> customObservationConvention) {
Neo4jVectorStore.Neo4jVectorStoreConfig config = Neo4jVectorStore.Neo4jVectorStoreConfig.builder()
.withDatabaseName(properties.getDatabaseName())
.withEmbeddingDimension(properties.getEmbeddingDimension())
@@ -50,7 +56,9 @@ public class Neo4jVectorStoreAutoConfiguration {
.withConstraintName(properties.getConstraintName())
.build();
return new Neo4jVectorStore(driver, embeddingModel, config, properties.isInitializeSchema());
return new Neo4jVectorStore(driver, embeddingModel, config, properties.isInitializeSchema(),
observationRegistry.getIfUnique(() -> ObservationRegistry.NOOP),
customObservationConvention.getIfAvailable(() -> null));
}
}

View File

@@ -26,6 +26,8 @@ import org.opensearch.client.transport.aws.AwsSdk2TransportOptions;
import org.opensearch.client.transport.httpclient5.ApacheHttpClient5TransportBuilder;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.vectorstore.OpenSearchVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.ObjectProvider;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
@@ -33,6 +35,8 @@ import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingClas
import org.springframework.boot.context.properties.EnableConfigurationProperties;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import io.micrometer.observation.ObservationRegistry;
import software.amazon.awssdk.auth.credentials.AwsBasicCredentials;
import software.amazon.awssdk.auth.credentials.StaticCredentialsProvider;
import software.amazon.awssdk.http.SdkHttpClient;
@@ -57,12 +61,14 @@ public class OpenSearchVectorStoreAutoConfiguration {
@Bean
@ConditionalOnMissingBean
OpenSearchVectorStore vectorStore(OpenSearchVectorStoreProperties properties, OpenSearchClient openSearchClient,
EmbeddingModel embeddingModel) {
EmbeddingModel embeddingModel, ObjectProvider<ObservationRegistry> observationRegistry,
ObjectProvider<VectorStoreObservationConvention> customObservationConvention) {
var indexName = Optional.ofNullable(properties.getIndexName()).orElse(OpenSearchVectorStore.DEFAULT_INDEX_NAME);
var mappingJson = Optional.ofNullable(properties.getMappingJson())
.orElse(OpenSearchVectorStore.DEFAULT_MAPPING_EMBEDDING_TYPE_KNN_VECTOR_DIMENSION_1536);
return new OpenSearchVectorStore(indexName, openSearchClient, embeddingModel, mappingJson,
properties.isInitializeSchema());
properties.isInitializeSchema(), observationRegistry.getIfUnique(() -> ObservationRegistry.NOOP),
customObservationConvention.getIfAvailable(() -> null));
}
@Configuration(proxyBeanMethods = false)

View File

@@ -19,6 +19,8 @@ import javax.sql.DataSource;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.vectorstore.OracleVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.ObjectProvider;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
@@ -27,9 +29,12 @@ import org.springframework.boot.context.properties.EnableConfigurationProperties
import org.springframework.context.annotation.Bean;
import org.springframework.jdbc.core.JdbcTemplate;
import io.micrometer.observation.ObservationRegistry;
/**
* @author Loïc Lefèvre
* @author Eddú Meléndez
* @author Christian Tzolov
*/
@AutoConfiguration(after = JdbcTemplateAutoConfiguration.class)
@ConditionalOnClass({ OracleVectorStore.class, DataSource.class, JdbcTemplate.class })
@@ -39,11 +44,13 @@ public class OracleVectorStoreAutoConfiguration {
@Bean
@ConditionalOnMissingBean
public OracleVectorStore vectorStore(JdbcTemplate jdbcTemplate, EmbeddingModel embeddingModel,
OracleVectorStoreProperties properties) {
OracleVectorStoreProperties properties, ObjectProvider<ObservationRegistry> observationRegistry,
ObjectProvider<VectorStoreObservationConvention> customObservationConvention) {
return new OracleVectorStore(jdbcTemplate, embeddingModel, properties.getTableName(), properties.getIndexType(),
properties.getDistanceType(), properties.getDimensions(), properties.getSearchAccuracy(),
properties.isInitializeSchema(), properties.isRemoveExistingVectorStoreTable(),
properties.isForcedNormalization());
properties.isForcedNormalization(), observationRegistry.getIfUnique(() -> ObservationRegistry.NOOP),
customObservationConvention.getIfAvailable(() -> null));
}
}

View File

@@ -44,7 +44,7 @@ public class PgVectorStoreAutoConfiguration {
@ConditionalOnMissingBean
public PgVectorStore vectorStore(JdbcTemplate jdbcTemplate, EmbeddingModel embeddingModel,
PgVectorStoreProperties properties, ObjectProvider<ObservationRegistry> observationRegistry,
ObjectProvider<VectorStoreObservationConvention> customSearchObservationConvention) {
ObjectProvider<VectorStoreObservationConvention> customObservationConvention) {
var initializeSchema = properties.isInitializeSchema();
@@ -57,7 +57,7 @@ public class PgVectorStoreAutoConfiguration {
.withIndexType(properties.getIndexType())
.withInitializeSchema(initializeSchema)
.withObservationRegistry(observationRegistry.getIfUnique(() -> ObservationRegistry.NOOP))
.withSearchObservationConvention(customSearchObservationConvention.getIfAvailable(() -> null))
.withSearchObservationConvention(customObservationConvention.getIfAvailable(() -> null))
.build();
}

View File

@@ -18,12 +18,16 @@ package org.springframework.ai.autoconfigure.vectorstore.pinecone;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.vectorstore.PineconeVectorStore;
import org.springframework.ai.vectorstore.PineconeVectorStore.PineconeVectorStoreConfig;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.ObjectProvider;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
import org.springframework.boot.context.properties.EnableConfigurationProperties;
import org.springframework.context.annotation.Bean;
import io.micrometer.observation.ObservationRegistry;
/**
* @author Christian Tzolov
*/
@@ -34,7 +38,9 @@ public class PineconeVectorStoreAutoConfiguration {
@Bean
@ConditionalOnMissingBean
public PineconeVectorStore vectorStore(EmbeddingModel embeddingModel, PineconeVectorStoreProperties properties) {
public PineconeVectorStore vectorStore(EmbeddingModel embeddingModel, PineconeVectorStoreProperties properties,
ObjectProvider<ObservationRegistry> observationRegistry,
ObjectProvider<VectorStoreObservationConvention> customObservationConvention) {
var config = PineconeVectorStoreConfig.builder()
.withApiKey(properties.getApiKey())
@@ -47,7 +53,9 @@ public class PineconeVectorStoreAutoConfiguration {
.withServerSideTimeout(properties.getServerSideTimeout())
.build();
return new PineconeVectorStore(config, embeddingModel);
return new PineconeVectorStore(config, embeddingModel,
observationRegistry.getIfUnique(() -> ObservationRegistry.NOOP),
customObservationConvention.getIfAvailable(() -> null));
}
}

View File

@@ -15,10 +15,13 @@
*/
package org.springframework.ai.autoconfigure.vectorstore.qdrant;
import io.micrometer.observation.ObservationRegistry;
import io.qdrant.client.QdrantClient;
import io.qdrant.client.QdrantGrpcClient;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.ai.vectorstore.qdrant.QdrantVectorStore;
import org.springframework.beans.factory.ObjectProvider;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
@@ -28,6 +31,7 @@ import org.springframework.context.annotation.Bean;
/**
* @author Anush Shetty
* @author Eddú Meléndez
* @author Christian Tzolov
* @since 0.8.1
*/
@AutoConfiguration
@@ -57,9 +61,11 @@ public class QdrantVectorStoreAutoConfiguration {
@Bean
@ConditionalOnMissingBean
public QdrantVectorStore vectorStore(EmbeddingModel embeddingModel, QdrantVectorStoreProperties properties,
QdrantClient qdrantClient) {
QdrantClient qdrantClient, ObjectProvider<ObservationRegistry> observationRegistry,
ObjectProvider<VectorStoreObservationConvention> customObservationConvention) {
return new QdrantVectorStore(qdrantClient, properties.getCollectionName(), embeddingModel,
properties.isInitializeSchema());
properties.isInitializeSchema(), observationRegistry.getIfUnique(() -> ObservationRegistry.NOOP),
customObservationConvention.getIfAvailable(() -> null));
}
static class PropertiesQdrantConnectionDetails implements QdrantConnectionDetails {

View File

@@ -18,6 +18,8 @@ package org.springframework.ai.autoconfigure.vectorstore.redis;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.vectorstore.RedisVectorStore;
import org.springframework.ai.vectorstore.RedisVectorStore.RedisVectorStoreConfig;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.ObjectProvider;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnBean;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
@@ -26,6 +28,8 @@ import org.springframework.boot.autoconfigure.data.redis.RedisAutoConfiguration;
import org.springframework.boot.context.properties.EnableConfigurationProperties;
import org.springframework.context.annotation.Bean;
import org.springframework.data.redis.connection.jedis.JedisConnectionFactory;
import io.micrometer.observation.ObservationRegistry;
import redis.clients.jedis.JedisPooled;
/**
@@ -41,7 +45,8 @@ public class RedisVectorStoreAutoConfiguration {
@Bean
@ConditionalOnMissingBean
public RedisVectorStore vectorStore(EmbeddingModel embeddingModel, RedisVectorStoreProperties properties,
JedisConnectionFactory jedisConnectionFactory) {
JedisConnectionFactory jedisConnectionFactory, ObjectProvider<ObservationRegistry> observationRegistry,
ObjectProvider<VectorStoreObservationConvention> customObservationConvention) {
var config = RedisVectorStoreConfig.builder()
.withIndexName(properties.getIndex())
@@ -50,7 +55,8 @@ public class RedisVectorStoreAutoConfiguration {
return new RedisVectorStore(config, embeddingModel,
new JedisPooled(jedisConnectionFactory.getHostName(), jedisConnectionFactory.getPort()),
properties.isInitializeSchema());
properties.isInitializeSchema(), observationRegistry.getIfUnique(() -> ObservationRegistry.NOOP),
customObservationConvention.getIfAvailable(() -> null));
}
}

View File

@@ -18,6 +18,8 @@ package org.springframework.ai.autoconfigure.vectorstore.typesense;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.vectorstore.TypesenseVectorStore;
import org.springframework.ai.vectorstore.TypesenseVectorStore.TypesenseVectorStoreConfig;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.ObjectProvider;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
@@ -27,6 +29,8 @@ import org.typesense.api.Client;
import org.typesense.api.Configuration;
import org.typesense.resources.Node;
import io.micrometer.observation.ObservationRegistry;
import java.time.Duration;
import java.util.ArrayList;
import java.util.List;
@@ -50,14 +54,17 @@ public class TypesenseVectorStoreAutoConfiguration {
@Bean
@ConditionalOnMissingBean
public TypesenseVectorStore vectorStore(Client typesenseClient, EmbeddingModel embeddingModel,
TypesenseVectorStoreProperties properties) {
TypesenseVectorStoreProperties properties, ObjectProvider<ObservationRegistry> observationRegistry,
ObjectProvider<VectorStoreObservationConvention> customObservationConvention) {
TypesenseVectorStoreConfig config = TypesenseVectorStoreConfig.builder()
.withCollectionName(properties.getCollectionName())
.withEmbeddingDimension(properties.getEmbeddingDimension())
.build();
return new TypesenseVectorStore(typesenseClient, embeddingModel, config, properties.isInitializeSchema());
return new TypesenseVectorStore(typesenseClient, embeddingModel, config, properties.isInitializeSchema(),
observationRegistry.getIfUnique(() -> ObservationRegistry.NOOP),
customObservationConvention.getIfAvailable(() -> null));
}
@Bean

View File

@@ -15,6 +15,7 @@
*/
package org.springframework.ai.autoconfigure.vectorstore.weaviate;
import io.micrometer.observation.ObservationRegistry;
import io.weaviate.client.Config;
import io.weaviate.client.WeaviateAuthClient;
import io.weaviate.client.WeaviateClient;
@@ -23,6 +24,8 @@ import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.vectorstore.WeaviateVectorStore;
import org.springframework.ai.vectorstore.WeaviateVectorStore.WeaviateVectorStoreConfig;
import org.springframework.ai.vectorstore.WeaviateVectorStore.WeaviateVectorStoreConfig.MetadataField;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.ObjectProvider;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
@@ -62,7 +65,8 @@ public class WeaviateVectorStoreAutoConfiguration {
@Bean
@ConditionalOnMissingBean
public WeaviateVectorStore vectorStore(EmbeddingModel embeddingModel, WeaviateClient weaviateClient,
WeaviateVectorStoreProperties properties) {
WeaviateVectorStoreProperties properties, ObjectProvider<ObservationRegistry> observationRegistry,
ObjectProvider<VectorStoreObservationConvention> customObservationConvention) {
WeaviateVectorStoreConfig.Builder configBuilder = WeaviateVectorStore.WeaviateVectorStoreConfig.builder()
.withObjectClass(properties.getObjectClass())
@@ -73,7 +77,9 @@ public class WeaviateVectorStoreAutoConfiguration {
.toList())
.withConsistencyLevel(properties.getConsistencyLevel());
return new WeaviateVectorStore(configBuilder.build(), embeddingModel, weaviateClient);
return new WeaviateVectorStore(configBuilder.build(), embeddingModel, weaviateClient,
observationRegistry.getIfUnique(() -> ObservationRegistry.NOOP),
customObservationConvention.getIfAvailable(() -> null));
}
static class PropertiesWeaviateConnectionDetails implements WeaviateConnectionDetails {

View File

@@ -15,6 +15,10 @@
*/
package org.springframework.ai.autoconfigure.vectorstore.azure;
import static org.assertj.core.api.Assertions.assertThat;
import static org.hamcrest.Matchers.hasSize;
import static org.springframework.ai.autoconfigure.vectorstore.observation.ObservationTestUtil.assertObservationRegistry;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.time.Duration;
@@ -26,21 +30,21 @@ import org.awaitility.Awaitility;
import org.junit.jupiter.api.BeforeAll;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.azure.AzureVectorStore;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.VectorStore;
import org.springframework.ai.vectorstore.azure.AzureVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.core.io.DefaultResourceLoader;
import static org.assertj.core.api.Assertions.assertThat;
import static org.hamcrest.Matchers.hasSize;
import io.micrometer.observation.tck.TestObservationRegistry;
/**
* @author Christian Tzolov
@@ -97,6 +101,7 @@ public class AzureVectorStoreAutoConfigurationIT {
assertThat(properties.getIndexName()).isEqualTo("my_test_index");
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
assertThat(vectorStore).isInstanceOf(AzureVectorStore.class);
@@ -106,6 +111,10 @@ public class AzureVectorStoreAutoConfigurationIT {
return vectorStore.similaritySearch(SearchRequest.query("Spring").withTopK(1));
}, hasSize(1));
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.AZURE,
VectorStoreObservationContext.Operation.ADD);
observationRegistry.clear();
List<Document> results = vectorStore.similaritySearch(SearchRequest.query("Spring").withTopK(1));
assertThat(results).hasSize(1);
@@ -116,18 +125,32 @@ public class AzureVectorStoreAutoConfigurationIT {
assertThat(resultDoc.getMetadata()).hasSize(2);
assertThat(resultDoc.getMetadata()).containsKeys("spring", "distance");
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.AZURE,
VectorStoreObservationContext.Operation.QUERY);
observationRegistry.clear();
// Remove all documents from the store
vectorStore.delete(documents.stream().map(doc -> doc.getId()).toList());
Awaitility.await().until(() -> {
return vectorStore.similaritySearch(SearchRequest.query("Spring").withTopK(1));
}, hasSize(0));
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.AZURE,
VectorStoreObservationContext.Operation.DELETE);
observationRegistry.clear();
});
}
@Configuration(proxyBeanMethods = false)
static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public EmbeddingModel embeddingModel() {
return new TransformersEmbeddingModel();

View File

@@ -15,31 +15,36 @@
*/
package org.springframework.ai.autoconfigure.vectorstore.cassandra;
import static org.assertj.core.api.Assertions.assertThat;
import static org.springframework.ai.autoconfigure.vectorstore.observation.ObservationTestUtil.assertObservationRegistry;
import java.util.List;
import java.util.Map;
import org.junit.jupiter.api.Test;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import org.testcontainers.containers.CassandraContainer;
import org.testcontainers.utility.DockerImageName;
import org.springframework.ai.ResourceUtils;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.VectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.cassandra.CassandraAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.testcontainers.containers.CassandraContainer;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import org.testcontainers.utility.DockerImageName;
import static org.assertj.core.api.Assertions.assertThat;
import io.micrometer.observation.tck.TestObservationRegistry;
/**
* @author Mick Semb Wever
* @author Christian Tzolov
* @since 1.0.0
*/
@Testcontainers
@@ -72,8 +77,13 @@ class CassandraVectorStoreAutoConfigurationIT {
.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.CASSANDRA,
VectorStoreObservationContext.Operation.ADD);
observationRegistry.clear();
List<Document> results = vectorStore.similaritySearch(SearchRequest.query("Spring").withTopK(1));
assertThat(results).hasSize(1);
@@ -82,17 +92,30 @@ class CassandraVectorStoreAutoConfigurationIT {
assertThat(resultDoc.getContent()).contains(
"Spring AI provides abstractions that serve as the foundation for developing AI applications.");
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.CASSANDRA,
VectorStoreObservationContext.Operation.QUERY);
observationRegistry.clear();
// Remove all documents from the store
vectorStore.delete(documents.stream().map(doc -> doc.getId()).toList());
results = vectorStore.similaritySearch(SearchRequest.query("Spring").withTopK(1));
assertThat(results).isEmpty();
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.CASSANDRA,
VectorStoreObservationContext.Operation.DELETE);
observationRegistry.clear();
});
}
@Configuration(proxyBeanMethods = false)
static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public EmbeddingModel embeddingModel() {
return new TransformersEmbeddingModel();

View File

@@ -23,17 +23,25 @@ import org.testcontainers.chromadb.ChromaDBContainer;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import io.micrometer.observation.tck.TestObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistryAssert;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.VectorStore;
import org.springframework.ai.vectorstore.observation.DefaultVectorStoreObservationConvention;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.HighCardinalityKeyNames;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import static org.assertj.core.api.Assertions.assertThat;
import static org.springframework.ai.autoconfigure.vectorstore.observation.ObservationTestUtil.assertObservationRegistry;
/**
* @author Christian Tzolov
@@ -60,6 +68,7 @@ public class ChromaVectorStoreAutoConfigurationIT {
contextRunner.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
var bgDocument = new Document("The World is Big and Salvation Lurks Around the Corner",
Map.of("country", "Bulgaria"));
@@ -68,29 +77,62 @@ public class ChromaVectorStoreAutoConfigurationIT {
vectorStore.add(List.of(bgDocument, nlDocument));
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.CHROMA,
VectorStoreObservationContext.Operation.ADD);
observationRegistry.clear();
var request = SearchRequest.query("The World").withTopK(5);
List<Document> results = vectorStore.similaritySearch(request);
assertThat(results).hasSize(2);
observationRegistry.clear();
results = vectorStore
.similaritySearch(request.withSimilarityThresholdAll().withFilterExpression("country == 'Bulgaria'"));
assertThat(results).hasSize(1);
assertThat(results.get(0).getId()).isEqualTo(bgDocument.getId());
observationRegistry.clear();
results = vectorStore.similaritySearch(
request.withSimilarityThresholdAll().withFilterExpression("country == 'Netherlands'"));
assertThat(results).hasSize(1);
assertThat(results.get(0).getId()).isEqualTo(nlDocument.getId());
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store chroma query")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY_METADATA_FILTER.asString(),
"Expression[type=EQ, left=Key[key=country], right=Value[value=Netherlands]]")
.hasBeenStarted()
.hasBeenStopped();
observationRegistry.clear();
// Remove all documents from the store
vectorStore.delete(List.of(bgDocument, nlDocument).stream().map(doc -> doc.getId()).toList());
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store chroma delete")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY_METADATA_FILTER.asString(), "none")
.hasBeenStarted()
.hasBeenStopped();
observationRegistry.clear();
});
}
@Configuration(proxyBeanMethods = false)
static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public EmbeddingModel embeddingModel() {
return new TransformersEmbeddingModel();

View File

@@ -15,31 +15,38 @@
*/
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.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;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.core.io.DefaultResourceLoader;
import org.testcontainers.elasticsearch.ElasticsearchContainer;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import static org.assertj.core.api.Assertions.assertThat;
import static org.hamcrest.Matchers.hasSize;
import static org.springframework.ai.autoconfigure.vectorstore.observation.ObservationTestUtil.assertObservationRegistry;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.util.List;
import java.util.Map;
import static org.assertj.core.api.Assertions.assertThat;
import static org.hamcrest.Matchers.hasSize;
import org.awaitility.Awaitility;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.autoconfigure.openai.OpenAiAutoConfiguration;
import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration;
import org.springframework.ai.document.Document;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.vectorstore.ElasticsearchVectorStore;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.SimilarityFunction;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.elasticsearch.ElasticsearchRestClientAutoConfiguration;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.core.io.DefaultResourceLoader;
import org.testcontainers.elasticsearch.ElasticsearchContainer;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import io.micrometer.observation.tck.TestObservationRegistry;
@Testcontainers
@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+")
@@ -59,6 +66,7 @@ class ElasticsearchVectorStoreAutoConfigurationIT {
.withConfiguration(AutoConfigurations.of(ElasticsearchRestClientAutoConfiguration.class,
ElasticsearchVectorStoreAutoConfiguration.class, RestClientAutoConfiguration.class,
SpringAiRetryAutoConfiguration.class, OpenAiAutoConfiguration.class))
.withUserConfiguration(Config.class)
.withPropertyValues("spring.elasticsearch.uris=" + elasticsearchContainer.getHttpHostAddress(),
"spring.ai.vectorstore.elasticsearch.initializeSchema=true",
"spring.ai.openai.api-key=" + System.getenv("OPENAI_API_KEY"));
@@ -70,14 +78,21 @@ class ElasticsearchVectorStoreAutoConfigurationIT {
this.contextRunner.run(context -> {
ElasticsearchVectorStore vectorStore = context.getBean(ElasticsearchVectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.ELASTICSEARCH,
VectorStoreObservationContext.Operation.ADD);
observationRegistry.clear();
Awaitility.await()
.until(() -> vectorStore
.similaritySearch(SearchRequest.query("Great Depression").withTopK(1).withSimilarityThreshold(0)),
hasSize(1));
observationRegistry.clear();
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("Great Depression").withTopK(1).withSimilarityThreshold(0));
@@ -89,9 +104,17 @@ class ElasticsearchVectorStoreAutoConfigurationIT {
assertThat(resultDoc.getMetadata()).containsKey("meta2");
assertThat(resultDoc.getMetadata()).containsKey("distance");
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.ELASTICSEARCH,
VectorStoreObservationContext.Operation.QUERY);
observationRegistry.clear();
// Remove all documents from the store
vectorStore.delete(documents.stream().map(Document::getId).toList());
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.ELASTICSEARCH,
VectorStoreObservationContext.Operation.DELETE);
observationRegistry.clear();
Awaitility.await()
.until(() -> vectorStore
.similaritySearch(SearchRequest.query("Great Depression").withTopK(1).withSimilarityThreshold(0)),
@@ -135,4 +158,14 @@ class ElasticsearchVectorStoreAutoConfigurationIT {
}
}
@Configuration(proxyBeanMethods = false)
static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
}
}

View File

@@ -18,17 +18,12 @@ package org.springframework.ai.autoconfigure.vectorstore.gemfire;
import static org.assertj.core.api.Assertions.assertThat;
import static org.hamcrest.Matchers.hasSize;
import static org.springframework.ai.autoconfigure.vectorstore.observation.ObservationTestUtil.assertObservationRegistry;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import com.fasterxml.jackson.databind.JsonNode;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.github.dockerjava.api.model.ExposedPort;
import com.github.dockerjava.api.model.PortBinding;
import com.github.dockerjava.api.model.Ports;
import com.vmware.gemfire.testcontainers.GemFireCluster;
import org.awaitility.Awaitility;
import org.junit.jupiter.api.AfterAll;
import org.junit.jupiter.api.Assertions;
@@ -37,19 +32,30 @@ import org.junit.jupiter.api.Test;
import org.springframework.ai.ResourceUtils;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.GemFireVectorStore;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.VectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import com.fasterxml.jackson.databind.JsonNode;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.github.dockerjava.api.model.ExposedPort;
import com.github.dockerjava.api.model.PortBinding;
import com.github.dockerjava.api.model.Ports;
import com.vmware.gemfire.testcontainers.GemFireCluster;
import io.micrometer.observation.tck.TestObservationRegistry;
/**
* @author Geet Rawat
* @author Christian Tzolov
*/
class GemFireVectorStoreAutoConfigurationIT {
private static GemFireCluster gemFireCluster;
@@ -116,6 +122,7 @@ class GemFireVectorStoreAutoConfigurationIT {
void ensureGemFireVectorStoreCustomConfiguration() {
this.contextRunner.run(context -> {
GemFireVectorStore store = context.getBean(GemFireVectorStore.class);
Assertions.assertNotNull(store);
assertThat(store.getIndexName()).isEqualTo(INDEX_NAME);
assertThat(store.getBeamWidth()).isEqualTo(BEAM_WIDTH);
@@ -138,14 +145,24 @@ class GemFireVectorStoreAutoConfigurationIT {
public void addAndSearchTest() {
contextRunner.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.GEMFIRE,
VectorStoreObservationContext.Operation.ADD);
Awaitility.await().until(() -> {
return vectorStore.similaritySearch(SearchRequest.query("Spring").withTopK(1));
}, hasSize(1));
observationRegistry.clear();
List<Document> results = vectorStore.similaritySearch(SearchRequest.query("Spring").withTopK(1));
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.GEMFIRE,
VectorStoreObservationContext.Operation.QUERY);
observationRegistry.clear();
assertThat(results).hasSize(1);
Document resultDoc = results.get(0);
assertThat(resultDoc.getId()).isEqualTo(documents.get(0).getId());
@@ -157,9 +174,13 @@ class GemFireVectorStoreAutoConfigurationIT {
// Remove all documents from the store
vectorStore.delete(documents.stream().map(doc -> doc.getId()).toList());
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.GEMFIRE,
VectorStoreObservationContext.Operation.DELETE);
Awaitility.await().until(() -> {
return vectorStore.similaritySearch(SearchRequest.query("Spring").withTopK(1));
}, hasSize(0));
observationRegistry.clear();
});
}
@@ -191,6 +212,11 @@ class GemFireVectorStoreAutoConfigurationIT {
@Configuration(proxyBeanMethods = false)
static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public EmbeddingModel embeddingModel() {
return new TransformersEmbeddingModel();

View File

@@ -15,26 +15,30 @@
*/
package org.springframework.ai.autoconfigure.vectorstore.milvus;
import static org.assertj.core.api.Assertions.assertThat;
import static org.springframework.ai.autoconfigure.vectorstore.observation.ObservationTestUtil.assertObservationRegistry;
import java.util.List;
import java.util.Map;
import org.junit.jupiter.api.Test;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import org.springframework.ai.ResourceUtils;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.VectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import org.testcontainers.milvus.MilvusContainer;
import static org.assertj.core.api.Assertions.assertThat;
import io.micrometer.observation.tck.TestObservationRegistry;
/**
* @author Christian Tzolov
@@ -68,8 +72,14 @@ public class MilvusVectorStoreAutoConfigurationIT {
"spring.ai.vectorstore.milvus.client.port=" + milvus.getMappedPort(19530))
.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.MILVUS,
VectorStoreObservationContext.Operation.ADD);
observationRegistry.clear();
List<Document> results = vectorStore.similaritySearch(SearchRequest.query("Spring").withTopK(1));
assertThat(results).hasSize(1);
@@ -80,17 +90,31 @@ public class MilvusVectorStoreAutoConfigurationIT {
assertThat(resultDoc.getMetadata()).hasSize(2);
assertThat(resultDoc.getMetadata()).containsKeys("spring", "distance");
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.MILVUS,
VectorStoreObservationContext.Operation.QUERY);
observationRegistry.clear();
// Remove all documents from the store
vectorStore.delete(documents.stream().map(doc -> doc.getId()).toList());
results = vectorStore.similaritySearch(SearchRequest.query("Spring").withTopK(1));
assertThat(results).hasSize(0);
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.MILVUS,
VectorStoreObservationContext.Operation.DELETE);
observationRegistry.clear();
});
}
@Configuration(proxyBeanMethods = false)
static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public EmbeddingModel embeddingModel() {
return new TransformersEmbeddingModel();

View File

@@ -15,33 +15,43 @@
*/
package org.springframework.ai.autoconfigure.vectorstore.mongo;
import org.junit.jupiter.api.Test;
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.SearchRequest;
import org.springframework.ai.vectorstore.VectorStore;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.data.mongo.MongoDataAutoConfiguration;
import org.springframework.boot.autoconfigure.mongo.MongoAutoConfiguration;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.testcontainers.containers.GenericContainer;
import org.testcontainers.containers.wait.strategy.Wait;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import static org.assertj.core.api.Assertions.assertThat;
import static org.springframework.ai.autoconfigure.vectorstore.observation.ObservationTestUtil.assertObservationRegistry;
import java.time.Duration;
import java.util.Collections;
import java.util.List;
import java.util.stream.Collectors;
import static org.assertj.core.api.Assertions.assertThat;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.autoconfigure.openai.OpenAiAutoConfiguration;
import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration;
import org.springframework.ai.document.Document;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.VectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.data.mongo.MongoDataAutoConfiguration;
import org.springframework.boot.autoconfigure.mongo.MongoAutoConfiguration;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.testcontainers.containers.GenericContainer;
import org.testcontainers.containers.wait.strategy.Wait;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import io.micrometer.observation.tck.TestObservationRegistry;
/**
* @author Eddú Meléndez
* @author Christian Tzolov
*/
@Testcontainers
@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+")
class MongoDBAtlasVectorStoreAutoConfigurationIT {
@Container
@@ -61,6 +71,7 @@ class MongoDBAtlasVectorStoreAutoConfigurationIT {
Collections.singletonMap("meta2", "meta2")));
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withUserConfiguration(Config.class)
.withConfiguration(AutoConfigurations.of(MongoAutoConfiguration.class, MongoDataAutoConfiguration.class,
MongoDBAtlasVectorStoreAutoConfiguration.class, RestClientAutoConfiguration.class,
SpringAiRetryAutoConfiguration.class, OpenAiAutoConfiguration.class))
@@ -79,8 +90,13 @@ class MongoDBAtlasVectorStoreAutoConfigurationIT {
contextRunner.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.MONGODB,
VectorStoreObservationContext.Operation.ADD);
observationRegistry.clear();
Thread.sleep(5000); // Await a second for the document to be indexed
List<Document> results = vectorStore.similaritySearch(SearchRequest.query("Great").withTopK(1));
@@ -92,12 +108,30 @@ class MongoDBAtlasVectorStoreAutoConfigurationIT {
"Great Depression Great Depression Great Depression Great Depression Great Depression Great Depression");
assertThat(resultDoc.getMetadata()).containsEntry("meta2", "meta2");
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.MONGODB,
VectorStoreObservationContext.Operation.QUERY);
observationRegistry.clear();
// Remove all documents from the store
vectorStore.delete(documents.stream().map(Document::getId).collect(Collectors.toList()));
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.MONGODB,
VectorStoreObservationContext.Operation.DELETE);
observationRegistry.clear();
List<Document> results2 = vectorStore.similaritySearch(SearchRequest.query("Great").withTopK(1));
assertThat(results2).isEmpty();
});
}
@Configuration(proxyBeanMethods = false)
static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
}
}

View File

@@ -15,32 +15,37 @@
*/
package org.springframework.ai.autoconfigure.vectorstore.neo4j;
import static org.assertj.core.api.Assertions.assertThat;
import static org.springframework.ai.autoconfigure.vectorstore.observation.ObservationTestUtil.assertObservationRegistry;
import java.util.List;
import java.util.Map;
import org.junit.jupiter.api.Test;
import org.springframework.ai.ResourceUtils;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.VectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.neo4j.Neo4jAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
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.ResourceUtils;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.VectorStore;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import static org.assertj.core.api.Assertions.assertThat;
import io.micrometer.observation.tck.TestObservationRegistry;
/**
* @author Jingzhou Ou
* @author Soby Chacko
* @author Christian Tzolov
*/
@Testcontainers
public class Neo4jVectorStoreAutoConfigurationIT {
@@ -48,7 +53,7 @@ public class Neo4jVectorStoreAutoConfigurationIT {
// Needs to be Neo4j 5.15+, because Neo4j 5.15 deprecated the used embedding storing
// function.
@Container
static Neo4jContainer<?> neo4jContainer = new Neo4jContainer<>(DockerImageName.parse("neo4j:5.15"))
static Neo4jContainer<?> neo4jContainer = new Neo4jContainer<>(DockerImageName.parse("neo4j:5.18"))
.withRandomPassword();
List<Document> documents = List.of(
@@ -76,8 +81,14 @@ public class Neo4jVectorStoreAutoConfigurationIT {
assertThat(properties.getIndexName()).isEqualTo("customIndexName");
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.NEO4J,
VectorStoreObservationContext.Operation.ADD);
observationRegistry.clear();
List<Document> results = vectorStore.similaritySearch(SearchRequest.query("Spring").withTopK(1));
assertThat(results).hasSize(1);
@@ -86,9 +97,17 @@ public class Neo4jVectorStoreAutoConfigurationIT {
assertThat(resultDoc.getContent()).contains(
"Spring AI provides abstractions that serve as the foundation for developing AI applications.");
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.NEO4J,
VectorStoreObservationContext.Operation.QUERY);
observationRegistry.clear();
// Remove all documents from the store
vectorStore.delete(documents.stream().map(doc -> doc.getId()).toList());
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.NEO4J,
VectorStoreObservationContext.Operation.DELETE);
observationRegistry.clear();
results = vectorStore.similaritySearch(SearchRequest.query("Spring").withTopK(1));
assertThat(results).isEmpty();
});
@@ -97,6 +116,11 @@ public class Neo4jVectorStoreAutoConfigurationIT {
@Configuration(proxyBeanMethods = false)
static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public EmbeddingModel embeddingModel() {
return new TransformersEmbeddingModel();

View File

@@ -0,0 +1,43 @@
/*
* Copyright 2024 - 2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (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.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.autoconfigure.vectorstore.observation;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.vectorstore.observation.DefaultVectorStoreObservationConvention;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import io.micrometer.observation.tck.TestObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistryAssert;
/**
* @author Christian Tzolov
* @since 1.0.0
*/
public class ObservationTestUtil {
public static void assertObservationRegistry(TestObservationRegistry observationRegistry, String kind,
VectorStoreProvider vectorStoreProvider, VectorStoreObservationContext.Operation operation) {
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo(kind + " " + vectorStoreProvider.value() + " " + operation.value())
.hasBeenStarted()
.hasBeenStopped();
}
}

View File

@@ -47,7 +47,7 @@ import static org.awaitility.Awaitility.await;
import static org.hamcrest.Matchers.hasSize;
@Testcontainers
class a {
class AwsOpenSearchVectorStoreAutoConfigurationIT {
@Container
private static final LocalStackContainer localstack = new LocalStackContainer(

View File

@@ -15,15 +15,26 @@
*/
package org.springframework.ai.autoconfigure.vectorstore.opensearch;
import static org.assertj.core.api.Assertions.assertThat;
import static org.hamcrest.Matchers.hasSize;
import static org.springframework.ai.autoconfigure.vectorstore.observation.ObservationTestUtil.assertObservationRegistry;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.util.List;
import java.util.Map;
import org.awaitility.Awaitility;
import org.junit.jupiter.api.Test;
import org.opensearch.testcontainers.OpensearchContainer;
import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.OpenSearchVectorStore;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.test.context.FilteredClassLoader;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
@@ -33,17 +44,11 @@ import org.springframework.core.io.DefaultResourceLoader;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import org.testcontainers.utility.DockerImageName;
import io.micrometer.observation.tck.TestObservationRegistry;
import software.amazon.awssdk.http.apache.ApacheHttpClient;
import software.amazon.awssdk.regions.Region;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.util.List;
import java.util.Map;
import static org.assertj.core.api.Assertions.assertThat;
import static org.hamcrest.Matchers.hasSize;
@Testcontainers
class OpenSearchVectorStoreAutoConfigurationIT {
@@ -82,17 +87,28 @@ class OpenSearchVectorStoreAutoConfigurationIT {
this.contextRunner.run(context -> {
OpenSearchVectorStore vectorStore = context.getBean(OpenSearchVectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.OPENSEARCH,
VectorStoreObservationContext.Operation.ADD);
Awaitility.await()
.until(() -> vectorStore
.similaritySearch(SearchRequest.query("Great Depression").withTopK(1).withSimilarityThreshold(0)),
hasSize(1));
observationRegistry.clear();
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("Great Depression").withTopK(1).withSimilarityThreshold(0));
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.OPENSEARCH,
VectorStoreObservationContext.Operation.QUERY);
observationRegistry.clear();
assertThat(results).hasSize(1);
Document resultDoc = results.get(0);
assertThat(resultDoc.getId()).isEqualTo(documents.get(2).getId());
@@ -104,6 +120,10 @@ class OpenSearchVectorStoreAutoConfigurationIT {
// Remove all documents from the store
vectorStore.delete(documents.stream().map(Document::getId).toList());
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.OPENSEARCH,
VectorStoreObservationContext.Operation.DELETE);
observationRegistry.clear();
Awaitility.await()
.until(() -> vectorStore
.similaritySearch(SearchRequest.query("Great Depression").withTopK(1).withSimilarityThreshold(0)),
@@ -124,6 +144,11 @@ class OpenSearchVectorStoreAutoConfigurationIT {
@Configuration(proxyBeanMethods = false)
static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public EmbeddingModel embeddingModel() {
return new TransformersEmbeddingModel();

View File

@@ -15,31 +15,35 @@
*/
package org.springframework.ai.autoconfigure.vectorstore.oracle;
import static org.assertj.core.api.Assertions.assertThat;
import static org.springframework.ai.autoconfigure.vectorstore.observation.ObservationTestUtil.assertObservationRegistry;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.util.List;
import java.util.Map;
import org.junit.jupiter.api.Test;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.VectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.jdbc.DataSourceAutoConfiguration;
import org.springframework.boot.autoconfigure.jdbc.JdbcTemplateAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.core.io.DefaultResourceLoader;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import org.testcontainers.oracle.OracleContainer;
import org.testcontainers.utility.MountableFile;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.VectorStore;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.jdbc.DataSourceAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.core.io.DefaultResourceLoader;
import static org.assertj.core.api.Assertions.assertThat;
import io.micrometer.observation.tck.TestObservationRegistry;
/**
* @author Christian Tzolov
@@ -76,9 +80,14 @@ public class OracleVectorStoreAutoConfigurationIT {
contextRunner.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.ORACLE,
VectorStoreObservationContext.Operation.ADD);
observationRegistry.clear();
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("What is Great Depression?").withTopK(1));
@@ -87,8 +96,17 @@ public class OracleVectorStoreAutoConfigurationIT {
assertThat(resultDoc.getId()).isEqualTo(documents.get(2).getId());
assertThat(resultDoc.getMetadata()).containsKeys("depression", "distance");
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.ORACLE,
VectorStoreObservationContext.Operation.QUERY);
observationRegistry.clear();
// Remove all documents from the store
vectorStore.delete(documents.stream().map(doc -> doc.getId()).toList());
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.ORACLE,
VectorStoreObservationContext.Operation.DELETE);
observationRegistry.clear();
results = vectorStore.similaritySearch(SearchRequest.query("Great Depression").withTopK(1));
assertThat(results).hasSize(0);
});
@@ -107,6 +125,11 @@ public class OracleVectorStoreAutoConfigurationIT {
@Configuration(proxyBeanMethods = false)
static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public EmbeddingModel embeddingModel() {
return new TransformersEmbeddingModel();

View File

@@ -16,6 +16,7 @@
package org.springframework.ai.autoconfigure.vectorstore.pgvector;
import static org.assertj.core.api.Assertions.assertThat;
import static org.springframework.ai.autoconfigure.vectorstore.observation.ObservationTestUtil.assertObservationRegistry;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
@@ -27,9 +28,11 @@ import org.junit.jupiter.params.ParameterizedTest;
import org.junit.jupiter.params.provider.ValueSource;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.PgVectorStore;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.jdbc.DataSourceAutoConfiguration;
import org.springframework.boot.autoconfigure.jdbc.JdbcTemplateAutoConfiguration;
@@ -43,6 +46,8 @@ import org.testcontainers.containers.PostgreSQLContainer;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import io.micrometer.observation.tck.TestObservationRegistry;
/**
* @author Christian Tzolov
* @author Muthukumaran Navaneethakrishnan
@@ -88,6 +93,7 @@ public class PgVectorStoreAutoConfigurationIT {
contextRunner.run(context -> {
PgVectorStore vectorStore = context.getBean(PgVectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
assertThat(isFullyQualifiedTableExists(context, PgVectorStore.DEFAULT_SCHEMA_NAME,
PgVectorStore.DEFAULT_TABLE_NAME))
@@ -95,6 +101,10 @@ public class PgVectorStoreAutoConfigurationIT {
vectorStore.add(documents);
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.PG_VECTOR,
VectorStoreObservationContext.Operation.ADD);
observationRegistry.clear();
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("What is Great Depression?").withTopK(1));
@@ -103,10 +113,19 @@ public class PgVectorStoreAutoConfigurationIT {
assertThat(resultDoc.getId()).isEqualTo(documents.get(2).getId());
assertThat(resultDoc.getMetadata()).containsKeys("depression", "distance");
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.PG_VECTOR,
VectorStoreObservationContext.Operation.QUERY);
observationRegistry.clear();
// Remove all documents from the store
vectorStore.delete(documents.stream().map(doc -> doc.getId()).toList());
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.PG_VECTOR,
VectorStoreObservationContext.Operation.DELETE);
results = vectorStore.similaritySearch(SearchRequest.query("Great Depression").withTopK(1));
assertThat(results).hasSize(0);
observationRegistry.clear();
});
}
@@ -142,6 +161,11 @@ public class PgVectorStoreAutoConfigurationIT {
@Configuration(proxyBeanMethods = false)
static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public EmbeddingModel embeddingModel() {
return new TransformersEmbeddingModel();

View File

@@ -17,6 +17,7 @@ package org.springframework.ai.autoconfigure.vectorstore.pinecone;
import static org.assertj.core.api.Assertions.assertThat;
import static org.hamcrest.Matchers.hasSize;
import static org.springframework.ai.autoconfigure.vectorstore.observation.ObservationTestUtil.assertObservationRegistry;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
@@ -31,15 +32,19 @@ import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.PineconeVectorStore;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.core.io.DefaultResourceLoader;
import io.micrometer.observation.tck.TestObservationRegistry;
/**
* @author Christian Tzolov
* @author Soby Chacko
@@ -85,12 +90,17 @@ public class PineconeVectorStoreAutoConfigurationIT {
contextRunner.run(context -> {
PineconeVectorStore vectorStore = context.getBean(PineconeVectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.PINECONE,
VectorStoreObservationContext.Operation.ADD);
Awaitility.await().until(() -> {
return vectorStore.similaritySearch(SearchRequest.query("Spring").withTopK(1));
}, hasSize(1));
observationRegistry.clear();
List<Document> results = vectorStore.similaritySearch(SearchRequest.query("Spring").withTopK(1));
@@ -102,9 +112,17 @@ public class PineconeVectorStoreAutoConfigurationIT {
assertThat(resultDoc.getMetadata()).hasSize(2);
assertThat(resultDoc.getMetadata()).containsKeys("spring", "customDistanceField");
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.PINECONE,
VectorStoreObservationContext.Operation.QUERY);
observationRegistry.clear();
// Remove all documents from the store
vectorStore.delete(documents.stream().map(doc -> doc.getId()).toList());
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.PINECONE,
VectorStoreObservationContext.Operation.DELETE);
observationRegistry.clear();
Awaitility.await().until(() -> {
return vectorStore.similaritySearch(SearchRequest.query("Spring").withTopK(1));
}, hasSize(0));
@@ -114,6 +132,11 @@ public class PineconeVectorStoreAutoConfigurationIT {
@Configuration(proxyBeanMethods = false)
static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public EmbeddingModel embeddingModel() {
return new TransformersEmbeddingModel();

View File

@@ -15,28 +15,32 @@
*/
package org.springframework.ai.autoconfigure.vectorstore.qdrant;
import static org.assertj.core.api.Assertions.assertThat;
import static org.springframework.ai.autoconfigure.vectorstore.observation.ObservationTestUtil.assertObservationRegistry;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.util.List;
import java.util.Map;
import org.junit.jupiter.api.Test;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import org.testcontainers.qdrant.QdrantContainer;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.VectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.core.io.DefaultResourceLoader;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import org.testcontainers.qdrant.QdrantContainer;
import static org.assertj.core.api.Assertions.assertThat;
import io.micrometer.observation.tck.TestObservationRegistry;
/**
* @author Christian Tzolov
@@ -67,9 +71,14 @@ public class QdrantVectorStoreAutoConfigurationIT {
contextRunner.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.QDRANT,
VectorStoreObservationContext.Operation.ADD);
observationRegistry.clear();
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("What is Great Depression?").withTopK(1));
@@ -78,10 +87,18 @@ public class QdrantVectorStoreAutoConfigurationIT {
assertThat(resultDoc.getId()).isEqualTo(documents.get(2).getId());
assertThat(resultDoc.getMetadata()).containsKeys("depression", "distance");
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.QDRANT,
VectorStoreObservationContext.Operation.QUERY);
observationRegistry.clear();
// Remove all documents from the store
vectorStore.delete(documents.stream().map(doc -> doc.getId()).toList());
results = vectorStore.similaritySearch(SearchRequest.query("Great Depression").withTopK(1));
assertThat(results).hasSize(0);
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.QDRANT,
VectorStoreObservationContext.Operation.DELETE);
observationRegistry.clear();
});
}
@@ -98,6 +115,11 @@ public class QdrantVectorStoreAutoConfigurationIT {
@Configuration(proxyBeanMethods = false)
static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public EmbeddingModel embeddingModel() {
return new TransformersEmbeddingModel();

View File

@@ -16,6 +16,7 @@
package org.springframework.ai.autoconfigure.vectorstore.redis;
import static org.assertj.core.api.Assertions.assertThat;
import static org.springframework.ai.autoconfigure.vectorstore.observation.ObservationTestUtil.assertObservationRegistry;
import java.util.List;
import java.util.Map;
@@ -24,9 +25,11 @@ import org.junit.jupiter.api.Test;
import org.springframework.ai.ResourceUtils;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.VectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.data.redis.RedisAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
@@ -37,10 +40,13 @@ import org.testcontainers.junit.jupiter.Testcontainers;
import com.redis.testcontainers.RedisStackContainer;
import io.micrometer.observation.tck.TestObservationRegistry;
/**
* @author Julien Ruaux
* @author Eddú Meléndez
* @author Soby Chacko
* @author Christian Tzolov
*/
@Testcontainers
class RedisVectorStoreAutoConfigurationIT {
@@ -66,8 +72,14 @@ class RedisVectorStoreAutoConfigurationIT {
void addAndSearch() {
contextRunner.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.REDIS,
VectorStoreObservationContext.Operation.ADD);
observationRegistry.clear();
List<Document> results = vectorStore.similaritySearch(SearchRequest.query("Spring").withTopK(1));
assertThat(results).hasSize(1);
@@ -76,9 +88,17 @@ class RedisVectorStoreAutoConfigurationIT {
assertThat(resultDoc.getContent()).contains(
"Spring AI provides abstractions that serve as the foundation for developing AI applications.");
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.REDIS,
VectorStoreObservationContext.Operation.QUERY);
observationRegistry.clear();
// Remove all documents from the store
vectorStore.delete(documents.stream().map(doc -> doc.getId()).toList());
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.REDIS,
VectorStoreObservationContext.Operation.DELETE);
observationRegistry.clear();
results = vectorStore.similaritySearch(SearchRequest.query("Spring").withTopK(1));
assertThat(results).isEmpty();
});
@@ -87,6 +107,11 @@ class RedisVectorStoreAutoConfigurationIT {
@Configuration(proxyBeanMethods = false)
static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public EmbeddingModel embeddingModel() {
return new TransformersEmbeddingModel();

View File

@@ -15,13 +15,22 @@
*/
package org.springframework.ai.autoconfigure.vectorstore.typesense;
import static org.assertj.core.api.Assertions.assertThat;
import static org.springframework.ai.autoconfigure.vectorstore.observation.ObservationTestUtil.assertObservationRegistry;
import java.time.Duration;
import java.util.List;
import java.util.Map;
import org.junit.jupiter.api.Test;
import org.springframework.ai.ResourceUtils;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.VectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
@@ -30,16 +39,13 @@ import org.testcontainers.containers.GenericContainer;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import java.time.Duration;
import java.util.List;
import java.util.Map;
import static org.assertj.core.api.Assertions.assertThat;
import io.micrometer.observation.tck.TestObservationRegistry;
/**
* @author Pablo Sanchidrian Herrera
* @author Eddú Meléndez
* @author Soby Chacko
* @author Christian Tzolov
*/
@Testcontainers
public class TypesenseVectorStoreAutoConfigurationIT {
@@ -71,8 +77,14 @@ public class TypesenseVectorStoreAutoConfigurationIT {
"spring.ai.vectorstore.typesense.client.port=" + typesenseContainer.getMappedPort(8108).toString())
.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.TYPESENSE,
VectorStoreObservationContext.Operation.ADD);
observationRegistry.clear();
List<Document> results = vectorStore.similaritySearch(SearchRequest.query("Spring").withTopK(1));
assertThat(results).hasSize(1);
@@ -83,8 +95,16 @@ public class TypesenseVectorStoreAutoConfigurationIT {
assertThat(resultDoc.getMetadata()).hasSize(2);
assertThat(resultDoc.getMetadata()).containsKeys("spring", "distance");
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.TYPESENSE,
VectorStoreObservationContext.Operation.QUERY);
observationRegistry.clear();
vectorStore.delete(documents.stream().map(doc -> doc.getId()).toList());
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.TYPESENSE,
VectorStoreObservationContext.Operation.DELETE);
observationRegistry.clear();
results = vectorStore.similaritySearch(SearchRequest.query("Spring").withTopK(1));
assertThat(results).hasSize(0);
});
@@ -93,6 +113,11 @@ public class TypesenseVectorStoreAutoConfigurationIT {
@Configuration(proxyBeanMethods = false)
static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public EmbeddingModel embeddingModel() {
return new TransformersEmbeddingModel();

View File

@@ -24,17 +24,22 @@ import org.testcontainers.junit.jupiter.Testcontainers;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.VectorStore;
import org.springframework.ai.vectorstore.WeaviateVectorStore.WeaviateVectorStoreConfig.MetadataField;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.testcontainers.weaviate.WeaviateContainer;
import io.micrometer.observation.tck.TestObservationRegistry;
import static org.assertj.core.api.Assertions.assertThat;
import static org.springframework.ai.autoconfigure.vectorstore.observation.ObservationTestUtil.assertObservationRegistry;
/**
* @author Christian Tzolov
@@ -73,6 +78,8 @@ public class WeaviateVectorStoreAutoConfigurationIT {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
var bgDocument = new Document("The World is Big and Salvation Lurks Around the Corner",
Map.of("country", "Bulgaria", "price", 3.14, "active", true, "year", 2020));
var nlDocument = new Document("The World is Big and Salvation Lurks Around the Corner",
@@ -80,6 +87,10 @@ public class WeaviateVectorStoreAutoConfigurationIT {
vectorStore.add(List.of(bgDocument, nlDocument));
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.WEAVIATE,
VectorStoreObservationContext.Operation.ADD);
observationRegistry.clear();
var request = SearchRequest.query("The World").withTopK(5);
List<Document> results = vectorStore.similaritySearch(request);
@@ -90,6 +101,9 @@ public class WeaviateVectorStoreAutoConfigurationIT {
assertThat(results).hasSize(1);
assertThat(results.get(0).getId()).isEqualTo(bgDocument.getId());
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.WEAVIATE,
VectorStoreObservationContext.Operation.QUERY);
results = vectorStore.similaritySearch(
request.withSimilarityThresholdAll().withFilterExpression("country == 'Netherlands'"));
assertThat(results).hasSize(1);
@@ -109,14 +123,24 @@ public class WeaviateVectorStoreAutoConfigurationIT {
assertThat(results).hasSize(1);
assertThat(results.get(0).getId()).isEqualTo(nlDocument.getId());
observationRegistry.clear();
// Remove all documents from the store
vectorStore.delete(List.of(bgDocument, nlDocument).stream().map(doc -> doc.getId()).toList());
assertObservationRegistry(observationRegistry, "vector_store", VectorStoreProvider.WEAVIATE,
VectorStoreObservationContext.Operation.DELETE);
});
}
@Configuration(proxyBeanMethods = false)
static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public EmbeddingModel embeddingModel() {
return new TransformersEmbeddingModel();

View File

@@ -18,6 +18,7 @@ package org.springframework.ai.testcontainers.service.connection.ollama;
import org.apache.commons.logging.Log;
import org.apache.commons.logging.LogFactory;
import org.junit.jupiter.api.BeforeAll;
import org.junit.jupiter.api.Disabled;
import org.junit.jupiter.api.Test;
import org.springframework.ai.autoconfigure.ollama.OllamaAutoConfiguration;
import org.springframework.ai.embedding.EmbeddingResponse;
@@ -42,6 +43,7 @@ import static org.assertj.core.api.Assertions.assertThat;
* @author Eddú Meléndez
*/
@SpringJUnitConfig
@Disabled("requires more memory than is often available on dev machines")
@Testcontainers
@TestPropertySource(properties = "spring.ai.ollama.embedding.options.model="
+ OllamaContainerConnectionDetailsFactoryTest.MODEL_NAME)

View File

@@ -15,23 +15,22 @@
*/
package org.springframework.ai.evaluation;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.chat.model.ChatModel;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.chat.prompt.PromptTemplate;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.SystemMessage;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.core.io.Resource;
import static org.assertj.core.api.Assertions.assertThat;
import static org.assertj.core.api.Assertions.fail;
import java.util.List;
import java.util.Map;
import static org.assertj.core.api.Assertions.assertThat;
import static org.assertj.core.api.Assertions.fail;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.SystemMessage;
import org.springframework.ai.chat.model.ChatModel;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.chat.prompt.PromptTemplate;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.core.io.Resource;
public class BasicEvaluationTest {

View File

@@ -80,6 +80,11 @@
<artifactId>awaitility</artifactId>
<scope>test</scope>
</dependency>
<dependency>
<groupId>io.micrometer</groupId>
<artifactId>micrometer-observation-test</artifactId>
<scope>test</scope>
</dependency>
</dependencies>

View File

@@ -15,6 +15,31 @@
*/
package org.springframework.ai.vectorstore.azure;
import java.util.ArrayList;
import java.util.Collections;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import java.util.stream.Collectors;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.model.EmbeddingUtils;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.observation.conventions.VectorStoreSimilarityMetric;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.filter.FilterExpressionConverter;
import org.springframework.ai.vectorstore.observation.AbstractObservationVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext.Builder;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.InitializingBean;
import org.springframework.util.Assert;
import org.springframework.util.CollectionUtils;
import org.springframework.util.StringUtils;
import com.alibaba.fastjson2.JSONObject;
import com.alibaba.fastjson2.TypeReference;
import com.azure.core.util.Context;
@@ -34,25 +59,8 @@ import com.azure.search.documents.models.IndexingResult;
import com.azure.search.documents.models.SearchOptions;
import com.azure.search.documents.models.VectorSearchOptions;
import com.azure.search.documents.models.VectorizedQuery;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.model.EmbeddingUtils;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.VectorStore;
import org.springframework.ai.vectorstore.filter.FilterExpressionConverter;
import org.springframework.beans.factory.InitializingBean;
import org.springframework.util.Assert;
import org.springframework.util.CollectionUtils;
import org.springframework.util.StringUtils;
import java.util.ArrayList;
import java.util.Collections;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import java.util.stream.Collectors;
import io.micrometer.observation.ObservationRegistry;
/**
* Uses Azure Cognitive Search as a backing vector store. Documents can be preloaded into
@@ -65,7 +73,7 @@ import java.util.stream.Collectors;
* @author Christian Tzolov
* @author Josh Long
*/
public class AzureVectorStore implements VectorStore, InitializingBean {
public class AzureVectorStore extends AbstractObservationVectorStore implements InitializingBean {
private static final Logger logger = LoggerFactory.getLogger(AzureVectorStore.class);
@@ -166,6 +174,25 @@ public class AzureVectorStore implements VectorStore, InitializingBean {
*/
public AzureVectorStore(SearchIndexClient searchIndexClient, EmbeddingModel embeddingModel,
boolean initializeSchema, List<MetadataField> filterMetadataFields) {
this(searchIndexClient, embeddingModel, initializeSchema, filterMetadataFields, ObservationRegistry.NOOP, null);
}
/**
* Constructs a new AzureCognitiveSearchVectorStore.
* @param searchIndexClient A pre-configured Azure {@link SearchIndexClient} that CRUD
* for Azure search indexes and factory for {@link SearchClient}.
* @param embeddingModel The client for embedding operations.
* @param filterMetadataFields List of metadata fields (as field name and type) that
* can be used in similarity search query filter expressions.
* @param observationRegistry The observation registry to use.
* @param customObservationConvention The optional, custom search observation
* convention to use.
*/
public AzureVectorStore(SearchIndexClient searchIndexClient, EmbeddingModel embeddingModel,
boolean initializeSchema, List<MetadataField> filterMetadataFields, ObservationRegistry observationRegistry,
VectorStoreObservationConvention customObservationConvention) {
super(observationRegistry, customObservationConvention);
Assert.notNull(embeddingModel, "The embedding model can not be null.");
Assert.notNull(searchIndexClient, "The search index client can not be null.");
@@ -208,7 +235,7 @@ public class AzureVectorStore implements VectorStore, InitializingBean {
}
@Override
public void add(List<Document> documents) {
public void doAdd(List<Document> documents) {
Assert.notNull(documents, "The document list should not be null.");
if (CollectionUtils.isEmpty(documents)) {
@@ -243,7 +270,7 @@ public class AzureVectorStore implements VectorStore, InitializingBean {
}
@Override
public Optional<Boolean> delete(List<String> documentIds) {
public Optional<Boolean> doDelete(List<String> documentIds) {
Assert.notNull(documentIds, "The document ID list should not be null.");
if (CollectionUtils.isEmpty(documentIds)) {
@@ -278,7 +305,7 @@ public class AzureVectorStore implements VectorStore, InitializingBean {
}
@Override
public List<Document> similaritySearch(SearchRequest request) {
public List<Document> doSimilaritySearch(SearchRequest request) {
Assert.notNull(request, "The search request must not be null.");
@@ -379,4 +406,13 @@ public class AzureVectorStore implements VectorStore, InitializingBean {
this.searchClient = this.searchIndexClient.getSearchClient(this.indexName);
}
@Override
public Builder createObservationContextBuilder(String operationName) {
return VectorStoreObservationContext.builder(VectorStoreProvider.AZURE.value(), operationName)
.withDimensions(this.embeddingModel.dimensions())
.withSimilarityMetric(this.initializeSchema ? VectorStoreSimilarityMetric.COSINE.value() : null)
.withIndexName(this.indexName);
}
}

View File

@@ -0,0 +1,183 @@
/*
* Copyright 2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (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.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.vectorstore.azure;
import static org.assertj.core.api.Assertions.assertThat;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.time.Duration;
import java.util.List;
import java.util.Map;
import java.util.concurrent.TimeUnit;
import org.awaitility.Awaitility;
import org.junit.jupiter.api.BeforeAll;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.VectorStore;
import org.springframework.ai.vectorstore.azure.AzureVectorStore.MetadataField;
import org.springframework.ai.vectorstore.observation.DefaultVectorStoreObservationConvention;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.HighCardinalityKeyNames;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.LowCardinalityKeyNames;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.boot.autoconfigure.EnableAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.core.io.DefaultResourceLoader;
import com.azure.core.credential.AzureKeyCredential;
import com.azure.search.documents.indexes.SearchIndexClient;
import com.azure.search.documents.indexes.SearchIndexClientBuilder;
import io.micrometer.observation.ObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistryAssert;
/**
* Integration tests for observation instrumentation AbstractObservationVectorStore in
* {@link AzureVectorStore}.
*
* @author Christian Tzolov
*/
@EnabledIfEnvironmentVariable(named = "AZURE_AI_SEARCH_API_KEY", matches = ".+")
@EnabledIfEnvironmentVariable(named = "AZURE_AI_SEARCH_ENDPOINT", matches = ".+")
public class AzureVectorStoreObservationIT {
List<Document> documents = List.of(
new Document(getText("classpath:/test/data/spring.ai.txt"), Map.of("meta1", "meta1")),
new Document(getText("classpath:/test/data/time.shelter.txt")),
new Document(getText("classpath:/test/data/great.depression.txt"), Map.of("meta2", "meta2")));
public static String getText(String uri) {
var resource = new DefaultResourceLoader().getResource(uri);
try {
return resource.getContentAsString(StandardCharsets.UTF_8);
}
catch (IOException e) {
throw new RuntimeException(e);
}
}
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withUserConfiguration(Config.class);
@BeforeAll
public static void beforeAll() {
Awaitility.setDefaultPollInterval(2, TimeUnit.SECONDS);
Awaitility.setDefaultPollDelay(Duration.ZERO);
Awaitility.setDefaultTimeout(Duration.ofMinutes(1));
}
@Test
void observationVectorStoreAddAndQueryOperations() {
contextRunner.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store azure add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(), "azure")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "384")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "cosine")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "none")
.hasBeenStarted()
.hasBeenStopped();
observationRegistry.clear();
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("What is Great Depression").withTopK(1));
assertThat(results).isNotEmpty();
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store azure query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(), "azure")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "What is Great Depression")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "384")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "cosine")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "1")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "0.0")
.hasBeenStarted()
.hasBeenStopped();
});
}
@SpringBootConfiguration
@EnableAutoConfiguration
public static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public SearchIndexClient searchIndexClient() {
return new SearchIndexClientBuilder().endpoint(System.getenv("AZURE_AI_SEARCH_ENDPOINT"))
.credential(new AzureKeyCredential(System.getenv("AZURE_AI_SEARCH_API_KEY")))
.buildClient();
}
@Bean
public VectorStore vectorStore(SearchIndexClient searchIndexClient, EmbeddingModel embeddingModel,
ObservationRegistry observationRegistry) {
var filterableMetaFields = List.of(MetadataField.text("country"), MetadataField.int64("year"),
MetadataField.date("activationDate"));
return new AzureVectorStore(searchIndexClient, embeddingModel, true, filterableMetaFields,
observationRegistry, null);
}
@Bean
public EmbeddingModel embeddingModel() {
return new TransformersEmbeddingModel();
}
}
}

View File

@@ -71,6 +71,11 @@
<scope>test</scope>
</dependency>
<dependency>
<groupId>io.micrometer</groupId>
<artifactId>micrometer-observation-test</artifactId>
<scope>test</scope>
</dependency>
</dependencies>
</project>

View File

@@ -29,14 +29,22 @@ import com.datastax.oss.driver.api.querybuilder.insert.InsertInto;
import com.datastax.oss.driver.api.querybuilder.insert.RegularInsert;
import com.datastax.oss.driver.shaded.guava.common.base.Preconditions;
import io.micrometer.observation.ObservationRegistry;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.model.EmbeddingUtils;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.observation.conventions.VectorStoreSimilarityMetric;
import org.springframework.ai.vectorstore.CassandraVectorStoreConfig.SchemaColumn;
import org.springframework.ai.vectorstore.filter.FilterExpressionConverter;
import org.springframework.ai.vectorstore.observation.AbstractObservationVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext.Builder;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import java.util.ArrayList;
import java.util.HashMap;
@@ -86,12 +94,13 @@ import java.util.concurrent.ConcurrentMap;
* also serve as a protecting throttle against your embedding model.
*
* @author Mick Semb Wever
* @author Christian Tzolov
* @see VectorStore
* @see org.springframework.ai.vectorstore.CassandraVectorStoreConfig
* @see EmbeddingModel
* @since 1.0.0
*/
public class CassandraVectorStore implements VectorStore, AutoCloseable {
public class CassandraVectorStore extends AbstractObservationVectorStore implements AutoCloseable {
/**
* Indexes are automatically created with COSINE. This can be changed manually via
@@ -127,11 +136,14 @@ public class CassandraVectorStore implements VectorStore, AutoCloseable {
private final Similarity similarity;
public static CassandraVectorStore create(CassandraVectorStoreConfig conf, EmbeddingModel embeddingModel) {
return new CassandraVectorStore(conf, embeddingModel);
public CassandraVectorStore(CassandraVectorStoreConfig conf, EmbeddingModel embeddingModel) {
this(conf, embeddingModel, ObservationRegistry.NOOP, null);
}
public CassandraVectorStore(CassandraVectorStoreConfig conf, EmbeddingModel embeddingModel) {
public CassandraVectorStore(CassandraVectorStoreConfig conf, EmbeddingModel embeddingModel,
ObservationRegistry observationRegistry, VectorStoreObservationConvention customObservationConvention) {
super(observationRegistry, customObservationConvention);
Preconditions.checkArgument(null != conf, "Config must not be null");
Preconditions.checkArgument(null != embeddingModel, "Embedding model must not be null");
@@ -156,7 +168,7 @@ public class CassandraVectorStore implements VectorStore, AutoCloseable {
}
@Override
public void add(List<Document> documents) {
public void doAdd(List<Document> documents) {
var futures = new CompletableFuture[documents.size()];
int i = 0;
@@ -194,7 +206,7 @@ public class CassandraVectorStore implements VectorStore, AutoCloseable {
}
@Override
public Optional<Boolean> delete(List<String> idList) {
public Optional<Boolean> doDelete(List<String> idList) {
CompletableFuture[] futures = new CompletableFuture[idList.size()];
int i = 0;
for (String id : idList) {
@@ -207,7 +219,7 @@ public class CassandraVectorStore implements VectorStore, AutoCloseable {
}
@Override
public List<Document> similaritySearch(SearchRequest request) {
public List<Document> doSimilaritySearch(SearchRequest request) {
Preconditions.checkArgument(request.getTopK() <= 1000);
var embedding = toFloatArray(this.embeddingModel.embed(request.getQuery()));
CqlVector<Float> cqlVector = CqlVector.newInstance(embedding);
@@ -366,4 +378,25 @@ public class CassandraVectorStore implements VectorStore, AutoCloseable {
return embeddingFloat;
}
@Override
public Builder createObservationContextBuilder(String operationName) {
return VectorStoreObservationContext.builder(VectorStoreProvider.CASSANDRA.value(), operationName)
.withDimensions(this.embeddingModel.dimensions())
.withCollectionName(this.conf.schema.table())
.withNamespace(this.conf.schema.keyspace())
.withSimilarityMetric(getSimilarityMetric())
.withIndexName(this.conf.schema.index());
}
private static Map<Similarity, VectorStoreSimilarityMetric> SIMILARITY_TYPE_MAPPING = Map.of(Similarity.COSINE,
VectorStoreSimilarityMetric.COSINE, Similarity.EUCLIDEAN, VectorStoreSimilarityMetric.EUCLIDEAN,
Similarity.DOT_PRODUCT, VectorStoreSimilarityMetric.DOT);
private String getSimilarityMetric() {
if (!SIMILARITY_TYPE_MAPPING.containsKey(this.similarity)) {
return this.similarity.name();
}
return SIMILARITY_TYPE_MAPPING.get(this.similarity).value();
}
}

View File

@@ -195,7 +195,7 @@ class CassandraRichSchemaVectorStoreIT {
contextRunner.run(context -> {
try (CassandraVectorStore store = CassandraVectorStore.create(
try (CassandraVectorStore store = new CassandraVectorStore(
storeBuilder(context, List.of()).withFixedThreadPoolExecutorSize(nThreads).build(),
context.getBean(EmbeddingModel.class))) {

View File

@@ -396,7 +396,7 @@ class CassandraVectorStoreIT {
.build();
conf.dropKeyspace();
return CassandraVectorStore.create(conf, embeddingModel);
return new CassandraVectorStore(conf, embeddingModel);
}
@Bean
@@ -432,7 +432,7 @@ class CassandraVectorStoreIT {
CassandraVectorStoreConfig.Builder builder) {
CassandraVectorStoreConfig conf = builder.build();
conf.dropKeyspace();
return CassandraVectorStore.create(conf, context.getBean(EmbeddingModel.class));
return new CassandraVectorStore(conf, context.getBean(EmbeddingModel.class));
}
}

View File

@@ -0,0 +1,187 @@
/*
* Copyright 2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (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.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.vectorstore;
import static org.assertj.core.api.Assertions.assertThat;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.util.List;
import java.util.Map;
import org.junit.jupiter.api.Test;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.CassandraVectorStoreConfig.SchemaColumn;
import org.springframework.ai.vectorstore.observation.DefaultVectorStoreObservationConvention;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.HighCardinalityKeyNames;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.LowCardinalityKeyNames;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.boot.autoconfigure.EnableAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.core.io.DefaultResourceLoader;
import org.testcontainers.containers.CassandraContainer;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import org.testcontainers.utility.DockerImageName;
import com.datastax.oss.driver.api.core.CqlSession;
import com.datastax.oss.driver.api.core.CqlSessionBuilder;
import com.datastax.oss.driver.api.core.type.DataTypes;
import io.micrometer.observation.ObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistryAssert;
/**
* @author Christian Tzolov
*/
@Testcontainers
public class CassandraVectorStoreObservationIT {
static final DockerImageName DEFAULT_IMAGE_NAME = DockerImageName.parse("cassandra");
@Container
static CassandraContainer cassandraContainer = new CassandraContainer(DEFAULT_IMAGE_NAME.withTag("5.0"));
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withUserConfiguration(Config.class);
List<Document> documents = List.of(
new Document(getText("classpath:/test/data/spring.ai.txt"), Map.of("meta1", "meta1")),
new Document(getText("classpath:/test/data/time.shelter.txt")),
new Document(getText("classpath:/test/data/great.depression.txt"), Map.of("meta2", "meta2")));
public static String getText(String uri) {
var resource = new DefaultResourceLoader().getResource(uri);
try {
return resource.getContentAsString(StandardCharsets.UTF_8);
}
catch (IOException e) {
throw new RuntimeException(e);
}
}
@Test
void observationVectorStoreAddAndQueryOperations() {
contextRunner.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store cassandra add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(), "cassandra")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "384")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "ai_vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "test_springframework")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "cosine")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "none")
.hasBeenStarted()
.hasBeenStopped();
observationRegistry.clear();
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("What is Great Depression").withTopK(1));
assertThat(results).isNotEmpty();
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store cassandra query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(), "cassandra")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "What is Great Depression")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "384")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "ai_vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "test_springframework")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "cosine")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "1")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "0.0")
.hasBeenStarted()
.hasBeenStopped();
});
}
@SpringBootConfiguration
@EnableAutoConfiguration
public static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public EmbeddingModel embeddingModel() {
return new TransformersEmbeddingModel();
}
@Bean
public CassandraVectorStore store(CqlSession cqlSession, EmbeddingModel embeddingModel,
ObservationRegistry observationRegistry) {
CassandraVectorStoreConfig conf = storeBuilder(cqlSession)
.addMetadataColumns(new SchemaColumn("meta1", DataTypes.TEXT),
new SchemaColumn("meta2", DataTypes.TEXT), new SchemaColumn("country", DataTypes.TEXT),
new SchemaColumn("year", DataTypes.SMALLINT))
.build();
conf.dropKeyspace();
return new CassandraVectorStore(conf, embeddingModel, observationRegistry, null);
}
@Bean
public CqlSession cqlSession() {
return new CqlSessionBuilder()
// comment next two lines out to connect to a local C* cluster
.addContactPoint(cassandraContainer.getContactPoint())
.withLocalDatacenter(cassandraContainer.getLocalDatacenter())
.build();
}
}
private static CassandraVectorStoreConfig.Builder storeBuilder(CqlSession cqlSession) {
return CassandraVectorStoreConfig.builder()
.withCqlSession(cqlSession)
.withKeyspaceName("test_" + CassandraVectorStoreConfig.DEFAULT_KEYSPACE_NAME);
}
}

View File

@@ -119,7 +119,7 @@ class WikiVectorStoreExample {
})
.build();
return CassandraVectorStore.create(conf, embeddingModel());
return new CassandraVectorStore(conf, embeddingModel());
}
@Bean

View File

@@ -1,6 +1,7 @@
<?xml version="1.0" encoding="UTF-8"?>
<project xmlns="http://maven.apache.org/POM/4.0.0"
xmlns:xsi="http://www.w3.org/2001/XMLSchema-instance" xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 http://maven.apache.org/maven-v4_0_0.xsd">
xmlns:xsi="http://www.w3.org/2001/XMLSchema-instance"
xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 http://maven.apache.org/maven-v4_0_0.xsd">
<modelVersion>4.0.0</modelVersion>
<parent>
<groupId>org.springframework.ai</groupId>
@@ -58,6 +59,18 @@
<scope>test</scope>
</dependency>
<dependency>
<groupId>org.springframework.ai</groupId>
<artifactId>spring-ai-test</artifactId>
<version>${project.parent.version}</version>
<scope>test</scope>
</dependency>
<dependency>
<groupId>io.micrometer</groupId>
<artifactId>micrometer-observation-test</artifactId>
<scope>test</scope>
</dependency>
</dependencies>
</project>
</project>

View File

@@ -27,12 +27,19 @@ import org.springframework.ai.chroma.ChromaApi.DeleteEmbeddingsRequest;
import org.springframework.ai.chroma.ChromaApi.Embedding;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.vectorstore.filter.FilterExpressionConverter;
import org.springframework.ai.vectorstore.observation.AbstractObservationVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext.Builder;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.InitializingBean;
import org.springframework.util.Assert;
import org.springframework.util.CollectionUtils;
import org.springframework.util.StringUtils;
import io.micrometer.observation.ObservationRegistry;
/**
* {@link ChromaVectorStore} is a concrete implementation of the {@link VectorStore}
* interface. It is responsible for adding, deleting, and searching documents based on
@@ -40,7 +47,7 @@ import org.springframework.util.StringUtils;
* embedding calculations. For more information about how it does this, see the official
* <a href="https://www.trychroma.com/">Chroma website</a>.
*/
public class ChromaVectorStore implements VectorStore, InitializingBean {
public class ChromaVectorStore extends AbstractObservationVectorStore implements InitializingBean {
public static final String DISTANCE_FIELD_NAME = "distance";
@@ -68,6 +75,15 @@ public class ChromaVectorStore implements VectorStore, InitializingBean {
public ChromaVectorStore(EmbeddingModel embeddingModel, ChromaApi chromaApi, String collectionName,
boolean initializeSchema) {
this(embeddingModel, chromaApi, collectionName, initializeSchema, ObservationRegistry.NOOP, null);
}
public ChromaVectorStore(EmbeddingModel embeddingModel, ChromaApi chromaApi, String collectionName,
boolean initializeSchema, ObservationRegistry observationRegistry,
VectorStoreObservationConvention customObservationConvention) {
super(observationRegistry, customObservationConvention);
this.embeddingModel = embeddingModel;
this.chromaApi = chromaApi;
this.collectionName = collectionName;
@@ -81,7 +97,7 @@ public class ChromaVectorStore implements VectorStore, InitializingBean {
}
@Override
public void add(List<Document> documents) {
public void doAdd(List<Document> documents) {
Assert.notNull(documents, "Documents must not be null");
if (CollectionUtils.isEmpty(documents)) {
return;
@@ -105,7 +121,7 @@ public class ChromaVectorStore implements VectorStore, InitializingBean {
}
@Override
public Optional<Boolean> delete(List<String> idList) {
public Optional<Boolean> doDelete(List<String> idList) {
Assert.notNull(idList, "Document id list must not be null");
List<String> deletedIds = this.chromaApi.deleteEmbeddings(this.collectionId,
new DeleteEmbeddingsRequest(idList));
@@ -113,7 +129,7 @@ public class ChromaVectorStore implements VectorStore, InitializingBean {
}
@Override
public List<Document> similaritySearch(SearchRequest request) {
public List<Document> doSimilaritySearch(SearchRequest request) {
String nativeFilterExpression = (request.getFilterExpression() != null)
? this.filterExpressionConverter.convertExpression(request.getFilterExpression()) : "";
@@ -149,6 +165,14 @@ public class ChromaVectorStore implements VectorStore, InitializingBean {
return responseDocuments;
}
public String getCollectionName() {
return this.collectionName;
}
public String getCollectionId() {
return this.collectionId;
}
@Override
public void afterPropertiesSet() throws Exception {
@@ -162,4 +186,12 @@ public class ChromaVectorStore implements VectorStore, InitializingBean {
this.collectionId = collection.id();
}
@Override
public Builder createObservationContextBuilder(String operationName) {
return VectorStoreObservationContext.builder(VectorStoreProvider.CHROMA.value(), operationName)
.withDimensions(this.embeddingModel.dimensions())
.withCollectionName(this.collectionName + ":" + this.collectionId)
.withFieldName(this.initializeSchema ? DISTANCE_FIELD_NAME : null);
}
}

View File

@@ -0,0 +1,175 @@
/*
* Copyright 2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (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.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.vectorstore;
import static org.assertj.core.api.Assertions.assertThat;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.util.List;
import java.util.Map;
import org.junit.jupiter.api.Test;
import org.springframework.ai.chroma.ChromaApi;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.openai.OpenAiEmbeddingModel;
import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.ai.vectorstore.observation.DefaultVectorStoreObservationConvention;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.HighCardinalityKeyNames;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.LowCardinalityKeyNames;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.boot.autoconfigure.EnableAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.core.io.DefaultResourceLoader;
import org.springframework.http.client.SimpleClientHttpRequestFactory;
import org.springframework.web.client.RestClient;
import org.testcontainers.chromadb.ChromaDBContainer;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import io.micrometer.observation.ObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistryAssert;
/**
* @author Christian Tzolov
*/
@Testcontainers
public class ChromaVectorStoreObservationIT {
@Container
static ChromaDBContainer chromaContainer = new ChromaDBContainer("ghcr.io/chroma-core/chroma:0.5.0");
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withUserConfiguration(Config.class);
List<Document> documents = List.of(
new Document(getText("classpath:/test/data/spring.ai.txt"), Map.of("meta1", "meta1")),
new Document(getText("classpath:/test/data/time.shelter.txt")),
new Document(getText("classpath:/test/data/great.depression.txt"), Map.of("meta2", "meta2")));
public static String getText(String uri) {
var resource = new DefaultResourceLoader().getResource(uri);
try {
return resource.getContentAsString(StandardCharsets.UTF_8);
}
catch (IOException e) {
throw new RuntimeException(e);
}
}
@Test
void observationVectorStoreAddAndQueryOperations() {
contextRunner.run(context -> {
ChromaVectorStore vectorStore = context.getBean(ChromaVectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store chroma add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.CHROMA.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "1536")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(),
"TestCollection:" + vectorStore.getCollectionId())
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "distance")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "none")
.hasBeenStarted()
.hasBeenStopped();
observationRegistry.clear();
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("What is Great Depression").withTopK(1));
assertThat(results).isNotEmpty();
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store chroma query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.CHROMA.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "What is Great Depression")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "1536")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(),
"TestCollection:" + vectorStore.getCollectionId())
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "distance")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "1")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "0.0")
.hasBeenStarted()
.hasBeenStopped();
});
}
@SpringBootConfiguration
@EnableAutoConfiguration
public static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public RestClient.Builder builder() {
return RestClient.builder().requestFactory(new SimpleClientHttpRequestFactory());
}
@Bean
public ChromaApi chromaApi(RestClient.Builder builder) {
return new ChromaApi(chromaContainer.getEndpoint(), builder);
}
@Bean
public VectorStore chromaVectorStore(EmbeddingModel embeddingModel, ChromaApi chromaApi,
ObservationRegistry observationRegistry) {
return new ChromaVectorStore(embeddingModel, chromaApi, "TestCollection", true, observationRegistry, null);
}
@Bean
public EmbeddingModel embeddingModel() {
return new OpenAiEmbeddingModel(new OpenAiApi(System.getenv("OPENAI_API_KEY")));
}
}
}

View File

@@ -71,6 +71,11 @@
<scope>test</scope>
</dependency>
<dependency>
<groupId>io.micrometer</groupId>
<artifactId>micrometer-observation-test</artifactId>
<scope>test</scope>
</dependency>
</dependencies>
</project>

View File

@@ -15,6 +15,35 @@
*/
package org.springframework.ai.vectorstore;
import static java.lang.Math.sqrt;
import java.io.IOException;
import java.util.List;
import java.util.Map;
import java.util.Objects;
import java.util.Optional;
import java.util.stream.Collectors;
import org.elasticsearch.client.RestClient;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.model.EmbeddingUtils;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.observation.conventions.VectorStoreSimilarityMetric;
import org.springframework.ai.vectorstore.filter.Filter;
import org.springframework.ai.vectorstore.filter.FilterExpressionConverter;
import org.springframework.ai.vectorstore.observation.AbstractObservationVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext.Builder;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.InitializingBean;
import org.springframework.util.Assert;
import com.fasterxml.jackson.databind.DeserializationFeature;
import com.fasterxml.jackson.databind.ObjectMapper;
import co.elastic.clients.elasticsearch.ElasticsearchClient;
import co.elastic.clients.elasticsearch.core.BulkRequest;
import co.elastic.clients.elasticsearch.core.BulkResponse;
@@ -23,27 +52,7 @@ import co.elastic.clients.elasticsearch.core.bulk.BulkResponseItem;
import co.elastic.clients.elasticsearch.core.search.Hit;
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.elasticsearch.client.RestClient;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.model.EmbeddingUtils;
import org.springframework.ai.vectorstore.filter.Filter;
import org.springframework.ai.vectorstore.filter.FilterExpressionConverter;
import org.springframework.beans.factory.InitializingBean;
import org.springframework.util.Assert;
import java.io.IOException;
import java.util.List;
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;
import io.micrometer.observation.ObservationRegistry;
/**
* The ElasticsearchVectorStore class implements the VectorStore interface and provides
@@ -58,9 +67,10 @@ import static org.springframework.ai.vectorstore.SimilarityFunction.l2_norm;
* @author Wei Jiang
* @author Laura Trotta
* @author Soby Chacko
* @author Christian Tzolov
* @since 1.0.0
*/
public class ElasticsearchVectorStore implements VectorStore, InitializingBean {
public class ElasticsearchVectorStore extends AbstractObservationVectorStore implements InitializingBean {
private static final Logger logger = LoggerFactory.getLogger(ElasticsearchVectorStore.class);
@@ -80,6 +90,15 @@ public class ElasticsearchVectorStore implements VectorStore, InitializingBean {
public ElasticsearchVectorStore(ElasticsearchVectorStoreOptions options, RestClient restClient,
EmbeddingModel embeddingModel, boolean initializeSchema) {
this(options, restClient, embeddingModel, initializeSchema, ObservationRegistry.NOOP, null);
}
public ElasticsearchVectorStore(ElasticsearchVectorStoreOptions options, RestClient restClient,
EmbeddingModel embeddingModel, boolean initializeSchema, ObservationRegistry observationRegistry,
VectorStoreObservationConvention customObservationConvention) {
super(observationRegistry, customObservationConvention);
this.initializeSchema = initializeSchema;
Objects.requireNonNull(embeddingModel, "RestClient must not be null");
Objects.requireNonNull(embeddingModel, "EmbeddingModel must not be null");
@@ -91,7 +110,7 @@ public class ElasticsearchVectorStore implements VectorStore, InitializingBean {
}
@Override
public void add(List<Document> documents) {
public void doAdd(List<Document> documents) {
BulkRequest.Builder bulkRequestBuilder = new BulkRequest.Builder();
for (Document document : documents) {
@@ -119,7 +138,7 @@ public class ElasticsearchVectorStore implements VectorStore, InitializingBean {
}
@Override
public Optional<Boolean> delete(List<String> idList) {
public Optional<Boolean> doDelete(List<String> idList) {
BulkRequest.Builder bulkRequestBuilder = new BulkRequest.Builder();
// We call operations on BulkRequest.Builder only if the index exists.
// For the index to be present, either it must be pre-created or set the
@@ -142,12 +161,12 @@ public class ElasticsearchVectorStore implements VectorStore, InitializingBean {
}
@Override
public List<Document> similaritySearch(SearchRequest searchRequest) {
public List<Document> doSimilaritySearch(SearchRequest searchRequest) {
Assert.notNull(searchRequest, "The search request must not be null.");
try {
float threshold = (float) searchRequest.getSimilarityThreshold();
// reverting l2_norm distance to its original value
if (options.getSimilarity().equals(l2_norm)) {
if (options.getSimilarity().equals(SimilarityFunction.l2_norm)) {
threshold = 1 - threshold;
}
final float finalThreshold = threshold;
@@ -230,4 +249,24 @@ public class ElasticsearchVectorStore implements VectorStore, InitializingBean {
}
}
@Override
public Builder createObservationContextBuilder(String operationName) {
return VectorStoreObservationContext.builder(VectorStoreProvider.ELASTICSEARCH.value(), operationName)
.withDimensions(this.embeddingModel.dimensions())
.withIndexName(this.options.getIndexName())
.withSimilarityMetric(getSimilarityMetric());
}
private static Map<SimilarityFunction, VectorStoreSimilarityMetric> SIMILARITY_TYPE_MAPPING = Map.of(
SimilarityFunction.cosine, VectorStoreSimilarityMetric.COSINE, SimilarityFunction.l2_norm,
VectorStoreSimilarityMetric.EUCLIDEAN, SimilarityFunction.dot_product, VectorStoreSimilarityMetric.DOT);
private String getSimilarityMetric() {
if (!SIMILARITY_TYPE_MAPPING.containsKey(this.options.getSimilarity())) {
return this.options.getSimilarity().name();
}
return SIMILARITY_TYPE_MAPPING.get(this.options.getSimilarity()).value();
}
}

View File

@@ -0,0 +1,217 @@
/*
* Copyright 2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (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.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.vectorstore;
import static org.assertj.core.api.Assertions.assertThat;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.time.Duration;
import java.util.List;
import java.util.Map;
import java.util.concurrent.TimeUnit;
import org.apache.http.HttpHost;
import org.awaitility.Awaitility;
import org.elasticsearch.client.RestClient;
import org.junit.jupiter.api.BeforeAll;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.openai.OpenAiEmbeddingModel;
import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.ai.vectorstore.observation.DefaultVectorStoreObservationConvention;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.HighCardinalityKeyNames;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.LowCardinalityKeyNames;
import org.springframework.boot.SpringBootConfiguration;
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.springframework.core.io.DefaultResourceLoader;
import org.testcontainers.elasticsearch.ElasticsearchContainer;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import com.fasterxml.jackson.databind.DeserializationFeature;
import com.fasterxml.jackson.databind.ObjectMapper;
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 io.micrometer.observation.ObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistryAssert;
import static org.hamcrest.Matchers.greaterThan;;
/**
* @author Christian Tzolov
*/
@Testcontainers
@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+")
public class ElasticsearchVectorStoreObservationIT {
@Container
private static final ElasticsearchContainer elasticsearchContainer = new ElasticsearchContainer(
"docker.elastic.co/elasticsearch/elasticsearch:8.13.3")
.withEnv("xpack.security.enabled", "false");
List<Document> documents = List.of(
new Document(getText("classpath:/test/data/spring.ai.txt"), Map.of("meta1", "meta1")),
new Document(getText("classpath:/test/data/time.shelter.txt")),
new Document(getText("classpath:/test/data/great.depression.txt"), Map.of("meta2", "meta2")));
public static String getText(String uri) {
var resource = new DefaultResourceLoader().getResource(uri);
try {
return resource.getContentAsString(StandardCharsets.UTF_8);
}
catch (IOException e) {
throw new RuntimeException(e);
}
}
private ApplicationContextRunner getContextRunner() {
return new ApplicationContextRunner().withUserConfiguration(Config.class);
}
@BeforeAll
public static void beforeAll() {
Awaitility.setDefaultPollInterval(2, TimeUnit.SECONDS);
Awaitility.setDefaultPollDelay(Duration.ZERO);
Awaitility.setDefaultTimeout(Duration.ofMinutes(1));
}
@BeforeEach
void cleanDatabase() {
getContextRunner().run(context -> {
// 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));
}
});
}
@Test
void observationVectorStoreAddAndQueryOperations() {
getContextRunner().run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store elasticsearch add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.ELASTICSEARCH.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "1536")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "cosine")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "none")
.hasBeenStarted()
.hasBeenStopped();
Awaitility.await()
.until(() -> vectorStore
.similaritySearch(SearchRequest.query("What is Great Depression").withSimilarityThresholdAll())
.size(), greaterThan(1));
observationRegistry.clear();
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("What is Great Depression").withTopK(1));
assertThat(results).isNotEmpty();
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store elasticsearch query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.ELASTICSEARCH.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "What is Great Depression")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "1536")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "cosine")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "1")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "0.0")
.hasBeenStarted()
.hasBeenStopped();
});
}
@SpringBootConfiguration
@EnableAutoConfiguration(exclude = { DataSourceAutoConfiguration.class })
public static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public ElasticsearchVectorStore vectorStoreDefault(EmbeddingModel embeddingModel, RestClient restClient,
ObservationRegistry observationRegistry) {
return new ElasticsearchVectorStore(new ElasticsearchVectorStoreOptions(), restClient, embeddingModel, true,
observationRegistry, null);
}
@Bean
public EmbeddingModel embeddingModel() {
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))));
}
}
}

View File

@@ -78,7 +78,11 @@
<version>3.0.0</version>
<scope>test</scope>
</dependency>
<dependency>
<groupId>io.micrometer</groupId>
<artifactId>micrometer-observation-test</artifactId>
<scope>test</scope>
</dependency>
</dependencies>
</project>

View File

@@ -23,15 +23,15 @@ import java.util.List;
import java.util.Map;
import java.util.Optional;
import com.fasterxml.jackson.annotation.JsonCreator;
import com.fasterxml.jackson.annotation.JsonInclude;
import com.fasterxml.jackson.annotation.JsonProperty;
import com.fasterxml.jackson.core.JsonProcessingException;
import com.fasterxml.jackson.databind.ObjectMapper;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.vectorstore.observation.AbstractObservationVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext.Builder;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.InitializingBean;
import org.springframework.http.HttpMethod;
import org.springframework.http.MediaType;
@@ -41,6 +41,14 @@ import org.springframework.web.reactive.function.client.WebClient;
import org.springframework.web.reactive.function.client.WebClientException;
import org.springframework.web.reactive.function.client.WebClientResponseException;
import org.springframework.web.util.UriComponentsBuilder;
import com.fasterxml.jackson.annotation.JsonCreator;
import com.fasterxml.jackson.annotation.JsonInclude;
import com.fasterxml.jackson.annotation.JsonProperty;
import com.fasterxml.jackson.core.JsonProcessingException;
import com.fasterxml.jackson.databind.ObjectMapper;
import io.micrometer.observation.ObservationRegistry;
import reactor.util.annotation.NonNull;
/**
@@ -48,8 +56,9 @@ import reactor.util.annotation.NonNull;
* deleting, and similarity searching of documents in a GemFire index.
*
* @author Geet Rawat
* @author Christian Tzolov
*/
public class GemFireVectorStore implements VectorStore, InitializingBean {
public class GemFireVectorStore extends AbstractObservationVectorStore implements InitializingBean {
private static final Logger logger = LoggerFactory.getLogger(GemFireVectorStore.class);
@@ -70,10 +79,29 @@ public class GemFireVectorStore implements VectorStore, InitializingBean {
* configuration.
* @param config the configuration for the GemFireVectorStore
* @param embeddingModel the embedding client used for generating embeddings
* @param initializeSchema whether to initialize the schema during initialization
*/
public GemFireVectorStore(GemFireVectorStoreConfig config, EmbeddingModel embeddingModel,
boolean initializeSchema) {
this(config, embeddingModel, initializeSchema, ObservationRegistry.NOOP, null);
}
/**
* Configures and initializes a GemFireVectorStore instance based on the provided
* configuration.
* @param config the configuration for the GemFireVectorStore
* @param embeddingModel the embedding client used for generating embeddings
* @param initializeSchema whether to initialize the schema during initialization
* @param observationRegistry the observation registry to use for recording
* observations
* @param customObservationConvention the custom observation convention to use for
* observing operations
*/
public GemFireVectorStore(GemFireVectorStoreConfig config, EmbeddingModel embeddingModel, boolean initializeSchema,
ObservationRegistry observationRegistry, VectorStoreObservationConvention customObservationConvention) {
super(observationRegistry, customObservationConvention);
Assert.notNull(config, "GemFireVectorStoreConfig must not be null");
Assert.notNull(embeddingModel, "EmbeddingModel must not be null");
this.initializeSchema = initializeSchema;
@@ -374,7 +402,7 @@ public class GemFireVectorStore implements VectorStore, InitializingBean {
}
@Override
public void add(List<Document> documents) {
public void doAdd(List<Document> documents) {
UploadRequest upload = new UploadRequest(documents.stream().map(document -> {
// Compute and assign an embedding to the document.
document.setEmbedding(this.embeddingModel.embed(document));
@@ -404,7 +432,7 @@ public class GemFireVectorStore implements VectorStore, InitializingBean {
}
@Override
public Optional<Boolean> delete(List<String> idList) {
public Optional<Boolean> doDelete(List<String> idList) {
try {
client.method(HttpMethod.DELETE)
.uri("/" + indexName + EMBEDDINGS)
@@ -421,7 +449,7 @@ public class GemFireVectorStore implements VectorStore, InitializingBean {
}
@Override
public List<Document> similaritySearch(SearchRequest request) {
public List<Document> doSimilaritySearch(SearchRequest request) {
if (request.hasFilterExpression()) {
throw new UnsupportedOperationException("GemFire currently does not support metadata filter expressions.");
}
@@ -507,4 +535,12 @@ public class GemFireVectorStore implements VectorStore, InitializingBean {
}
}
@Override
public Builder createObservationContextBuilder(String operationName) {
return VectorStoreObservationContext.builder(VectorStoreProvider.GEMFIRE.value(), operationName)
.withDimensions(this.embeddingModel.dimensions())
.withIndexName(this.indexName)
.withFieldName(EMBEDDINGS);
}
}

View File

@@ -0,0 +1,210 @@
/*
* Copyright 2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (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.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.vectorstore;
import static org.assertj.core.api.Assertions.assertThat;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.util.List;
import java.util.Map;
import org.awaitility.Awaitility;
import org.junit.jupiter.api.AfterAll;
import org.junit.jupiter.api.BeforeAll;
import org.junit.jupiter.api.Test;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.observation.DefaultVectorStoreObservationConvention;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.HighCardinalityKeyNames;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.LowCardinalityKeyNames;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.boot.autoconfigure.EnableAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.core.io.DefaultResourceLoader;
import com.github.dockerjava.api.model.ExposedPort;
import com.github.dockerjava.api.model.PortBinding;
import com.github.dockerjava.api.model.Ports;
import com.vmware.gemfire.testcontainers.GemFireCluster;
import io.micrometer.observation.ObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistryAssert;
import static java.util.concurrent.TimeUnit.MINUTES;
import static org.hamcrest.Matchers.hasSize;
/**
* @author Christian Tzolov
*/
public class GemFireVectorStoreObservationIT {
public static final String INDEX_NAME = "spring-ai-index1";
private static GemFireCluster gemFireCluster;
private static final int HTTP_SERVICE_PORT = 9090;
private static final int LOCATOR_COUNT = 1;
private static final int SERVER_COUNT = 1;
@AfterAll
public static void stopGemFireCluster() {
gemFireCluster.close();
}
@BeforeAll
public static void startGemFireCluster() {
Ports.Binding hostPort = Ports.Binding.bindPort(HTTP_SERVICE_PORT);
ExposedPort exposedPort = new ExposedPort(HTTP_SERVICE_PORT);
PortBinding mappedPort = new PortBinding(hostPort, exposedPort);
gemFireCluster = new GemFireCluster("gemfire/gemfire-all:10.1-jdk17", LOCATOR_COUNT, SERVER_COUNT);
gemFireCluster.withConfiguration(GemFireCluster.SERVER_GLOB,
container -> container.withExposedPorts(HTTP_SERVICE_PORT)
.withCreateContainerCmdModifier(cmd -> cmd.getHostConfig().withPortBindings(mappedPort)));
gemFireCluster.withGemFireProperty(GemFireCluster.SERVER_GLOB, "http-service-port",
Integer.toString(HTTP_SERVICE_PORT));
gemFireCluster.acceptLicense().start();
System.setProperty("spring.data.gemfire.pool.locators",
String.format("localhost[%d]", gemFireCluster.getLocatorPort()));
}
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withUserConfiguration(Config.class);
List<Document> documents = List.of(
new Document(getText("classpath:/test/data/spring.ai.txt"), Map.of("meta1", "meta1")),
new Document(getText("classpath:/test/data/time.shelter.txt")),
new Document(getText("classpath:/test/data/great.depression.txt"), Map.of("meta2", "meta2")));
public static String getText(String uri) {
var resource = new DefaultResourceLoader().getResource(uri);
try {
return resource.getContentAsString(StandardCharsets.UTF_8);
}
catch (IOException e) {
throw new RuntimeException(e);
}
}
@Test
void observationVectorStoreAddAndQueryOperations() {
contextRunner.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store gemfire add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.GEMFIRE.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "384")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "/embeddings")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.INDEX_NAME.asString(), INDEX_NAME)
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "none")
.hasBeenStarted()
.hasBeenStopped();
Awaitility.await()
.atMost(1, MINUTES)
.until(() -> vectorStore
.similaritySearch(SearchRequest.query("Great Depression").withTopK(5).withSimilarityThresholdAll()),
hasSize(3));
observationRegistry.clear();
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("What is Great Depression").withTopK(1));
assertThat(results).isNotEmpty();
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store gemfire query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.GEMFIRE.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "What is Great Depression")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "384")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "/embeddings")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.INDEX_NAME.asString(), INDEX_NAME)
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "1")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "0.0")
.hasBeenStarted()
.hasBeenStopped();
});
}
@SpringBootConfiguration
@EnableAutoConfiguration
public static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public GemFireVectorStoreConfig gemfireVectorStoreConfig() {
return new GemFireVectorStoreConfig().setHost("localhost")
.setPort(HTTP_SERVICE_PORT)
.setIndexName(INDEX_NAME);
}
@Bean
public GemFireVectorStore vectorStore(GemFireVectorStoreConfig config, EmbeddingModel embeddingModel,
ObservationRegistry observationRegistry) {
return new GemFireVectorStore(config, embeddingModel, true, observationRegistry, null);
}
@Bean
public EmbeddingModel embeddingModel() {
return new TransformersEmbeddingModel();
}
}
}

View File

@@ -1,6 +1,7 @@
<?xml version="1.0" encoding="UTF-8"?>
<project xmlns="http://maven.apache.org/POM/4.0.0"
xmlns:xsi="http://www.w3.org/2001/XMLSchema-instance" xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 http://maven.apache.org/maven-v4_0_0.xsd">
xmlns:xsi="http://www.w3.org/2001/XMLSchema-instance"
xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 http://maven.apache.org/maven-v4_0_0.xsd">
<modelVersion>4.0.0</modelVersion>
<parent>
<groupId>org.springframework.ai</groupId>
@@ -71,6 +72,10 @@
<version>${parent.version}</version>
<scope>test</scope>
</dependency>
<dependency>
<groupId>io.micrometer</groupId>
<artifactId>micrometer-observation-test</artifactId>
<scope>test</scope>
</dependency>
</dependencies>
</project>
</project>

View File

@@ -16,11 +16,20 @@
package org.springframework.ai.vectorstore;
import com.fasterxml.jackson.core.JsonProcessingException;
import io.micrometer.observation.ObservationRegistry;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.model.EmbeddingUtils;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.observation.conventions.VectorStoreSimilarityMetric;
import org.springframework.ai.vectorstore.observation.AbstractObservationVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext.Builder;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import java.util.Collections;
import java.util.List;
@@ -54,12 +63,13 @@ import java.util.stream.Collectors;
* generated by <code>org.springframework.ai.embedding.EmbeddingModel</code>
*
* @author Rahul Mittal
* @author Christian Tzolov
* @see <a href=
* "https://help.sap.com/docs/hana-cloud-database/sap-hana-cloud-sap-hana-database-vector-engine-guide/introduction">SAP
* HANA Database Vector Engine Guide</a>
* @since 1.0.0
*/
public class HanaCloudVectorStore implements VectorStore {
public class HanaCloudVectorStore extends AbstractObservationVectorStore {
private static final Logger logger = LoggerFactory.getLogger(HanaCloudVectorStore.class);
@@ -71,13 +81,23 @@ public class HanaCloudVectorStore implements VectorStore {
public HanaCloudVectorStore(HanaVectorRepository<? extends HanaVectorEntity> repository,
EmbeddingModel embeddingModel, HanaCloudVectorStoreConfig config) {
this(repository, embeddingModel, config, ObservationRegistry.NOOP, null);
}
public HanaCloudVectorStore(HanaVectorRepository<? extends HanaVectorEntity> repository,
EmbeddingModel embeddingModel, HanaCloudVectorStoreConfig config, ObservationRegistry observationRegistry,
VectorStoreObservationConvention customObservationConvention) {
super(observationRegistry, customObservationConvention);
this.repository = repository;
this.embeddingModel = embeddingModel;
this.config = config;
}
@Override
public void add(List<Document> documents) {
public void doAdd(List<Document> documents) {
int count = 1;
for (Document document : documents) {
logger.info("[{}/{}] Calling EmbeddingModel for document id = {}", count++, documents.size(),
@@ -90,7 +110,7 @@ public class HanaCloudVectorStore implements VectorStore {
}
@Override
public Optional<Boolean> delete(List<String> idList) {
public Optional<Boolean> doDelete(List<String> idList) {
int deleteCount = repository.deleteEmbeddingsById(config.getTableName(), idList);
logger.info("{} embeddings deleted", deleteCount);
return Optional.of(deleteCount == idList.size());
@@ -108,7 +128,7 @@ public class HanaCloudVectorStore implements VectorStore {
}
@Override
public List<Document> similaritySearch(SearchRequest request) {
public List<Document> doSimilaritySearch(SearchRequest request) {
if (request.hasFilterExpression()) {
throw new UnsupportedOperationException(
"SAPHanaVectorEngine does not support metadata filter expressions yet.");
@@ -144,4 +164,13 @@ public class HanaCloudVectorStore implements VectorStore {
.collect(Collectors.joining(", ")) + "]";
}
@Override
public Builder createObservationContextBuilder(String operationName) {
return VectorStoreObservationContext.builder(VectorStoreProvider.HANA.value(), operationName)
.withDimensions(this.embeddingModel.dimensions())
.withCollectionName(this.config.getTableName())
.withSimilarityMetric(VectorStoreSimilarityMetric.COSINE.value());
}
}

View File

@@ -0,0 +1,195 @@
/*
* Copyright 2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (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.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.vectorstore;
import static org.assertj.core.api.Assertions.assertThat;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.util.List;
import java.util.Map;
import javax.sql.DataSource;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.openai.OpenAiEmbeddingModel;
import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.ai.vectorstore.observation.DefaultVectorStoreObservationConvention;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.HighCardinalityKeyNames;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.LowCardinalityKeyNames;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.boot.autoconfigure.EnableAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.core.io.DefaultResourceLoader;
import org.springframework.jdbc.datasource.DriverManagerDataSource;
import org.springframework.orm.jpa.JpaVendorAdapter;
import org.springframework.orm.jpa.LocalContainerEntityManagerFactoryBean;
import org.springframework.orm.jpa.vendor.HibernateJpaVendorAdapter;
import io.micrometer.observation.ObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistryAssert;
/**
* @author Christian Tzolov
*/
@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+")
@EnabledIfEnvironmentVariable(named = "HANA_DATASOURCE_URL", matches = ".+")
@EnabledIfEnvironmentVariable(named = "HANA_DATASOURCE_USERNAME", matches = ".+")
@EnabledIfEnvironmentVariable(named = "HANA_DATASOURCE_PASSWORD", matches = ".+")
public class HanaVectorStoreObservationIT {
List<Document> documents = List.of(
new Document(getText("classpath:/test/data/spring.ai.txt"), Map.of("meta1", "meta1")),
new Document(getText("classpath:/test/data/time.shelter.txt")),
new Document(getText("classpath:/test/data/great.depression.txt"), Map.of("meta2", "meta2")));
public static String getText(String uri) {
var resource = new DefaultResourceLoader().getResource(uri);
try {
return resource.getContentAsString(StandardCharsets.UTF_8);
}
catch (IOException e) {
throw new RuntimeException(e);
}
}
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withUserConfiguration(Config.class);
@Test
void observationVectorStoreAddAndQueryOperations() {
contextRunner.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store hana add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.HANA.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "1536")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "CRICKET_WORLD_CUP")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "cosine")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "none")
.hasBeenStarted()
.hasBeenStopped();
observationRegistry.clear();
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("What is Great Depression").withTopK(1));
assertThat(results).isNotEmpty();
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store hana query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.HANA.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "What is Great Depression")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "1536")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "CRICKET_WORLD_CUP")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "cosine")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "1")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "0.0")
.hasBeenStarted()
.hasBeenStopped();
});
}
@SpringBootConfiguration
@EnableAutoConfiguration
public static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public VectorStore hanaCloudVectorStore(CricketWorldCupRepository cricketWorldCupRepository,
EmbeddingModel embeddingModel, ObservationRegistry observationRegistry) {
return new HanaCloudVectorStore(cricketWorldCupRepository, embeddingModel,
HanaCloudVectorStoreConfig.builder().tableName("CRICKET_WORLD_CUP").topK(1).build(),
observationRegistry, null);
}
@Bean
public CricketWorldCupRepository cricketWorldCupRepository() {
return new CricketWorldCupRepository();
}
@Bean
public DataSource dataSource() {
DriverManagerDataSource dataSource = new DriverManagerDataSource();
dataSource.setDriverClassName("com.sap.db.jdbc.Driver");
dataSource.setUrl(System.getenv("HANA_DATASOURCE_URL"));
dataSource.setUsername(System.getenv("HANA_DATASOURCE_USERNAME"));
dataSource.setPassword(System.getenv("HANA_DATASOURCE_PASSWORD"));
return dataSource;
}
@Bean
public LocalContainerEntityManagerFactoryBean entityManagerFactory() {
LocalContainerEntityManagerFactoryBean em = new LocalContainerEntityManagerFactoryBean();
em.setDataSource(dataSource());
em.setPackagesToScan("org.springframework.ai.vectorstore");
JpaVendorAdapter vendorAdapter = new HibernateJpaVendorAdapter();
em.setJpaVendorAdapter(vendorAdapter);
return em;
}
@Bean
public EmbeddingModel embeddingModel() {
return new OpenAiEmbeddingModel(new OpenAiApi(System.getenv("OPENAI_API_KEY")));
}
}
}

View File

@@ -74,6 +74,11 @@
<scope>test</scope>
</dependency>
<dependency>
<groupId>io.micrometer</groupId>
<artifactId>micrometer-observation-test</artifactId>
<scope>test</scope>
</dependency>
</dependencies>
</project>

View File

@@ -17,6 +17,7 @@ package org.springframework.ai.vectorstore;
import java.util.ArrayList;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import java.util.stream.Collectors;
@@ -25,13 +26,19 @@ import org.slf4j.LoggerFactory;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.model.EmbeddingUtils;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.observation.conventions.VectorStoreSimilarityMetric;
import org.springframework.ai.vectorstore.filter.FilterExpressionConverter;
import org.springframework.ai.vectorstore.observation.AbstractObservationVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.InitializingBean;
import org.springframework.util.Assert;
import org.springframework.util.StringUtils;
import com.alibaba.fastjson.JSONObject;
import io.micrometer.observation.ObservationRegistry;
import io.milvus.client.MilvusServiceClient;
import io.milvus.common.clientenum.ConsistencyLevelEnum;
import io.milvus.grpc.DataType;
@@ -61,7 +68,7 @@ import io.milvus.response.SearchResultsWrapper;
/**
* @author Christian Tzolov
*/
public class MilvusVectorStore implements VectorStore, InitializingBean {
public class MilvusVectorStore extends AbstractObservationVectorStore implements InitializingBean {
private static final Logger logger = LoggerFactory.getLogger(MilvusVectorStore.class);
@@ -250,6 +257,15 @@ public class MilvusVectorStore implements VectorStore, InitializingBean {
public MilvusVectorStore(MilvusServiceClient milvusClient, EmbeddingModel embeddingModel,
MilvusVectorStoreConfig config, boolean initializeSchema) {
this(milvusClient, embeddingModel, config, initializeSchema, ObservationRegistry.NOOP, null);
}
public MilvusVectorStore(MilvusServiceClient milvusClient, EmbeddingModel embeddingModel,
MilvusVectorStoreConfig config, boolean initializeSchema, ObservationRegistry observationRegistry,
VectorStoreObservationConvention customObservationConvention) {
super(observationRegistry, customObservationConvention);
this.initializeSchema = initializeSchema;
Assert.notNull(milvusClient, "MilvusServiceClient must not be null");
@@ -261,7 +277,7 @@ public class MilvusVectorStore implements VectorStore, InitializingBean {
}
@Override
public void add(List<Document> documents) {
public void doAdd(List<Document> documents) {
Assert.notNull(documents, "Documents must not be null");
@@ -300,7 +316,7 @@ public class MilvusVectorStore implements VectorStore, InitializingBean {
}
@Override
public Optional<Boolean> delete(List<String> idList) {
public Optional<Boolean> doDelete(List<String> idList) {
Assert.notNull(idList, "Document id list must not be null");
String deleteExpression = String.format("%s in [%s]", DOC_ID_FIELD_NAME,
@@ -320,7 +336,7 @@ public class MilvusVectorStore implements VectorStore, InitializingBean {
}
@Override
public List<Document> similaritySearch(SearchRequest request) {
public List<Document> doSimilaritySearch(SearchRequest request) {
String nativeFilterExpressions = (request.getFilterExpression() != null)
? this.filterExpressionConverter.convertExpression(request.getFilterExpression()) : "";
@@ -520,4 +536,27 @@ public class MilvusVectorStore implements VectorStore, InitializingBean {
}
}
@Override
public org.springframework.ai.vectorstore.observation.VectorStoreObservationContext.Builder createObservationContextBuilder(
String operationName) {
return VectorStoreObservationContext.builder(VectorStoreProvider.MILVUS.value(), operationName)
.withDimensions(this.embeddingModel.dimensions())
.withCollectionName(this.config.collectionName)
.withIndexName(this.config.indexType.name())
.withSimilarityMetric(getSimilarityMetric())
.withNamespace(this.config.databaseName);
}
private static Map<MetricType, VectorStoreSimilarityMetric> SIMILARITY_TYPE_MAPPING = Map.of(MetricType.COSINE,
VectorStoreSimilarityMetric.COSINE, MetricType.L2, VectorStoreSimilarityMetric.EUCLIDEAN, MetricType.IP,
VectorStoreSimilarityMetric.DOT);
private String getSimilarityMetric() {
if (!SIMILARITY_TYPE_MAPPING.containsKey(this.config.metricType)) {
return this.config.metricType.name();
}
return SIMILARITY_TYPE_MAPPING.get(this.config.metricType).value();
}
}

View File

@@ -0,0 +1,179 @@
/*
* Copyright 2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (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.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.vectorstore;
import static org.assertj.core.api.Assertions.assertThat;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.util.List;
import java.util.Map;
import org.junit.jupiter.api.Test;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.openai.OpenAiEmbeddingModel;
import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.ai.vectorstore.MilvusVectorStore.MilvusVectorStoreConfig;
import org.springframework.ai.vectorstore.observation.DefaultVectorStoreObservationConvention;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.HighCardinalityKeyNames;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.LowCardinalityKeyNames;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.boot.autoconfigure.EnableAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.core.io.DefaultResourceLoader;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import org.testcontainers.milvus.MilvusContainer;
import io.micrometer.observation.ObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistryAssert;
import io.milvus.client.MilvusServiceClient;
import io.milvus.param.ConnectParam;
import io.milvus.param.IndexType;
import io.milvus.param.MetricType;
/**
* @author Christian Tzolov
*/
@Testcontainers
public class MilvusVectorStoreObservationIT {
@Container
private static MilvusContainer milvusContainer = new MilvusContainer("milvusdb/milvus:v2.3.8");
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withUserConfiguration(Config.class);
List<Document> documents = List.of(
new Document(getText("classpath:/test/data/spring.ai.txt"), Map.of("meta1", "meta1")),
new Document(getText("classpath:/test/data/time.shelter.txt")),
new Document(getText("classpath:/test/data/great.depression.txt"), Map.of("meta2", "meta2")));
public static String getText(String uri) {
var resource = new DefaultResourceLoader().getResource(uri);
try {
return resource.getContentAsString(StandardCharsets.UTF_8);
}
catch (IOException e) {
throw new RuntimeException(e);
}
}
@Test
void observationVectorStoreAddAndQueryOperations() {
contextRunner.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store milvus add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.MILVUS.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "1536")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "test_vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "default")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "cosine")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "none")
.hasBeenStarted()
.hasBeenStopped();
observationRegistry.clear();
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("What is Great Depression").withTopK(1));
assertThat(results).isNotEmpty();
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store milvus query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.MILVUS.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "What is Great Depression")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "1536")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "test_vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "default")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "cosine")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "1")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "0.0")
.hasBeenStarted()
.hasBeenStopped();
});
}
@SpringBootConfiguration
@EnableAutoConfiguration
public static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public VectorStore vectorStore(MilvusServiceClient milvusClient, EmbeddingModel embeddingModel,
ObservationRegistry observationRegistry) {
MilvusVectorStoreConfig config = MilvusVectorStoreConfig.builder()
.withCollectionName("test_vector_store")
.withDatabaseName("default")
.withIndexType(IndexType.IVF_FLAT)
.withMetricType(MetricType.COSINE)
.build();
return new MilvusVectorStore(milvusClient, embeddingModel, config, true, observationRegistry, null);
}
@Bean
public MilvusServiceClient milvusClient() {
return new MilvusServiceClient(ConnectParam.newBuilder()
.withAuthorization("minioadmin", "minioadmin")
.withUri(milvusContainer.getEndpoint())
.build());
}
@Bean
public EmbeddingModel embeddingModel() {
return new OpenAiEmbeddingModel(new OpenAiApi(System.getenv("OPENAI_API_KEY")));
}
}
}

View File

@@ -54,5 +54,11 @@
<artifactId>junit-jupiter</artifactId>
<scope>test</scope>
</dependency>
<dependency>
<groupId>io.micrometer</groupId>
<artifactId>micrometer-observation-test</artifactId>
<scope>test</scope>
</dependency>
</dependencies>
</project>

View File

@@ -15,16 +15,21 @@
*/
package org.springframework.ai.vectorstore;
import static org.springframework.data.mongodb.core.query.Criteria.where;
import java.util.ArrayList;
import java.util.Collections;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import com.mongodb.MongoCommandException;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.model.EmbeddingUtils;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.vectorstore.observation.AbstractObservationVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.InitializingBean;
import org.springframework.data.mongodb.UncategorizedMongoDbException;
import org.springframework.data.mongodb.core.MongoTemplate;
@@ -33,14 +38,17 @@ import org.springframework.data.mongodb.core.query.Criteria;
import org.springframework.data.mongodb.core.query.Query;
import org.springframework.util.Assert;
import static org.springframework.data.mongodb.core.query.Criteria.where;
import com.mongodb.MongoCommandException;
import io.micrometer.observation.ObservationRegistry;
/**
* @author Chris Smith
* @author Soby Chacko
* @author Christian Tzolov
* @since 1.0.0
*/
public class MongoDBAtlasVectorStore implements VectorStore, InitializingBean {
public class MongoDBAtlasVectorStore extends AbstractObservationVectorStore implements InitializingBean {
public static final String ID_FIELD_NAME = "_id";
@@ -79,6 +87,15 @@ public class MongoDBAtlasVectorStore implements VectorStore, InitializingBean {
public MongoDBAtlasVectorStore(MongoTemplate mongoTemplate, EmbeddingModel embeddingModel,
MongoDBVectorStoreConfig config, boolean initializeSchema) {
this(mongoTemplate, embeddingModel, config, initializeSchema, ObservationRegistry.NOOP, null);
}
public MongoDBAtlasVectorStore(MongoTemplate mongoTemplate, EmbeddingModel embeddingModel,
MongoDBVectorStoreConfig config, boolean initializeSchema, ObservationRegistry observationRegistry,
VectorStoreObservationConvention customObservationConvention) {
super(observationRegistry, customObservationConvention);
this.mongoTemplate = mongoTemplate;
this.embeddingModel = embeddingModel;
this.config = config;
@@ -156,7 +173,7 @@ public class MongoDBAtlasVectorStore implements VectorStore, InitializingBean {
}
@Override
public void add(List<Document> documents) {
public void doAdd(List<Document> documents) {
for (Document document : documents) {
float[] embedding = this.embeddingModel.embed(document);
document.setEmbedding(embedding);
@@ -165,7 +182,7 @@ public class MongoDBAtlasVectorStore implements VectorStore, InitializingBean {
}
@Override
public Optional<Boolean> delete(List<String> idList) {
public Optional<Boolean> doDelete(List<String> idList) {
Query query = new Query(where(ID_FIELD_NAME).in(idList));
var deleteRes = this.mongoTemplate.remove(query, this.config.collectionName);
@@ -180,7 +197,7 @@ public class MongoDBAtlasVectorStore implements VectorStore, InitializingBean {
}
@Override
public List<Document> similaritySearch(SearchRequest request) {
public List<Document> doSimilaritySearch(SearchRequest request) {
String nativeFilterExpressions = (request.getFilterExpression() != null)
? this.filterExpressionConverter.convertExpression(request.getFilterExpression()) : "";
@@ -299,4 +316,14 @@ public class MongoDBAtlasVectorStore implements VectorStore, InitializingBean {
}
@Override
public VectorStoreObservationContext.Builder createObservationContextBuilder(String operationName) {
return VectorStoreObservationContext.builder(VectorStoreProvider.MONGODB.value(), operationName)
.withDimensions(this.embeddingModel.dimensions())
.withCollectionName(this.config.collectionName)
.withFieldName(this.config.pathName)
.withIndexName(this.config.vectorIndexName);
}
}

View File

@@ -44,7 +44,7 @@ import static org.assertj.core.api.Assertions.assertThat;
* @author Chris Smith
*/
@Testcontainers
@Disabled("Disabled due to https://github.com/spring-projects/spring-ai/issues/698")
// @Disabled("Disabled due to https://github.com/spring-projects/spring-ai/issues/698")
class MongoDBAtlasVectorStoreIT {
@Container

View File

@@ -0,0 +1,187 @@
/*
* Copyright 2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (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.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.vectorstore;
import static org.assertj.core.api.Assertions.assertThat;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.util.List;
import java.util.Map;
import java.util.Vector;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.openai.OpenAiEmbeddingModel;
import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.ai.vectorstore.observation.DefaultVectorStoreObservationConvention;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.HighCardinalityKeyNames;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.LowCardinalityKeyNames;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.boot.autoconfigure.EnableAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.core.io.DefaultResourceLoader;
import org.springframework.data.mongodb.core.MongoTemplate;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import com.mongodb.client.MongoClient;
import io.micrometer.observation.ObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistryAssert;
/**
* @author Christian Tzolov
*/
@Testcontainers
@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+")
public class MongoDbVectorStoreObservationIT {
@Container
private static MongoDBAtlasContainer container = new MongoDBAtlasContainer();
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withUserConfiguration(Config.class)
.withPropertyValues("spring.data.mongodb.database=springaisample",
String.format("spring.data.mongodb.uri=" + container.getConnectionString()));
List<Document> documents = List.of(
new Document(getText("classpath:/test/data/spring.ai.txt"), Map.of("meta1", "meta1")),
new Document(getText("classpath:/test/data/time.shelter.txt")),
new Document(getText("classpath:/test/data/great.depression.txt"), Map.of("meta2", "meta2")));
public static String getText(String uri) {
var resource = new DefaultResourceLoader().getResource(uri);
try {
return resource.getContentAsString(StandardCharsets.UTF_8);
}
catch (IOException e) {
throw new RuntimeException(e);
}
}
@BeforeEach
public void beforeEach() {
contextRunner.run(context -> {
MongoTemplate mongoTemplate = context.getBean(MongoTemplate.class);
mongoTemplate.getCollection("vector_store").deleteMany(new org.bson.Document());
});
}
@Test
void observationVectorStoreAddAndQueryOperations() {
contextRunner.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store mongodb add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.MONGODB.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "1536")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "embedding")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.INDEX_NAME.asString(), "vector_index")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "none")
.hasBeenStarted()
.hasBeenStopped();
observationRegistry.clear();
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("What is Great Depression").withTopK(1));
assertThat(results).isNotEmpty();
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store mongodb query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.MONGODB.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "What is Great Depression")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "1536")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "embedding")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.INDEX_NAME.asString(), "vector_index")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "1")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "0.0")
.hasBeenStarted()
.hasBeenStopped();
});
}
@SpringBootConfiguration
@EnableAutoConfiguration
public static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public VectorStore vectorStore(MongoTemplate mongoTemplate, EmbeddingModel embeddingModel,
ObservationRegistry observationRegistry) {
return new MongoDBAtlasVectorStore(mongoTemplate, embeddingModel,
MongoDBAtlasVectorStore.MongoDBVectorStoreConfig.builder()
.withMetadataFieldsToFilter(List.of("country", "year"))
.build(),
true, observationRegistry, null);
}
@Bean
public MongoTemplate mongoTemplate(MongoClient mongoClient) {
return new MongoTemplate(mongoClient, "springaisample");
}
@Bean
public EmbeddingModel embeddingModel() {
return new OpenAiEmbeddingModel(new OpenAiApi(System.getenv("OPENAI_API_KEY")));
}
}
}

View File

@@ -75,6 +75,18 @@
<scope>test</scope>
</dependency>
<dependency>
<groupId>org.springframework.ai</groupId>
<artifactId>spring-ai-test</artifactId>
<version>${project.parent.version}</version>
<scope>test</scope>
</dependency>
<dependency>
<groupId>io.micrometer</groupId>
<artifactId>micrometer-observation-test</artifactId>
<scope>test</scope>
</dependency>
</dependencies>

View File

@@ -27,15 +27,23 @@ import org.neo4j.driver.SessionConfig;
import org.neo4j.driver.Values;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.observation.conventions.VectorStoreSimilarityMetric;
import org.springframework.ai.vectorstore.filter.Neo4jVectorFilterExpressionConverter;
import org.springframework.ai.vectorstore.observation.AbstractObservationVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.InitializingBean;
import org.springframework.util.Assert;
import io.micrometer.observation.ObservationRegistry;
/**
* @author Gerrit Meier
* @author Michael Simons
* @author Christian Tzolov
*/
public class Neo4jVectorStore implements VectorStore, InitializingBean {
public class Neo4jVectorStore extends AbstractObservationVectorStore implements InitializingBean {
/**
* An enum to configure the distance function used in the Neo4j vector index.
@@ -277,6 +285,15 @@ public class Neo4jVectorStore implements VectorStore, InitializingBean {
public Neo4jVectorStore(Driver driver, EmbeddingModel embeddingModel, Neo4jVectorStoreConfig config,
boolean initializeSchema) {
this(driver, embeddingModel, config, initializeSchema, ObservationRegistry.NOOP, null);
}
public Neo4jVectorStore(Driver driver, EmbeddingModel embeddingModel, Neo4jVectorStoreConfig config,
boolean initializeSchema, ObservationRegistry observationRegistry,
VectorStoreObservationConvention customObservationConvention) {
super(observationRegistry, customObservationConvention);
this.initializeSchema = initializeSchema;
Assert.notNull(driver, "Neo4j driver must not be null");
@@ -289,7 +306,7 @@ public class Neo4jVectorStore implements VectorStore, InitializingBean {
}
@Override
public void add(List<Document> documents) {
public void doAdd(List<Document> documents) {
var rows = documents.stream().map(this::documentToRecord).toList();
@@ -311,7 +328,7 @@ public class Neo4jVectorStore implements VectorStore, InitializingBean {
}
@Override
public Optional<Boolean> delete(List<String> idList) {
public Optional<Boolean> doDelete(List<String> idList) {
try (var session = this.driver.session(this.config.sessionConfig)) {
@@ -327,7 +344,7 @@ public class Neo4jVectorStore implements VectorStore, InitializingBean {
}
@Override
public List<Document> similaritySearch(SearchRequest request) {
public List<Document> doSimilaritySearch(SearchRequest request) {
Assert.isTrue(request.getTopK() > 0, "The number of documents to returned must be greater than zero");
Assert.isTrue(request.getSimilarityThreshold() >= 0 && request.getSimilarityThreshold() <= 1,
"The similarity score is bounded between 0 and 1; least to most similar respectively.");
@@ -412,4 +429,24 @@ public class Neo4jVectorStore implements VectorStore, InitializingBean {
Map.copyOf(metaData));
}
@Override
public VectorStoreObservationContext.Builder createObservationContextBuilder(String operationName) {
return VectorStoreObservationContext.builder(VectorStoreProvider.NEO4J.value(), operationName)
.withDimensions(this.embeddingModel.dimensions())
.withIndexName(this.config.indexName)
.withSimilarityMetric(getSimilarityMetric());
}
private static Map<Neo4jDistanceType, VectorStoreSimilarityMetric> SIMILARITY_TYPE_MAPPING = Map.of(
Neo4jDistanceType.COSINE, VectorStoreSimilarityMetric.COSINE, Neo4jDistanceType.EUCLIDEAN,
VectorStoreSimilarityMetric.EUCLIDEAN);
private String getSimilarityMetric() {
if (!SIMILARITY_TYPE_MAPPING.containsKey(this.config.distanceType)) {
return this.config.distanceType.name();
}
return SIMILARITY_TYPE_MAPPING.get(this.config.distanceType).value();
}
}

View File

@@ -0,0 +1,183 @@
/*
* Copyright 2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (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.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.vectorstore;
import static org.assertj.core.api.Assertions.assertThat;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.util.List;
import java.util.Map;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.neo4j.driver.AuthTokens;
import org.neo4j.driver.Driver;
import org.neo4j.driver.GraphDatabase;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.openai.OpenAiEmbeddingModel;
import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.ai.vectorstore.observation.DefaultVectorStoreObservationConvention;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.HighCardinalityKeyNames;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.LowCardinalityKeyNames;
import org.springframework.boot.SpringBootConfiguration;
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.springframework.core.io.DefaultResourceLoader;
import org.testcontainers.containers.Neo4jContainer;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import org.testcontainers.utility.DockerImageName;
import io.micrometer.observation.ObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistryAssert;
/**
* @author Christian Tzolov
*/
@Testcontainers
@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+")
public class Neo4jVectorStoreObservationIT {
@Container
static Neo4jContainer<?> neo4jContainer = new Neo4jContainer<>(DockerImageName.parse("neo4j:5.18"))
.withRandomPassword();
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withUserConfiguration(Config.class);
List<Document> documents = List.of(
new Document(getText("classpath:/test/data/spring.ai.txt"), Map.of("meta1", "meta1")),
new Document(getText("classpath:/test/data/time.shelter.txt")),
new Document(getText("classpath:/test/data/great.depression.txt"), Map.of("meta2", "meta2")));
public static String getText(String uri) {
var resource = new DefaultResourceLoader().getResource(uri);
try {
return resource.getContentAsString(StandardCharsets.UTF_8);
}
catch (IOException e) {
throw new RuntimeException(e);
}
}
@BeforeEach
void cleanDatabase() {
this.contextRunner
.run(context -> context.getBean(Driver.class).executableQuery("MATCH (n) DETACH DELETE n").execute());
}
@Test
void observationVectorStoreAddAndQueryOperations() {
contextRunner.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store neo4j add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.NEO4J.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "1536")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "cosine")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "none")
.hasBeenStarted()
.hasBeenStopped();
observationRegistry.clear();
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("What is Great Depression").withTopK(1));
assertThat(results).isNotEmpty();
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store neo4j query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.NEO4J.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "What is Great Depression")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "1536")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "cosine")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "1")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "0.0")
.hasBeenStarted()
.hasBeenStopped();
});
}
@SpringBootConfiguration
@EnableAutoConfiguration(exclude = { DataSourceAutoConfiguration.class })
public static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public VectorStore vectorStore(Driver driver, EmbeddingModel embeddingModel,
ObservationRegistry observationRegistry) {
return new Neo4jVectorStore(driver, embeddingModel, Neo4jVectorStore.Neo4jVectorStoreConfig.defaultConfig(),
true, observationRegistry, null);
}
@Bean
public Driver driver() {
return GraphDatabase.driver(neo4jContainer.getBoltUrl(),
AuthTokens.basic("neo4j", neo4jContainer.getAdminPassword()));
}
@Bean
public EmbeddingModel embeddingModel() {
return new OpenAiEmbeddingModel(new OpenAiApi(System.getenv("OPENAI_API_KEY")));
}
}
}

View File

@@ -80,6 +80,12 @@
<scope>test</scope>
</dependency>
<dependency>
<groupId>io.micrometer</groupId>
<artifactId>micrometer-observation-test</artifactId>
<scope>test</scope>
</dependency>
</dependencies>
</project>

View File

@@ -31,11 +31,19 @@ import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.observation.conventions.VectorStoreSimilarityMetric;
import org.springframework.ai.vectorstore.filter.Filter;
import org.springframework.ai.vectorstore.filter.FilterExpressionConverter;
import org.springframework.ai.vectorstore.observation.AbstractObservationVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext.Builder;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.InitializingBean;
import org.springframework.util.Assert;
import io.micrometer.observation.ObservationRegistry;
import java.io.IOException;
import java.io.StringReader;
import java.util.List;
@@ -46,9 +54,10 @@ import java.util.stream.Collectors;
/**
* @author Jemin Huh
* @author Soby Chacko
* @author Christian Tzolov
* @since 1.0.0
*/
public class OpenSearchVectorStore implements VectorStore, InitializingBean {
public class OpenSearchVectorStore extends AbstractObservationVectorStore implements InitializingBean {
public static final String COSINE_SIMILARITY_FUNCTION = "cosinesimil";
@@ -94,6 +103,15 @@ public class OpenSearchVectorStore implements VectorStore, InitializingBean {
public OpenSearchVectorStore(String index, OpenSearchClient openSearchClient, EmbeddingModel embeddingModel,
String mappingJson, boolean initializeSchema) {
this(index, openSearchClient, embeddingModel, mappingJson, initializeSchema, ObservationRegistry.NOOP, null);
}
public OpenSearchVectorStore(String index, OpenSearchClient openSearchClient, EmbeddingModel embeddingModel,
String mappingJson, boolean initializeSchema, ObservationRegistry observationRegistry,
VectorStoreObservationConvention customObservationConvention) {
super(observationRegistry, customObservationConvention);
Objects.requireNonNull(embeddingModel, "RestClient must not be null");
Objects.requireNonNull(embeddingModel, "EmbeddingModel must not be null");
this.openSearchClient = openSearchClient;
@@ -113,7 +131,7 @@ public class OpenSearchVectorStore implements VectorStore, InitializingBean {
}
@Override
public void add(List<Document> documents) {
public void doAdd(List<Document> documents) {
BulkRequest.Builder bulkRequestBuilder = new BulkRequest.Builder();
for (Document document : documents) {
if (Objects.isNull(document.getEmbedding()) || document.getEmbedding().length == 0) {
@@ -127,7 +145,7 @@ public class OpenSearchVectorStore implements VectorStore, InitializingBean {
}
@Override
public Optional<Boolean> delete(List<String> idList) {
public Optional<Boolean> doDelete(List<String> idList) {
BulkRequest.Builder bulkRequestBuilder = new BulkRequest.Builder();
for (String id : idList)
bulkRequestBuilder.operations(op -> op.delete(idx -> idx.index(this.index).id(id)));
@@ -144,7 +162,7 @@ public class OpenSearchVectorStore implements VectorStore, InitializingBean {
}
@Override
public List<Document> similaritySearch(SearchRequest searchRequest) {
public List<Document> doSimilaritySearch(SearchRequest searchRequest) {
Assert.notNull(searchRequest, "The search request must not be null.");
return similaritySearch(this.embeddingModel.embed(searchRequest.getQuery()), searchRequest.getTopK(),
searchRequest.getSimilarityThreshold(), searchRequest.getFilterExpression());
@@ -240,4 +258,23 @@ public class OpenSearchVectorStore implements VectorStore, InitializingBean {
}
}
@Override
public Builder createObservationContextBuilder(String operationName) {
return VectorStoreObservationContext.builder(VectorStoreProvider.OPENSEARCH.value(), operationName)
.withDimensions(this.embeddingModel.dimensions())
.withSimilarityMetric(getSimilarityFunction())
.withIndexName(this.index);
}
private String getSimilarityFunction() {
if ("cosinesimil".equalsIgnoreCase(this.similarityFunction)) {
return VectorStoreSimilarityMetric.COSINE.value();
}
else if ("l2".equalsIgnoreCase(this.similarityFunction)) {
return VectorStoreSimilarityMetric.EUCLIDEAN.value();
}
return this.similarityFunction;
}
}

View File

@@ -0,0 +1,215 @@
/*
* Copyright 2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (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.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.vectorstore;
import static org.assertj.core.api.Assertions.assertThat;
import java.io.IOException;
import java.net.URISyntaxException;
import java.nio.charset.StandardCharsets;
import java.time.Duration;
import java.util.List;
import java.util.Map;
import java.util.concurrent.TimeUnit;
import org.apache.hc.core5.http.HttpHost;
import org.awaitility.Awaitility;
import org.junit.jupiter.api.BeforeAll;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.opensearch.client.opensearch.OpenSearchClient;
import org.opensearch.client.transport.httpclient5.ApacheHttpClient5TransportBuilder;
import org.opensearch.testcontainers.OpensearchContainer;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.openai.OpenAiEmbeddingModel;
import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.ai.vectorstore.observation.DefaultVectorStoreObservationConvention;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.HighCardinalityKeyNames;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.LowCardinalityKeyNames;
import org.springframework.boot.SpringBootConfiguration;
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.springframework.core.io.DefaultResourceLoader;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import org.testcontainers.utility.DockerImageName;
import io.micrometer.observation.ObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistryAssert;
import static org.hamcrest.Matchers.hasSize;
/**
* @author Christian Tzolov
*/
@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+")
@Testcontainers
public class OpenSearchVectorStoreObservationIT {
@Container
private static final OpensearchContainer<?> opensearchContainer = new OpensearchContainer<>(
DockerImageName.parse("opensearchproject/opensearch:2.13.0"));
List<Document> documents = List.of(
new Document(getText("classpath:/test/data/spring.ai.txt"), Map.of("meta1", "meta1")),
new Document(getText("classpath:/test/data/time.shelter.txt")),
new Document(getText("classpath:/test/data/great.depression.txt"), Map.of("meta2", "meta2")));
public static String getText(String uri) {
var resource = new DefaultResourceLoader().getResource(uri);
try {
return resource.getContentAsString(StandardCharsets.UTF_8);
}
catch (IOException e) {
throw new RuntimeException(e);
}
}
private ApplicationContextRunner getContextRunner() {
return new ApplicationContextRunner().withUserConfiguration(Config.class);
}
@BeforeAll
public static void beforeAll() {
Awaitility.setDefaultPollInterval(2, TimeUnit.SECONDS);
Awaitility.setDefaultPollDelay(Duration.ZERO);
Awaitility.setDefaultTimeout(Duration.ofMinutes(1));
}
@BeforeEach
void cleanDatabase() {
getContextRunner().run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
vectorStore.delete(List.of("_all"));
});
}
@Test
void observationVectorStoreAddAndQueryOperations() {
getContextRunner().run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store opensearch add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.OPENSEARCH.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "1536")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "cosine")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "none")
.hasBeenStarted()
.hasBeenStopped();
Awaitility.await()
.until(() -> vectorStore
.similaritySearch(SearchRequest.query("Great Depression").withTopK(1).withSimilarityThreshold(0)),
hasSize(1));
observationRegistry.clear();
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("What is Great Depression").withTopK(1));
assertThat(results).isNotEmpty();
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store opensearch query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.OPENSEARCH.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "What is Great Depression")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "1536")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "cosine")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "1")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "0.0")
.hasBeenStarted()
.hasBeenStopped();
observationRegistry.clear();
vectorStore.delete(documents.stream().map(Document::getId).toList());
Awaitility.await()
.until(() -> vectorStore
.similaritySearch(SearchRequest.query("Great Depression").withTopK(1).withSimilarityThreshold(0)),
hasSize(0));
});
}
@SpringBootConfiguration
@EnableAutoConfiguration(exclude = { DataSourceAutoConfiguration.class })
static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public OpenSearchVectorStore vectorStore(EmbeddingModel embeddingModel,
ObservationRegistry observationRegistry) {
try {
return new OpenSearchVectorStore(OpenSearchVectorStore.DEFAULT_INDEX_NAME,
new OpenSearchClient(ApacheHttpClient5TransportBuilder
.builder(HttpHost.create(opensearchContainer.getHttpHostAddress()))
.build()),
embeddingModel, OpenSearchVectorStore.DEFAULT_MAPPING_EMBEDDING_TYPE_KNN_VECTOR_DIMENSION_1536,
true, observationRegistry, null);
}
catch (URISyntaxException e) {
throw new RuntimeException(e);
}
}
@Bean
public EmbeddingModel embeddingModel() {
return new OpenAiEmbeddingModel(new OpenAiApi(System.getenv("OPENAI_API_KEY")));
}
}
}

View File

@@ -95,7 +95,11 @@
<artifactId>junit-jupiter</artifactId>
<scope>test</scope>
</dependency>
<dependency>
<groupId>io.micrometer</groupId>
<artifactId>micrometer-observation-test</artifactId>
<scope>test</scope>
</dependency>
</dependencies>
</project>

View File

@@ -15,22 +15,8 @@
*/
package org.springframework.ai.vectorstore;
import oracle.jdbc.OracleType;
import oracle.sql.VECTOR;
import oracle.sql.json.OracleJsonFactory;
import oracle.sql.json.OracleJsonGenerator;
import oracle.sql.json.OracleJsonObject;
import oracle.sql.json.OracleJsonValue;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.vectorstore.filter.FilterExpressionConverter;
import org.springframework.beans.factory.InitializingBean;
import org.springframework.jdbc.core.BatchPreparedStatementSetter;
import org.springframework.jdbc.core.JdbcTemplate;
import org.springframework.jdbc.core.RowMapper;
import org.springframework.util.StringUtils;
import static org.springframework.ai.vectorstore.OracleVectorStore.OracleVectorStoreDistanceType.DOT;
import static org.springframework.jdbc.core.StatementCreatorUtils.setParameterValue;
import java.io.ByteArrayOutputStream;
import java.sql.PreparedStatement;
@@ -44,8 +30,30 @@ import java.util.List;
import java.util.Map;
import java.util.Optional;
import static org.springframework.ai.vectorstore.OracleVectorStore.OracleVectorStoreDistanceType.DOT;
import static org.springframework.jdbc.core.StatementCreatorUtils.setParameterValue;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.observation.conventions.VectorStoreSimilarityMetric;
import org.springframework.ai.vectorstore.filter.FilterExpressionConverter;
import org.springframework.ai.vectorstore.observation.AbstractObservationVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext.Builder;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.InitializingBean;
import org.springframework.jdbc.core.BatchPreparedStatementSetter;
import org.springframework.jdbc.core.JdbcTemplate;
import org.springframework.jdbc.core.RowMapper;
import org.springframework.util.StringUtils;
import io.micrometer.observation.ObservationRegistry;
import oracle.jdbc.OracleType;
import oracle.sql.VECTOR;
import oracle.sql.json.OracleJsonFactory;
import oracle.sql.json.OracleJsonGenerator;
import oracle.sql.json.OracleJsonObject;
import oracle.sql.json.OracleJsonValue;
/**
* <p>
@@ -69,8 +77,9 @@ import static org.springframework.jdbc.core.StatementCreatorUtils.setParameterVa
* </ul>
*
* @author Loïc Lefèvre
* @author Christian Tzolov
*/
public class OracleVectorStore implements VectorStore, InitializingBean {
public class OracleVectorStore extends AbstractObservationVectorStore implements InitializingBean {
private static final Logger logger = LoggerFactory.getLogger(OracleVectorStore.class);
@@ -126,7 +135,7 @@ public class OracleVectorStore implements VectorStore, InitializingBean {
public enum OracleVectorStoreDistanceType {
/**
* Default metric. It calculates the cosine distane between two vectors.
* Default metric. It calculates the cosine distance between two vectors.
*/
COSINE,
@@ -220,6 +229,18 @@ public class OracleVectorStore implements VectorStore, InitializingBean {
OracleVectorStoreIndexType indexType, OracleVectorStoreDistanceType distanceType, int dimensions,
int searchAccuracy, boolean initializeSchema, boolean removeExistingVectorStoreTable,
boolean forcedNormalization) {
this(jdbcTemplate, embeddingModel, tableName, indexType, distanceType, dimensions, searchAccuracy,
initializeSchema, removeExistingVectorStoreTable, forcedNormalization, ObservationRegistry.NOOP, null);
}
public OracleVectorStore(JdbcTemplate jdbcTemplate, EmbeddingModel embeddingModel, String tableName,
OracleVectorStoreIndexType indexType, OracleVectorStoreDistanceType distanceType, int dimensions,
int searchAccuracy, boolean initializeSchema, boolean removeExistingVectorStoreTable,
boolean forcedNormalization, ObservationRegistry observationRegistry,
VectorStoreObservationConvention customObservationConvention) {
super(observationRegistry, customObservationConvention);
if (dimensions != DEFAULT_DIMENSIONS) {
if (dimensions <= 0) {
throw new RuntimeException("Number of dimensions must be strictly positive");
@@ -251,7 +272,7 @@ public class OracleVectorStore implements VectorStore, InitializingBean {
}
@Override
public void add(final List<Document> documents) {
public void doAdd(final List<Document> documents) {
this.jdbcTemplate.batchUpdate(getIngestStatement(), new BatchPreparedStatementSetter() {
@Override
public void setValues(PreparedStatement ps, int i) throws SQLException {
@@ -366,7 +387,7 @@ public class OracleVectorStore implements VectorStore, InitializingBean {
}
@Override
public Optional<Boolean> delete(final List<String> idList) {
public Optional<Boolean> doDelete(final List<String> idList) {
final String sql = String.format("delete from %s where id=?", tableName);
final int[] argTypes = { Types.VARCHAR };
@@ -429,7 +450,7 @@ public class OracleVectorStore implements VectorStore, InitializingBean {
}
@Override
public List<Document> similaritySearch(SearchRequest request) {
public List<Document> doSimilaritySearch(SearchRequest request) {
try {
// From the provided query, generate a vector using the embedding model
final VECTOR embeddingVector = toVECTOR(embeddingModel.embed(request.getQuery()));
@@ -599,4 +620,24 @@ public class OracleVectorStore implements VectorStore, InitializingBean {
return tableName;
}
@Override
public Builder createObservationContextBuilder(String operationName) {
return VectorStoreObservationContext.builder(VectorStoreProvider.ORACLE.value(), operationName)
.withDimensions(this.embeddingModel.dimensions())
.withCollectionName(this.getTableName())
.withFieldName(getSimilarityMetric());
}
private static Map<OracleVectorStoreDistanceType, VectorStoreSimilarityMetric> SIMILARITY_TYPE_MAPPING = Map.of(
OracleVectorStoreDistanceType.COSINE, VectorStoreSimilarityMetric.COSINE,
OracleVectorStoreDistanceType.EUCLIDEAN, VectorStoreSimilarityMetric.EUCLIDEAN,
OracleVectorStoreDistanceType.DOT, VectorStoreSimilarityMetric.DOT);
private String getSimilarityMetric() {
if (!SIMILARITY_TYPE_MAPPING.containsKey(this.distanceType)) {
return this.distanceType.name();
}
return SIMILARITY_TYPE_MAPPING.get(this.distanceType).value();
}
}

View File

@@ -0,0 +1,206 @@
/*
* Copyright 2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (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.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.vectorstore;
import static org.assertj.core.api.Assertions.assertThat;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.util.List;
import java.util.Map;
import javax.sql.DataSource;
import org.junit.jupiter.api.Test;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.OracleVectorStore.OracleVectorStoreDistanceType;
import org.springframework.ai.vectorstore.observation.DefaultVectorStoreObservationConvention;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.HighCardinalityKeyNames;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.LowCardinalityKeyNames;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.boot.autoconfigure.EnableAutoConfiguration;
import org.springframework.boot.autoconfigure.jdbc.DataSourceAutoConfiguration;
import org.springframework.boot.autoconfigure.jdbc.DataSourceProperties;
import org.springframework.boot.context.properties.ConfigurationProperties;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.ApplicationContext;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Primary;
import org.springframework.core.io.DefaultResourceLoader;
import org.springframework.jdbc.core.JdbcTemplate;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import org.testcontainers.oracle.OracleContainer;
import org.testcontainers.utility.MountableFile;
import io.micrometer.observation.ObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistryAssert;
import oracle.jdbc.pool.OracleDataSource;
/**
* @author Christian Tzolov
*/
@Testcontainers
public class OracleVectorStoreObservationIT {
@Container
static OracleContainer oracle23aiContainer = new OracleContainer("gvenzl/oracle-free:23-slim")
.withCopyFileToContainer(MountableFile.forClasspathResource("/initialize.sql"),
"/container-entrypoint-initdb.d/initialize.sql");
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withUserConfiguration(Config.class)
.withPropertyValues("test.spring.ai.vectorstore.oracle.dimensions=384",
// JdbcTemplate configuration
String.format("app.datasource.url=%s", oracle23aiContainer.getJdbcUrl()),
String.format("app.datasource.username=%s", oracle23aiContainer.getUsername()),
String.format("app.datasource.password=%s", oracle23aiContainer.getPassword()),
"app.datasource.type=oracle.jdbc.pool.OracleDataSource");
List<Document> documents = List.of(
new Document(getText("classpath:/test/data/spring.ai.txt"), Map.of("meta1", "meta1")),
new Document(getText("classpath:/test/data/time.shelter.txt")),
new Document(getText("classpath:/test/data/great.depression.txt"), Map.of("meta2", "meta2")));
public static String getText(String uri) {
var resource = new DefaultResourceLoader().getResource(uri);
try {
return resource.getContentAsString(StandardCharsets.UTF_8);
}
catch (IOException e) {
throw new RuntimeException(e);
}
}
private static void dropTable(ApplicationContext context, String tableName) {
JdbcTemplate jdbcTemplate = context.getBean(JdbcTemplate.class);
jdbcTemplate.execute("DROP TABLE IF EXISTS " + tableName + " PURGE");
}
@Test
void observationVectorStoreAddAndQueryOperations() {
contextRunner.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store oracle add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.ORACLE.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "384")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(),
OracleVectorStore.DEFAULT_TABLE_NAME)
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "cosine")
.hasBeenStarted()
.hasBeenStopped();
observationRegistry.clear();
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("What is Great Depression").withTopK(1));
assertThat(results).isNotEmpty();
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store oracle query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.ORACLE.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "What is Great Depression")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "384")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(),
OracleVectorStore.DEFAULT_TABLE_NAME)
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "cosine")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "1")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "0.0")
.hasBeenStarted()
.hasBeenStopped();
dropTable(context, ((OracleVectorStore) vectorStore).getTableName());
});
}
@SpringBootConfiguration
@EnableAutoConfiguration(exclude = { DataSourceAutoConfiguration.class })
public static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public VectorStore vectorStore(JdbcTemplate jdbcTemplate, EmbeddingModel embeddingModel,
ObservationRegistry observationRegistry) {
return new OracleVectorStore(jdbcTemplate, embeddingModel, OracleVectorStore.DEFAULT_TABLE_NAME,
OracleVectorStore.OracleVectorStoreIndexType.IVF, OracleVectorStoreDistanceType.COSINE, 384,
OracleVectorStore.DEFAULT_SEARCH_ACCURACY, true, true, true, observationRegistry, null);
}
@Bean
public JdbcTemplate myJdbcTemplate(DataSource dataSource) {
return new JdbcTemplate(dataSource);
}
@Bean
@Primary
@ConfigurationProperties("app.datasource")
public DataSourceProperties dataSourceProperties() {
return new DataSourceProperties();
}
@Bean
public OracleDataSource dataSource(DataSourceProperties dataSourceProperties) {
return dataSourceProperties.initializeDataSourceBuilder().type(OracleDataSource.class).build();
}
@Bean
public EmbeddingModel embeddingModel() {
return new TransformersEmbeddingModel();
}
}
}

View File

@@ -139,10 +139,9 @@ public class PgVectorStore extends AbstractObservationVectorStore implements Ini
private PgVectorStore(String schemaName, String vectorTableName, boolean vectorTableValidationsEnabled,
JdbcTemplate jdbcTemplate, EmbeddingModel embeddingModel, int dimensions, PgDistanceType distanceType,
boolean removeExistingVectorStoreTable, PgIndexType createIndexMethod, boolean initializeSchema,
ObservationRegistry observationRegistry,
VectorStoreObservationConvention customSearchObservationConvention) {
ObservationRegistry observationRegistry, VectorStoreObservationConvention customObservationConvention) {
super(observationRegistry, customSearchObservationConvention);
super(observationRegistry, customObservationConvention);
this.vectorTableName = (null == vectorTableName || vectorTableName.isEmpty()) ? DEFAULT_TABLE_NAME
: vectorTableName.trim();
@@ -554,9 +553,8 @@ public class PgVectorStore extends AbstractObservationVectorStore implements Ini
return this;
}
public Builder withSearchObservationConvention(
VectorStoreObservationConvention customSearchObservationConvention) {
this.searchObservationConvention = customSearchObservationConvention;
public Builder withSearchObservationConvention(VectorStoreObservationConvention customObservationConvention) {
this.searchObservationConvention = customObservationConvention;
return this;
}

View File

@@ -122,6 +122,7 @@ public class PgVectorObservationIT {
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "cosine")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "none")
.hasBeenStarted()
.hasBeenStopped();
@@ -149,6 +150,7 @@ public class PgVectorObservationIT {
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "cosine")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "1")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "0.0")
.hasBeenStarted()
.hasBeenStopped();

View File

@@ -101,6 +101,11 @@
<scope>test</scope>
</dependency>
<dependency>
<groupId>io.micrometer</groupId>
<artifactId>micrometer-observation-test</artifactId>
<scope>test</scope>
</dependency>
</dependencies>
</project>

View File

@@ -21,11 +21,25 @@ import java.util.List;
import java.util.Map;
import java.util.Optional;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.model.EmbeddingUtils;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.vectorstore.filter.FilterExpressionConverter;
import org.springframework.ai.vectorstore.filter.converter.PineconeFilterExpressionConverter;
import org.springframework.ai.vectorstore.observation.AbstractObservationVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.util.Assert;
import org.springframework.util.StringUtils;
import com.fasterxml.jackson.core.type.TypeReference;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.google.protobuf.Struct;
import com.google.protobuf.Value;
import com.google.protobuf.util.JsonFormat;
import io.micrometer.observation.ObservationRegistry;
import io.pinecone.PineconeClient;
import io.pinecone.PineconeClientConfig;
import io.pinecone.PineconeConnection;
@@ -36,14 +50,6 @@ import io.pinecone.proto.QueryResponse;
import io.pinecone.proto.UpsertRequest;
import io.pinecone.proto.Vector;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.model.EmbeddingUtils;
import org.springframework.ai.vectorstore.filter.FilterExpressionConverter;
import org.springframework.ai.vectorstore.filter.converter.PineconeFilterExpressionConverter;
import org.springframework.util.Assert;
import org.springframework.util.StringUtils;
/**
* A VectorStore implementation backed by Pinecone, a cloud-based vector database. This
* store supports creating, updating, deleting, and similarity searching of documents in a
@@ -52,7 +58,7 @@ import org.springframework.util.StringUtils;
* @author Christian Tzolov
* @author Adam Bchouti
*/
public class PineconeVectorStore implements VectorStore {
public class PineconeVectorStore extends AbstractObservationVectorStore {
public static final String CONTENT_FIELD_NAME = "document_content";
@@ -252,6 +258,19 @@ public class PineconeVectorStore implements VectorStore {
* @param embeddingModel The client for embedding operations.
*/
public PineconeVectorStore(PineconeVectorStoreConfig config, EmbeddingModel embeddingModel) {
this(config, embeddingModel, ObservationRegistry.NOOP, null);
}
/**
* Constructs a new PineconeVectorStore.
* @param config The configuration for the store.
* @param embeddingModel The client for embedding operations.
* @param observationRegistry The registry for observations.
* @param customObservationConvention The custom observation convention.
*/
public PineconeVectorStore(PineconeVectorStoreConfig config, EmbeddingModel embeddingModel,
ObservationRegistry observationRegistry, VectorStoreObservationConvention customObservationConvention) {
super(observationRegistry, customObservationConvention);
Assert.notNull(config, "PineconeVectorStoreConfig must not be null");
Assert.notNull(embeddingModel, "EmbeddingModel must not be null");
@@ -294,7 +313,7 @@ public class PineconeVectorStore implements VectorStore {
* @param documents The list of documents to be added.
*/
@Override
public void add(List<Document> documents) {
public void doAdd(List<Document> documents) {
add(documents, this.pineconeNamespace);
}
@@ -352,7 +371,7 @@ public class PineconeVectorStore implements VectorStore {
* @return An optional boolean indicating the deletion status.
*/
@Override
public Optional<Boolean> delete(List<String> documentIds) {
public Optional<Boolean> doDelete(List<String> documentIds) {
return delete(documentIds, this.pineconeNamespace);
}
@@ -390,7 +409,7 @@ public class PineconeVectorStore implements VectorStore {
}
@Override
public List<Document> similaritySearch(SearchRequest request) {
public List<Document> doSimilaritySearch(SearchRequest request) {
return similaritySearch(request, this.pineconeNamespace);
}
@@ -424,4 +443,13 @@ public class PineconeVectorStore implements VectorStore {
}
}
@Override
public VectorStoreObservationContext.Builder createObservationContextBuilder(String operationName) {
return VectorStoreObservationContext.builder(VectorStoreProvider.PINECONE.value(), operationName)
.withDimensions(this.embeddingModel.dimensions())
.withNamespace(this.pineconeNamespace)
.withFieldName(this.pineconeContentFieldName);
}
}

View File

@@ -0,0 +1,202 @@
/*
* Copyright 2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (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.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.vectorstore;
import static org.assertj.core.api.Assertions.assertThat;
import static org.hamcrest.Matchers.hasSize;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.util.List;
import java.util.Map;
import java.util.concurrent.TimeUnit;
import org.awaitility.Awaitility;
import org.awaitility.Duration;
import org.junit.jupiter.api.BeforeAll;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.PineconeVectorStore.PineconeVectorStoreConfig;
import org.springframework.ai.vectorstore.observation.DefaultVectorStoreObservationConvention;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.HighCardinalityKeyNames;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.LowCardinalityKeyNames;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.boot.autoconfigure.EnableAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.core.io.DefaultResourceLoader;
import io.micrometer.observation.ObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistryAssert;
/**
* @author Christian Tzolov
*/
@EnabledIfEnvironmentVariable(named = "PINECONE_API_KEY", matches = ".+")
public class PineconeVectorStoreObservationIT {
private static final String PINECONE_ENVIRONMENT = "gcp-starter";
private static final String PINECONE_PROJECT_ID = "814621f";
private static final String PINECONE_INDEX_NAME = "spring-ai-test-index";
// NOTE: Leave it empty as for free tier as later doesn't support namespaces.
private static final String PINECONE_NAMESPACE = "";
private static final String CUSTOM_CONTENT_FIELD_NAME = "article";
List<Document> documents = List.of(
new Document(getText("classpath:/test/data/spring.ai.txt"), Map.of("meta1", "meta1")),
new Document(getText("classpath:/test/data/time.shelter.txt")),
new Document(getText("classpath:/test/data/great.depression.txt"), Map.of("meta2", "meta2")));
public static String getText(String uri) {
var resource = new DefaultResourceLoader().getResource(uri);
try {
return resource.getContentAsString(StandardCharsets.UTF_8);
}
catch (IOException e) {
throw new RuntimeException(e);
}
}
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withUserConfiguration(Config.class);
@BeforeAll
public static void beforeAll() {
Awaitility.setDefaultPollInterval(2, TimeUnit.SECONDS);
Awaitility.setDefaultPollDelay(Duration.ZERO);
Awaitility.setDefaultTimeout(Duration.ONE_MINUTE);
}
@Test
void observationVectorStoreAddAndQueryOperations() {
contextRunner.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store pinecone add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.PINECONE.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "384")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "article")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "none")
.hasBeenStarted()
.hasBeenStopped();
Awaitility.await().until(() -> {
return vectorStore.similaritySearch(SearchRequest.query("What is Great Depression").withTopK(1));
}, hasSize(1));
observationRegistry.clear();
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("What is Great Depression").withTopK(1));
assertThat(results).isNotEmpty();
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store pinecone query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.PINECONE.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "What is Great Depression")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "384")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "article")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "1")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "0.0")
.hasBeenStarted()
.hasBeenStopped();
// Remove all documents from the store
vectorStore.delete(documents.stream().map(doc -> doc.getId()).toList());
Awaitility.await().until(() -> {
return vectorStore.similaritySearch(SearchRequest.query("Hello").withTopK(1));
}, hasSize(0));
});
}
@SpringBootConfiguration
@EnableAutoConfiguration
public static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public PineconeVectorStoreConfig pineconeVectorStoreConfig() {
return PineconeVectorStoreConfig.builder()
.withApiKey(System.getenv("PINECONE_API_KEY"))
.withEnvironment(PINECONE_ENVIRONMENT)
.withProjectId(PINECONE_PROJECT_ID)
.withIndexName(PINECONE_INDEX_NAME)
.withNamespace(PINECONE_NAMESPACE)
.withContentFieldName(CUSTOM_CONTENT_FIELD_NAME)
.build();
}
@Bean
public VectorStore vectorStore(PineconeVectorStoreConfig config, EmbeddingModel embeddingModel,
ObservationRegistry observationRegistry) {
return new PineconeVectorStore(config, embeddingModel, observationRegistry, null);
}
@Bean
public EmbeddingModel embeddingModel() {
return new TransformersEmbeddingModel();
}
}
}

View File

@@ -75,5 +75,18 @@
<artifactId>junit-jupiter</artifactId>
<scope>test</scope>
</dependency>
<dependency>
<groupId>org.springframework.ai</groupId>
<artifactId>spring-ai-test</artifactId>
<version>${parent.version}</version>
<scope>test</scope>
</dependency>
<dependency>
<groupId>io.micrometer</groupId>
<artifactId>micrometer-observation-test</artifactId>
<scope>test</scope>
</dependency>
</dependencies>
</project>

View File

@@ -29,11 +29,16 @@ import java.util.concurrent.ExecutionException;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.model.EmbeddingUtils;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.observation.conventions.VectorStoreSimilarityMetric;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.VectorStore;
import org.springframework.ai.vectorstore.observation.AbstractObservationVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.InitializingBean;
import org.springframework.util.Assert;
import io.micrometer.observation.ObservationRegistry;
import io.qdrant.client.QdrantClient;
import io.qdrant.client.grpc.Collections.Distance;
import io.qdrant.client.grpc.Collections.VectorParams;
@@ -55,7 +60,7 @@ import io.qdrant.client.grpc.Points.UpdateStatus;
* @author Josh Long
* @since 0.8.1
*/
public class QdrantVectorStore implements VectorStore, InitializingBean {
public class QdrantVectorStore extends AbstractObservationVectorStore implements InitializingBean {
private static final String CONTENT_FIELD_NAME = "doc_content";
@@ -152,9 +157,28 @@ public class QdrantVectorStore implements VectorStore, InitializingBean {
* @param qdrantClient A {@link QdrantClient} instance for interfacing with Qdrant.
* @param collectionName The name of the collection to use in Qdrant.
* @param embeddingModel The client for embedding operations.
* @param initializeSchema A boolean indicating whether to initialize the schema.
*/
public QdrantVectorStore(QdrantClient qdrantClient, String collectionName, EmbeddingModel embeddingModel,
boolean initializeSchema) {
this(qdrantClient, collectionName, embeddingModel, initializeSchema, ObservationRegistry.NOOP, null);
}
/**
* Constructs a new QdrantVectorStore.
* @param qdrantClient A {@link QdrantClient} instance for interfacing with Qdrant.
* @param collectionName The name of the collection to use in Qdrant.
* @param embeddingModel The client for embedding operations.
* @param initializeSchema A boolean indicating whether to initialize the schema.
* @param observationRegistry The observation registry to use.
* @param customObservationConvention The custom search observation convention to use.
*/
public QdrantVectorStore(QdrantClient qdrantClient, String collectionName, EmbeddingModel embeddingModel,
boolean initializeSchema, ObservationRegistry observationRegistry,
VectorStoreObservationConvention customObservationConvention) {
super(observationRegistry, customObservationConvention);
Assert.notNull(qdrantClient, "QdrantClient must not be null");
Assert.notNull(collectionName, "collectionName must not be null");
Assert.notNull(embeddingModel, "EmbeddingModel must not be null");
@@ -170,7 +194,7 @@ public class QdrantVectorStore implements VectorStore, InitializingBean {
* @param documents The list of documents to be added.
*/
@Override
public void add(List<Document> documents) {
public void doAdd(List<Document> documents) {
try {
List<PointStruct> points = documents.stream().map(document -> {
// Compute and assign an embedding to the document.
@@ -196,7 +220,7 @@ public class QdrantVectorStore implements VectorStore, InitializingBean {
* @return An optional boolean indicating the deletion status.
*/
@Override
public Optional<Boolean> delete(List<String> documentIds) {
public Optional<Boolean> doDelete(List<String> documentIds) {
try {
List<PointId> ids = documentIds.stream().map(id -> id(UUID.fromString(id))).toList();
var result = this.qdrantClient.deleteAsync(this.collectionName, ids)
@@ -216,7 +240,7 @@ public class QdrantVectorStore implements VectorStore, InitializingBean {
* @return A list of documents that are similar to the query.
*/
@Override
public List<Document> similaritySearch(SearchRequest request) {
public List<Document> doSimilaritySearch(SearchRequest request) {
try {
Filter filter = (request.getFilterExpression() != null)
? this.filterExpressionConverter.convertExpression(request.getFilterExpression())
@@ -307,4 +331,13 @@ public class QdrantVectorStore implements VectorStore, InitializingBean {
}
}
@Override
public VectorStoreObservationContext.Builder createObservationContextBuilder(String operationName) {
return VectorStoreObservationContext.builder(VectorStoreProvider.QDRANT.value(), operationName)
.withDimensions(this.embeddingModel.dimensions())
.withCollectionName(this.collectionName);
}
}

View File

@@ -0,0 +1,209 @@
/*
* Copyright 2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (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.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.vectorstore.qdrant;
import static org.assertj.core.api.Assertions.assertThat;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.util.List;
import java.util.Map;
import java.util.concurrent.ExecutionException;
import org.junit.jupiter.api.BeforeAll;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.azure.openai.AzureOpenAiEmbeddingModel;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.VectorStore;
import org.springframework.ai.vectorstore.observation.DefaultVectorStoreObservationConvention;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.HighCardinalityKeyNames;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.LowCardinalityKeyNames;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.boot.autoconfigure.EnableAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.core.io.DefaultResourceLoader;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import org.testcontainers.qdrant.QdrantContainer;
import com.azure.ai.openai.OpenAIClient;
import com.azure.ai.openai.OpenAIClientBuilder;
import com.azure.core.credential.AzureKeyCredential;
import io.micrometer.observation.ObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistryAssert;
import io.qdrant.client.QdrantClient;
import io.qdrant.client.QdrantGrpcClient;
import io.qdrant.client.grpc.Collections.Distance;
import io.qdrant.client.grpc.Collections.VectorParams;
/**
* @author Christian Tzolov
*/
@Testcontainers
@EnabledIfEnvironmentVariable(named = "AZURE_AI_SEARCH_API_KEY", matches = ".+")
@EnabledIfEnvironmentVariable(named = "AZURE_AI_SEARCH_ENDPOINT", matches = ".+")
public class QdrantVectorStoreObservationIT {
private static final String COLLECTION_NAME = "test_collection";
private static final int EMBEDDING_DIMENSION = 1536;
@Container
static QdrantContainer qdrantContainer = new QdrantContainer("qdrant/qdrant:v1.9.2");
List<Document> documents = List.of(
new Document(getText("classpath:/test/data/spring.ai.txt"), Map.of("meta1", "meta1")),
new Document(getText("classpath:/test/data/time.shelter.txt")),
new Document(getText("classpath:/test/data/great.depression.txt"), Map.of("meta2", "meta2")));
public static String getText(String uri) {
var resource = new DefaultResourceLoader().getResource(uri);
try {
return resource.getContentAsString(StandardCharsets.UTF_8);
}
catch (IOException e) {
throw new RuntimeException(e);
}
}
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withUserConfiguration(Config.class);
@BeforeAll
static void setup() throws InterruptedException, ExecutionException {
String host = qdrantContainer.getHost();
int port = qdrantContainer.getGrpcPort();
QdrantClient client = new QdrantClient(QdrantGrpcClient.newBuilder(host, port, false).build());
client
.createCollectionAsync(COLLECTION_NAME,
VectorParams.newBuilder().setDistance(Distance.Cosine).setSize(EMBEDDING_DIMENSION).build())
.get();
client.close();
}
@Test
void observationVectorStoreAddAndQueryOperations() {
contextRunner.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store qdrant add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.QDRANT.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "1536")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "test_collection")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "none")
.hasBeenStarted()
.hasBeenStopped();
observationRegistry.clear();
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("What is Great Depression").withTopK(1));
assertThat(results).isNotEmpty();
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store qdrant query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.QDRANT.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "What is Great Depression")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "1536")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "test_collection")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "1")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "0.0")
.hasBeenStarted()
.hasBeenStopped();
});
}
@SpringBootConfiguration
@EnableAutoConfiguration
public static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public QdrantClient qdrantClient() {
String host = qdrantContainer.getHost();
int port = qdrantContainer.getGrpcPort();
QdrantClient qdrantClient = new QdrantClient(QdrantGrpcClient.newBuilder(host, port, false).build());
return qdrantClient;
}
@Bean
public VectorStore qdrantVectorStore(EmbeddingModel embeddingModel, QdrantClient qdrantClient,
ObservationRegistry observationRegistry) {
return new QdrantVectorStore(qdrantClient, COLLECTION_NAME, embeddingModel, true, observationRegistry,
null);
}
@Bean
public OpenAIClient openAIClient() {
return new OpenAIClientBuilder().credential(new AzureKeyCredential(System.getenv("AZURE_OPENAI_API_KEY")))
.endpoint(System.getenv("AZURE_OPENAI_ENDPOINT"))
.buildClient();
}
@Bean
public AzureOpenAiEmbeddingModel azureEmbeddingModel(OpenAIClient openAIClient) {
return new AzureOpenAiEmbeddingModel(openAIClient);
}
}
}

View File

@@ -79,6 +79,11 @@
<scope>test</scope>
</dependency>
<dependency>
<groupId>io.micrometer</groupId>
<artifactId>micrometer-observation-test</artifactId>
<scope>test</scope>
</dependency>
</dependencies>

View File

@@ -30,10 +30,16 @@ import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.vectorstore.filter.FilterExpressionConverter;
import org.springframework.ai.vectorstore.observation.AbstractObservationVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.InitializingBean;
import org.springframework.util.Assert;
import org.springframework.util.CollectionUtils;
import io.micrometer.observation.ObservationRegistry;
import redis.clients.jedis.JedisPooled;
import redis.clients.jedis.Pipeline;
import redis.clients.jedis.json.Path2;
@@ -72,7 +78,7 @@ import redis.clients.jedis.search.schemafields.VectorField.VectorAlgorithm;
* @see RedisVectorStoreConfig
* @see EmbeddingModel
*/
public class RedisVectorStore implements VectorStore, InitializingBean {
public class RedisVectorStore extends AbstractObservationVectorStore implements InitializingBean {
public enum Algorithm {
@@ -274,6 +280,15 @@ public class RedisVectorStore implements VectorStore, InitializingBean {
public RedisVectorStore(RedisVectorStoreConfig config, EmbeddingModel embeddingModel, JedisPooled jedis,
boolean initializeSchema) {
this(config, embeddingModel, jedis, initializeSchema, ObservationRegistry.NOOP, null);
}
public RedisVectorStore(RedisVectorStoreConfig config, EmbeddingModel embeddingModel, JedisPooled jedis,
boolean initializeSchema, ObservationRegistry observationRegistry,
VectorStoreObservationConvention customObservationConvention) {
super(observationRegistry, customObservationConvention);
Assert.notNull(config, "Config must not be null");
Assert.notNull(embeddingModel, "Embedding model must not be null");
this.initializeSchema = initializeSchema;
@@ -289,7 +304,7 @@ public class RedisVectorStore implements VectorStore, InitializingBean {
}
@Override
public void add(List<Document> documents) {
public void doAdd(List<Document> documents) {
try (Pipeline pipeline = this.jedis.pipelined()) {
for (Document document : documents) {
var embedding = this.embeddingModel.embed(document);
@@ -318,7 +333,7 @@ public class RedisVectorStore implements VectorStore, InitializingBean {
}
@Override
public Optional<Boolean> delete(List<String> idList) {
public Optional<Boolean> doDelete(List<String> idList) {
try (Pipeline pipeline = this.jedis.pipelined()) {
for (String id : idList) {
pipeline.jsonDel(key(id));
@@ -336,7 +351,7 @@ public class RedisVectorStore implements VectorStore, InitializingBean {
}
@Override
public List<Document> similaritySearch(SearchRequest request) {
public List<Document> doSimilaritySearch(SearchRequest request) {
Assert.isTrue(request.getTopK() > 0, "The number of documents to returned must be greater than zero");
Assert.isTrue(request.getSimilarityThreshold() >= 0 && request.getSimilarityThreshold() <= 1,
@@ -457,13 +472,15 @@ public class RedisVectorStore implements VectorStore, InitializingBean {
return JSON_PATH_PREFIX + field;
}
private static float[] toFloatArray(List<Float> embedding) {
float[] embeddingFloat = new float[embedding.size()];
int i = 0;
for (Float d : embedding) {
embeddingFloat[i++] = d.floatValue();
}
return embeddingFloat;
@Override
public VectorStoreObservationContext.Builder createObservationContextBuilder(String operationName) {
return VectorStoreObservationContext.builder(VectorStoreProvider.REDIS.value(), operationName)
.withDimensions(this.embeddingModel.dimensions())
.withFieldName(this.config.embeddingFieldName)
.withSimilarityMetric(vectorAlgorithm().name())
.withIndexName(this.config.indexName);
}
}

View File

@@ -0,0 +1,184 @@
/*
* Copyright 2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (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.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.vectorstore;
import static org.assertj.core.api.Assertions.assertThat;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.util.List;
import java.util.Map;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.RedisVectorStore.MetadataField;
import org.springframework.ai.vectorstore.RedisVectorStore.RedisVectorStoreConfig;
import org.springframework.ai.vectorstore.observation.DefaultVectorStoreObservationConvention;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.HighCardinalityKeyNames;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.LowCardinalityKeyNames;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.EnableAutoConfiguration;
import org.springframework.boot.autoconfigure.data.redis.RedisAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.core.io.DefaultResourceLoader;
import org.springframework.data.redis.connection.jedis.JedisConnectionFactory;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import com.redis.testcontainers.RedisStackContainer;
import io.micrometer.observation.ObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistryAssert;
import redis.clients.jedis.JedisPooled;
/**
* @author Christian Tzolov
*/
@Testcontainers
public class RedisVectorStoreObservationIT {
@Container
static RedisStackContainer redisContainer = new RedisStackContainer(
RedisStackContainer.DEFAULT_IMAGE_NAME.withTag(RedisStackContainer.DEFAULT_TAG));
List<Document> documents = List.of(
new Document(getText("classpath:/test/data/spring.ai.txt"), Map.of("meta1", "meta1")),
new Document(getText("classpath:/test/data/time.shelter.txt")),
new Document(getText("classpath:/test/data/great.depression.txt"), Map.of("meta2", "meta2")));
public static String getText(String uri) {
var resource = new DefaultResourceLoader().getResource(uri);
try {
return resource.getContentAsString(StandardCharsets.UTF_8);
}
catch (IOException e) {
throw new RuntimeException(e);
}
}
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withConfiguration(AutoConfigurations.of(RedisAutoConfiguration.class))
.withUserConfiguration(Config.class)
.withPropertyValues("spring.data.redis.url=" + redisContainer.getRedisURI());
@BeforeEach
void cleanDatabase() {
this.contextRunner.run(context -> context.getBean(RedisVectorStore.class).getJedis().flushAll());
}
@Test
void observationVectorStoreAddAndQueryOperations() {
contextRunner.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store redis add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.REDIS.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "384")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "embedding")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "HNSW")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.INDEX_NAME.asString(), "spring-ai-index")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "none")
.hasBeenStarted()
.hasBeenStopped();
observationRegistry.clear();
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("What is Great Depression").withTopK(1));
assertThat(results).isNotEmpty();
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store redis query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.REDIS.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "What is Great Depression")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "384")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "embedding")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "HNSW")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.INDEX_NAME.asString(), "spring-ai-index")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "1")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "0.0")
.hasBeenStarted()
.hasBeenStopped();
});
}
@SpringBootConfiguration
@EnableAutoConfiguration
public static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public RedisVectorStore vectorStore(EmbeddingModel embeddingModel,
JedisConnectionFactory jedisConnectionFactory, ObservationRegistry observationRegistry) {
return new RedisVectorStore(
RedisVectorStoreConfig.builder()
.withMetadataFields(MetadataField.tag("meta1"), MetadataField.tag("meta2"),
MetadataField.tag("country"), MetadataField.numeric("year"))
.build(),
embeddingModel,
new JedisPooled(jedisConnectionFactory.getHostName(), jedisConnectionFactory.getPort()), true,
observationRegistry, null);
}
@Bean
public EmbeddingModel embeddingModel() {
return new TransformersEmbeddingModel();
}
}
}

View File

@@ -62,6 +62,13 @@
<scope>test</scope>
</dependency>
<dependency>
<groupId>io.micrometer</groupId>
<artifactId>micrometer-observation-test</artifactId>
<scope>test</scope>
</dependency>
</dependencies>
</project>

View File

@@ -16,7 +16,6 @@
package org.springframework.ai.vectorstore;
import java.util.Arrays;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
@@ -28,7 +27,12 @@ import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.observation.conventions.VectorStoreSimilarityMetric;
import org.springframework.ai.vectorstore.filter.FilterExpressionConverter;
import org.springframework.ai.vectorstore.observation.AbstractObservationVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.beans.factory.InitializingBean;
import org.springframework.util.Assert;
import org.typesense.api.Client;
@@ -42,11 +46,14 @@ import org.typesense.model.MultiSearchCollectionParameters;
import org.typesense.model.MultiSearchResult;
import org.typesense.model.MultiSearchSearchesParameter;
import io.micrometer.observation.ObservationRegistry;
/**
* @author Pablo Sanchidrian Herrera
* @author Soby Chacko
* @author Christian Tzolov
*/
public class TypesenseVectorStore implements VectorStore, InitializingBean {
public class TypesenseVectorStore extends AbstractObservationVectorStore implements InitializingBean {
private static final Logger logger = LoggerFactory.getLogger(TypesenseVectorStore.class);
@@ -154,6 +161,16 @@ public class TypesenseVectorStore implements VectorStore, InitializingBean {
public TypesenseVectorStore(Client client, EmbeddingModel embeddingModel, TypesenseVectorStoreConfig config,
boolean initializeSchema) {
this(client, embeddingModel, config, initializeSchema, ObservationRegistry.NOOP, null);
}
public TypesenseVectorStore(Client client, EmbeddingModel embeddingModel, TypesenseVectorStoreConfig config,
boolean initializeSchema, ObservationRegistry observationRegistry,
VectorStoreObservationConvention customObservationConvention) {
super(observationRegistry, customObservationConvention);
Assert.notNull(client, "Typesense must not be null");
Assert.notNull(embeddingModel, "EmbeddingModel must not be null");
@@ -164,7 +181,7 @@ public class TypesenseVectorStore implements VectorStore, InitializingBean {
}
@Override
public void add(List<Document> documents) {
public void doAdd(List<Document> documents) {
Assert.notNull(documents, "Documents must not be null");
List<HashMap<String, Object>> documentList = documents.stream().map(document -> {
@@ -194,7 +211,7 @@ public class TypesenseVectorStore implements VectorStore, InitializingBean {
}
@Override
public Optional<Boolean> delete(List<String> idList) {
public Optional<Boolean> doDelete(List<String> idList) {
DeleteDocumentsParameters deleteDocumentsParameters = new DeleteDocumentsParameters();
deleteDocumentsParameters.filterBy(DOC_ID_FIELD_NAME + ":=[" + String.join(",", idList) + "]");
@@ -217,7 +234,7 @@ public class TypesenseVectorStore implements VectorStore, InitializingBean {
}
@Override
public List<Document> similaritySearch(SearchRequest request) {
public List<Document> doSimilaritySearch(SearchRequest request) {
Assert.notNull(request.getQuery(), "Query string must not be null");
String nativeFilterExpressions = (request.getFilterExpression() != null)
@@ -361,4 +378,14 @@ public class TypesenseVectorStore implements VectorStore, InitializingBean {
}
@Override
public VectorStoreObservationContext.Builder createObservationContextBuilder(String operationName) {
return VectorStoreObservationContext.builder(VectorStoreProvider.TYPESENSE.value(), operationName)
.withDimensions(this.embeddingModel.dimensions())
.withCollectionName(this.config.collectionName)
.withFieldName(EMBEDDING_FIELD_NAME)
.withSimilarityMetric(VectorStoreSimilarityMetric.COSINE.value());
}
}

View File

@@ -0,0 +1,185 @@
/*
* Copyright 2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (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.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.vectorstore;
import static org.assertj.core.api.Assertions.assertThat;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.time.Duration;
import java.util.ArrayList;
import java.util.List;
import java.util.Map;
import org.junit.jupiter.api.Test;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.TypesenseVectorStore.TypesenseVectorStoreConfig;
import org.springframework.ai.vectorstore.observation.DefaultVectorStoreObservationConvention;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.HighCardinalityKeyNames;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.LowCardinalityKeyNames;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.boot.autoconfigure.EnableAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.core.io.DefaultResourceLoader;
import org.testcontainers.containers.GenericContainer;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import org.typesense.api.Client;
import org.typesense.api.Configuration;
import org.typesense.resources.Node;
import io.micrometer.observation.ObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistryAssert;
/**
* @author Christian Tzolov
*/
@Testcontainers
public class TypesenseVectorStoreObservationIT {
@Container
private static GenericContainer<?> typesenseContainer = new GenericContainer<>("typesense/typesense:26.0")
.withExposedPorts(8108)
.withCommand("--data-dir", "/tmp", "--api-key=xyz", "--enable-cors");
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withUserConfiguration(Config.class);
List<Document> documents = List.of(
new Document(getText("classpath:/test/data/spring.ai.txt"), Map.of("meta1", "meta1")),
new Document(getText("classpath:/test/data/time.shelter.txt")),
new Document(getText("classpath:/test/data/great.depression.txt"), Map.of("meta2", "meta2")));
public static String getText(String uri) {
var resource = new DefaultResourceLoader().getResource(uri);
try {
return resource.getContentAsString(StandardCharsets.UTF_8);
}
catch (IOException e) {
throw new RuntimeException(e);
}
}
@Test
void observationVectorStoreAddAndQueryOperations() {
contextRunner.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store typesense add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.TYPESENSE.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "384")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "test_vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "embedding")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "cosine")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.INDEX_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "none")
.hasBeenStarted()
.hasBeenStopped();
observationRegistry.clear();
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("What is Great Depression").withTopK(1));
assertThat(results).isNotEmpty();
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store typesense query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.TYPESENSE.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "What is Great Depression")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "384")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "test_vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "embedding")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "cosine")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.INDEX_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "1")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "0.0")
.hasBeenStarted()
.hasBeenStopped();
});
}
@SpringBootConfiguration
@EnableAutoConfiguration
public static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public VectorStore vectorStore(Client client, EmbeddingModel embeddingModel,
ObservationRegistry observationRegistry) {
TypesenseVectorStoreConfig config = TypesenseVectorStoreConfig.builder()
.withCollectionName("test_vector_store")
.withEmbeddingDimension(embeddingModel.dimensions())
.build();
return new TypesenseVectorStore(client, embeddingModel, config, true, observationRegistry, null);
}
@Bean
public Client typesenseClient() {
List<Node> nodes = new ArrayList<>();
nodes
.add(new Node("http", typesenseContainer.getHost(), typesenseContainer.getMappedPort(8108).toString()));
Configuration configuration = new Configuration(nodes, Duration.ofSeconds(5), "xyz");
return new Client(configuration);
}
@Bean
public EmbeddingModel embeddingModel() {
return new TransformersEmbeddingModel();
}
}
}

View File

@@ -80,6 +80,12 @@
<scope>test</scope>
</dependency>
<dependency>
<groupId>io.micrometer</groupId>
<artifactId>micrometer-observation-test</artifactId>
<scope>test</scope>
</dependency>
</dependencies>
</project>

View File

@@ -23,8 +23,23 @@ import java.util.Map;
import java.util.Optional;
import java.util.stream.Collectors;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.model.EmbeddingUtils;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.vectorstore.WeaviateVectorStore.WeaviateVectorStoreConfig.ConsistentLevel;
import org.springframework.ai.vectorstore.WeaviateVectorStore.WeaviateVectorStoreConfig.MetadataField;
import org.springframework.ai.vectorstore.observation.AbstractObservationVectorStore;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationContext;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationConvention;
import org.springframework.util.Assert;
import org.springframework.util.CollectionUtils;
import org.springframework.util.StringUtils;
import com.fasterxml.jackson.core.JsonProcessingException;
import com.fasterxml.jackson.databind.ObjectMapper;
import io.micrometer.observation.ObservationRegistry;
import io.weaviate.client.WeaviateClient;
import io.weaviate.client.base.Result;
import io.weaviate.client.base.WeaviateErrorMessage;
@@ -42,16 +57,6 @@ import io.weaviate.client.v1.graphql.query.builder.GetBuilder.GetBuilderBuilder;
import io.weaviate.client.v1.graphql.query.fields.Field;
import io.weaviate.client.v1.graphql.query.fields.Fields;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.model.EmbeddingUtils;
import org.springframework.ai.vectorstore.WeaviateVectorStore.WeaviateVectorStoreConfig.ConsistentLevel;
import org.springframework.ai.vectorstore.WeaviateVectorStore.WeaviateVectorStoreConfig.MetadataField;
import org.springframework.beans.factory.InitializingBean;
import org.springframework.util.Assert;
import org.springframework.util.CollectionUtils;
import org.springframework.util.StringUtils;
/**
* A VectorStore implementation backed by Weaviate vector database.
*
@@ -65,7 +70,7 @@ import org.springframework.util.StringUtils;
* @author Josh Long
* @author Soby Chacko
*/
public class WeaviateVectorStore implements VectorStore {
public class WeaviateVectorStore extends AbstractObservationVectorStore {
public static final String DOCUMENT_METADATA_DISTANCE_KEY_NAME = "distance";
@@ -281,9 +286,27 @@ public class WeaviateVectorStore implements VectorStore {
* Constructs a new WeaviateVectorStore.
* @param vectorStoreConfig The configuration for the store.
* @param embeddingModel The client for embedding operations.
* @param weaviateClient The client for Weaviate operations.
*/
public WeaviateVectorStore(WeaviateVectorStoreConfig vectorStoreConfig, EmbeddingModel embeddingModel,
WeaviateClient weaviateClient) {
this(vectorStoreConfig, embeddingModel, weaviateClient, ObservationRegistry.NOOP, null);
}
/**
* Constructs a new WeaviateVectorStore.
* @param vectorStoreConfig The configuration for the store.
* @param embeddingModel The client for embedding operations.
* @param weaviateClient The client for Weaviate operations.
* @param observationRegistry The registry for observations.
* @param customObservationConvention The custom observation convention.
*/
public WeaviateVectorStore(WeaviateVectorStoreConfig vectorStoreConfig, EmbeddingModel embeddingModel,
WeaviateClient weaviateClient, ObservationRegistry observationRegistry,
VectorStoreObservationConvention customObservationConvention) {
super(observationRegistry, customObservationConvention);
Assert.notNull(vectorStoreConfig, "WeaviateVectorStoreConfig must not be null");
Assert.notNull(embeddingModel, "EmbeddingModel must not be null");
@@ -318,7 +341,7 @@ public class WeaviateVectorStore implements VectorStore {
}
@Override
public void add(List<Document> documents) {
public void doAdd(List<Document> documents) {
if (CollectionUtils.isEmpty(documents)) {
return;
@@ -395,7 +418,7 @@ public class WeaviateVectorStore implements VectorStore {
}
@Override
public Optional<Boolean> delete(List<String> documentIds) {
public Optional<Boolean> doDelete(List<String> documentIds) {
Result<BatchDeleteResponse> result = this.weaviateClient.batch()
.objectsBatchDeleter()
@@ -421,7 +444,7 @@ public class WeaviateVectorStore implements VectorStore {
}
@Override
public List<Document> similaritySearch(SearchRequest request) {
public List<Document> doSimilaritySearch(SearchRequest request) {
float[] embedding = this.embeddingModel.embed(request.getQuery());
@@ -518,4 +541,12 @@ public class WeaviateVectorStore implements VectorStore {
return document;
}
@Override
public VectorStoreObservationContext.Builder createObservationContextBuilder(String operationName) {
return VectorStoreObservationContext.builder(VectorStoreProvider.WEAVIATE.value(), operationName)
.withDimensions(this.embeddingModel.dimensions())
.withCollectionName(this.weaviateObjectClass);
}
}

View File

@@ -0,0 +1,174 @@
/*
* Copyright 2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (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.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.vectorstore;
import static org.assertj.core.api.Assertions.assertThat;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.util.List;
import java.util.Map;
import org.junit.jupiter.api.Test;
import org.springframework.ai.document.Document;
import org.springframework.ai.embedding.EmbeddingModel;
import org.springframework.ai.observation.conventions.VectorStoreProvider;
import org.springframework.ai.transformers.TransformersEmbeddingModel;
import org.springframework.ai.vectorstore.WeaviateVectorStore.WeaviateVectorStoreConfig;
import org.springframework.ai.vectorstore.observation.DefaultVectorStoreObservationConvention;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.HighCardinalityKeyNames;
import org.springframework.ai.vectorstore.observation.VectorStoreObservationDocumentation.LowCardinalityKeyNames;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.boot.autoconfigure.EnableAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.core.io.DefaultResourceLoader;
import org.testcontainers.containers.wait.strategy.Wait;
import org.testcontainers.junit.jupiter.Container;
import org.testcontainers.junit.jupiter.Testcontainers;
import org.testcontainers.weaviate.WeaviateContainer;
import io.micrometer.observation.ObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistry;
import io.micrometer.observation.tck.TestObservationRegistryAssert;
import io.weaviate.client.WeaviateClient;
/**
* @author Christian Tzolov
*/
@Testcontainers
public class WeaviateVectorStoreObservationIT {
@Container
static WeaviateContainer weaviateContainer = new WeaviateContainer("semitechnologies/weaviate:1.25.4")
.waitingFor(Wait.forHttp("/v1/.well-known/ready").forPort(8080));
List<Document> documents = List.of(
new Document(getText("classpath:/test/data/spring.ai.txt"), Map.of("meta1", "meta1")),
new Document(getText("classpath:/test/data/time.shelter.txt")),
new Document(getText("classpath:/test/data/great.depression.txt"), Map.of("meta2", "meta2")));
public static String getText(String uri) {
var resource = new DefaultResourceLoader().getResource(uri);
try {
return resource.getContentAsString(StandardCharsets.UTF_8);
}
catch (IOException e) {
throw new RuntimeException(e);
}
}
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withUserConfiguration(Config.class);
@Test
void observationVectorStoreAddAndQueryOperations() {
contextRunner.run(context -> {
VectorStore vectorStore = context.getBean(VectorStore.class);
TestObservationRegistry observationRegistry = context.getBean(TestObservationRegistry.class);
vectorStore.add(documents);
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store weaviate add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "add")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.WEAVIATE.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "384")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "SpringAiWeaviate")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.INDEX_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "none")
.hasBeenStarted()
.hasBeenStopped();
observationRegistry.clear();
List<Document> results = vectorStore
.similaritySearch(SearchRequest.query("What is Great Depression").withTopK(1));
assertThat(results).isNotEmpty();
TestObservationRegistryAssert.assertThat(observationRegistry)
.doesNotHaveAnyRemainingCurrentObservation()
.hasObservationWithNameEqualTo(DefaultVectorStoreObservationConvention.DEFAULT_NAME)
.that()
.hasContextualNameEqualTo("vector_store weaviate query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_OPERATION_NAME.asString(), "query")
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.DB_SYSTEM.asString(),
VectorStoreProvider.WEAVIATE.value())
.hasLowCardinalityKeyValue(LowCardinalityKeyNames.SPRING_AI_KIND.asString(), "vector_store")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.QUERY.asString(), "What is Great Depression")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.DIMENSIONS.asString(), "384")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.COLLECTION_NAME.asString(), "SpringAiWeaviate")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.NAMESPACE.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.FIELD_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_METRIC.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.INDEX_NAME.asString(), "none")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.TOP_K.asString(), "1")
.hasHighCardinalityKeyValue(HighCardinalityKeyNames.SIMILARITY_THRESHOLD.asString(), "0.0")
.hasBeenStarted()
.hasBeenStopped();
// vectorStore.delete(documents.stream().map(Document::getId).toList());
});
}
@SpringBootConfiguration
@EnableAutoConfiguration
public static class Config {
@Bean
public TestObservationRegistry observationRegistry() {
return TestObservationRegistry.create();
}
@Bean
public WeaviateVectorStore vectorStore(EmbeddingModel embeddingModel, ObservationRegistry observationRegistry) {
WeaviateClient weaviateClient = new WeaviateClient(
new io.weaviate.client.Config("http", weaviateContainer.getHttpHostAddress()));
WeaviateVectorStoreConfig config = WeaviateVectorStore.WeaviateVectorStoreConfig.builder()
.withConsistencyLevel(WeaviateVectorStoreConfig.ConsistentLevel.ONE)
.build();
return new WeaviateVectorStore(config, embeddingModel, weaviateClient, observationRegistry, null);
}
@Bean
public EmbeddingModel embeddingModel() {
return new TransformersEmbeddingModel();
}
}
}