From 0e1e594d13ec7434fe4a76b8249a2d7be36d81cf Mon Sep 17 00:00:00 2001 From: Mark Paluch Date: Wed, 16 Apr 2025 10:03:08 +0200 Subject: [PATCH] Explore returning Search Results. Closes #1573 --- .../cassandra/core/CassandraOperations.java | 29 +-- .../cassandra/core/CassandraTemplate.java | 3 +- .../data/cassandra/core/StatementFactory.java | 5 +- .../core/convert/CassandraConverters.java | 27 --- .../cassandra/core/convert/QueryMapper.java | 10 +- .../data/cassandra/core/query/ColumnName.java | 2 +- .../data/cassandra/core/query/Columns.java | 62 ++++-- .../core/query/SerializationUtils.java | 2 +- .../query/AbstractCassandraQuery.java | 13 +- .../query/AbstractReactiveCassandraQuery.java | 4 + .../repository/query/BindingContext.java | 62 ++++-- .../query/CassandraParameterAccessor.java | 24 ++ .../repository/query/CassandraParameters.java | 34 ++- .../CassandraParametersParameterAccessor.java | 39 ++++ .../query/CassandraQueryCreator.java | 58 ++++- .../query/CassandraQueryExecution.java | 22 +- .../query/ConvertingParameterAccessor.java | 29 +++ .../query/QueryStatementCreator.java | 64 +++++- .../ReactiveCassandraQueryExecution.java | 45 ++++ .../core/StatementFactoryUnitTests.java | 5 +- .../core/convert/QueryMapperUnitTests.java | 43 ++-- .../ReactiveVectorSearchIntegrationTests.java | 202 +++++++++++++++++ .../VectorSearchIntegrationTests.java | 207 ++++++++++++++++++ .../query/CassandraParametersUnitTests.java | 14 +- .../query/StubParameterAccessor.java | 29 +++ .../test/util/CassandraDelegate.java | 2 +- 26 files changed, 892 insertions(+), 144 deletions(-) create mode 100644 spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/ReactiveVectorSearchIntegrationTests.java create mode 100644 spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/VectorSearchIntegrationTests.java diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/CassandraOperations.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/CassandraOperations.java index 4edb52905..a7a970cd2 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/CassandraOperations.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/CassandraOperations.java @@ -20,6 +20,7 @@ import java.util.List; import java.util.stream.Stream; import org.jspecify.annotations.Nullable; + import org.springframework.dao.DataAccessException; import org.springframework.data.cassandra.core.convert.CassandraConverter; import org.springframework.data.cassandra.core.cql.CqlOperations; @@ -92,7 +93,7 @@ public interface CassandraOperations extends FluentCassandraOperations { /** * The table name used for the specified class by this template. * - * @param entityClass The entity type must not be {@literal null}. + * @param entityClass the entity type must not be {@literal null}. * @return the {@link CqlIdentifier} */ CqlIdentifier getTableName(Class entityClass); @@ -105,7 +106,7 @@ public interface CassandraOperations extends FluentCassandraOperations { * Execute a {@code SELECT} query and convert the resulting items to a {@link List} of entities. * * @param cql must not be {@literal null}. - * @param entityClass The entity type must not be {@literal null}. + * @param entityClass the entity type must not be {@literal null}. * @return the converted results * @throws DataAccessException if there is any problem executing the query. */ @@ -129,7 +130,7 @@ public interface CassandraOperations extends FluentCassandraOperations { * Execute a {@code SELECT} query and convert the resulting item to an entity. * * @param cql must not be {@literal null}. - * @param entityClass The entity type must not be {@literal null}. + * @param entityClass the entity type must not be {@literal null}. * @return the converted object or {@literal null}. * @throws DataAccessException if there is any problem executing the query. */ @@ -154,7 +155,7 @@ public interface CassandraOperations extends FluentCassandraOperations { * Execute a {@code SELECT} query and convert the resulting items to a {@link List} of entities. * * @param statement must not be {@literal null}. - * @param entityClass The entity type must not be {@literal null}. + * @param entityClass the entity type must not be {@literal null}. * @return the converted results * @throws DataAccessException if there is any problem executing the query. */ @@ -165,7 +166,7 @@ public interface CassandraOperations extends FluentCassandraOperations { * query translates the effective {@link Statement#getPageSize()} to the page size. * * @param statement the CQL statement, must not be {@literal null}. - * @param entityClass The entity type must not be {@literal null}. + * @param entityClass the entity type must not be {@literal null}. * @return the converted results * @throws DataAccessException if there is any problem executing the query. * @since 2.0 @@ -190,7 +191,7 @@ public interface CassandraOperations extends FluentCassandraOperations { * Execute a {@code SELECT} query and convert the resulting item to an entity. * * @param statement must not be {@literal null}. - * @param entityClass The entity type must not be {@literal null}. + * @param entityClass the entity type must not be {@literal null}. * @return the converted object or {@literal null}. * @throws DataAccessException if there is any problem executing the query. */ @@ -204,7 +205,7 @@ public interface CassandraOperations extends FluentCassandraOperations { * Execute a {@code SELECT} query and convert the resulting items to a {@link List} of entities. * * @param query must not be {@literal null}. - * @param entityClass The entity type must not be {@literal null}. + * @param entityClass the entity type must not be {@literal null}. * @return the converted results * @throws DataAccessException if there is any problem executing the query. * @since 2.0 @@ -215,7 +216,7 @@ public interface CassandraOperations extends FluentCassandraOperations { * Execute a {@code SELECT} query with paging and convert the result set to a {@link Slice} of entities. * * @param query the query object used to create a CQL statement, must not be {@literal null}. - * @param entityClass The entity type must not be {@literal null}. + * @param entityClass the entity type must not be {@literal null}. * @return the converted results * @throws DataAccessException if there is any problem executing the query. * @since 2.0 @@ -241,7 +242,7 @@ public interface CassandraOperations extends FluentCassandraOperations { * Execute a {@code SELECT} query and convert the resulting item to an entity. * * @param query must not be {@literal null}. - * @param entityClass The entity type must not be {@literal null}. + * @param entityClass the entity type must not be {@literal null}. * @return the converted object or {@literal null}. * @throws DataAccessException if there is any problem executing the query. * @since 2.0 @@ -253,7 +254,7 @@ public interface CassandraOperations extends FluentCassandraOperations { * * @param query must not be {@literal null}. * @param update must not be {@literal null}. - * @param entityClass The entity type must not be {@literal null}. + * @param entityClass the entity type must not be {@literal null}. * @throws DataAccessException if there is any problem executing the query. */ boolean update(Query query, Update update, Class entityClass) throws DataAccessException; @@ -262,7 +263,7 @@ public interface CassandraOperations extends FluentCassandraOperations { * Remove entities (rows)/columns from the table by {@link Query}. * * @param query must not be {@literal null}. - * @param entityClass The entity type must not be {@literal null}. + * @param entityClass the entity type must not be {@literal null}. * @throws DataAccessException if there is any problem executing the query. */ boolean delete(Query query, Class entityClass) throws DataAccessException; @@ -322,7 +323,7 @@ public interface CassandraOperations extends FluentCassandraOperations { * @param id the Id value. For single primary keys it's the plain value. For composite primary keys either the * {@link org.springframework.data.cassandra.core.mapping.PrimaryKeyClass} or * {@link org.springframework.data.cassandra.core.mapping.MapId}. Must not be {@literal null}. - * @param entityClass The entity type must not be {@literal null}. + * @param entityClass the entity type must not be {@literal null}. * @return the converted object or {@literal null}. * @throws DataAccessException if there is any problem executing the query. */ @@ -407,7 +408,7 @@ public interface CassandraOperations extends FluentCassandraOperations { * @param id the Id value. For single primary keys it's the plain value. For composite primary keys either the * {@link org.springframework.data.cassandra.core.mapping.PrimaryKeyClass} or * {@link org.springframework.data.cassandra.core.mapping.MapId}. Must not be {@literal null}. - * @param entityClass The entity type must not be {@literal null}. + * @param entityClass the entity type must not be {@literal null}. * @throws DataAccessException if there is any problem executing the query. */ boolean deleteById(Object id, Class entityClass) throws DataAccessException; @@ -415,7 +416,7 @@ public interface CassandraOperations extends FluentCassandraOperations { /** * Execute a {@code TRUNCATE} query to remove all entities of a given class. * - * @param entityClass The entity type must not be {@literal null}. + * @param entityClass the entity type must not be {@literal null}. * @throws DataAccessException if there is any problem executing the query. */ void truncate(Class entityClass) throws DataAccessException; diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/CassandraTemplate.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/CassandraTemplate.java index 300bc1c7b..163f37098 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/CassandraTemplate.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/CassandraTemplate.java @@ -354,7 +354,8 @@ public class CassandraTemplate implements CassandraOperations, ApplicationEventP Assert.notNull(statement, "Statement must not be null"); Assert.notNull(entityClass, "Entity type must not be null"); - return doSelect(statement, entityClass, getTableName(entityClass), entityClass, QueryResultConverter.entity()); + return doSelect(statement, entityClass, EntityQueryUtils.getTableName(statement), entityClass, + QueryResultConverter.entity()); } List doSelect(Statement statement, Class entityClass, CqlIdentifier tableName, Class returnType, 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 79dbfd6fc..cab05be15 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 @@ -767,12 +767,11 @@ public class StatementFactory { .stream().map(param -> { if (param instanceof ColumnSelector s) { - - return com.datastax.oss.driver.api.querybuilder.select.Selector.column(s.getExpression()); + return com.datastax.oss.driver.api.querybuilder.select.Selector.column(s.getIdentifier()); } if (param instanceof CqlIdentifier i) { - return com.datastax.oss.driver.api.querybuilder.select.Selector.column(i.toString()); + return com.datastax.oss.driver.api.querybuilder.select.Selector.column(i); } return new SimpleSelector(param.toString()); diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/convert/CassandraConverters.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/convert/CassandraConverters.java index c9049544d..c7738e069 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/convert/CassandraConverters.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/convert/CassandraConverters.java @@ -73,11 +73,9 @@ public abstract class CassandraConverters { converters.add(VectorToFloatArrayConverter.INSTANCE); converters.add(VectorToDoubleArrayConverter.INSTANCE); - converters.add(VectorToFloatListConverter.INSTANCE); converters.add(FloatArrayToVectorConverter.INSTANCE); converters.add(DoubleArrayToVectorConverter.INSTANCE); - converters.add(NumberListToVectorConverter.INSTANCE); converters.add(VectorToCqlVectorConverter.INSTANCE); converters.add(CqlVectorToVectorConverter.INSTANCE); @@ -307,17 +305,6 @@ public abstract class CassandraConverters { } } - @ReadingConverter - public enum NumberListToVectorConverter implements Converter, CqlVector> { - - INSTANCE; - - @Override - public CqlVector convert(List source) { - return CqlVector.newInstance(source); - } - } - @ReadingConverter public enum VectorToFloatArrayConverter implements Converter, float[]> { @@ -348,18 +335,4 @@ public abstract class CassandraConverters { } } - @ReadingConverter - public enum VectorToFloatListConverter implements Converter, List> { - - INSTANCE; - - @Override - public List convert(CqlVector source) { - List values = new ArrayList<>(source.size()); - for (int i = 0; i < source.size(); i++) { - values.add(source.get(i).floatValue()); - } - return values; - } - } } diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/convert/QueryMapper.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/convert/QueryMapper.java index 03d72d8f0..b95c71bf7 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/convert/QueryMapper.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/convert/QueryMapper.java @@ -199,7 +199,7 @@ public class QueryMapper { Field field = createPropertyField(entity, column); - columns.getSelector(column).ifPresent(selector -> { + columns.getSelector(column).forEach(selector -> { List mappedColumnNames = getCqlIdentifier(column, field); @@ -301,8 +301,12 @@ public class QueryMapper { Field field = createPropertyField(entity, column); field.getProperty().ifPresent(seen::add); - columns.getSelector(column).filter(selector -> selector instanceof ColumnSelector) - .ifPresent(columnSelector -> columnNames.addAll(getCqlIdentifier(column, field))); + columns.getSelector(column).forEach(columnSelector -> { + + if (columnSelector instanceof ColumnSelector) { + columnNames.addAll(getCqlIdentifier(column, field)); + } + }); } if (columns.isEmpty()) { diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/query/ColumnName.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/query/ColumnName.java index 3a3837b31..6fccd44c7 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/query/ColumnName.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/query/ColumnName.java @@ -169,7 +169,7 @@ public abstract class ColumnName { @Override public String toCql() { - return this.cqlIdentifier.toString(); + return this.cqlIdentifier.asInternal(); } @Override diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/query/Columns.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/query/Columns.java index a63cc5f28..9e57f6786 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/query/Columns.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/query/Columns.java @@ -15,6 +15,7 @@ */ package org.springframework.data.cassandra.core.query; +import java.util.ArrayList; import java.util.Arrays; import java.util.Collections; import java.util.Iterator; @@ -26,10 +27,11 @@ import java.util.Map.Entry; import java.util.Optional; import java.util.function.Function; +import org.jspecify.annotations.Nullable; + import org.springframework.data.cassandra.core.convert.CassandraVector; import org.springframework.data.cassandra.core.mapping.SimilarityFunction; import org.springframework.data.domain.Vector; -import org.jspecify.annotations.Nullable; import org.springframework.util.Assert; import org.springframework.util.ObjectUtils; import org.springframework.util.StringUtils; @@ -51,9 +53,9 @@ import com.datastax.oss.driver.api.core.data.CqlVector; */ public class Columns implements Iterable { - private final Map columns; + private final Map> columns; - private Columns(Map columns) { + private Columns(Map> columns) { this.columns = Collections.unmodifiableMap(columns); } @@ -76,10 +78,10 @@ public class Columns implements Iterable { Assert.notNull(columnNames, "Column names must not be null"); - Map columns = new LinkedHashMap<>(columnNames.length, 1); + Map> columns = new LinkedHashMap<>(columnNames.length, 1); for (String columnName : columnNames) { - columns.put(ColumnName.from(columnName), ColumnSelector.from(columnName)); + columns.put(ColumnName.from(columnName), new ArrayList<>(List.of(ColumnSelector.from(columnName)))); } return new Columns(columns); @@ -95,10 +97,10 @@ public class Columns implements Iterable { Assert.notNull(columnNames, "Column names must not be null"); - Map columns = new LinkedHashMap<>(columnNames.length, 1); + Map> columns = new LinkedHashMap<>(columnNames.length, 1); for (CqlIdentifier cqlId : columnNames) { - columns.put(ColumnName.from(cqlId), ColumnSelector.from(cqlId)); + columns.put(ColumnName.from(cqlId), new ArrayList<>(List.of(ColumnSelector.from(cqlId)))); } return new Columns(columns); @@ -209,8 +211,8 @@ public class Columns implements Iterable { */ private Columns select(ColumnName columnName, Selector selector) { - Map result = new LinkedHashMap<>(this.columns); - result.put(columnName, selector); + Map> result = new LinkedHashMap<>(this.columns); + result.computeIfAbsent(columnName, it -> new ArrayList<>()).add(selector); return new Columns(result); } @@ -231,9 +233,10 @@ public class Columns implements Iterable { */ public Columns and(Columns columns) { - Map result = new LinkedHashMap<>(this.columns); + Map> result = new LinkedHashMap<>(this.columns); - result.putAll(columns.columns); + columns.columns + .forEach((col, selectors) -> result.computeIfAbsent(col, columnName -> new ArrayList<>()).addAll(selectors)); return new Columns(result); } @@ -247,11 +250,12 @@ public class Columns implements Iterable { * @param columnName must not be {@literal null}. * @return the {@link Optional} {@link Selector} for {@link ColumnName}. */ - public Optional getSelector(ColumnName columnName) { + public List getSelector(ColumnName columnName) { Assert.notNull(columnName, "ColumnName must not be null"); - return Optional.ofNullable(this.columns.get(columnName)); + List selectors = this.columns.get(columnName); + return selectors == null ? List.of() : selectors; } @Override @@ -280,7 +284,7 @@ public class Columns implements Iterable { @Override public String toString() { - Iterator> iterator = this.columns.entrySet().iterator(); + Iterator>> iterator = this.columns.entrySet().iterator(); StringBuilder builder = toString(iterator); if (builder.isEmpty()) { @@ -290,24 +294,26 @@ public class Columns implements Iterable { return builder.toString(); } - private StringBuilder toString(Iterator> iterator) { + private StringBuilder toString(Iterator>> iterator) { StringBuilder builder = new StringBuilder(); boolean first = true; while (iterator.hasNext()) { - Entry entry = iterator.next(); + Entry> entry = iterator.next(); - Selector expression = entry.getValue(); + for (Selector selector : entry.getValue()) { - if (first) { - first = false; - } else { - builder.append(", "); + if (first) { + first = false; + } else { + builder.append(", "); + } + + builder.append(selector); } - builder.append(expression.toString()); } return builder; @@ -340,6 +346,9 @@ public class Columns implements Iterable { */ Selector as(CqlIdentifier alias); + /** + * @return the expression that forms this selection. + */ String getExpression(); Optional getAlias(); @@ -410,6 +419,10 @@ public class Columns implements Iterable { return alias; } + public CqlIdentifier getIdentifier() { + return columnName.getCqlIdentifier().orElseGet(() -> CqlIdentifier.fromCql(columnName.toCql())); + } + @Override public String getExpression() { return columnName.toCql(); @@ -662,8 +675,9 @@ public class Columns implements Iterable { @Override public Selector using(SimilarityFunction similarityFunction) { - return FunctionCall.from("similarity_" + similarityFunction.name().toLowerCase(Locale.ROOT), columnName, - vector).as(columnName); + return FunctionCall + .from("similarity_" + similarityFunction.name().toLowerCase(Locale.ROOT), columnName, vector) + .as(columnName); } }; } diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/query/SerializationUtils.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/query/SerializationUtils.java index 7d49a1b79..61373f73e 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/query/SerializationUtils.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/query/SerializationUtils.java @@ -57,7 +57,7 @@ abstract class SerializationUtils { CriteriaDefinition.Predicate predicate = criteria.getPredicate(); return String.format("%s %s", criteria.getColumnName(), - predicate.getOperator().toCql(serializeToCqlSafely(predicate.getValue()))); + predicate.getOperator().toCql(predicate != null ? serializeToCqlSafely(predicate.getValue()) : "")); } /** diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/AbstractCassandraQuery.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/AbstractCassandraQuery.java index a8d99be73..66adea760 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/AbstractCassandraQuery.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/AbstractCassandraQuery.java @@ -16,17 +16,10 @@ package org.springframework.data.cassandra.repository.query; import org.jspecify.annotations.Nullable; + import org.springframework.core.convert.converter.Converter; import org.springframework.data.cassandra.core.CassandraOperations; -import org.springframework.data.cassandra.repository.query.CassandraQueryExecution.CollectionExecution; -import org.springframework.data.cassandra.repository.query.CassandraQueryExecution.ExistsExecution; -import org.springframework.data.cassandra.repository.query.CassandraQueryExecution.ResultProcessingConverter; -import org.springframework.data.cassandra.repository.query.CassandraQueryExecution.ResultProcessingExecution; -import org.springframework.data.cassandra.repository.query.CassandraQueryExecution.ResultSetQuery; -import org.springframework.data.cassandra.repository.query.CassandraQueryExecution.SingleEntityExecution; -import org.springframework.data.cassandra.repository.query.CassandraQueryExecution.SlicedExecution; -import org.springframework.data.cassandra.repository.query.CassandraQueryExecution.StreamExecution; -import org.springframework.data.cassandra.repository.query.CassandraQueryExecution.WindowExecution; +import org.springframework.data.cassandra.repository.query.CassandraQueryExecution.*; import org.springframework.data.repository.query.ParameterAccessor; import org.springframework.data.repository.query.RepositoryQuery; import org.springframework.data.repository.query.ResultProcessor; @@ -120,6 +113,8 @@ public abstract class AbstractCassandraQuery extends CassandraRepositoryQuerySup return new SlicedExecution(getOperations(), parameterAccessor.getPageable()); } else if (getQueryMethod().isScrollQuery()) { return new WindowExecution(getOperations(), parameterAccessor.getScrollPosition(), parameterAccessor.getLimit()); + } else if (getQueryMethod().isSearchQuery()) { + return new SearchExecution(getOperations(), parameterAccessor); } else if (getQueryMethod().isCollectionQuery()) { return new CollectionExecution(getOperations()); } else if (getQueryMethod().isResultSetQuery()) { diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/AbstractReactiveCassandraQuery.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/AbstractReactiveCassandraQuery.java index 47478b7fd..1872474ab 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/AbstractReactiveCassandraQuery.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/AbstractReactiveCassandraQuery.java @@ -18,6 +18,7 @@ package org.springframework.data.cassandra.repository.query; import reactor.core.publisher.Mono; import org.reactivestreams.Publisher; + import org.springframework.core.convert.converter.Converter; import org.springframework.data.cassandra.ReactiveResultSet; import org.springframework.data.cassandra.core.CassandraOperations; @@ -26,6 +27,7 @@ import org.springframework.data.cassandra.repository.query.ReactiveCassandraQuer import org.springframework.data.cassandra.repository.query.ReactiveCassandraQueryExecution.ExistsExecution; import org.springframework.data.cassandra.repository.query.ReactiveCassandraQueryExecution.ResultProcessingConverter; import org.springframework.data.cassandra.repository.query.ReactiveCassandraQueryExecution.ResultProcessingExecution; +import org.springframework.data.cassandra.repository.query.ReactiveCassandraQueryExecution.SearchExecution; import org.springframework.data.cassandra.repository.query.ReactiveCassandraQueryExecution.SingleEntityExecution; import org.springframework.data.cassandra.repository.query.ReactiveCassandraQueryExecution.SlicedExecution; import org.springframework.data.cassandra.repository.query.ReactiveCassandraQueryExecution.WindowExecution; @@ -126,6 +128,8 @@ public abstract class AbstractReactiveCassandraQuery extends CassandraRepository } else if (getQueryMethod().isScrollQuery()) { return new WindowExecution(getReactiveCassandraOperations(), parameterAccessor.getScrollPosition(), parameterAccessor.getLimit()); + } else if (getQueryMethod().isSearchQuery()) { + return new SearchExecution(getReactiveCassandraOperations(), parameterAccessor); } else if (getQueryMethod().isCollectionQuery()) { return new CollectionExecution(getReactiveCassandraOperations()); } else if (isCountQuery()) { diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/BindingContext.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/BindingContext.java index 378795ab7..028f1bd0b 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/BindingContext.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/BindingContext.java @@ -20,9 +20,9 @@ import java.util.Collections; import java.util.List; import org.jspecify.annotations.Nullable; + +import org.springframework.data.domain.Limit; import org.springframework.data.mapping.model.ValueExpressionEvaluator; -import org.springframework.data.repository.query.Parameter; -import org.springframework.data.repository.query.ParameterAccessor; import org.springframework.util.Assert; /** @@ -35,7 +35,7 @@ class BindingContext { private final CassandraParameters parameters; - private final ParameterAccessor parameterAccessor; + private final CassandraParameterAccessor parameterAccessor; private final List bindings; @@ -44,7 +44,7 @@ class BindingContext { /** * Create new {@link BindingContext}. */ - BindingContext(CassandraParameters parameters, ParameterAccessor parameterAccessor, + BindingContext(CassandraParameters parameters, CassandraParameterAccessor parameterAccessor, List bindings, ValueExpressionEvaluator evaluator) { this.parameters = parameters; @@ -75,8 +75,9 @@ class BindingContext { List parameters = new ArrayList<>(bindings.size()); for (ParameterBinding binding : bindings) { + Object parameterValueForBinding = getParameterValueForBinding(binding); - parameters.add(parameterValueForBinding); + parameters.add(binding.prepareValue(parameterValueForBinding)); } return parameters; @@ -95,20 +96,20 @@ class BindingContext { } return binding.isNamed() - ? parameterAccessor.getBindableValue(getParameterIndex(parameters, binding.getRequiredParameterName())) + ? parameterAccessor.getValue(getParameterIndex(parameters, binding.getRequiredParameterName())) : parameterAccessor.getBindableValue(binding.getParameterIndex()); } private int getParameterIndex(CassandraParameters parameters, String parameterName) { - return parameters.stream() // - .filter(cassandraParameter -> cassandraParameter // - .getName().filter(s -> s.equals(parameterName)) // - .isPresent()) // - .mapToInt(Parameter::getIndex) // - .findFirst() // - .orElseThrow(() -> new IllegalArgumentException( - String.format("Invalid parameter name; Cannot resolve parameter [%s]", parameterName))); + for (CassandraParameters.CassandraParameter parameter : parameters) { + if (parameter.getName().filter(s -> s.equals(parameterName)).isPresent()) { + return parameter.getIndex(); + } + } + + throw new IllegalArgumentException( + String.format("Invalid parameter name; Cannot resolve parameter [%s]", parameterName)); } /** @@ -129,31 +130,31 @@ class BindingContext { this.parameterName = parameterName; } - static ParameterBinding expression(String expression, boolean quoted) { + public static ParameterBinding expression(String expression, boolean quoted) { return new ParameterBinding(-1, expression, null); } - static ParameterBinding indexed(int parameterIndex) { + public static ParameterBinding indexed(int parameterIndex) { return new ParameterBinding(parameterIndex, null, null); } - static ParameterBinding named(String name) { + public static ParameterBinding named(String name) { return new ParameterBinding(-1, null, name); } - boolean isNamed() { + public boolean isNamed() { return (parameterName != null); } - int getParameterIndex() { + public int getParameterIndex() { return parameterIndex; } - String getParameter() { + public String getParameter() { return ("?" + (isExpression() ? "expr" : "") + parameterIndex); } - String getRequiredExpression() { + public String getRequiredExpression() { Assert.state(expression != null, "ParameterBinding is not an expression"); return expression; @@ -169,5 +170,24 @@ class BindingContext { return parameterName; } + + /** + * Prepare a value before binding it to the query. + * + * @param value + * @return + */ + public @Nullable Object prepareValue(@Nullable Object value) { + + if (value == null) { + return value; + } + + if (value instanceof Limit limit) { + return limit.max(); + } + + return value; + } } } diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/CassandraParameterAccessor.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/CassandraParameterAccessor.java index 022de0a25..8d3c8a77e 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/CassandraParameterAccessor.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/CassandraParameterAccessor.java @@ -16,9 +16,11 @@ package org.springframework.data.cassandra.repository.query; import org.jspecify.annotations.Nullable; + import org.springframework.data.cassandra.core.cql.QueryOptions; import org.springframework.data.cassandra.core.mapping.CassandraType; import org.springframework.data.cassandra.core.query.CassandraScrollPosition; +import org.springframework.data.domain.ScoringFunction; import org.springframework.data.repository.query.ParameterAccessor; import com.datastax.oss.driver.api.core.type.DataType; @@ -65,6 +67,17 @@ public interface CassandraParameterAccessor extends ParameterAccessor { */ Class getParameterType(int index); + /** + * Get the value of the parameter at the given index. In contrast to {@link #getBindableValue(int)}, this method has + * access to all parameters. + * + * @param parameterIndex + * @return + * @since 5.0 + */ + @Nullable + Object getValue(int parameterIndex); + /** * Returns the raw parameter values of the underlying query method. * @@ -76,6 +89,16 @@ public interface CassandraParameterAccessor extends ParameterAccessor { @Override CassandraScrollPosition getScrollPosition(); + /** + * Returns the {@link ScoringFunction} from a {@link org.springframework.data.domain.Score} or + * {@link org.springframework.data.domain.Range} of scores if such a parameter is declared. + * + * @return the scoring function or {@literal null} if none is provided. + * @since 5.0 + */ + @Nullable + ScoringFunction getScoringFunction(); + /** * Returns the {@link QueryOptions} associated with the associated Repository query method. * @@ -85,4 +108,5 @@ public interface CassandraParameterAccessor extends ParameterAccessor { @Nullable QueryOptions getQueryOptions(); + } diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/CassandraParameters.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/CassandraParameters.java index bd9fc6c7c..da24b3b10 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/CassandraParameters.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/CassandraParameters.java @@ -29,6 +29,9 @@ import org.springframework.core.annotation.AnnotatedElementUtils; import org.springframework.data.cassandra.core.cql.QueryOptions; import org.springframework.data.cassandra.core.mapping.CassandraType; import org.springframework.data.cassandra.repository.query.CassandraParameters.CassandraParameter; +import org.springframework.data.domain.Range; +import org.springframework.data.domain.Score; +import org.springframework.data.domain.ScoringFunction; import org.springframework.data.repository.query.Parameter; import org.springframework.data.repository.query.Parameters; import org.springframework.data.repository.query.ParametersSource; @@ -46,6 +49,7 @@ import org.springframework.data.util.TypeInformation; public class CassandraParameters extends Parameters { private final @Nullable Integer queryOptionsIndex; + private final @Nullable Integer scoringFunctionIndex; /** * Create a new {@link CassandraParameters} instance from the given {@link Method}. @@ -58,18 +62,23 @@ public class CassandraParameters extends Parameters originals, @Nullable Integer queryOptionsIndex) { + private CassandraParameters(List originals, @Nullable Integer queryOptionsIndex, + @Nullable Integer scoringFunctionIndex) { super(originals); this.queryOptionsIndex = queryOptionsIndex; + this.scoringFunctionIndex = scoringFunctionIndex; } @Override protected CassandraParameters createFrom(List parameters) { - return new CassandraParameters(parameters, queryOptionsIndex); + return new CassandraParameters(parameters, queryOptionsIndex, scoringFunctionIndex); } /** @@ -82,6 +91,16 @@ public class CassandraParameters extends Parameters parameterType; + private final boolean isScoreRange; + private final boolean isScoringFunction; CassandraParameter(MethodParameter parameter, TypeInformation domainType) { @@ -104,12 +125,17 @@ public class CassandraParameters extends Parameters range = getScoreRange(); + + if (range != null) { + + if (range.getLowerBound().isBounded()) { + return range.getLowerBound().getValue().get().getFunction(); + } + + if (range.getUpperBound().isBounded()) { + return range.getUpperBound().getValue().get().getFunction(); + } + } + + int scoringFunctionIndex = getParameters().getScoringFunctionIndex(); + if (scoringFunctionIndex != -1) { + return (ScoringFunction) getValue(scoringFunctionIndex); + } + + return null; + } + @Override public Limit getLimit() { diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/CassandraQueryCreator.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/CassandraQueryCreator.java index e475a285d..eebb31904 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/CassandraQueryCreator.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/CassandraQueryCreator.java @@ -32,9 +32,12 @@ import org.springframework.data.cassandra.core.query.Criteria; import org.springframework.data.cassandra.core.query.CriteriaDefinition; import org.springframework.data.cassandra.core.query.Filter; import org.springframework.data.cassandra.core.query.Query; +import org.springframework.data.cassandra.core.query.VectorSort; import org.springframework.data.domain.Range; import org.springframework.data.domain.Sort; +import org.springframework.data.domain.Vector; import org.springframework.data.mapping.PersistentPropertyPath; +import org.springframework.data.mapping.PropertyPath; import org.springframework.data.mapping.context.MappingContext; import org.springframework.data.repository.query.parser.AbstractQueryCreator; import org.springframework.data.repository.query.parser.Part; @@ -56,6 +59,8 @@ class CassandraQueryCreator extends AbstractQueryCreator { private final MappingContext mappingContext; private final QueryBuilder queryBuilder = new QueryBuilder(); + private final CassandraParameterAccessor parameterAccessor; + private final PartTree tree; /** * Create a new {@link CassandraQueryCreator} from the given {@link PartTree}, {@link ConvertingParameterAccessor} and @@ -72,6 +77,8 @@ class CassandraQueryCreator extends AbstractQueryCreator { Assert.notNull(mappingContext, "CassandraMappingContext must not be null"); + this.tree = tree; + this.parameterAccessor = parameterAccessor; this.mappingContext = mappingContext; } @@ -102,9 +109,6 @@ class CassandraQueryCreator extends AbstractQueryCreator { .getPersistentPropertyPath(part.getProperty()); CassandraPersistentProperty property = path.getLeafProperty(); - - Assert.state(property != null && path.toDotPath() != null, "Leaf property must not be null"); - Object filterOrCriteria = from(part, property, Criteria.where(path.toDotPath()), iterator); if (filterOrCriteria instanceof CriteriaDefinition) { @@ -115,10 +119,12 @@ class CassandraQueryCreator extends AbstractQueryCreator { } @Override - protected Filter and(Part part, Filter base, Iterator iterator) { + protected Filter and(Part part, @Nullable Filter base, Iterator iterator) { - for (CriteriaDefinition criterion : base) { - getQueryBuilder().and(criterion); + if (base != null) { + for (CriteriaDefinition criterion : base) { + getQueryBuilder().and(criterion); + } } return create(part, iterator); @@ -139,7 +145,8 @@ class CassandraQueryCreator extends AbstractQueryCreator { } } - Query query = getQueryBuilder().create(sort); + Query query = sort.isUnsorted() && parameterAccessor.getVector() != null ? getQueryBuilder().create(getVectorSort()) + : getQueryBuilder().create(sort); if (LOG.isDebugEnabled()) { LOG.debug(String.format("Created query [%s]", query)); @@ -148,10 +155,29 @@ class CassandraQueryCreator extends AbstractQueryCreator { return query; } + private Sort getVectorSort() { + return VectorSort.ann(getVectorProperty().toDotPath(), parameterAccessor.getVector()); + } + + PropertyPath getVectorProperty() { + + for (PartTree.OrPart parts : tree) { + for (Part part : parts) { + + if (part.getType() == Type.NEAR || part.getType() == Type.WITHIN) { + return part.getProperty(); + } + } + } + + throw new IllegalArgumentException("No Near/Within property found"); + } + /** * Returns a {@link Filter} or {@link CriteriaDefinition} object representing the criterion for a {@link Part}. */ - private Object from(Part part, CassandraPersistentProperty property, Criteria where, Iterator parameters) { + private @Nullable Object from(Part part, CassandraPersistentProperty property, Criteria where, + Iterator parameters) { Type type = part.getType(); @@ -182,10 +208,24 @@ class CassandraQueryCreator extends AbstractQueryCreator { return where.is(false); case SIMPLE_PROPERTY: return where.is(parameters.next()); + + case NEAR: + case WITHIN: + + Object next = parameters.next(); + + if (!(next instanceof Vector)) { + + throw new IllegalArgumentException("Expected a Vector for Near/Within keyword but got [%s]" + .formatted(next == null ? "null" : next.getClass())); + } + + return null; default: throw new InvalidDataAccessApiUsageException( String.format("Unsupported keyword [%s] in part [%s]", type, part)); } + } /** @@ -280,7 +320,7 @@ class CassandraQueryCreator extends AbstractQueryCreator { */ static class QueryBuilder { - private List criterias = new ArrayList<>(); + private final List criterias = new ArrayList<>(); CriteriaDefinition and(CriteriaDefinition clause) { criterias.add(clause); diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/CassandraQueryExecution.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/CassandraQueryExecution.java index 5cd99b883..151c79f47 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/CassandraQueryExecution.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/CassandraQueryExecution.java @@ -18,6 +18,7 @@ package org.springframework.data.cassandra.repository.query; import java.util.List; import org.jspecify.annotations.Nullable; + import org.springframework.core.convert.converter.Converter; import org.springframework.dao.IncorrectResultSizeDataAccessException; import org.springframework.data.cassandra.core.CassandraOperations; @@ -28,6 +29,11 @@ import org.springframework.data.cassandra.core.query.CassandraScrollPosition; import org.springframework.data.convert.DtoInstantiatingConverter; import org.springframework.data.domain.Limit; import org.springframework.data.domain.Pageable; +import org.springframework.data.domain.Score; +import org.springframework.data.domain.ScoringFunction; +import org.springframework.data.domain.SearchResult; +import org.springframework.data.domain.SearchResults; +import org.springframework.data.domain.Similarity; import org.springframework.data.domain.Slice; import org.springframework.data.domain.SliceImpl; import org.springframework.data.mapping.context.MappingContext; @@ -180,8 +186,9 @@ interface CassandraQueryExecution { ScoringFunction function = accessor.getScoringFunction(); - List> results = operations.select(statement, type, (o, row) -> { + List> results = operations.query(statement).as(type).map((row, reader) -> { + Object o = reader.get(); if (row.getColumnDefinitions().contains("__score__")) { return new SearchResult<>(o, getScore(row, "__score__", function)); } @@ -189,17 +196,20 @@ interface CassandraQueryExecution { if (row.getColumnDefinitions().contains("score")) { return new SearchResult<>(o, getScore(row, "score", function)); } - return new SearchResult<>(o, 0); - }); - return new SearchResults(results); + return new SearchResult<>(o, Similarity.of(0)); + }).all(); + + return new SearchResults<>(results); } - private Score getScore(Row row, String columnName, ScoringFunction function) { + private Score getScore(Row row, String columnName, @Nullable ScoringFunction function) { Object object = row.getObject(columnName); - return Score.of(((Number) object).doubleValue(), function); + return Similarity.raw(((Number) object).doubleValue(), + function == null ? ScoringFunction.unspecified() : function); } + } /** diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/ConvertingParameterAccessor.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/ConvertingParameterAccessor.java index 003de8677..53f539826 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/ConvertingParameterAccessor.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/ConvertingParameterAccessor.java @@ -18,6 +18,7 @@ package org.springframework.data.cassandra.repository.query; import java.util.Iterator; import org.jspecify.annotations.Nullable; + import org.springframework.data.cassandra.core.convert.CassandraConverter; import org.springframework.data.cassandra.core.cql.QueryOptions; import org.springframework.data.cassandra.core.mapping.CassandraType; @@ -25,7 +26,10 @@ import org.springframework.data.cassandra.core.query.CassandraScrollPosition; import org.springframework.data.domain.Limit; import org.springframework.data.domain.Pageable; import org.springframework.data.domain.Range; +import org.springframework.data.domain.Score; +import org.springframework.data.domain.ScoringFunction; import org.springframework.data.domain.Sort; +import org.springframework.data.domain.Vector; import com.datastax.oss.driver.api.core.type.DataType; @@ -54,6 +58,26 @@ class ConvertingParameterAccessor implements CassandraParameterAccessor { return delegate.getScrollPosition(); } + @Override + public ScoringFunction getScoringFunction() { + return delegate.getScoringFunction(); + } + + @Override + public @Nullable Vector getVector() { + return delegate.getVector(); + } + + @Override + public @Nullable Score getScore() { + return delegate.getScore(); + } + + @Override + public @Nullable Range getScoreRange() { + return delegate.getScoreRange(); + } + @Override public Pageable getPageable() { return this.delegate.getPageable(); @@ -113,6 +137,11 @@ class ConvertingParameterAccessor implements CassandraParameterAccessor { return this.delegate.getValues(); } + @Override + public @Nullable Object getValue(int parameterIndex) { + return potentiallyConvert(parameterIndex, this.delegate.getValue(parameterIndex)); + } + @Nullable Object potentiallyConvert(int index, @Nullable Object bindableValue) { diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/QueryStatementCreator.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/QueryStatementCreator.java index e023eea7b..0b868264a 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/QueryStatementCreator.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/QueryStatementCreator.java @@ -15,11 +15,16 @@ */ package org.springframework.data.cassandra.repository.query; +import java.util.ArrayList; +import java.util.List; +import java.util.Map; import java.util.Optional; import java.util.function.Function; import org.apache.commons.logging.Log; import org.apache.commons.logging.LogFactory; +import org.jspecify.annotations.Nullable; + import org.springframework.data.cassandra.core.StatementFactory; import org.springframework.data.cassandra.core.cql.QueryExtractorDelegate; import org.springframework.data.cassandra.core.cql.QueryOptions; @@ -27,10 +32,15 @@ import org.springframework.data.cassandra.core.cql.QueryOptions.QueryOptionsBuil import org.springframework.data.cassandra.core.cql.QueryOptionsUtil; import org.springframework.data.cassandra.core.mapping.CassandraPersistentEntity; import org.springframework.data.cassandra.core.mapping.CassandraPersistentProperty; +import org.springframework.data.cassandra.core.mapping.SimilarityFunction; import org.springframework.data.cassandra.core.query.Columns; import org.springframework.data.cassandra.core.query.Query; import org.springframework.data.cassandra.repository.Query.Idempotency; import org.springframework.data.domain.Limit; +import org.springframework.data.domain.ScoringFunction; +import org.springframework.data.domain.Vector; +import org.springframework.data.domain.VectorScoringFunctions; +import org.springframework.data.mapping.PropertyPath; import org.springframework.data.mapping.context.MappingContext; import org.springframework.data.mapping.model.ValueExpressionEvaluator; import org.springframework.data.repository.query.QueryCreationException; @@ -50,6 +60,11 @@ import com.datastax.oss.driver.api.core.cql.Statement; */ class QueryStatementCreator { + private static final Map SIMILARITY_FUNCTIONS = Map.of( + VectorScoringFunctions.COSINE, SimilarityFunction.COSINE, // + VectorScoringFunctions.EUCLIDEAN, SimilarityFunction.EUCLIDEAN, // + VectorScoringFunctions.DOT_PRODUCT, SimilarityFunction.DOT_PRODUCT); + private static final Log LOG = LogFactory.getLog(QueryStatementCreator.class); private final CassandraQueryMethod queryMethod; @@ -81,9 +96,28 @@ class QueryStatementCreator { ReturnedType returnedType = processor.withDynamicProjection(parameterAccessor).getReturnedType(); + Columns columns = null; if (returnedType.needsCustomConstruction()) { + columns = Columns.from(returnedType.getInputProperties().toArray(new String[0])); + } else if (queryMethod.isSearchQuery()) { + columns = getColumns(returnedType.getReturnedType()); + } + + if (columns != null) { + + if (queryMethod.isSearchQuery()) { + + CassandraQueryCreator queryCreator = new CassandraQueryCreator(tree, parameterAccessor, this.mappingContext); + + PropertyPath vectorProperty = queryCreator.getVectorProperty(); + Vector vector = parameterAccessor.getVector(); + SimilarityFunction similarityFunction = getSimilarityFunction(parameterAccessor.getScoringFunction()); + + columns = columns.select(vectorProperty.toDotPath(), + selectorBuilder -> selectorBuilder.similarity(vector).using(similarityFunction).as("\"__score__\"")); + } + - Columns columns = Columns.from(returnedType.getInputProperties().toArray(new String[0])); query = query.columns(columns); } @@ -99,6 +133,34 @@ class QueryStatementCreator { return doWithQuery(parameterAccessor, tree, function); } + private Columns getColumns(Class returnedType) { + + CassandraPersistentEntity entity = mappingContext.getRequiredPersistentEntity(returnedType); + List names = new ArrayList<>(); + for (CassandraPersistentProperty property : entity) { + names.add(property.getName()); + } + + return Columns.from(names.toArray(new String[0])); + } + + private SimilarityFunction getSimilarityFunction(@Nullable ScoringFunction function) { + + if (function == null) { + throw new IllegalStateException( + "Cannot determine ScoringFunction. No ScoringFunction, Score/Similarity or bounded Score Range parameters provided."); + } + + SimilarityFunction similarityFunction = SIMILARITY_FUNCTIONS.get(function); + + if (similarityFunction == null) { + throw new IllegalArgumentException( + "Cannot determine SimilarityFunction from ScoreFunction '%s'".formatted(function)); + } + + return similarityFunction; + } + /** * Create a {@literal COUNT} {@link Statement} from a {@link PartTree} and apply query options. * diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/ReactiveCassandraQueryExecution.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/ReactiveCassandraQueryExecution.java index fe4476e54..ec155af10 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/ReactiveCassandraQueryExecution.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/query/ReactiveCassandraQueryExecution.java @@ -20,7 +20,9 @@ import reactor.core.publisher.Mono; import java.util.List; +import org.jspecify.annotations.Nullable; import org.reactivestreams.Publisher; + import org.springframework.core.convert.converter.Converter; import org.springframework.dao.IncorrectResultSizeDataAccessException; import org.springframework.data.cassandra.core.ReactiveCassandraOperations; @@ -31,6 +33,10 @@ import org.springframework.data.cassandra.core.query.CassandraScrollPosition; import org.springframework.data.convert.DtoInstantiatingConverter; import org.springframework.data.domain.Limit; import org.springframework.data.domain.Pageable; +import org.springframework.data.domain.Score; +import org.springframework.data.domain.ScoringFunction; +import org.springframework.data.domain.SearchResult; +import org.springframework.data.domain.Similarity; import org.springframework.data.domain.Slice; import org.springframework.data.domain.SliceImpl; import org.springframework.data.mapping.context.MappingContext; @@ -152,6 +158,45 @@ interface ReactiveCassandraQueryExecution { } + final class SearchExecution implements ReactiveCassandraQueryExecution { + + private final ReactiveCassandraOperations operations; + private final CassandraParameterAccessor accessor; + + public SearchExecution(ReactiveCassandraOperations operations, CassandraParameterAccessor accessor) { + + this.operations = operations; + this.accessor = accessor; + } + + @Override + public Publisher execute(Statement statement, Class type) { + + ScoringFunction function = accessor.getScoringFunction(); + + return operations.query(statement).as(type).map((row, reader) -> { + + Object entity = reader.get(); + if (row.getColumnDefinitions().contains("__score__")) { + return new SearchResult<>(entity, getScore(row, "__score__", function)); + } + + if (row.getColumnDefinitions().contains("score")) { + return new SearchResult<>(entity, getScore(row, "score", function)); + } + return new SearchResult<>(entity, 0); + }).all(); + } + + private Score getScore(Row row, String columnName, @Nullable ScoringFunction function) { + + Object object = row.getObject(columnName); + return Similarity.raw(((Number) object).doubleValue(), + function == null ? ScoringFunction.unspecified() : function); + } + + } + /** * {@link ReactiveCassandraQueryExecution} to return a single entity. * diff --git a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/core/StatementFactoryUnitTests.java b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/core/StatementFactoryUnitTests.java index 20dbc69bc..f3c573392 100644 --- a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/core/StatementFactoryUnitTests.java +++ b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/core/StatementFactoryUnitTests.java @@ -921,16 +921,14 @@ class StatementFactoryUnitTests { withVector.id = "foo"; withVector.vector = CqlVector.newInstance(1.2f, 1.3f); withVector.array = new float[] { 2.2f, 2.3f }; - withVector.list = Arrays.asList(3.2f, 3.3f); SimpleStatement statement = statementFactory.update(withVector, WriteOptions.empty()) .build(ParameterHandling.BY_NAME); - assertThat(statement.getQuery()).isEqualTo("UPDATE withvector SET vector=:p0, array=:p1, list=:p2 WHERE id=:p3"); + assertThat(statement.getQuery()).isEqualTo("UPDATE withvector SET vector=:p0, array=:p1 WHERE id=:p2"); assertThat(statement.getNamedValues().get(CqlIdentifier.fromCql("p0"))).isInstanceOf(CqlVector.class).hasToString("[1.2, 1.3]"); assertThat(statement.getNamedValues().get(CqlIdentifier.fromCql("p1"))).isInstanceOf(CqlVector.class).hasToString("[2.2, 2.3]"); - assertThat(statement.getNamedValues().get(CqlIdentifier.fromCql("p2"))).isInstanceOf(CqlVector.class).hasToString("[3.2, 3.3]"); } @Test // GH-1504 @@ -990,6 +988,5 @@ class StatementFactoryUnitTests { @VectorType(dimensions = 12) CqlVector vector; @VectorType(dimensions = 12) float[] array; - @VectorType(dimensions = 12) List list; } } diff --git a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/core/convert/QueryMapperUnitTests.java b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/core/convert/QueryMapperUnitTests.java index 65ae478f4..2e60d9cb2 100644 --- a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/core/convert/QueryMapperUnitTests.java +++ b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/core/convert/QueryMapperUnitTests.java @@ -21,7 +21,6 @@ import static org.springframework.data.domain.Sort.Order.*; import java.time.LocalDate; import java.time.LocalTime; -import java.util.Arrays; import java.util.Collection; import java.util.Collections; import java.util.Currency; @@ -65,6 +64,7 @@ import org.springframework.data.convert.ValueConverter; import org.springframework.data.domain.Sort; import org.springframework.data.domain.Sort.Direction; import org.springframework.data.domain.Sort.Order; +import org.springframework.data.domain.Vector; import org.springframework.lang.Nullable; import com.datastax.oss.driver.api.core.CqlIdentifier; @@ -368,6 +368,31 @@ public class QueryMapperUnitTests { assertThat(mappedObject).contains(CqlIdentifier.fromCql("first_name")); } + @Test // + void shouldMapMultipleColumnNames() { + + Columns columnNames = Columns.from("array").select("array", + selectorBuilder -> selectorBuilder.similarity(Vector.of(1, 2)).cosine().as("score")); + + List mappedObject = queryMapper.getMappedColumnNames(columnNames, + mappingContext.getRequiredPersistentEntity(WithVector.class)); + + assertThat(mappedObject).contains(CqlIdentifier.fromCql("array")); + } + + @Test // + void shouldMapMultipleSelectorsNames() { + + Columns columnNames = Columns.from("array").select("array", + selectorBuilder -> selectorBuilder.similarity(Vector.of(1, 2)).cosine().as("score")); + + List mappedObject = queryMapper.getMappedSelectors(columnNames, + mappingContext.getRequiredPersistentEntity(WithVector.class)); + + assertThat(mappedObject).extracting(Selector::toString).contains("array", + "similarity_cosine(array, [1.0, 2.0]) AS score"); + } + @Test // DATACASS-523 @SuppressWarnings("all") void shouldMapTuple() { @@ -440,7 +465,7 @@ public class QueryMapperUnitTests { assertThat(mappedObject.iterator().next().getPredicate().getValue()).isEqualTo(42L); } - @Test // + @Test // GH-1504 void shouldConvertVectorValues() { Filter filter = Filter.from(Criteria.where("array").is(new float[] { 1.1f, 2.2f })); @@ -452,23 +477,13 @@ public class QueryMapperUnitTests { assertThat(mappedObject.iterator().next().getPredicate().getValue()).isEqualTo(new float[] { 1.1f, 2.2f }); } - @Test // GH-1504 - void shouldConvertVectorValuesFromList() { - - Filter filter = Filter.from(Criteria.where("list").is(Arrays.asList(1.1f, 2.2f))); - - Filter mappedObject = this.queryMapper.getMappedObject(filter, - this.mappingContext.getRequiredPersistentEntity(WithVector.class)); - - assertThat(mappedObject.iterator().next().getColumnName()).isEqualTo(ColumnName.from("list")); - assertThat(mappedObject.iterator().next().getPredicate().getValue()).isEqualTo(CqlVector.newInstance(1.1f, 2.2f)); - } @Test // GH-1504 void shouldConvertVectorSelectorFunction() { Columns columns = Columns.empty(); - Columns.FunctionCall similarity = Columns.FunctionCall.from("similarity_cosine", CqlIdentifier.fromCql("array"), Arrays.asList(1.1f, 2.2f)); + Columns.FunctionCall similarity = Columns.FunctionCall.from("similarity_cosine", CqlIdentifier.fromCql("array"), + CqlVector.newInstance(1.1f, 2.2f)); Query query = Query.empty().columns(columns.select("array", similarity)); diff --git a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/ReactiveVectorSearchIntegrationTests.java b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/ReactiveVectorSearchIntegrationTests.java new file mode 100644 index 000000000..cb2471629 --- /dev/null +++ b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/ReactiveVectorSearchIntegrationTests.java @@ -0,0 +1,202 @@ +/* + * Copyright 2016-2025 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.data.cassandra.repository; + +import static org.assertj.core.api.Assertions.*; + +import reactor.core.publisher.Flux; +import reactor.test.StepVerifier; + +import java.util.Collections; +import java.util.List; +import java.util.Set; +import java.util.UUID; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; + +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.context.annotation.ComponentScan; +import org.springframework.context.annotation.Configuration; +import org.springframework.context.annotation.FilterType; +import org.springframework.data.annotation.Id; +import org.springframework.data.annotation.PersistenceCreator; +import org.springframework.data.cassandra.config.SchemaAction; +import org.springframework.data.cassandra.core.mapping.SaiIndexed; +import org.springframework.data.cassandra.core.mapping.Table; +import org.springframework.data.cassandra.core.mapping.VectorType; +import org.springframework.data.cassandra.repository.config.EnableReactiveCassandraRepositories; +import org.springframework.data.cassandra.repository.support.AbstractSpringDataEmbeddedCassandraIntegrationTest; +import org.springframework.data.cassandra.repository.support.IntegrationTestConfig; +import org.springframework.data.domain.Limit; +import org.springframework.data.domain.ScoringFunction; +import org.springframework.data.domain.SearchResult; +import org.springframework.data.domain.Similarity; +import org.springframework.data.domain.Vector; +import org.springframework.data.domain.VectorScoringFunctions; +import org.springframework.data.repository.reactive.ReactiveCrudRepository; +import org.springframework.test.context.junit.jupiter.SpringJUnitConfig; + +/** + * Integration tests for Vector Search using reactive repositories. + * + * @author Mark Paluch + */ +@SpringJUnitConfig +class ReactiveVectorSearchIntegrationTests extends AbstractSpringDataEmbeddedCassandraIntegrationTest { + + Vector VECTOR = Vector.of(0.2001f, 0.32345f, 0.43456f, 0.54567f, 0.65678f); + + @Configuration + @EnableReactiveCassandraRepositories(basePackageClasses = ReactiveVectorSearchRepository.class, + considerNestedRepositories = true, + includeFilters = @ComponentScan.Filter(classes = ReactiveVectorSearchRepository.class, + type = FilterType.ASSIGNABLE_TYPE)) + public static class Config extends IntegrationTestConfig { + + @Override + protected Set> getInitialEntitySet() { + return Collections.singleton(WithVectorFields.class); + } + + @Override + public SchemaAction getSchemaAction() { + return SchemaAction.RECREATE_DROP_UNUSED; + } + } + + @Autowired ReactiveVectorSearchRepository repository; + + @BeforeEach + void setUp() { + + repository.deleteAll().as(StepVerifier::create).verifyComplete(); + + WithVectorFields w1 = new WithVectorFields("de", "one", Vector.of(0.1001f, 0.22345f, 0.33456f, 0.44567f, 0.55678f)); + WithVectorFields w2 = new WithVectorFields("de", "two", Vector.of(0.2001f, 0.32345f, 0.43456f, 0.54567f, 0.65678f)); + WithVectorFields w3 = new WithVectorFields("en", "three", + Vector.of(0.9001f, 0.82345f, 0.73456f, 0.64567f, 0.55678f)); + WithVectorFields w4 = new WithVectorFields("de", "four", + Vector.of(0.9001f, 0.92345f, 0.93456f, 0.94567f, 0.95678f)); + + repository.saveAll(List.of(w1, w2, w3, w4)).as(StepVerifier::create).expectNextCount(4).verifyComplete(); + } + + @Test // GH- + void shouldConsiderScoringFunction() { + + Vector vector = Vector.of(0.9f, 0.54f, 0.12f, 0.1f, 0.95f); + + List> results = repository + .searchByEmbeddingNear(vector, VectorScoringFunctions.COSINE, Limit.of(100)).collectList().block(); + + assertThat(results).hasSize(4); + for (SearchResult result : results) { + assertThat(result.getScore()).isInstanceOf(Similarity.class); + assertThat(result.getScore().getValue()).isNotCloseTo(0d, offset(0.1d)); + } + + results = repository.searchByEmbeddingNear(VECTOR, VectorScoringFunctions.EUCLIDEAN, Limit.of(100)).collectList() + .block(); + + assertThat(results).hasSize(4); + for (SearchResult result : results) { + assertThat(result.getScore()).isInstanceOf(Similarity.class); + assertThat(result.getScore().getValue()).isNotCloseTo(0.3d, offset(0.1d)); + } + } + + @Test // GH- + void shouldRunAnnotatedSearchByVector() { + + List> results = repository.searchAnnotatedByEmbeddingNear(VECTOR, Limit.of(100)) + .collectList().block(); + + assertThat(results).hasSize(4); + for (SearchResult result : results) { + assertThat(result.getScore()).isInstanceOf(Similarity.class); + assertThat(result.getScore().getValue()).isNotCloseTo(0d, offset(0.1d)); + } + } + + @Test // GH- + void shouldFindByVector() { + + List result = repository.findByEmbeddingNear(VECTOR, Limit.of(100)).collectList().block(); + + assertThat(result).hasSize(4); + } + + interface ReactiveVectorSearchRepository extends ReactiveCrudRepository { + + Flux> searchByEmbeddingNear(Vector embedding, ScoringFunction function, Limit limit); + + Flux findByEmbeddingNear(Vector embedding, Limit limit); + + @Query("SELECT id,description,country,similarity_cosine(embedding,:embedding) AS score FROM withvectorfields ORDER BY embedding ANN OF :embedding LIMIT :limit") + Flux> searchAnnotatedByEmbeddingNear(Vector embedding, Limit limit); + + } + + @Table + static class WithVectorFields { + + @Id String id; + String country; + String description; + + @VectorType(dimensions = 5) + @SaiIndexed Vector embedding; + + @PersistenceCreator + public WithVectorFields(String id, String country, String description, Vector embedding) { + this.id = id; + this.country = country; + this.description = description; + this.embedding = embedding; + } + + public WithVectorFields(String country, String description, Vector embedding) { + this.id = UUID.randomUUID().toString(); + this.country = country; + this.description = description; + this.embedding = embedding; + } + + public String getId() { + return id; + } + + public String getCountry() { + return country; + } + + public String getDescription() { + return description; + } + + public Vector getEmbedding() { + return embedding; + } + + @Override + public String toString() { + return "WithVectorFields{" + "id='" + id + '\'' + ", country='" + country + '\'' + ", description='" + description + + '\'' + '}'; + } + } + +} diff --git a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/VectorSearchIntegrationTests.java b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/VectorSearchIntegrationTests.java new file mode 100644 index 000000000..1fd73c71f --- /dev/null +++ b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/VectorSearchIntegrationTests.java @@ -0,0 +1,207 @@ +/* + * Copyright 2016-2025 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.data.cassandra.repository; + +import static org.assertj.core.api.Assertions.*; + +import java.util.Collections; +import java.util.List; +import java.util.Set; +import java.util.UUID; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; + +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.context.annotation.ComponentScan; +import org.springframework.context.annotation.Configuration; +import org.springframework.context.annotation.FilterType; +import org.springframework.data.annotation.Id; +import org.springframework.data.annotation.PersistenceCreator; +import org.springframework.data.cassandra.config.SchemaAction; +import org.springframework.data.cassandra.core.mapping.SaiIndexed; +import org.springframework.data.cassandra.core.mapping.Table; +import org.springframework.data.cassandra.core.mapping.VectorType; +import org.springframework.data.cassandra.repository.config.EnableCassandraRepositories; +import org.springframework.data.cassandra.repository.support.AbstractSpringDataEmbeddedCassandraIntegrationTest; +import org.springframework.data.cassandra.repository.support.IntegrationTestConfig; +import org.springframework.data.domain.Limit; +import org.springframework.data.domain.ScoringFunction; +import org.springframework.data.domain.SearchResult; +import org.springframework.data.domain.SearchResults; +import org.springframework.data.domain.Similarity; +import org.springframework.data.domain.Vector; +import org.springframework.data.domain.VectorScoringFunctions; +import org.springframework.data.repository.CrudRepository; +import org.springframework.data.repository.query.QueryCreationException; +import org.springframework.test.context.junit.jupiter.SpringJUnitConfig; + +/** + * Integration tests for Vector Search using repositories. + * + * @author Mark Paluch + */ +@SpringJUnitConfig +class VectorSearchIntegrationTests extends AbstractSpringDataEmbeddedCassandraIntegrationTest { + + Vector VECTOR = Vector.of(0.2001f, 0.32345f, 0.43456f, 0.54567f, 0.65678f); + + @Configuration + @EnableCassandraRepositories(basePackageClasses = VectorSearchRepository.class, considerNestedRepositories = true, + includeFilters = @ComponentScan.Filter(classes = VectorSearchRepository.class, type = FilterType.ASSIGNABLE_TYPE)) + public static class Config extends IntegrationTestConfig { + + @Override + protected Set> getInitialEntitySet() { + return Collections.singleton(WithVectorFields.class); + } + + @Override + public SchemaAction getSchemaAction() { + return SchemaAction.RECREATE_DROP_UNUSED; + } + } + + @Autowired VectorSearchRepository repository; + + @BeforeEach + void setUp() { + + repository.deleteAll(); + + WithVectorFields w1 = new WithVectorFields("de", "one", Vector.of(0.1001f, 0.22345f, 0.33456f, 0.44567f, 0.55678f)); + WithVectorFields w2 = new WithVectorFields("de", "two", Vector.of(0.2001f, 0.32345f, 0.43456f, 0.54567f, 0.65678f)); + WithVectorFields w3 = new WithVectorFields("en", "three", + Vector.of(0.9001f, 0.82345f, 0.73456f, 0.64567f, 0.55678f)); + WithVectorFields w4 = new WithVectorFields("de", "four", + Vector.of(0.9001f, 0.92345f, 0.93456f, 0.94567f, 0.95678f)); + + repository.saveAll(List.of(w1, w2, w3, w4)); + } + + @Test // GH- + void searchWithoutScoringFunctionShouldFail() { + assertThatExceptionOfType(QueryCreationException.class) + .isThrownBy(() -> repository.searchByEmbeddingNear(VECTOR, Limit.of(100))); + } + + @Test // GH- + void shouldConsiderScoringFunction() { + + SearchResults results = repository.searchByEmbeddingNear(VECTOR, + ScoringFunction.dotProduct(), Limit.of(100)); + + assertThat(results).hasSize(4); + + for (SearchResult result : results) { + assertThat(result.getScore()).isInstanceOf(Similarity.class); + assertThat(result.getScore().getValue()).isNotCloseTo(0d, offset(0.1d)); + } + + results = repository.searchByEmbeddingNear(VECTOR, VectorScoringFunctions.EUCLIDEAN, Limit.of(100)); + + assertThat(results).hasSize(4); + + for (SearchResult result : results) { + + assertThat(result.getScore()).isInstanceOf(Similarity.class); + assertThat(result.getScore().getValue()).isNotCloseTo(0.3d, offset(0.1d)); + } + } + + @Test // GH- + void shouldRunAnnotatedSearchByVector() { + + SearchResults results = repository.searchAnnotatedByEmbeddingNear(VECTOR, Limit.of(100)); + + assertThat(results).hasSize(4); + for (SearchResult result : results) { + assertThat(result.getScore()).isInstanceOf(Similarity.class); + assertThat(result.getScore().getValue()).isNotCloseTo(0d, offset(0.1d)); + } + } + + @Test // GH- + void shouldFindByVector() { + + List result = repository.findByEmbeddingNear(VECTOR, Limit.of(100)); + + assertThat(result).hasSize(4); + } + + interface VectorSearchRepository extends CrudRepository { + + SearchResults searchByEmbeddingNear(Vector embedding, ScoringFunction function, Limit limit); + + SearchResults searchByEmbeddingNear(Vector embedding, Limit limit); + + List findByEmbeddingNear(Vector embedding, Limit limit); + + @Query("SELECT id,description,country,similarity_cosine(embedding,:embedding) AS score FROM withvectorfields ORDER BY embedding ANN OF :embedding LIMIT :limit") + SearchResults searchAnnotatedByEmbeddingNear(Vector embedding, Limit limit); + + } + + @Table + static class WithVectorFields { + + @Id String id; + String country; + String description; + + @VectorType(dimensions = 5) + @SaiIndexed Vector embedding; + + @PersistenceCreator + public WithVectorFields(String id, String country, String description, Vector embedding) { + this.id = id; + this.country = country; + this.description = description; + this.embedding = embedding; + } + + public WithVectorFields(String country, String description, Vector embedding) { + this.id = UUID.randomUUID().toString(); + this.country = country; + this.description = description; + this.embedding = embedding; + } + + public String getId() { + return id; + } + + public String getCountry() { + return country; + } + + public String getDescription() { + return description; + } + + public Vector getEmbedding() { + return embedding; + } + + @Override + public String toString() { + return "WithVectorFields{" + "id='" + id + '\'' + ", country='" + country + '\'' + ", description='" + description + + '\'' + '}'; + } + + } + +} diff --git a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/query/CassandraParametersUnitTests.java b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/query/CassandraParametersUnitTests.java index 292de9b81..dfd04f7e5 100755 --- a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/query/CassandraParametersUnitTests.java +++ b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/query/CassandraParametersUnitTests.java @@ -26,10 +26,11 @@ import org.junit.jupiter.api.Test; import org.junit.jupiter.api.extension.ExtendWith; import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; + import org.springframework.data.cassandra.core.mapping.CassandraType; import org.springframework.data.cassandra.domain.Person; +import org.springframework.data.domain.ScoringFunction; import org.springframework.data.repository.Repository; -import org.springframework.data.repository.core.support.DefaultRepositoryMetadata; import org.springframework.data.repository.query.ParametersSource; /** @@ -89,6 +90,15 @@ class CassandraParametersUnitTests { assertThat(cassandraParameters.getParameter(0).getCassandraType().type()).isEqualTo(Name.BOOLEAN); } + @Test // GH- + void considersScoringFunctionIndex() throws Exception { + + Method method = PersonRepository.class.getMethod("findByObject", ScoringFunction.class); + CassandraParameters cassandraParameters = new CassandraParameters(ParametersSource.of(method)); + + assertThat(cassandraParameters.getScoringFunctionIndex()).isEqualTo(0); + } + interface PersonRepository extends Repository { Person findByFirstname(String firstname); @@ -97,6 +107,8 @@ class CassandraParametersUnitTests { Person findByObject(Object firstname); + Person findByObject(ScoringFunction firstname); + Person findByAnnotatedObject(@CassandraType(type = Name.TIME) Object firstname); Person findByComposedAnnotationObject(@ComposedCassandraTypeAnnotation Object firstname); diff --git a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/query/StubParameterAccessor.java b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/query/StubParameterAccessor.java index dafd0a650..3245ade93 100644 --- a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/query/StubParameterAccessor.java +++ b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/query/StubParameterAccessor.java @@ -24,7 +24,11 @@ import org.springframework.data.cassandra.core.cql.QueryOptions; import org.springframework.data.cassandra.core.mapping.CassandraType; import org.springframework.data.cassandra.core.query.CassandraScrollPosition; import org.springframework.data.domain.Pageable; +import org.springframework.data.domain.Range; +import org.springframework.data.domain.Score; +import org.springframework.data.domain.ScoringFunction; import org.springframework.data.domain.Sort; +import org.springframework.data.domain.Vector; import org.springframework.data.repository.query.ParameterAccessor; import com.datastax.oss.driver.api.core.type.DataType; @@ -77,6 +81,26 @@ class StubParameterAccessor implements CassandraParameterAccessor { return null; } + @Override + public ScoringFunction getScoringFunction() { + return null; + } + + @Override + public @Nullable Vector getVector() { + return null; + } + + @Override + public @Nullable Score getScore() { + return null; + } + + @Override + public @Nullable Range getScoreRange() { + return null; + } + @Override public Pageable getPageable() { return null; @@ -117,4 +141,9 @@ class StubParameterAccessor implements CassandraParameterAccessor { public Object[] getValues() { return new Object[0]; } + + @Override + public @Nullable Object getValue(int parameterIndex) { + return null; + } } diff --git a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/test/util/CassandraDelegate.java b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/test/util/CassandraDelegate.java index 2636125d6..477a889b9 100644 --- a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/test/util/CassandraDelegate.java +++ b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/test/util/CassandraDelegate.java @@ -290,7 +290,7 @@ class CassandraDelegate { if (container == null) { container = getCassandraDockerImageName().map(CassandraContainer::new) - .orElseGet(() -> new CassandraContainer("cassandra:5.0.3")); + .orElseGet(() -> new CassandraContainer("cassandra:5.0.3")).withReuse(true); container.withEnv("MAX_HEAP_SIZE", "1500M"); container.withEnv("HEAP_NEWSIZE", "300M");