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 8972e6711..8d794f30a 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 @@ -944,63 +944,65 @@ public class StatementFactory { () -> new IllegalArgumentException(String.format("Unknown operator [%s]", predicate.getOperator()))); ColumnRelationBuilder column = Relation.column(columnName); + Object value = predicate.getValue(); + switch (predicateOperator) { case EQ: - return column.isEqualTo(factory.create(predicate.getValue())); + return column.isEqualTo(factory.create(value)); case NE: - return column.isNotEqualTo(factory.create(predicate.getValue())); + return column.isNotEqualTo(factory.create(value)); case GT: - return column.isGreaterThan(factory.create(predicate.getValue())); + return column.isGreaterThan(factory.create(value)); case GTE: - return column.isGreaterThanOrEqualTo(factory.create(predicate.getValue())); + return column.isGreaterThanOrEqualTo(factory.create(value)); case LT: - return column.isLessThan(factory.create(predicate.getValue())); + return column.isLessThan(factory.create(value)); case LTE: - return column.isLessThanOrEqualTo(factory.create(predicate.getValue())); + return column.isLessThanOrEqualTo(factory.create(value)); case IN: - if (predicate.getValue() instanceof List - || (predicate.getValue() != null && predicate.getValue().getClass().isArray())) { - Term term = factory.create(predicate.getValue()); - if (term instanceof BindMarker) { - return column.in((BindMarker) term); + if (isCollectionLike(value)) { + + if (factory.canBindCollection()) { + Term term = factory.create(value); + return term instanceof BindMarker ? column.in((BindMarker) term) : column.in(term); } - return column.in(toLiterals(predicate.getValue())); + return column.in(toLiterals(value)); } - return column.in(factory.create(predicate.getValue())); + return column.in(factory.create(value)); case LIKE: - return column.like(factory.create(predicate.getValue())); + return column.like(factory.create(value)); case IS_NOT_NULL: return column.isNotNull(); case CONTAINS: - Assert.state(predicate.getValue() != null, + Assert.state(value != null, () -> String.format("CONTAINS value for column %s is null", columnName)); - return column.contains(factory.create(predicate.getValue())); + return column.contains(factory.create(value)); case CONTAINS_KEY: - Assert.state(predicate.getValue() != null, + Assert.state(value != null, () -> String.format("CONTAINS KEY value for column %s is null", columnName)); - return column.containsKey(factory.create(predicate.getValue())); + return column.containsKey(factory.create(value)); } throw new IllegalArgumentException( - String.format("Criteria %s %s %s not supported", columnName, predicate.getOperator(), predicate.getValue())); + String.format("Criteria %s %s %s not supported", columnName, predicate.getOperator(), value)); } private static Condition toCondition(CriteriaDefinition criteriaDefinition, TermFactory factory) { @@ -1014,43 +1016,45 @@ public class StatementFactory { () -> new IllegalArgumentException(String.format("Unknown operator [%s]", predicate.getOperator()))); ConditionBuilder column = Condition.column(columnName); + Object value = predicate.getValue(); + switch (predicateOperator) { case EQ: - return column.isEqualTo(factory.create(predicate.getValue())); + return column.isEqualTo(factory.create(value)); case NE: - return column.isNotEqualTo(factory.create(predicate.getValue())); + return column.isNotEqualTo(factory.create(value)); case GT: - return column.isGreaterThan(factory.create(predicate.getValue())); + return column.isGreaterThan(factory.create(value)); case GTE: - return column.isGreaterThanOrEqualTo(factory.create(predicate.getValue())); + return column.isGreaterThanOrEqualTo(factory.create(value)); case LT: - return column.isLessThan(factory.create(predicate.getValue())); + return column.isLessThan(factory.create(value)); case LTE: - return column.isLessThanOrEqualTo(factory.create(predicate.getValue())); + return column.isLessThanOrEqualTo(factory.create(value)); case IN: - if (predicate.getValue() instanceof List - || (predicate.getValue() != null && predicate.getValue().getClass().isArray())) { - Term term = factory.create(predicate.getValue()); - if (term instanceof BindMarker) { - return column.in((BindMarker) term); + if (isCollectionLike(value)) { + + if (factory.canBindCollection()) { + Term term = factory.create(value); + return term instanceof BindMarker ? column.in((BindMarker) term) : column.in(term); } - return column.in(toLiterals(predicate.getValue())); + return column.in(toLiterals(value)); } - return column.in(factory.create(predicate.getValue())); + return column.in(factory.create(value)); } throw new IllegalArgumentException(String.format("Criteria %s %s %s not supported for IF Conditions", columnName, - predicate.getOperator(), predicate.getValue())); + predicate.getOperator(), value)); } static List toLiterals(@Nullable Object arrayOrList) { @@ -1084,6 +1088,10 @@ public class StatementFactory { return Collections.emptyList(); } + private static boolean isCollectionLike(@Nullable Object value) { + return value instanceof List || (value != null && value.getClass().isArray()); + } + static class SimpleSelector implements com.datastax.oss.driver.api.querybuilder.select.Selector { private final String selector; diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/cql/util/StatementBuilder.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/cql/util/StatementBuilder.java index e7d5d134a..cee14e0ac 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/cql/util/StatementBuilder.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/cql/util/StatementBuilder.java @@ -206,7 +206,17 @@ public class StatementBuilder { if (parameterHandling == ParameterHandling.INLINE) { - TermFactory termFactory = value -> toLiteralTerms(value, codecRegistry); + TermFactory termFactory = new TermFactory() { + @Override + public Term create(@Nullable Object value) { + return toLiteralTerms(value, codecRegistry); + } + + @Override + public boolean canBindCollection() { + return false; + } + }; for (BuilderRunnable runnable : queryActions) { statement = runnable.run(statement, termFactory); diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/cql/util/TermFactory.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/cql/util/TermFactory.java index ca80b6d16..bd2c0365e 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/cql/util/TermFactory.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/core/cql/util/TermFactory.java @@ -15,10 +15,10 @@ */ package org.springframework.data.cassandra.core.cql.util; -import com.datastax.oss.driver.api.querybuilder.term.Term; - import org.springframework.lang.Nullable; +import com.datastax.oss.driver.api.querybuilder.term.Term; + /** * Factory for {@link Term} objects encapsulating a binding {@code value}. Classes implementing this factory interface * may return inline terms to render values as part of the query string, or bind markers to supply parameters on @@ -39,4 +39,14 @@ public interface TermFactory { * @return the {@link Term} for the given {@code value}. */ Term create(@Nullable Object value); + + /** + * Check whether the term factory accepts {@link java.util.Collection} values to be created as {@link Term}. + * + * @return {@code true} whether the term factory can {@link java.util.Collection} values. + * @since 3.2.6 + */ + default boolean canBindCollection() { + return true; + } } 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 80cbf2e5d..cfd9df270 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 @@ -19,6 +19,7 @@ import static org.assertj.core.api.Assertions.*; import static org.springframework.data.domain.Sort.Direction.*; import java.time.Duration; +import java.util.Arrays; import java.util.Collections; import java.util.List; import java.util.Map; @@ -43,6 +44,7 @@ import org.springframework.data.cassandra.core.query.Update; import org.springframework.data.cassandra.domain.Group; import org.springframework.data.domain.Sort; +import com.datastax.oss.driver.api.core.CqlIdentifier; import com.datastax.oss.driver.api.core.DefaultConsistencyLevel; import com.datastax.oss.driver.api.core.cql.SimpleStatement; import com.datastax.oss.driver.api.querybuilder.delete.Delete; @@ -171,14 +173,54 @@ class StatementFactoryUnitTests { .isEqualTo("SELECT * FROM group LIMIT 10 ALLOW FILTERING"); } - @Test - void shouldMapSelectInQuery() { + @Test // GH-1172 + void shouldMapSelectInQueryAsInlineValue() { - Query query = Query.query(Criteria.where("foo").in("bar")); - - StatementBuilder select = statementFactory.select(Query.query(Criteria.where("foo").in("bar")), + groupEntity); assertThat(select.build(ParameterHandling.INLINE).getQuery()).isEqualTo("SELECT * FROM group WHERE foo IN ('bar')"); + + select = statementFactory.select(Query.query(Criteria.where("foo").in("bar", "baz")), groupEntity); + + assertThat(select.build(ParameterHandling.INLINE).getQuery()) + .isEqualTo("SELECT * FROM group WHERE foo IN ('bar','baz')"); + } + + @Test // GH-1172 + void shouldMapSelectInQueryAsByIndexValue() { + + StatementBuilder select = statementFactory.select(Query.query(Criteria.where("foo").in("bar")), + groupEntity); + SimpleStatement statement = select.build(ParameterHandling.BY_NAME); + + assertThat(statement.getQuery()).isEqualTo("SELECT * FROM group WHERE foo IN :p0"); + assertThat(statement.getNamedValues()).hasSize(1).containsEntry(CqlIdentifier.fromCql("p0"), + Collections.singletonList("bar")); + + select = statementFactory.select(Query.query(Criteria.where("foo").in("bar", "baz")), groupEntity); + statement = select.build(ParameterHandling.BY_NAME); + + assertThat(statement.getQuery()).isEqualTo("SELECT * FROM group WHERE foo IN :p0"); + assertThat(statement.getNamedValues()).hasSize(1).containsEntry(CqlIdentifier.fromCql("p0"), + (Arrays.asList("bar", "baz"))); } @Test // DATACASS-343