From daf131be2cc668fe8f1b536b3025dff14e1d231f Mon Sep 17 00:00:00 2001 From: mck Date: Wed, 24 Apr 2024 14:12:49 +0200 Subject: [PATCH] Fix column creation, when adding additional normal and embedding colums), and make index name unique (for when there are multiple vector indexes in the same keyspace) And change stream to for-loop when converting List to Float[] for performance --- .../CassandraVectorStoreProperties.java | 2 +- .../CassandraVectorStorePropertiesTests.java | 2 +- .../ai/vectorstore/CassandraVectorStore.java | 11 ++++++++++- .../vectorstore/CassandraVectorStoreConfig.java | 16 ++++++++++++---- .../CassandraRichSchemaVectorStoreIT.java | 7 ++++++- .../resources/test_wiki_partial_4_schema.cql | 10 ++++++++++ 6 files changed, 40 insertions(+), 8 deletions(-) create mode 100644 vector-stores/spring-ai-cassandra/src/test/resources/test_wiki_partial_4_schema.cql diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/cassandra/CassandraVectorStoreProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/cassandra/CassandraVectorStoreProperties.java index 1f2433100..27af7605e 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/cassandra/CassandraVectorStoreProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/cassandra/CassandraVectorStoreProperties.java @@ -33,7 +33,7 @@ public class CassandraVectorStoreProperties { private String table = CassandraVectorStoreConfig.DEFAULT_TABLE_NAME; - private String indexName = CassandraVectorStoreConfig.DEFAULT_INDEX_NAME; + private String indexName = null; private String contentColumnName = CassandraVectorStoreConfig.DEFAULT_CONTENT_COLUMN_NAME; diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/cassandra/CassandraVectorStorePropertiesTests.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/cassandra/CassandraVectorStorePropertiesTests.java index ca5c678d7..c0d5ad04a 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/cassandra/CassandraVectorStorePropertiesTests.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/cassandra/CassandraVectorStorePropertiesTests.java @@ -34,7 +34,7 @@ class CassandraVectorStorePropertiesTests { assertThat(props.getTable()).isEqualTo(CassandraVectorStoreConfig.DEFAULT_TABLE_NAME); assertThat(props.getContentColumnName()).isEqualTo(CassandraVectorStoreConfig.DEFAULT_CONTENT_COLUMN_NAME); assertThat(props.getEmbeddingColumnName()).isEqualTo(CassandraVectorStoreConfig.DEFAULT_EMBEDDING_COLUMN_NAME); - assertThat(props.getIndexName()).isEqualTo(CassandraVectorStoreConfig.DEFAULT_INDEX_NAME); + assertThat(props.getIndexName()).isNull(); assertThat(props.getDisallowSchemaCreation()).isFalse(); assertThat(props.getFixedThreadPoolExecutorSize()) .isEqualTo(CassandraVectorStoreConfig.DEFAULT_ADD_CONCURRENCY); 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 4e2732d58..c7422cc5b 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 @@ -206,7 +206,7 @@ public final class CassandraVectorStore implements VectorStore, InitializingBean @Override public List similaritySearch(SearchRequest request) { Preconditions.checkArgument(request.getTopK() <= 1000); - var embedding = this.embeddingClient.embed(request.getQuery()).stream().map(Double::floatValue).toList(); + var embedding = toFloatArray(this.embeddingClient.embed(request.getQuery())); CqlVector cqlVector = CqlVector.newInstance(embedding); String whereClause = ""; @@ -350,4 +350,13 @@ public final class CassandraVectorStore implements VectorStore, InitializingBean return this.conf.primaryKeyTranslator.apply(primaryKeyValues); } + private static Float[] toFloatArray(List embeddingDouble) { + Float[] embeddingFloat = new Float[embeddingDouble.size()]; + int i = 0; + for (Double d : embeddingDouble) { + embeddingFloat[i++] = d.floatValue(); + } + return embeddingFloat; + } + } 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 91356e5e0..32a76d508 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 @@ -44,9 +44,12 @@ import com.datastax.oss.driver.api.querybuilder.schema.CreateTable; import com.datastax.oss.driver.api.querybuilder.schema.CreateTableStart; import com.datastax.oss.driver.shaded.guava.common.annotations.VisibleForTesting; import com.datastax.oss.driver.shaded.guava.common.base.Preconditions; + import org.slf4j.Logger; import org.slf4j.LoggerFactory; +import org.springframework.lang.Nullable; + /** * Configuration for the Cassandra vector store. * @@ -71,7 +74,7 @@ public final class CassandraVectorStoreConfig implements AutoCloseable { public static final String DEFAULT_ID_NAME = "id"; - public static final String DEFAULT_INDEX_NAME = "embedding_index"; + public static final String DEFAULT_INDEX_SUFFIX = "idx"; public static final String DEFAULT_CONTENT_COLUMN_NAME = "content"; @@ -186,7 +189,7 @@ public final class CassandraVectorStoreConfig implements AutoCloseable { private List clusteringKeys = List.of(); - private String indexName = DEFAULT_INDEX_NAME; + private String indexName = null; private String contentColumnName = DEFAULT_CONTENT_COLUMN_NAME; @@ -257,6 +260,8 @@ public final class CassandraVectorStoreConfig implements AutoCloseable { return this; } + /** defaults (if null) to '__idx' **/ + @Nullable public Builder withIndexName(String name) { this.indexName = name; return this; @@ -324,6 +329,9 @@ public final class CassandraVectorStoreConfig implements AutoCloseable { } public CassandraVectorStoreConfig build() { + if (null == this.indexName) { + this.indexName = String.format("%s_%s_%s", this.table, this.embeddingColumnName, DEFAULT_INDEX_SUFFIX); + } for (SchemaColumn metadata : this.metadataColumns) { Preconditions.checkArgument( @@ -530,7 +538,7 @@ public final class CassandraVectorStoreConfig implements AutoCloseable { // special case for embedding column, bc JAVA-3118, as above StringBuilder alterTableStmt = new StringBuilder(((BuildableQuery) alterTable).asCql()); if (newColumns.isEmpty() && !addContent) { - alterTableStmt.append(" ADD "); + alterTableStmt.append(" ADD ("); } else { alterTableStmt.setLength(alterTableStmt.length() - 1); @@ -539,7 +547,7 @@ public final class CassandraVectorStoreConfig implements AutoCloseable { alterTableStmt.append(this.schema.embedding) .append(" vector"); + .append(">)"); logger.debug("Executing {}", alterTableStmt.toString()); this.session.execute(alterTableStmt.toString()); diff --git a/vector-stores/spring-ai-cassandra/src/test/java/org/springframework/ai/vectorstore/CassandraRichSchemaVectorStoreIT.java b/vector-stores/spring-ai-cassandra/src/test/java/org/springframework/ai/vectorstore/CassandraRichSchemaVectorStoreIT.java index 2dc59afae..868c79fbb 100644 --- a/vector-stores/spring-ai-cassandra/src/test/java/org/springframework/ai/vectorstore/CassandraRichSchemaVectorStoreIT.java +++ b/vector-stores/spring-ai-cassandra/src/test/java/org/springframework/ai/vectorstore/CassandraRichSchemaVectorStoreIT.java @@ -134,7 +134,8 @@ class CassandraRichSchemaVectorStoreIT { @Test void ensureSchemaPartialCreation() { this.contextRunner.run(context -> { - for (int i = 0; i < 4; ++i) { + int PARTIAL_FILES = 5; + for (int i = 0; i < PARTIAL_FILES; ++i) { executeCqlFile(context, format("test_wiki_partial_%d_schema.cql", i)); var wrapper = createStore(context, List.of(), false, false); try { @@ -148,6 +149,10 @@ class CassandraRichSchemaVectorStoreIT { wrapper.store().close(); } } + // make sure there's not more files to test + Assertions.assertThrows(IOException.class, () -> { + executeCqlFile(context, format("test_wiki_partial_%d_schema.cql", PARTIAL_FILES)); + }); }); } diff --git a/vector-stores/spring-ai-cassandra/src/test/resources/test_wiki_partial_4_schema.cql b/vector-stores/spring-ai-cassandra/src/test/resources/test_wiki_partial_4_schema.cql new file mode 100644 index 000000000..68b4583c4 --- /dev/null +++ b/vector-stores/spring-ai-cassandra/src/test/resources/test_wiki_partial_4_schema.cql @@ -0,0 +1,10 @@ +CREATE KEYSPACE IF NOT EXISTS test_wikidata WITH replication = {'class': 'SimpleStrategy', 'replication_factor': 1}; + +CREATE TABLE IF NOT EXISTS test_wikidata.articles ( + wiki text, + language text, + title text, + chunk_no int, + messages text, + PRIMARY KEY ((wiki, language, title), chunk_no) +); \ No newline at end of file