diff --git a/vector-stores/spring-ai-cassandra/src/main/java/org/springframework/ai/vectorstore/CassandraVectorStore.java b/vector-stores/spring-ai-cassandra/src/main/java/org/springframework/ai/vectorstore/CassandraVectorStore.java index c7422cc5b..a9532d850 100644 --- a/vector-stores/spring-ai-cassandra/src/main/java/org/springframework/ai/vectorstore/CassandraVectorStore.java +++ b/vector-stores/spring-ai-cassandra/src/main/java/org/springframework/ai/vectorstore/CassandraVectorStore.java @@ -160,11 +160,9 @@ public final class CassandraVectorStore implements VectorStore, InitializingBean futures[i++] = CompletableFuture.runAsync(() -> { List 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(), diff --git a/vector-stores/spring-ai-cassandra/src/main/java/org/springframework/ai/vectorstore/CassandraVectorStoreConfig.java b/vector-stores/spring-ai-cassandra/src/main/java/org/springframework/ai/vectorstore/CassandraVectorStoreConfig.java index 528da6c05..288894648 100644 --- a/vector-stores/spring-ai-cassandra/src/main/java/org/springframework/ai/vectorstore/CassandraVectorStoreConfig.java +++ b/vector-stores/spring-ai-cassandra/src/main/java/org/springframework/ai/vectorstore/CassandraVectorStoreConfig.java @@ -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 diff --git a/vector-stores/spring-ai-cassandra/src/test/java/org/springframework/ai/vectorstore/CassandraVectorStoreIT.java b/vector-stores/spring-ai-cassandra/src/test/java/org/springframework/ai/vectorstore/CassandraVectorStoreIT.java index d1fac3901..27ed4246d 100644 --- a/vector-stores/spring-ai-cassandra/src/test/java/org/springframework/ai/vectorstore/CassandraVectorStoreIT.java +++ b/vector-stores/spring-ai-cassandra/src/test/java/org/springframework/ai/vectorstore/CassandraVectorStoreIT.java @@ -68,13 +68,14 @@ class CassandraVectorStoreIT { private final ApplicationContextRunner contextRunner = new ApplicationContextRunner() .withUserConfiguration(TestApplication.class); - List 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 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 documents = documents(); store.add(documents); + for (Document d : documents) { + assertThat(d.getEmbedding()).satisfiesAnyOf(e -> assertThat(e).isNotNull(), + e -> assertThat(e).isNotEmpty()); + } List 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 documents = documents(); + store.add(documents); + for (Document d : documents) { + assertThat(d.getEmbedding()).satisfiesAnyOf(e -> assertThat(e).isNotNull(), + e -> assertThat(e).isNotEmpty()); + } + + List 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 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));