Add option to CassandraVectorStore to return embeddings in documents from similarity searches
This commit is contained in:
@@ -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(),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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));
|
||||
|
||||
Reference in New Issue
Block a user