diff --git a/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/repository/support/SimpleMongoRepository.java b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/repository/support/SimpleMongoRepository.java index fbfbd0da7..9fca7fb69 100644 --- a/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/repository/support/SimpleMongoRepository.java +++ b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/repository/support/SimpleMongoRepository.java @@ -20,7 +20,9 @@ import static org.springframework.data.mongodb.core.query.Criteria.*; import java.io.Serializable; import java.util.ArrayList; import java.util.Collections; +import java.util.HashSet; import java.util.List; +import java.util.Set; import org.springframework.data.domain.Page; import org.springframework.data.domain.PageImpl; @@ -171,10 +173,23 @@ public class SimpleMongoRepository implements Paging * @see org.springframework.data.repository.CrudRepository#findAll() */ public List findAll() { - return findAll(new Query()); } + /* + * (non-Javadoc) + * @see org.springframework.data.repository.CrudRepository#findAll(java.lang.Iterable) + */ + public Iterable findAll(Iterable ids) { + + Set parameters = new HashSet(); + for (ID id : ids) { + parameters.add(id); + } + + return findAll(new Query(new Criteria(entityInformation.getIdAttribute()).in(parameters))); + } + /* * (non-Javadoc) * @see org.springframework.data.repository.PagingAndSortingRepository#findAll(org.springframework.data.domain.Pageable) diff --git a/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/repository/AbstractPersonRepositoryIntegrationTests.java b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/repository/AbstractPersonRepositoryIntegrationTests.java index 4679a2678..19b32df3b 100644 --- a/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/repository/AbstractPersonRepositoryIntegrationTests.java +++ b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/repository/AbstractPersonRepositoryIntegrationTests.java @@ -90,6 +90,14 @@ public abstract class AbstractPersonRepositoryIntegrationTests { assertThat(result.containsAll(all), is(true)); } + @Test + public void findsAllWithGivenIds() { + + Iterable result = repository.findAll(Arrays.asList(dave.id, boyd.id)); + assertThat(result, hasItems(dave, boyd)); + assertThat(result, not(hasItems(oliver, carter, stefan, leroi, alicia))); + } + @Test public void deletesPersonCorrectly() throws Exception {