diff --git a/spring-batch-core-tests/src/test/java/org/springframework/batch/core/test/step/FaultTolerantStepFactoryBeanRollbackTests.java b/spring-batch-core-tests/src/test/java/org/springframework/batch/core/test/step/FaultTolerantStepFactoryBeanRollbackTests.java new file mode 100644 index 000000000..c64283cf1 --- /dev/null +++ b/spring-batch-core-tests/src/test/java/org/springframework/batch/core/test/step/FaultTolerantStepFactoryBeanRollbackTests.java @@ -0,0 +1,279 @@ +package org.springframework.batch.core.test.step; + +import static org.junit.Assert.assertEquals; + +import java.sql.ResultSet; +import java.sql.SQLException; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Collection; +import java.util.Collections; +import java.util.HashMap; +import java.util.List; +import java.util.Map; + +import javax.sql.DataSource; + +import org.apache.commons.logging.Log; +import org.apache.commons.logging.LogFactory; +import org.junit.Before; +import org.junit.Test; +import org.junit.runner.RunWith; +import org.springframework.batch.core.BatchStatus; +import org.springframework.batch.core.JobExecution; +import org.springframework.batch.core.JobParameters; +import org.springframework.batch.core.Step; +import org.springframework.batch.core.StepExecution; +import org.springframework.batch.core.repository.JobRepository; +import org.springframework.batch.core.step.item.FaultTolerantStepFactoryBean; +import org.springframework.batch.item.ItemProcessor; +import org.springframework.batch.item.ItemReader; +import org.springframework.batch.item.ItemWriter; +import org.springframework.batch.item.ParseException; +import org.springframework.batch.item.UnexpectedInputException; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.jdbc.core.simple.ParameterizedRowMapper; +import org.springframework.jdbc.core.simple.SimpleJdbcTemplate; +import org.springframework.scheduling.concurrent.ThreadPoolTaskExecutor; +import org.springframework.test.context.ContextConfiguration; +import org.springframework.test.context.junit4.SpringJUnit4ClassRunner; +import org.springframework.test.jdbc.SimpleJdbcTestUtils; +import org.springframework.transaction.PlatformTransactionManager; +import org.springframework.util.Assert; + +/** + * Tests for {@link FaultTolerantStepFactoryBean}. + */ +@ContextConfiguration(locations = "/simple-job-launcher-context.xml") +@RunWith(SpringJUnit4ClassRunner.class) +public class FaultTolerantStepFactoryBeanRollbackTests { + + private static final int MAX_COUNT = 1000; + + private final Log logger = LogFactory.getLog(getClass()); + + private FaultTolerantStepFactoryBean factory; + + private SkipReaderStub reader; + + private SkipProcessorStub processor; + + private SkipWriterStub writer; + + private JobExecution jobExecution; + + private StepExecution stepExecution; + + @Autowired + private DataSource dataSource; + + @Autowired + private JobRepository repository; + + @Autowired + private PlatformTransactionManager transactionManager; + + @SuppressWarnings("unchecked") + @Before + public void setUp() throws Exception { + + reader = new SkipReaderStub(); + writer = new SkipWriterStub(dataSource); + processor = new SkipProcessorStub(dataSource); + + factory = new FaultTolerantStepFactoryBean(); + + factory.setBeanName("stepName"); + factory.setTransactionManager(transactionManager); + factory.setJobRepository(repository); + factory.setCommitInterval(3); + factory.setSkipLimit(10); + ThreadPoolTaskExecutor taskExecutor = new ThreadPoolTaskExecutor(); + taskExecutor.setCorePoolSize(3); + taskExecutor.setMaxPoolSize(6); + taskExecutor.setQueueCapacity(0); + taskExecutor.afterPropertiesSet(); + factory.setTaskExecutor(taskExecutor); + + factory.setSkippableExceptionClasses(getExceptionMap(Exception.class)); + + } + + @Test + public void testUpdatesNoRollback() throws Exception { + + SimpleJdbcTemplate jdbcTemplate = new SimpleJdbcTemplate(dataSource); + + writer.write(Arrays.asList("foo", "bar")); + processor.process("spam"); + assertEquals(3, SimpleJdbcTestUtils.countRowsInTable(jdbcTemplate, "ERROR_LOG")); + + writer.clear(); + processor.clear(); + assertEquals(0, SimpleJdbcTestUtils.countRowsInTable(jdbcTemplate, "ERROR_LOG")); + + } + + @Test + public void testMultithreadedSkipInWriter() throws Throwable { + + jobExecution = repository.createJobExecution("skipJob", new JobParameters()); + + for (int i = 0; i < MAX_COUNT; i++) { + + SimpleJdbcTemplate jdbcTemplate = new SimpleJdbcTemplate(dataSource); + assertEquals(0, SimpleJdbcTestUtils.countRowsInTable(jdbcTemplate, "ERROR_LOG")); + + try { + + reader.clear(); + reader.setItems("1", "2", "3", "4", "5"); + factory.setItemReader(reader); + writer.clear(); + factory.setItemWriter(writer); + processor.clear(); + factory.setItemProcessor(processor); + + writer.setFailures("1", "2", "3", "4", "5"); + + Step step = (Step) factory.getObject(); + + stepExecution = jobExecution.createStepExecution(factory.getName()); + repository.add(stepExecution); + step.execute(stepExecution); + assertEquals(BatchStatus.COMPLETED, stepExecution.getStatus()); + + assertEquals("[]", writer.getCommitted().toString()); + assertEquals("[]", processor.getCommitted().toString()); + assertEquals(5, stepExecution.getSkipCount()); + + } + catch (Throwable e) { + logger.info("Failed on iteration " + i + " of " + MAX_COUNT); + throw e; + } + + } + + } + + private Map, Boolean> getExceptionMap(Class... args) { + Map, Boolean> map = new HashMap, Boolean>(); + for (Class arg : args) { + map.put(arg, true); + } + return map; + } + + private static class SkipReaderStub implements ItemReader { + + private String[] items; + + private int counter = -1; + + public SkipReaderStub() throws Exception { + super(); + } + + public void setItems(String... items) { + Assert.isTrue(counter < 0, "Items cannot be set once reading has started"); + this.items = items; + } + + public void clear() { + counter = -1; + } + + public synchronized String read() throws Exception, UnexpectedInputException, ParseException { + counter++; + if (counter >= items.length) { + return null; + } + String item = items[counter]; + return item; + } + } + + private static class SkipWriterStub implements ItemWriter { + + private List written = new ArrayList(); + + private Collection failures = Collections.emptySet(); + + private SimpleJdbcTemplate jdbcTemplate; + + public SkipWriterStub(DataSource dataSource) { + jdbcTemplate = new SimpleJdbcTemplate(dataSource); + } + + public void setFailures(String... failures) { + this.failures = Arrays.asList(failures); + } + + public List getCommitted() { + return jdbcTemplate.query("SELECT MESSAGE from ERROR_LOG where STEP_NAME='written'", + new ParameterizedRowMapper() { + public String mapRow(ResultSet rs, int rowNum) throws SQLException { + return rs.getString(1); + } + }); + } + + public void clear() { + written.clear(); + jdbcTemplate.update("DELETE FROM ERROR_LOG where STEP_NAME='written'"); + } + + public void write(List items) throws Exception { + for (String item : items) { + written.add(item); + jdbcTemplate.update("INSERT INTO ERROR_LOG (MESSAGE, STEP_NAME) VALUES (?, ?)", item, "written"); + checkFailure(item); + } + } + + private void checkFailure(String item) { + if (failures.contains(item)) { + throw new RuntimeException("Planned failure"); + } + } + } + + private static class SkipProcessorStub implements ItemProcessor { + + private final Log logger = LogFactory.getLog(getClass()); + + private List processed = new ArrayList(); + + private SimpleJdbcTemplate jdbcTemplate; + + /** + * @param dataSource + */ + public SkipProcessorStub(DataSource dataSource) { + jdbcTemplate = new SimpleJdbcTemplate(dataSource); + } + + public List getCommitted() { + return jdbcTemplate.query("SELECT MESSAGE from ERROR_LOG where STEP_NAME='processed'", + new ParameterizedRowMapper() { + public String mapRow(ResultSet rs, int rowNum) throws SQLException { + return rs.getString(1); + } + }); + } + + public void clear() { + processed.clear(); + jdbcTemplate.update("DELETE FROM ERROR_LOG where STEP_NAME='processed'"); + } + + public String process(String item) throws Exception { + processed.add(item); + logger.info("Processed item: "+item); + jdbcTemplate.update("INSERT INTO ERROR_LOG (MESSAGE, STEP_NAME) VALUES (?, ?)", item, "processed"); + return item; + } + } + +} diff --git a/spring-batch-core-tests/src/test/resources/log4j.properties b/spring-batch-core-tests/src/test/resources/log4j.properties index dae3702b2..bdb054942 100644 --- a/spring-batch-core-tests/src/test/resources/log4j.properties +++ b/spring-batch-core-tests/src/test/resources/log4j.properties @@ -9,5 +9,9 @@ log4j.category.org.apache.activemq=ERROR log4j.category.org.springframework.jdbc=INFO log4j.category.org.springframework.jms=INFO log4j.category.org.springframework.batch=INFO +#log4j.category.org.springframework.batch.core.scope=DEBUG +log4j.category.org.springframework.batch.core.step.item=INFO +log4j.category.org.springframework.batch.core.step=DEBUG +log4j.category.org.springframework.batch.core.test=DEBUG log4j.category.org.springframework.retry=INFO # log4j.category.org.springframework.beans.factory.config=TRACE diff --git a/spring-batch-core/src/main/java/org/springframework/batch/core/step/ThreadStepInterruptionPolicy.java b/spring-batch-core/src/main/java/org/springframework/batch/core/step/ThreadStepInterruptionPolicy.java index c1f72fd19..d551350de 100644 --- a/spring-batch-core/src/main/java/org/springframework/batch/core/step/ThreadStepInterruptionPolicy.java +++ b/spring-batch-core/src/main/java/org/springframework/batch/core/step/ThreadStepInterruptionPolicy.java @@ -30,10 +30,8 @@ import org.springframework.batch.core.StepExecution; */ public class ThreadStepInterruptionPolicy implements StepInterruptionPolicy { - protected static final Log logger = LogFactory - .getLog(ThreadStepInterruptionPolicy.class); + protected static final Log logger = LogFactory.getLog(ThreadStepInterruptionPolicy.class); - /** * Returns if the current job lifecycle has been interrupted by checking if * the current thread is interrupted. @@ -43,7 +41,7 @@ public class ThreadStepInterruptionPolicy implements StepInterruptionPolicy { if (isInterrupted(stepExecution)) { throw new JobInterruptedException("Job interrupted status detected."); } - + } /** @@ -51,9 +49,15 @@ public class ThreadStepInterruptionPolicy implements StepInterruptionPolicy { * @return true if the job has been interrupted */ private boolean isInterrupted(StepExecution stepExecution) { - boolean interrupted = (Thread.currentThread().isInterrupted() || stepExecution.isTerminateOnly()); - if(interrupted){ - logger.error("Step interrupted"); + boolean interrupted = Thread.currentThread().isInterrupted(); + if (interrupted) { + logger.info("Step interrupted through Thread API"); + } + else { + interrupted = stepExecution.isTerminateOnly(); + if (interrupted) { + logger.info("Step interrupted through StepExecution"); + } } return interrupted; } diff --git a/spring-batch-core/src/main/java/org/springframework/batch/core/step/item/Chunk.java b/spring-batch-core/src/main/java/org/springframework/batch/core/step/item/Chunk.java index 780e68929..bf0f80d01 100644 --- a/spring-batch-core/src/main/java/org/springframework/batch/core/step/item/Chunk.java +++ b/spring-batch-core/src/main/java/org/springframework/batch/core/step/item/Chunk.java @@ -239,6 +239,11 @@ public class Chunk implements Iterable { iterator.remove(); } + @Override + public String toString() { + return String.format("[items=%s, skips=%s]", items, skips); + } + } } \ No newline at end of file diff --git a/spring-batch-infrastructure/src/main/java/org/springframework/batch/support/transaction/TransactionAwareProxyFactory.java b/spring-batch-infrastructure/src/main/java/org/springframework/batch/support/transaction/TransactionAwareProxyFactory.java index abf6b8b3a..4bc720eb8 100644 --- a/spring-batch-infrastructure/src/main/java/org/springframework/batch/support/transaction/TransactionAwareProxyFactory.java +++ b/spring-batch-infrastructure/src/main/java/org/springframework/batch/support/transaction/TransactionAwareProxyFactory.java @@ -165,9 +165,9 @@ public class TransactionAwareProxyFactory { private class TargetSynchronization extends TransactionSynchronizationAdapter { - T cache; + private final T cache; - Object key; + private final Object key; public TargetSynchronization(Object key, T cache) { super();