diff --git a/spring-jdbc/src/main/java/org/springframework/jdbc/support/incrementer/AbstractIdentityColumnMaxValueIncrementer.java b/spring-jdbc/src/main/java/org/springframework/jdbc/support/incrementer/AbstractIdentityColumnMaxValueIncrementer.java index fd65f94a3e8..cbf7432f221 100644 --- a/spring-jdbc/src/main/java/org/springframework/jdbc/support/incrementer/AbstractIdentityColumnMaxValueIncrementer.java +++ b/spring-jdbc/src/main/java/org/springframework/jdbc/support/incrementer/AbstractIdentityColumnMaxValueIncrementer.java @@ -96,8 +96,7 @@ public abstract class AbstractIdentityColumnMaxValueIncrementer extends Abstract try { stmt = con.createStatement(); DataSourceUtils.applyTransactionTimeout(stmt, getDataSource()); - this.valueCache = new long[getCacheSize()]; - this.nextValueIndex = 0; + long[] newValues = new long[getCacheSize()]; for (int i = 0; i < getCacheSize(); i++) { stmt.executeUpdate(getIncrementStatement()); ResultSet rs = stmt.executeQuery(getIdentityStatement()); @@ -105,13 +104,16 @@ public abstract class AbstractIdentityColumnMaxValueIncrementer extends Abstract if (!rs.next()) { throw new DataAccessResourceFailureException("Identity statement failed after inserting"); } - this.valueCache[i] = rs.getLong(1); + newValues[i] = rs.getLong(1); } finally { JdbcUtils.closeResultSet(rs); } } - stmt.executeUpdate(getDeleteStatement(this.valueCache)); + stmt.executeUpdate(getDeleteStatement(newValues)); + // Only expose the new values once the entire range has been obtained + this.valueCache = newValues; + this.nextValueIndex = 0; } catch (SQLException ex) { throw new DataAccessResourceFailureException("Could not increment identity", ex); diff --git a/spring-jdbc/src/test/java/org/springframework/jdbc/support/incrementer/DataFieldMaxValueIncrementerTests.java b/spring-jdbc/src/test/java/org/springframework/jdbc/support/incrementer/DataFieldMaxValueIncrementerTests.java index ff33ec17818..66e73fcee75 100644 --- a/spring-jdbc/src/test/java/org/springframework/jdbc/support/incrementer/DataFieldMaxValueIncrementerTests.java +++ b/spring-jdbc/src/test/java/org/springframework/jdbc/support/incrementer/DataFieldMaxValueIncrementerTests.java @@ -30,6 +30,7 @@ import org.springframework.dao.DataAccessResourceFailureException; import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.Assertions.assertThatExceptionOfType; import static org.mockito.BDDMockito.given; +import static org.mockito.BDDMockito.willReturn; import static org.mockito.BDDMockito.willThrow; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.times; @@ -106,6 +107,115 @@ class DataFieldMaxValueIncrementerTests { verify(connection, times(2)).close(); } + @Test + void hsqlMaxValueIncrementerWithIncrementFailure() throws SQLException { + given(dataSource.getConnection()).willReturn(connection); + given(connection.createStatement()).willReturn(statement); + willThrow(new SQLException("Cannot insert")).willReturn(1) + .given(statement).executeUpdate("insert into myseq values(null)"); + given(statement.executeQuery("select max(identity()) from myseq")).willReturn(resultSet); + given(resultSet.next()).willReturn(true); + given(resultSet.getLong(1)).willReturn(1L, 2L); + + HsqlMaxValueIncrementer incrementer = new HsqlMaxValueIncrementer(); + incrementer.setDataSource(dataSource); + incrementer.setIncrementerName("myseq"); + incrementer.setColumnName("seq"); + incrementer.setCacheSize(2); + incrementer.afterPropertiesSet(); + + assertThatExceptionOfType(DataAccessResourceFailureException.class) + .isThrownBy(incrementer::nextLongValue); + assertThat(incrementer.nextLongValue()).isEqualTo(1); + assertThat(incrementer.nextLongValue()).isEqualTo(2); + + verify(statement, times(3)).executeUpdate("insert into myseq values(null)"); + verify(statement).executeUpdate("delete from myseq where seq < 2"); + verify(connection, times(2)).close(); + } + + @Test + void hsqlMaxValueIncrementerWithDeleteFailure() throws SQLException { + given(dataSource.getConnection()).willReturn(connection); + given(connection.createStatement()).willReturn(statement); + willThrow(new SQLException("Cannot delete")).willReturn(1) + .given(statement).executeUpdate("delete from myseq where seq < 2"); + given(statement.executeQuery("select max(identity()) from myseq")).willReturn(resultSet); + given(resultSet.next()).willReturn(true); + given(resultSet.getLong(1)).willReturn(1L, 2L, 3L, 4L); + + HsqlMaxValueIncrementer incrementer = new HsqlMaxValueIncrementer(); + incrementer.setDataSource(dataSource); + incrementer.setIncrementerName("myseq"); + incrementer.setColumnName("seq"); + incrementer.setCacheSize(2); + incrementer.afterPropertiesSet(); + + assertThatExceptionOfType(DataAccessResourceFailureException.class) + .isThrownBy(incrementer::nextLongValue); + assertThat(incrementer.nextLongValue()).isEqualTo(3); + assertThat(incrementer.nextLongValue()).isEqualTo(4); + + verify(statement, times(4)).executeUpdate("insert into myseq values(null)"); + verify(statement).executeUpdate("delete from myseq where seq < 4"); + verify(connection, times(2)).close(); + } + + @Test + void hsqlMaxValueIncrementerWithEmptyIdentityResult() throws SQLException { + given(dataSource.getConnection()).willReturn(connection); + given(connection.createStatement()).willReturn(statement); + given(statement.executeQuery("select max(identity()) from myseq")).willReturn(resultSet); + given(resultSet.next()).willReturn(false).willReturn(true); + given(resultSet.getLong(1)).willReturn(1L, 2L); + + HsqlMaxValueIncrementer incrementer = new HsqlMaxValueIncrementer(); + incrementer.setDataSource(dataSource); + incrementer.setIncrementerName("myseq"); + incrementer.setColumnName("seq"); + incrementer.setCacheSize(2); + incrementer.afterPropertiesSet(); + + assertThatExceptionOfType(DataAccessResourceFailureException.class) + .isThrownBy(incrementer::nextLongValue) + .withMessage("Identity statement failed after inserting") + .withNoCause(); + assertThat(incrementer.nextLongValue()).isEqualTo(1); + assertThat(incrementer.nextLongValue()).isEqualTo(2); + + verify(statement, times(3)).executeUpdate("insert into myseq values(null)"); + verify(statement).executeUpdate("delete from myseq where seq < 2"); + verify(connection, times(2)).close(); + } + + @Test + void hsqlMaxValueIncrementerWithPartialIncrementFailure() throws SQLException { + given(dataSource.getConnection()).willReturn(connection); + given(connection.createStatement()).willReturn(statement); + willReturn(1).willThrow(new SQLException("Cannot insert")).willReturn(1) + .given(statement).executeUpdate("insert into myseq values(null)"); + given(statement.executeQuery("select max(identity()) from myseq")).willReturn(resultSet); + given(resultSet.next()).willReturn(true); + given(resultSet.getLong(1)).willReturn(1L, 2L, 3L); + + HsqlMaxValueIncrementer incrementer = new HsqlMaxValueIncrementer(); + incrementer.setDataSource(dataSource); + incrementer.setIncrementerName("myseq"); + incrementer.setColumnName("seq"); + incrementer.setCacheSize(2); + incrementer.afterPropertiesSet(); + + assertThatExceptionOfType(DataAccessResourceFailureException.class) + .isThrownBy(incrementer::nextLongValue); + // Values obtained before the partial failure must not be served + assertThat(incrementer.nextLongValue()).isEqualTo(2); + assertThat(incrementer.nextLongValue()).isEqualTo(3); + + verify(statement, times(4)).executeUpdate("insert into myseq values(null)"); + verify(statement).executeUpdate("delete from myseq where seq < 3"); + verify(connection, times(2)).close(); + } + @Test void hsqlMaxValueIncrementerWithDeleteSpecificValues() throws SQLException { given(dataSource.getConnection()).willReturn(connection);