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 006cde0ef..4d20abe1f 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 @@ -78,6 +78,8 @@ import com.datastax.oss.driver.api.querybuilder.relation.Relation; import com.datastax.oss.driver.api.querybuilder.select.Select; import com.datastax.oss.driver.api.querybuilder.term.Term; import com.datastax.oss.driver.api.querybuilder.update.Assignment; +import com.datastax.oss.driver.api.querybuilder.update.OngoingAssignment; +import com.datastax.oss.driver.api.querybuilder.update.UpdateStart; import com.datastax.oss.driver.api.querybuilder.update.UpdateWithAssignments; /** @@ -268,7 +270,7 @@ public class StatementFactory { CassandraPersistentEntity persistentEntity = cassandraConverter.getMappingContext() .getRequiredPersistentEntity(objectToInsert.getClass()); - return insert(options, options, persistentEntity, persistentEntity.getTableName()); + return insert(objectToInsert, options, persistentEntity, persistentEntity.getTableName()); } /** @@ -303,7 +305,7 @@ public class StatementFactory { StatementBuilder builder = StatementBuilder .of(QueryBuilder.insertInto(tableName).valuesByIds(Collections.emptyMap())).bind((statement, factory) -> { - Map values = createTerms(insertNulls, object); + Map values = createTerms(insertNulls, object, factory); return statement.valuesByIds(values); }).apply(statement -> (RegularInsert) addWriteOptions(statement, options)); @@ -315,7 +317,8 @@ public class StatementFactory { return builder; } - private static Map createTerms(boolean insertNulls, Map object) { + private static Map createTerms(boolean insertNulls, Map object, + TermFactory factory) { Map values = new LinkedHashMap<>(object.size()); @@ -324,7 +327,7 @@ public class StatementFactory { if (o == null && !insertNulls) { return; } - values.put(cqlIdentifier, QueryBuilder.literal(o)); + values.put(cqlIdentifier, factory.create(o)); }); return values; } @@ -470,11 +473,12 @@ public class StatementFactory { * @param tableName must not be {@literal null}. * @return the delete builder. */ - StatementBuilder deleteById(Object id, EntityWriter entityWriter, CqlIdentifier tableName) { + StatementBuilder deleteById(Object id, CassandraPersistentEntity persistentEntity, + CqlIdentifier tableName) { Where where = new Where(); - entityWriter.write(id, where); + cassandraConverter.write(id, where, persistentEntity); return StatementBuilder.of(QueryBuilder.deleteFrom(tableName).where()).bind((statement, factory) -> { return statement.where(toRelations(where, factory)); @@ -640,6 +644,8 @@ public class StatementFactory { } select.onBuild(statementBuilder -> { + + query.getPagingState().ifPresent(statementBuilder::setPagingState); query.getQueryOptions().ifPresent(it -> QueryOptionsUtil.addQueryOptions(statementBuilder, it)); }); @@ -713,19 +719,24 @@ public class StatementFactory { private static StatementBuilder update(CqlIdentifier table, Update mappedUpdate, Filter filter) { - UpdateWithAssignments withAssignments = (UpdateWithAssignments) QueryBuilder.update(table); + UpdateStart updateStart = QueryBuilder.update(table); - for (AssignmentOp assignmentOp : mappedUpdate.getUpdateOperations()) { - withAssignments.set(getAssignment(assignmentOp)); - } + return StatementBuilder.of((com.datastax.oss.driver.api.querybuilder.update.Update) updateStart) + .bind((statement, factory) -> { - return StatementBuilder.of(withAssignments.where()).bind((statement, factory) -> { + List assignments = mappedUpdate.getUpdateOperations().stream() + .map(assignmentOp -> getAssignment(assignmentOp, factory)).collect(Collectors.toList()); - List relations = filter.stream().map(criteriaDefinition -> toClause(criteriaDefinition, factory)) - .collect(Collectors.toList()); + return (com.datastax.oss.driver.api.querybuilder.update.Update) ((OngoingAssignment) statement) + .set(assignments); - return statement.where(relations); - }); + }).bind((statement, factory) -> { + + List relations = filter.stream().map(criteriaDefinition -> toClause(criteriaDefinition, factory)) + .collect(Collectors.toList()); + + return statement.where(relations); + }); } static Iterable toRelations(Where where, TermFactory factory) { @@ -770,39 +781,39 @@ public class StatementFactory { }); } - private static Assignment getAssignment(AssignmentOp assignmentOp) { + private static Assignment getAssignment(AssignmentOp assignmentOp, TermFactory termFactory) { if (assignmentOp instanceof SetOp) { - return getAssignment((SetOp) assignmentOp); + return getAssignment((SetOp) assignmentOp, termFactory); } if (assignmentOp instanceof RemoveOp) { - return getAssignment((RemoveOp) assignmentOp); + return getAssignment((RemoveOp) assignmentOp, termFactory); } if (assignmentOp instanceof IncrOp) { - return getAssignment((IncrOp) assignmentOp); + return getAssignment((IncrOp) assignmentOp, termFactory); } if (assignmentOp instanceof AddToOp) { - return getAssignment((AddToOp) assignmentOp); + return getAssignment((AddToOp) assignmentOp, termFactory); } if (assignmentOp instanceof AddToMapOp) { - return getAssignment((AddToMapOp) assignmentOp); + return getAssignment((AddToMapOp) assignmentOp, termFactory); } throw new IllegalArgumentException(String.format("UpdateOp %s not supported", assignmentOp)); } - private static Assignment getAssignment(IncrOp incrOp) { + private static Assignment getAssignment(IncrOp incrOp, TermFactory termFactory) { return incrOp.getValue().intValue() > 0 - ? Assignment.increment(incrOp.toCqlIdentifier(), QueryBuilder.literal(Math.abs(incrOp.getValue().intValue()))) - : Assignment.decrement(incrOp.toCqlIdentifier(), QueryBuilder.literal(Math.abs(incrOp.getValue().intValue()))); + ? Assignment.increment(incrOp.toCqlIdentifier(), termFactory.create(Math.abs(incrOp.getValue().intValue()))) + : Assignment.decrement(incrOp.toCqlIdentifier(), termFactory.create(Math.abs(incrOp.getValue().intValue()))); } - private static Assignment getAssignment(SetOp updateOp) { + private static Assignment getAssignment(SetOp updateOp, TermFactory termFactory) { if (updateOp instanceof SetAtIndexOp) { SetAtIndexOp op = (SetAtIndexOp) updateOp; @@ -813,40 +824,44 @@ public class StatementFactory { if (updateOp instanceof SetAtKeyOp) { SetAtKeyOp op = (SetAtKeyOp) updateOp; - return Assignment.setMapValue(op.toCqlIdentifier(), QueryBuilder.literal(op.getKey()), - QueryBuilder.literal(op.getValue())); + return Assignment.setMapValue(op.toCqlIdentifier(), termFactory.create(op.getKey()), + termFactory.create(op.getValue())); } - return Assignment.setColumn(updateOp.toCqlIdentifier(), QueryBuilder.literal(updateOp.getValue())); + return Assignment.setColumn(updateOp.toCqlIdentifier(), termFactory.create(updateOp.getValue())); } - private static Assignment getAssignment(RemoveOp updateOp) { + private static Assignment getAssignment(RemoveOp updateOp, TermFactory termFactory) { if (updateOp.getValue() instanceof Set) { - return Assignment.removeSetElement(updateOp.toCqlIdentifier(), QueryBuilder.literal(updateOp.getValue())); + + + return Assignment.removeSetElement(updateOp.toCqlIdentifier(), termFactory.create(updateOp.getValue())); } if (updateOp.getValue() instanceof List) { - return Assignment.removeListElement(updateOp.toCqlIdentifier(), QueryBuilder.literal(updateOp.getValue())); + + + return Assignment.removeListElement(updateOp.toCqlIdentifier(), termFactory.create(updateOp.getValue())); } - return Assignment.remove(updateOp.toCqlIdentifier(), QueryBuilder.literal(updateOp.getValue())); + return Assignment.remove(updateOp.toCqlIdentifier(), termFactory.create(updateOp.getValue())); } @SuppressWarnings("unchecked") - private static Assignment getAssignment(AddToOp updateOp) { + private static Assignment getAssignment(AddToOp updateOp, TermFactory termFactory) { if (updateOp.getValue() instanceof Set) { - return Assignment.appendSetElement(updateOp.toCqlIdentifier(), QueryBuilder.literal(updateOp.getValue())); + return Assignment.appendSetElement(updateOp.toCqlIdentifier(), termFactory.create(updateOp.getValue())); } return Mode.PREPEND.equals(updateOp.getMode()) - ? Assignment.prependListElement(updateOp.toCqlIdentifier(), QueryBuilder.literal(updateOp.getValue())) - : Assignment.appendListElement(updateOp.getColumnName().toCql(), QueryBuilder.literal(updateOp.getValue())); + ? Assignment.prependListElement(updateOp.toCqlIdentifier(), termFactory.create(updateOp.getValue())) + : Assignment.appendListElement(updateOp.getColumnName().toCql(), termFactory.create(updateOp.getValue())); } - private static Assignment getAssignment(AddToMapOp updateOp) { - return Assignment.append(updateOp.toCqlIdentifier(), QueryBuilder.literal(updateOp.getValue())); + private static Assignment getAssignment(AddToMapOp updateOp, TermFactory termFactory) { + return Assignment.append(updateOp.toCqlIdentifier(), termFactory.create(updateOp.getValue())); } private StatementBuilder delete(List columnNames, CqlIdentifier from, Filter filter) { diff --git a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/core/cql/util/StatementBuilderUnitTests.java b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/core/cql/util/StatementBuilderUnitTests.java index 766cfdd29..32e7f4e0b 100644 --- a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/core/cql/util/StatementBuilderUnitTests.java +++ b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/core/cql/util/StatementBuilderUnitTests.java @@ -17,6 +17,8 @@ package org.springframework.data.cassandra.core.cql.util; import static org.assertj.core.api.Assertions.*; +import java.util.Collections; + import org.junit.Test; import com.datastax.oss.driver.api.core.CqlIdentifier; @@ -69,6 +71,39 @@ public class StatementBuilderUnitTests { assertThat(statement.getPositionalValues()).containsOnly("bar"); } + @Test // DATACASS-656 + public void shouldBindList() { + + SimpleStatement statement = StatementBuilder.of(QueryBuilder.selectFrom("person").all()) + .bind((select, factory) -> select + .where(Relation.column("foo").isEqualTo(factory.create(Collections.singletonList("value"))))) + .build(StatementBuilder.ParameterHandling.INLINE); + + assertThat(statement.getQuery()).isEqualTo("SELECT * FROM person WHERE foo=['value']"); + } + + @Test // DATACASS-656 + public void shouldBindSet() { + + SimpleStatement statement = StatementBuilder.of(QueryBuilder.selectFrom("person").all()) + .bind((select, factory) -> select + .where(Relation.column("foo").isEqualTo(factory.create(Collections.singleton("value"))))) + .build(StatementBuilder.ParameterHandling.INLINE); + + assertThat(statement.getQuery()).isEqualTo("SELECT * FROM person WHERE foo={'value'}"); + } + + @Test // DATACASS-656 + public void shouldBindMap() { + + SimpleStatement statement = StatementBuilder.of(QueryBuilder.selectFrom("person").all()) + .bind((select, factory) -> select + .where(Relation.column("foo").isEqualTo(factory.create(Collections.singletonMap("key", "value"))))) + .build(StatementBuilder.ParameterHandling.INLINE); + + assertThat(statement.getQuery()).isEqualTo("SELECT * FROM person WHERE foo={'key':'value'}"); + } + @Test // DATACASS-656 public void shouldBindByName() {