From 25352b86bf10af3a5b354e9e7d33f5e58ba8515c Mon Sep 17 00:00:00 2001 From: Mark Paluch Date: Wed, 21 Aug 2024 10:28:01 +0200 Subject: [PATCH] Add support for Cassandra Vector search. Closes #1504 --- .../data/cassandra/core/StatementFactory.java | 70 +++-- .../core/convert/CassandraConverters.java | 140 ++++++++- .../convert/CassandraCustomConversions.java | 1 - .../core/convert/CassandraVector.java | 106 +++++++ .../convert/DefaultColumnTypeResolver.java | 11 + .../convert/IndexSpecificationFactory.java | 34 ++- .../convert/MappingCassandraConverter.java | 6 +- .../cassandra/core/convert/QueryMapper.java | 39 ++- .../cassandra/core/convert/SchemaFactory.java | 15 + .../cql/keyspace/ColumnSpecification.java | 20 +- .../mapping/CassandraSimpleTypeHolder.java | 2 + .../cassandra/core/mapping/CassandraType.java | 9 +- .../{SAIIndexed.java => SaiIndexed.java} | 15 +- .../core/mapping/SimilarityFunction.java | 27 ++ .../cassandra/core/mapping/VectorType.java | 47 +++ .../data/cassandra/core/query/Columns.java | 268 +++++++++++++++--- .../data/cassandra/core/query/Query.java | 23 +- .../data/cassandra/core/query/VectorSort.java | 81 ++++++ .../CassandraParametersParameterAccessor.java | 1 + ...ersistentEntitySchemaCreatorUnitTests.java | 18 +- ...CassandraVectorSearchIntegrationTests.java | 166 +++++++++++ .../core/StatementFactoryUnitTests.java | 75 ++++- .../IndexSpecificationFactoryUnitTests.java | 48 +++- .../MappingCassandraConverterUnitTests.java | 1 + .../core/convert/QueryMapperUnitTests.java | 58 +++- .../core/convert/SchemaFactoryUnitTests.java | 51 ++++ .../CassandraSimpleTypeHolderUnitTests.java | 2 +- .../repository/support/SchemaTestUtils.java | 9 +- .../ROOT/examples/VectorSearchExample.java | 41 +++ .../examples/mapping/PersonWithIndexes.java | 4 + .../ROOT/pages/cassandra/template.adoc | 53 +++- .../modules/ROOT/pages/object-mapping.adoc | 13 +- 32 files changed, 1307 insertions(+), 147 deletions(-) create mode 100644 spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/convert/CassandraVector.java rename spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/mapping/{SAIIndexed.java => SaiIndexed.java} (93%) create mode 100644 spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/mapping/SimilarityFunction.java create mode 100644 spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/mapping/VectorType.java create mode 100644 spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/query/VectorSort.java create mode 100644 spring-data-cassandra/src/test/java/org/springframework/data/cassandra/core/CassandraVectorSearchIntegrationTests.java create mode 100644 src/main/antora/modules/ROOT/examples/VectorSearchExample.java diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/StatementFactory.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/StatementFactory.java index 56b864d8d..4a1724f4e 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/StatementFactory.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/StatementFactory.java @@ -59,6 +59,7 @@ import org.springframework.data.cassandra.core.query.Update.RemoveOp; import org.springframework.data.cassandra.core.query.Update.SetAtIndexOp; import org.springframework.data.cassandra.core.query.Update.SetAtKeyOp; import org.springframework.data.cassandra.core.query.Update.SetOp; +import org.springframework.data.cassandra.core.query.VectorSort; import org.springframework.data.convert.EntityWriter; import org.springframework.data.domain.Sort; import org.springframework.data.mapping.PersistentProperty; @@ -72,6 +73,7 @@ import org.springframework.util.Assert; import org.springframework.util.ClassUtils; import com.datastax.oss.driver.api.core.CqlIdentifier; +import com.datastax.oss.driver.api.core.data.CqlVector; import com.datastax.oss.driver.api.core.metadata.schema.ClusteringOrder; import com.datastax.oss.driver.api.querybuilder.BindMarker; import com.datastax.oss.driver.api.querybuilder.QueryBuilder; @@ -688,24 +690,30 @@ public class StatementFactory { private StatementBuilder builder = StatementBuilder.of((Select) QueryBuilder.selectFrom(from), + cassandraConverter.getCodecRegistry()); - if (selectors.isEmpty()) { - select = QueryBuilder.selectFrom(getKeyspace(entity, from), from).all(); - } else { + builder.bind((statement, factory) -> { - List mappedSelectors = new ArrayList<>( - selectors.size()); - for (Selector selector : selectors) { - com.datastax.oss.driver.api.querybuilder.select.Selector orElseGet = selector.getAlias() - .map(it -> getSelection(selector).as(it)).orElseGet(() -> getSelection(selector)); - mappedSelectors.add(orElseGet); + Select select; + + if (selectors.isEmpty()) { + select = QueryBuilder.selectFrom(getKeyspace(entity, from), from).all(); + } else { + + List mappedSelectors = new ArrayList<>( + selectors.size()); + for (Selector selector : selectors) { + com.datastax.oss.driver.api.querybuilder.select.Selector orElseGet = selector.getAlias() + .map(it -> getSelection(selector, factory).as(it)).orElseGet(() -> getSelection(selector, factory)); + mappedSelectors.add(orElseGet); + } + + select = QueryBuilder.selectFrom(getKeyspace(entity, from), from).selectors(mappedSelectors); } - select = QueryBuilder.selectFrom(getKeyspace(entity, from), from).selectors(mappedSelectors); - } - - StatementBuilder select = statementFactory.select(Query.query(where("foo").in("bar")), - groupEntity); + StatementBuilder select = statementFactory.select(Query.query(where("foo").in("bar")), - groupEntity); + StatementBuilder select = statementFactory.select(Query.query(where("foo").in("bar")), - groupEntity); + StatementBuilder