diff --git a/spring-batch-infrastructure/src/main/java/org/springframework/batch/item/data/RepositoryItemReader.java b/spring-batch-infrastructure/src/main/java/org/springframework/batch/item/data/RepositoryItemReader.java index 650abf22e..ea59a4b29 100644 --- a/spring-batch-infrastructure/src/main/java/org/springframework/batch/item/data/RepositoryItemReader.java +++ b/spring-batch-infrastructure/src/main/java/org/springframework/batch/item/data/RepositoryItemReader.java @@ -216,6 +216,11 @@ public class RepositoryItemReader extends AbstractItemCountingItemStreamItemR @Override protected void doClose() throws Exception { + synchronized (lock) { + current = 0; + page = 0; + results = null; + } } private Sort convertToSort(Map sorts) { diff --git a/spring-batch-infrastructure/src/test/java/org/springframework/batch/item/data/RepositoryItemReaderTests.java b/spring-batch-infrastructure/src/test/java/org/springframework/batch/item/data/RepositoryItemReaderTests.java index d7d5c24a3..83031a852 100644 --- a/spring-batch-infrastructure/src/test/java/org/springframework/batch/item/data/RepositoryItemReaderTests.java +++ b/spring-batch-infrastructure/src/test/java/org/springframework/batch/item/data/RepositoryItemReaderTests.java @@ -185,7 +185,7 @@ public class RepositoryItemReaderTests { public void testJumpToItem() throws Exception { reader.setPageSize(100); ArgumentCaptor pageRequestContainer = ArgumentCaptor.forClass(PageRequest.class); - when(repository.findAll(pageRequestContainer.capture())).thenReturn(new PageImpl(new ArrayList(){{ + when(repository.findAll(pageRequestContainer.capture())).thenReturn(new PageImpl(new ArrayList() {{ add(new Object()); }})); @@ -241,7 +241,7 @@ public class RepositoryItemReaderTests { reader.setPageSize(2); PageRequest request = new PageRequest(1, 2, new Sort(Direction.ASC, "id")); - when(repository.findAll(request)).thenReturn(new PageImpl(new ArrayList(){{ + when(repository.findAll(request)).thenReturn(new PageImpl(new ArrayList() {{ add("3"); add("4"); }})); @@ -274,7 +274,7 @@ public class RepositoryItemReaderTests { }})); request = new PageRequest(2, 2, new Sort(Direction.ASC, "id")); - when(repository.findAll(request)).thenReturn(new PageImpl(new ArrayList(){{ + when(repository.findAll(request)).thenReturn(new PageImpl(new ArrayList() {{ add("5"); add("6"); }})); @@ -294,6 +294,36 @@ public class RepositoryItemReaderTests { assertEquals("6", reader.read()); } + @Test + public void testResetOfPage() throws Exception { + reader.setPageSize(2); + + PageRequest request = new PageRequest(0, 2, new Sort(Direction.ASC, "id")); + when(repository.findAll(request)).thenReturn(new PageImpl(new ArrayList(){{ + add("1"); + add("2"); + }})); + + request = new PageRequest(1, 2, new Sort(Direction.ASC, "id")); + when(repository.findAll(request)).thenReturn(new PageImpl(new ArrayList() {{ + add("3"); + add("4"); + }})); + + ExecutionContext executionContext = new ExecutionContext(); + reader.open(executionContext); + + Object result = reader.read(); + reader.close(); + + assertEquals("1", result); + + reader.open(executionContext); + assertEquals("1", reader.read()); + assertEquals("2", reader.read()); + assertEquals("3", reader.read()); + } + public static interface TestRepository extends PagingAndSortingRepository { Page findFirstNames(Pageable pageable); }