diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/CassandraRepository.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/CassandraRepository.java index bf36c227b..ee0846a76 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/CassandraRepository.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/CassandraRepository.java @@ -54,8 +54,14 @@ public interface CassandraRepository extends CrudRepository { @Override List findAll(); - /* (non-Javadoc) - * @see org.springframework.data.repository.CrudRepository#findAllById(java.lang.Iterable) + /** + * {@inheritDoc} + *

+ * Note: Cassandra supports single-field {@code IN} queries only. When using {@link MapId} with multiple components, + * use {@link #findById(Object)}. + * + * @throws org.springframework.dao.InvalidDataAccessApiUsageException thrown when using {@link MapId} with multiple + * key components. */ @Override List findAllById(Iterable ids); diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/ReactiveCassandraRepository.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/ReactiveCassandraRepository.java index 7a99dc0cd..e314bf93b 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/ReactiveCassandraRepository.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/ReactiveCassandraRepository.java @@ -60,4 +60,21 @@ public interface ReactiveCassandraRepository extends ReactiveCrudReposito * @return the saved entity */ Flux insert(Publisher entities); + + /** + * {@inheritDoc} + *

+ * Note: Cassandra supports single-field {@code IN} queries only. Fetches each row individually when using + * {@link MapId} with multiple components. + */ + @Override + Flux findAllById(Iterable iterable); + + /** + * {@inheritDoc} + *

+ * Fetches each row individually. + */ + @Override + Flux findAllById(Publisher publisher); } diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/support/FindByIdQuery.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/support/FindByIdQuery.java new file mode 100644 index 000000000..f4a60f11a --- /dev/null +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/support/FindByIdQuery.java @@ -0,0 +1,118 @@ +/* + * Copyright 2019 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.support; + +import java.util.ArrayList; +import java.util.Iterator; +import java.util.List; + +import org.springframework.dao.InvalidDataAccessApiUsageException; +import org.springframework.data.cassandra.core.mapping.MapId; +import org.springframework.lang.Nullable; +import org.springframework.util.Assert; + +/** + * Value object representing a Find by ID query supporting also {@link MapId}. + * + * @author Mark Paluch + * @since 2.2 + */ +class FindByIdQuery { + + private final @Nullable String idProperty; + private final List idCollection; + + private FindByIdQuery(@Nullable String idProperty, List idCollection) { + + this.idProperty = idProperty; + this.idCollection = idCollection; + } + + /** + * Create a new {@link FindByIdQuery} given {@link Iterable} of {@code ID}s. Id's can be either scalar values or + * {@link MapId}s. In case of the latter, this method discovers the {@link #getIdProperty() Id property name}. + * + * @param ids + * @return + */ + static FindByIdQuery forIds(Iterable ids) { + + Assert.notNull(ids, "The given Iterable of ids must not be null"); + + List idCollection = new ArrayList<>(); + String idField = null; + + for (Object id : ids) { + + if (id instanceof MapId) { + + MapId mapId = (MapId) id; + Iterator iterator = mapId.keySet().iterator(); + + if (mapId.size() > 1) { + throw new InvalidDataAccessApiUsageException("MapId with multiple keys are not supported"); + } + + if (!iterator.hasNext()) { + throw new InvalidDataAccessApiUsageException("MapId is empty"); + } else { + + idField = iterator.next(); + idCollection.add(mapId.get(idField)); + } + } else { + idCollection.add(id); + } + } + + return new FindByIdQuery(idField, idCollection); + } + + /** + * Check if the {@link Iterable} of {@code ID}s contains composite keys. + * + * @param ids + * @return + */ + static boolean hasCompositeKeys(Iterable ids) { + + Assert.notNull(ids, "The given Iterable of ids must not be null"); + + for (Object id : ids) { + + if (id instanceof MapId) { + + MapId mapId = (MapId) id; + Iterator iterator = mapId.keySet().iterator(); + + if (mapId.size() > 1) { + return true; + } + } + } + + return false; + } + + @Nullable + String getIdProperty() { + return idProperty; + } + + List getIdCollection() { + return idCollection; + } +} diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/support/SimpleCassandraRepository.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/support/SimpleCassandraRepository.java index 3ab1e45b9..c013ba6a2 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/support/SimpleCassandraRepository.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/support/SimpleCassandraRepository.java @@ -15,9 +15,10 @@ */ package org.springframework.data.cassandra.repository.support; -import static org.springframework.data.cassandra.core.query.Criteria.where; +import static org.springframework.data.cassandra.core.query.Criteria.*; import java.util.ArrayList; +import java.util.Collections; import java.util.List; import java.util.Optional; @@ -32,8 +33,6 @@ import org.springframework.data.cassandra.repository.query.CassandraEntityInform import org.springframework.data.domain.Pageable; import org.springframework.data.domain.Slice; import org.springframework.data.mapping.context.AbstractMappingContext; -import org.springframework.data.util.StreamUtils; -import org.springframework.data.util.Streamable; import org.springframework.util.Assert; import com.datastax.driver.core.querybuilder.Insert; @@ -161,7 +160,11 @@ public class SimpleCassandraRepository implements CassandraRepository implements CassandraRepository idCollection = Streamable.of(ids).stream().collect(StreamUtils.toUnmodifiableList()); + FindByIdQuery mapIdQuery = FindByIdQuery.forIds(ids); + List idCollection = mapIdQuery.getIdCollection(); + String idField = mapIdQuery.getIdProperty(); - return this.operations.select(Query.query(where(this.entityInformation.getIdAttribute()).in(idCollection)), - this.entityInformation.getJavaType()); + if (idCollection.isEmpty()) { + return Collections.emptyList(); + } + + if (idField == null) { + idField = this.entityInformation.getIdAttribute(); + } + + return this.operations.select(Query.query(where(idField).in(idCollection)), this.entityInformation.getJavaType()); } /* (non-Javadoc) diff --git a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/support/SimpleReactiveCassandraRepository.java b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/support/SimpleReactiveCassandraRepository.java index 72228b94c..e730d5291 100644 --- a/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/support/SimpleReactiveCassandraRepository.java +++ b/spring-data-cassandra/src/main/java/org/springframework/data/cassandra/repository/support/SimpleReactiveCassandraRepository.java @@ -15,15 +15,21 @@ */ package org.springframework.data.cassandra.repository.support; +import static org.springframework.data.cassandra.core.query.Criteria.*; + import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; +import java.util.List; + import org.reactivestreams.Publisher; + import org.springframework.data.cassandra.core.EntityWriteResult; import org.springframework.data.cassandra.core.InsertOptions; import org.springframework.data.cassandra.core.ReactiveCassandraOperations; import org.springframework.data.cassandra.core.mapping.BasicCassandraPersistentEntity; import org.springframework.data.cassandra.core.mapping.CassandraPersistentProperty; +import org.springframework.data.cassandra.core.query.Query; import org.springframework.data.cassandra.repository.ReactiveCassandraRepository; import org.springframework.data.cassandra.repository.query.CassandraEntityInformation; import org.springframework.data.mapping.context.AbstractMappingContext; @@ -229,11 +235,27 @@ public class SimpleReactiveCassandraRepository implements ReactiveCassand * @see org.springframework.data.repository.reactive.ReactiveCrudRepository#findAllById(java.lang.Iterable) */ @Override - public Flux findAllById(Iterable iterable) { + public Flux findAllById(Iterable ids) { - Assert.notNull(iterable, "The given Iterable of ids must not be null"); + Assert.notNull(ids, "The given Iterable of ids must not be null"); - return findAllById(Flux.fromIterable(iterable)); + if (FindByIdQuery.hasCompositeKeys(ids)) { + return findAllById(Flux.fromIterable(ids)); + } + + FindByIdQuery query = FindByIdQuery.forIds(ids); + List idCollection = query.getIdCollection(); + String idField = query.getIdProperty(); + + if (idCollection.isEmpty()) { + return Flux.empty(); + } + + if (idField == null) { + idField = this.entityInformation.getIdAttribute(); + } + + return this.operations.select(Query.query(where(idField).in(idCollection)), this.entityInformation.getJavaType()); } /* diff --git a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/mapid/RepositoryMapIdIntegrationTests.java b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/mapid/RepositoryMapIdIntegrationTests.java index 20b4c8e23..0dae85316 100755 --- a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/mapid/RepositoryMapIdIntegrationTests.java +++ b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/mapid/RepositoryMapIdIntegrationTests.java @@ -18,12 +18,18 @@ package org.springframework.data.cassandra.repository.mapid; import static org.assertj.core.api.Assertions.*; import static org.springframework.data.cassandra.core.mapping.BasicMapId.*; +import java.util.Collections; +import java.util.List; + import org.junit.Before; import org.junit.Test; import org.junit.runner.RunWith; + import org.springframework.beans.factory.annotation.Autowired; import org.springframework.context.annotation.Configuration; +import org.springframework.dao.InvalidDataAccessApiUsageException; import org.springframework.data.cassandra.core.CassandraOperations; +import org.springframework.data.cassandra.core.mapping.BasicMapId; import org.springframework.data.cassandra.core.mapping.MapId; import org.springframework.data.cassandra.repository.config.EnableCassandraRepositories; import org.springframework.data.cassandra.repository.support.AbstractSpringDataEmbeddedCassandraIntegrationTest; @@ -52,13 +58,13 @@ public class RepositoryMapIdIntegrationTests extends AbstractSpringDataEmbeddedC } @Autowired CassandraOperations template; - @Autowired SinglePrimaryKecColumnRepository singlePrimaryKecColumnRepository; + @Autowired SinglePrimaryKecColumnRepository singlePrimaryKeyColumnRepository; @Autowired MultiPrimaryKeyColumnsRepository multiPrimaryKeyColumnsRepository; @Before public void before() { assertThat(template).isNotNull(); - assertThat(singlePrimaryKecColumnRepository).isNotNull(); + assertThat(singlePrimaryKeyColumnRepository).isNotNull(); assertThat(multiPrimaryKeyColumnsRepository).isNotNull(); } @@ -68,28 +74,38 @@ public class RepositoryMapIdIntegrationTests extends AbstractSpringDataEmbeddedC // insert SinglePrimaryKeyColumn inserted = new SinglePrimaryKeyColumn(uuid()); inserted.setValue(uuid()); - SinglePrimaryKeyColumn saved = singlePrimaryKecColumnRepository.save(inserted); + SinglePrimaryKeyColumn saved = singlePrimaryKeyColumnRepository.save(inserted); assertThat(inserted).isSameAs(saved); // select MapId id = id("key", saved.getKey()); - SinglePrimaryKeyColumn selected = singlePrimaryKecColumnRepository.findById(id).get(); + SinglePrimaryKeyColumn selected = singlePrimaryKeyColumnRepository.findById(id).get(); assertThat(saved).isNotSameAs(selected); assertThat(selected.getKey()).isEqualTo(saved.getKey()); assertThat(selected.getValue()).isEqualTo(saved.getValue()); + List allById = singlePrimaryKeyColumnRepository.findAllById(Collections.singletonList(id)); + + assertThat(allById).containsOnly(saved); + // update selected.setValue(uuid()); - SinglePrimaryKeyColumn updated = singlePrimaryKecColumnRepository.save(selected); + SinglePrimaryKeyColumn updated = singlePrimaryKeyColumnRepository.save(selected); assertThat(selected).isSameAs(updated); - selected = singlePrimaryKecColumnRepository.findById(id).get(); + selected = singlePrimaryKeyColumnRepository.findById(id).get(); assertThat(updated).isNotSameAs(selected); assertThat(selected.getValue()).isEqualTo(updated.getValue()); // delete - singlePrimaryKecColumnRepository.delete(selected); - assertThat(singlePrimaryKecColumnRepository.findById(id)).isEmpty(); + singlePrimaryKeyColumnRepository.delete(selected); + assertThat(singlePrimaryKeyColumnRepository.findById(id)).isEmpty(); + } + + @Test // DATACASS-661 + public void findAllByIdRejectsEmptyMapId() { + assertThatThrownBy(() -> multiPrimaryKeyColumnsRepository.findAllById(Collections.singletonList(BasicMapId.id()))) + .isInstanceOf(InvalidDataAccessApiUsageException.class); } @Test @@ -109,6 +125,9 @@ public class RepositoryMapIdIntegrationTests extends AbstractSpringDataEmbeddedC assertThat(selected.getKey1()).isEqualTo(saved.getKey1()); assertThat(selected.getValue()).isEqualTo(saved.getValue()); + assertThatThrownBy(() -> multiPrimaryKeyColumnsRepository.findAllById(Collections.singletonList(id))) + .isInstanceOf(InvalidDataAccessApiUsageException.class); + // update selected.setValue(uuid()); MultiPrimaryKeyColumns updated = multiPrimaryKeyColumnsRepository.save(selected); diff --git a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/mapid/SinglePrimaryKeyColumn.java b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/mapid/SinglePrimaryKeyColumn.java index 2276eb0e5..0f72452e3 100644 --- a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/mapid/SinglePrimaryKeyColumn.java +++ b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/mapid/SinglePrimaryKeyColumn.java @@ -15,6 +15,8 @@ */ package org.springframework.data.cassandra.repository.mapid; +import lombok.EqualsAndHashCode; + import org.springframework.data.cassandra.core.cql.PrimaryKeyType; import org.springframework.data.cassandra.core.mapping.Column; import org.springframework.data.cassandra.core.mapping.PrimaryKeyColumn; @@ -24,6 +26,7 @@ import org.springframework.data.cassandra.core.mapping.Table; * @author Matthew T. Adams */ @Table +@EqualsAndHashCode public class SinglePrimaryKeyColumn { @PrimaryKeyColumn(ordinal = 0, type = PrimaryKeyType.PARTITIONED) String key; diff --git a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/support/SimpleCassandraRepositoryMapIdIntegrationTests.java b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/support/SimpleCassandraRepositoryMapIdIntegrationTests.java new file mode 100644 index 000000000..22f01d1a1 --- /dev/null +++ b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/support/SimpleCassandraRepositoryMapIdIntegrationTests.java @@ -0,0 +1,89 @@ +/* + * Copyright 2019 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.support; + +import static org.assertj.core.api.Assertions.*; + +import java.util.Collections; + +import org.junit.Before; +import org.junit.Test; + +import org.springframework.dao.InvalidDataAccessApiUsageException; +import org.springframework.data.cassandra.core.CassandraTemplate; +import org.springframework.data.cassandra.core.mapping.BasicMapId; +import org.springframework.data.cassandra.core.mapping.MapId; +import org.springframework.data.cassandra.domain.TypeWithMapId; +import org.springframework.data.cassandra.domain.User; +import org.springframework.data.cassandra.test.util.AbstractKeyspaceCreatingIntegrationTest; + +/** + * Integration tests for {@link SimpleCassandraRepository} using MapId. + * + * @author Mark Paluch + */ +public class SimpleCassandraRepositoryMapIdIntegrationTests extends AbstractKeyspaceCreatingIntegrationTest { + + private SimpleCassandraRepository simple; + + private SimpleCassandraRepository composite; + + @Before + @SuppressWarnings("unchecked") + public void setUp() { + + CassandraTemplate template = new CassandraTemplate(this.session); + SchemaTestUtils.potentiallyCreateTableFor(User.class, template); + SchemaTestUtils.potentiallyCreateTableFor(TypeWithMapId.class, template); + + SchemaTestUtils.truncate(TypeWithMapId.class, template); + SchemaTestUtils.truncate(TypeWithMapId.class, template); + + simple = new SimpleCassandraRepository<>(new MappingCassandraEntityInformation( + template.getConverter().getMappingContext().getRequiredPersistentEntity(User.class), template.getConverter()), + template); + + composite = new SimpleCassandraRepository<>(new MappingCassandraEntityInformation( + template.getConverter().getMappingContext().getRequiredPersistentEntity(TypeWithMapId.class), + template.getConverter()), template); + } + + @Test // DATACASS-661 + public void shouldFindByIdWithSimpleKey() { + + User user = new User(); + user.setId("heisenberg"); + user.setFirstname("Walter"); + user.setLastname("White"); + + simple.save(user); + + assertThat(simple.findAllById(Collections.singletonList(BasicMapId.id("id", user.getId())))).hasSize(1); + } + + @Test // DATACASS-661 + public void shouldFindByIdWithCompositeKey() { + + TypeWithMapId withMapId = new TypeWithMapId(); + withMapId.setFirstname("Walter"); + withMapId.setLastname("White"); + + composite.save(withMapId); + + assertThatThrownBy(() -> composite.findAllById(Collections.singletonList(withMapId.getMapId()))) + .isInstanceOf(InvalidDataAccessApiUsageException.class); + } +} diff --git a/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/support/SimpleReactiveCassandraRepositoryMapIdIntegrationTests.java b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/support/SimpleReactiveCassandraRepositoryMapIdIntegrationTests.java new file mode 100644 index 000000000..5586ed4a0 --- /dev/null +++ b/spring-data-cassandra/src/test/java/org/springframework/data/cassandra/repository/support/SimpleReactiveCassandraRepositoryMapIdIntegrationTests.java @@ -0,0 +1,104 @@ +/* + * Copyright 2019 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.support; + +import reactor.test.StepVerifier; + +import java.util.Collections; + +import org.junit.Before; +import org.junit.Test; + +import org.springframework.data.cassandra.core.CassandraTemplate; +import org.springframework.data.cassandra.core.ReactiveCassandraTemplate; +import org.springframework.data.cassandra.core.cql.session.DefaultBridgedReactiveSession; +import org.springframework.data.cassandra.core.mapping.BasicMapId; +import org.springframework.data.cassandra.core.mapping.MapId; +import org.springframework.data.cassandra.domain.TypeWithMapId; +import org.springframework.data.cassandra.domain.User; +import org.springframework.data.cassandra.test.util.AbstractKeyspaceCreatingIntegrationTest; + +/** + * Integration tests for {@link SimpleReactiveCassandraRepository} using MapId. + * + * @author Mark Paluch + */ +public class SimpleReactiveCassandraRepositoryMapIdIntegrationTests extends AbstractKeyspaceCreatingIntegrationTest { + + private SimpleReactiveCassandraRepository simple; + + private SimpleReactiveCassandraRepository composite; + + @Before + @SuppressWarnings("unchecked") + public void setUp() { + + CassandraTemplate template = new CassandraTemplate(this.session); + SchemaTestUtils.potentiallyCreateTableFor(User.class, template); + SchemaTestUtils.potentiallyCreateTableFor(TypeWithMapId.class, template); + + SchemaTestUtils.truncate(TypeWithMapId.class, template); + SchemaTestUtils.truncate(TypeWithMapId.class, template); + + ReactiveCassandraTemplate reactiveTemplate = new ReactiveCassandraTemplate( + new DefaultBridgedReactiveSession(this.session)); + + simple = new SimpleReactiveCassandraRepository<>(new MappingCassandraEntityInformation( + template.getConverter().getMappingContext().getRequiredPersistentEntity(User.class), template.getConverter()), + reactiveTemplate); + + composite = new SimpleReactiveCassandraRepository<>(new MappingCassandraEntityInformation( + template.getConverter().getMappingContext().getRequiredPersistentEntity(TypeWithMapId.class), + template.getConverter()), reactiveTemplate); + } + + @Test // DATACASS-661 + public void shouldFindByIdWithSimpleKey() { + + User user = new User(); + user.setId("heisenberg"); + user.setFirstname("Walter"); + user.setLastname("White"); + + simple.save(user) // + .as(StepVerifier::create) // + .expectNextCount(1) // + .verifyComplete(); + + simple.findAllById(Collections.singletonList(BasicMapId.id("id", user.getId()))) // + .as(StepVerifier::create) // + .expectNextCount(1) // + .verifyComplete(); + } + + @Test // DATACASS-661 + public void shouldFindByIdWithCompositeKey() { + + TypeWithMapId withMapId = new TypeWithMapId(); + withMapId.setFirstname("Walter"); + withMapId.setLastname("White"); + + composite.save(withMapId) // + .as(StepVerifier::create) // + .expectNextCount(1) // + .verifyComplete(); + + composite.findAllById(Collections.singletonList(withMapId.getMapId())) // + .as(StepVerifier::create) // + .expectNextCount(1) // + .verifyComplete(); + } +}