Add option to CassandraVectorStore to return embeddings in documents from similarity searches

This commit is contained in:
mck
2024-05-03 13:20:45 +02:00
committed by Christian Tzolov
parent 6e40575977
commit 8a5f9dfb22
3 changed files with 89 additions and 20 deletions

View File

@@ -160,11 +160,9 @@ public final class CassandraVectorStore implements VectorStore, InitializingBean
futures[i++] = CompletableFuture.runAsync(() -> {
List<Object> primaryKeyValues = this.conf.documentIdTranslator.apply(d.getId());
var embedding = (null != d.getEmbedding() && !d.getEmbedding().isEmpty() ? d.getEmbedding()
: this.embeddingClient.embed(d))
.stream()
.map(Double::floatValue)
.toList();
if (null == d.getEmbedding() || d.getEmbedding().isEmpty()) {
d.setEmbedding(this.embeddingClient.embed(d));
}
BoundStatementBuilder builder = prepareAddStatement(d.getMetadata().keySet()).boundStatementBuilder();
for (int k = 0; k < primaryKeyValues.size(); ++k) {
@@ -173,7 +171,9 @@ public final class CassandraVectorStore implements VectorStore, InitializingBean
}
builder = builder.setString(this.conf.schema.content(), d.getContent())
.setVector(this.conf.schema.embedding(), CqlVector.newInstance(embedding), Float.class);
.setVector(this.conf.schema.embedding(),
CqlVector.newInstance(d.getEmbedding().stream().map(Double::floatValue).toList()),
Float.class);
for (var metadataColumn : this.conf.schema.metadataColumns()
.stream()
@@ -235,8 +235,15 @@ public final class CassandraVectorStore implements VectorStore, InitializingBean
docFields.put(metadata.name(), value);
}
}
Document doc = new Document(getDocumentId(row), row.getString(this.conf.schema.content()), docFields);
documents.add(new Document(getDocumentId(row), row.getString(this.conf.schema.content()), docFields));
if (this.conf.returnEmbeddings) {
doc.setEmbedding(row.getVector(this.conf.schema.embedding(), Float.class)
.stream()
.map(Float::doubleValue)
.toList());
}
documents.add(doc);
}
return documents;
}
@@ -328,6 +335,9 @@ public final class CassandraVectorStore implements VectorStore, InitializingBean
for (var m : this.conf.schema.metadataColumns()) {
extraSelectFields.append(',').append(m.name());
}
if (this.conf.returnEmbeddings) {
extraSelectFields.append(',').append(this.conf.schema.embedding());
}
// java-driver-query-builder doesn't support orderByAnnOf yet
String query = String.format(QUERY_FORMAT, similarityFunction, ids.toString(), this.conf.schema.content(),

View File

@@ -132,6 +132,8 @@ public final class CassandraVectorStoreConfig implements AutoCloseable {
final boolean disallowSchemaChanges;
final boolean returnEmbeddings;
final DocumentIdTranslator documentIdTranslator;
final PrimaryKeyTranslator primaryKeyTranslator;
@@ -148,6 +150,7 @@ public final class CassandraVectorStoreConfig implements AutoCloseable {
builder.contentColumnName, builder.embeddingColumnName, builder.indexName, builder.metadataColumns);
this.disallowSchemaChanges = builder.disallowSchemaCreation;
this.returnEmbeddings = builder.returnEmbeddings;
this.documentIdTranslator = builder.documentIdTranslator;
this.primaryKeyTranslator = builder.primaryKeyTranslator;
this.executor = Executors.newFixedThreadPool(builder.fixedThreadPoolExecutorSize);
@@ -199,6 +202,8 @@ public final class CassandraVectorStoreConfig implements AutoCloseable {
private boolean disallowSchemaCreation = false;
private boolean returnEmbeddings = false;
private int fixedThreadPoolExecutorSize = DEFAULT_ADD_CONCURRENCY;
private DocumentIdTranslator documentIdTranslator = (String id) -> List.of(id);
@@ -308,6 +313,11 @@ public final class CassandraVectorStoreConfig implements AutoCloseable {
return this;
}
public Builder returnEmbeddings() {
this.returnEmbeddings = true;
return this;
}
/**
* Executor to use when adding documents. The hotspot is the call to the
* embeddingClient. For remote transformers you probably want a higher value to

View File

@@ -68,13 +68,14 @@ class CassandraVectorStoreIT {
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withUserConfiguration(TestApplication.class);
List<Document> documents = List.of(
new Document("1", getText("classpath:/test/data/spring.ai.txt"), Map.of("meta1", "meta1")),
new Document("2", getText("classpath:/test/data/time.shelter.txt"), Map.of()),
new Document("3", getText("classpath:/test/data/great.depression.txt"),
Map.of("meta2", "meta2", "something_extra", "blue")));
private static List<Document> documents() {
return List.of(new Document("1", getText("classpath:/test/data/spring.ai.txt"), Map.of("meta1", "meta1")),
new Document("2", getText("classpath:/test/data/time.shelter.txt"), Map.of()),
new Document("3", getText("classpath:/test/data/great.depression.txt"),
Map.of("meta2", "meta2", "something_extra", "blue")));
}
public static String getText(String uri) {
private static String getText(String uri) {
var resource = new DefaultResourceLoader().getResource(uri);
try {
return resource.getContentAsString(StandardCharsets.UTF_8);
@@ -99,13 +100,21 @@ class CassandraVectorStoreIT {
contextRunner.run(context -> {
try (CassandraVectorStore store = createTestStore(context, new SchemaColumn("meta1", DataTypes.TEXT),
new SchemaColumn("meta2", DataTypes.TEXT))) {
List<Document> documents = documents();
store.add(documents);
for (Document d : documents) {
assertThat(d.getEmbedding()).satisfiesAnyOf(e -> assertThat(e).isNotNull(),
e -> assertThat(e).isNotEmpty());
}
List<Document> results = store.similaritySearch(SearchRequest.query("Spring").withTopK(1));
assertThat(results).hasSize(1);
Document resultDoc = results.get(0);
assertThat(resultDoc.getId()).isEqualTo(documents.get(0).getId());
assertThat(resultDoc.getId()).isEqualTo(documents().get(0).getId());
assertThat(resultDoc.getEmbedding()).satisfiesAnyOf(e -> assertThat(e).isNull(),
e -> assertThat(e).isEmpty());
assertThat(resultDoc.getContent()).contains(
"Spring AI provides abstractions that serve as the foundation for developing AI applications.");
@@ -114,7 +123,43 @@ class CassandraVectorStoreIT {
assertThat(resultDoc.getMetadata()).containsKeys("meta1", CassandraVectorStore.SIMILARITY_FIELD_NAME);
// Remove all documents from the store
store.delete(documents.stream().map(doc -> doc.getId()).toList());
store.delete(documents().stream().map(doc -> doc.getId()).toList());
results = store.similaritySearch(SearchRequest.query("Spring").withTopK(1));
assertThat(results).isEmpty();
}
});
}
@Test
void addAndSearchReturnEmbeddings() {
contextRunner.run(context -> {
CassandraVectorStoreConfig.Builder builder = storeBuilder(context.getBean(CqlSession.class))
.returnEmbeddings();
try (CassandraVectorStore store = createTestStore(context, builder)) {
List<Document> documents = documents();
store.add(documents);
for (Document d : documents) {
assertThat(d.getEmbedding()).satisfiesAnyOf(e -> assertThat(e).isNotNull(),
e -> assertThat(e).isNotEmpty());
}
List<Document> results = store.similaritySearch(SearchRequest.query("Spring").withTopK(1));
assertThat(results).hasSize(1);
Document resultDoc = results.get(0);
assertThat(resultDoc.getId()).isEqualTo(documents().get(0).getId());
assertThat(resultDoc.getEmbedding()).isNotEmpty();
assertThat(resultDoc.getContent()).contains(
"Spring AI provides abstractions that serve as the foundation for developing AI applications.");
assertThat(resultDoc.getMetadata()).hasSize(1);
assertThat(resultDoc.getMetadata()).containsKey(CassandraVectorStore.SIMILARITY_FIELD_NAME);
// Remove all documents from the store
store.delete(documents().stream().map(doc -> doc.getId()).toList());
results = store.similaritySearch(SearchRequest.query("Spring").withTopK(1));
assertThat(results).isEmpty();
@@ -309,7 +354,7 @@ class CassandraVectorStoreIT {
void searchWithThreshold() {
contextRunner.run(context -> {
try (CassandraVectorStore store = context.getBean(CassandraVectorStore.class)) {
store.add(documents);
store.add(documents());
List<Document> fullResult = store
.similaritySearch(SearchRequest.query("Spring").withTopK(5).withSimilarityThresholdAll());
@@ -327,7 +372,7 @@ class CassandraVectorStoreIT {
assertThat(results).hasSize(1);
Document resultDoc = results.get(0);
assertThat(resultDoc.getId()).isEqualTo(documents.get(0).getId());
assertThat(resultDoc.getId()).isEqualTo(documents().get(0).getId());
assertThat(resultDoc.getContent()).contains(
"Spring AI provides abstractions that serve as the foundation for developing AI applications.");
@@ -370,17 +415,21 @@ class CassandraVectorStoreIT {
}
static CassandraVectorStoreConfig.Builder storeBuilder(CqlSession cqlSession) {
private static CassandraVectorStoreConfig.Builder storeBuilder(CqlSession cqlSession) {
return CassandraVectorStoreConfig.builder()
.withCqlSession(cqlSession)
.withKeyspaceName("test_" + CassandraVectorStoreConfig.DEFAULT_KEYSPACE_NAME);
}
private CassandraVectorStore createTestStore(ApplicationContext context, SchemaColumn... metadataFields) {
private static CassandraVectorStore createTestStore(ApplicationContext context, SchemaColumn... metadataFields) {
CassandraVectorStoreConfig.Builder builder = storeBuilder(context.getBean(CqlSession.class))
.addMetadataColumns(metadataFields);
return createTestStore(context, builder);
}
private static CassandraVectorStore createTestStore(ApplicationContext context,
CassandraVectorStoreConfig.Builder builder) {
CassandraVectorStoreConfig conf = builder.build();
conf.dropKeyspace();
return new CassandraVectorStore(conf, context.getBean(EmbeddingClient.class));