From 2c979743667a7754e5c43ffbe4d42ed036db7f53 Mon Sep 17 00:00:00 2001 From: Robert McNees Date: Tue, 28 Mar 2023 08:37:30 -0400 Subject: [PATCH] Allow access to update counts in JdbcBatchItemWriter Resolves #3829 --- .../item/database/JdbcBatchItemWriter.java | 10 +++++++++ .../JdbcBatchItemWriterClassicTests.java | 21 +++++++++++++++++++ 2 files changed, 31 insertions(+) diff --git a/spring-batch-infrastructure/src/main/java/org/springframework/batch/item/database/JdbcBatchItemWriter.java b/spring-batch-infrastructure/src/main/java/org/springframework/batch/item/database/JdbcBatchItemWriter.java index e81ed4c49..2ac732e67 100644 --- a/spring-batch-infrastructure/src/main/java/org/springframework/batch/item/database/JdbcBatchItemWriter.java +++ b/spring-batch-infrastructure/src/main/java/org/springframework/batch/item/database/JdbcBatchItemWriter.java @@ -207,7 +207,17 @@ public class JdbcBatchItemWriter implements ItemWriter, InitializingBean { } } } + + processUpdateCounts(updateCounts); } } + /** + * Extension point to post process the update counts for each item. + * @param updateCounts the array of update counts for each item + */ + protected void processUpdateCounts(int[] updateCounts) { + // No Op + } + } diff --git a/spring-batch-infrastructure/src/test/java/org/springframework/batch/item/database/JdbcBatchItemWriterClassicTests.java b/spring-batch-infrastructure/src/test/java/org/springframework/batch/item/database/JdbcBatchItemWriterClassicTests.java index bc5eb0cc7..87470d0bf 100644 --- a/spring-batch-infrastructure/src/test/java/org/springframework/batch/item/database/JdbcBatchItemWriterClassicTests.java +++ b/spring-batch-infrastructure/src/test/java/org/springframework/batch/item/database/JdbcBatchItemWriterClassicTests.java @@ -22,6 +22,7 @@ import java.util.List; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; +import org.mockito.Mockito; import org.springframework.batch.item.Chunk; import org.springframework.dao.DataAccessException; @@ -35,6 +36,7 @@ import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertThrows; import static org.junit.jupiter.api.Assertions.assertTrue; import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.spy; import static org.mockito.Mockito.when; /** @@ -139,4 +141,23 @@ class JdbcBatchItemWriterClassicTests { assertTrue(list.contains("foo")); } + @Test + void testProcessUpdateCountsIsCalled() throws Exception { + JdbcBatchItemWriter customWriter = spy(new JdbcBatchItemWriter<>()); + + customWriter.setSql("SQL"); + customWriter.setJdbcTemplate(new NamedParameterJdbcTemplate(jdbcTemplate)); + customWriter.setItemPreparedStatementSetter((item, ps) -> list.add(item)); + customWriter.afterPropertiesSet(); + + ps.addBatch(); + int[] updateCounts = { 123 }; + when(ps.executeBatch()).thenReturn(updateCounts); + customWriter.write(Chunk.of("bar")); + assertEquals(2, list.size()); + assertTrue(list.contains("SQL")); + + Mockito.verify(customWriter, Mockito.times(1)).processUpdateCounts(updateCounts); + } + }