Set first batch args once while executing batch with generated keys

`batchArgs[0]` in `NamedParameterJdbcTemplate` and `indexedBatch[0]` in
`DefaultJdbcClient` are already set by
`pscf::newPreparedStatementCreator`, so it is unnecessary to set them
again via `BatchPreparedStatementSetter::setValues`.

Closes gh-37308

Signed-off-by: Yanming Zhou <zhouyanming@gmail.com>
This commit is contained in:
Yanming Zhou
2026-09-21 14:58:52 +02:00
committed by GitHub
parent 5772133ba7
commit 6e5efdc89b
5 changed files with 140 additions and 0 deletions
@@ -69,6 +69,7 @@ import org.springframework.util.ConcurrentLruCache;
*
* @author Thomas Risberg
* @author Juergen Hoeller
* @author Yanming Zhou
* @since 2.0
* @see NamedParameterJdbcOperations
* @see SqlParameterSource
@@ -423,6 +424,10 @@ public class NamedParameterJdbcTemplate implements NamedParameterJdbcOperations
return getJdbcOperations().batchUpdate(psc, new BatchPreparedStatementSetter() {
@Override
public void setValues(PreparedStatement ps, int i) throws SQLException {
if (i == 0) {
// batchArgs[0] is already set by pscf.newPreparedStatementCreator()
return;
}
@Nullable Object[] values = NamedParameterUtils.buildValueArray(parsedSql, batchArgs[i], null);
pscf.newPreparedStatementSetter(values).setValues(ps);
}
@@ -445,6 +445,10 @@ final class DefaultJdbcClient implements JdbcClient {
return classicOps.batchUpdate(psc, new BatchPreparedStatementSetter() {
@Override
public void setValues(PreparedStatement ps, int i) throws SQLException {
if (i == 0) {
// indexedBatch[0] is already set by pscf.newPreparedStatementCreator()
return;
}
pscf.newPreparedStatementSetter(indexedBatch.get(i)).setValues(ps);
}
@@ -21,6 +21,7 @@ import java.sql.DatabaseMetaData;
import java.sql.PreparedStatement;
import java.sql.ResultSet;
import java.sql.SQLException;
import java.sql.Statement;
import java.sql.Types;
import java.util.ArrayList;
import java.util.Arrays;
@@ -43,10 +44,12 @@ import org.springframework.jdbc.core.JdbcOperations;
import org.springframework.jdbc.core.JdbcTemplate;
import org.springframework.jdbc.core.PreparedStatementCallback;
import org.springframework.jdbc.core.SqlParameterValue;
import org.springframework.jdbc.support.GeneratedKeyHolder;
import static org.assertj.core.api.Assertions.assertThat;
import static org.assertj.core.api.Assertions.assertThatIllegalArgumentException;
import static org.mockito.ArgumentMatchers.anyString;
import static org.mockito.ArgumentMatchers.eq;
import static org.mockito.BDDMockito.given;
import static org.mockito.Mockito.atLeastOnce;
import static org.mockito.Mockito.inOrder;
@@ -60,6 +63,7 @@ import static org.mockito.Mockito.verify;
* @author Chris Beams
* @author Nikita Khateev
* @author Fedor Bobin
* @author Yanming Zhou
*/
class NamedParameterJdbcTemplateTests {
@@ -70,6 +74,9 @@ class NamedParameterJdbcTemplateTests {
private static final String SELECT_NO_PARAMETERS =
"select id, forename from custmr";
private static final String INSERT_NAMED_PARAMETERS =
"insert into custmr(forename,country) values (:forename,:country)";
private static final String UPDATE_NAMED_PARAMETERS =
"update seat_status set booking_id = null where performance_id = :perfId and price_band_id = :priceId";
private static final String UPDATE_NAMED_PARAMETERS_PARSED =
@@ -580,4 +587,30 @@ class NamedParameterJdbcTemplateTests {
verify(connection, atLeastOnce()).close();
}
@Test
void batchUpdateWithGeneratedKeys() throws Exception {
final SqlParameterSource[] batchArgs = new SqlParameterSource[2];
batchArgs[0] = new MapSqlParameterSource(Map.of("forename", "foo", "country", "UK"));
batchArgs[1] = new MapSqlParameterSource(Map.of("forename", "bar", "country", "US"));
final int[] rowsAffected = new int[] {1, 1};
given(connection.prepareStatement(anyString(), eq(Statement.RETURN_GENERATED_KEYS))).willReturn(preparedStatement);
given(preparedStatement.executeBatch()).willReturn(rowsAffected);
given(connection.getMetaData()).willReturn(databaseMetaData);
namedParameterTemplate = new NamedParameterJdbcTemplate(new JdbcTemplate(dataSource, false));
int[] actualRowsAffected = namedParameterTemplate.batchUpdate(INSERT_NAMED_PARAMETERS, batchArgs, new GeneratedKeyHolder());
assertThat(actualRowsAffected.length).as("executed 2 updates").isEqualTo(2);
assertThat(actualRowsAffected[0]).isEqualTo(rowsAffected[0]);
assertThat(actualRowsAffected[1]).isEqualTo(rowsAffected[1]);
verify(connection).prepareStatement("insert into custmr(forename,country) values (?,?)", Statement.RETURN_GENERATED_KEYS);
verify(preparedStatement).setString(1, "foo");
verify(preparedStatement).setString(2, "UK");
verify(preparedStatement).setString(1, "bar");
verify(preparedStatement).setString(2, "US");
verify(preparedStatement, times(2)).addBatch();
verify(preparedStatement, atLeastOnce()).close();
verify(connection, atLeastOnce()).close();
}
}
@@ -47,6 +47,7 @@ import static org.mockito.Mockito.verify;
/**
* @author Juergen Hoeller
* @author Yanming Zhou
* @since 6.1
*/
class JdbcClientIndexedParameterTests {
@@ -362,6 +363,30 @@ class JdbcClientIndexedParameterTests {
verify(connection).close();
}
@Test
void batchUpdateWithGeneratedKeys() throws SQLException {
given(resultSetMetaData.getColumnCount()).willReturn(1);
given(resultSetMetaData.getColumnLabel(1)).willReturn("1");
given(resultSet.getMetaData()).willReturn(resultSetMetaData);
given(resultSet.next()).willReturn(true, false);
given(resultSet.getObject(1)).willReturn(11);
given(preparedStatement.executeUpdate()).willReturn(1);
given(preparedStatement.getGeneratedKeys()).willReturn(resultSet);
given(connection.prepareStatement(INSERT_GENERATE_KEYS, PreparedStatement.RETURN_GENERATED_KEYS))
.willReturn(preparedStatement);
KeyHolder generatedKeyHolder = new GeneratedKeyHolder();
int[] rowsAffected = client.sql(INSERT_GENERATE_KEYS).batch().param("rod").add().update(generatedKeyHolder);
assertThat(rowsAffected).isEqualTo(new int[] { 1 });
assertThat(generatedKeyHolder.getKeyList()).hasSize(1);
assertThat(generatedKeyHolder.getKey()).isEqualTo(11);
verify(preparedStatement).setString(1, "rod");
verify(resultSet).close();
verify(preparedStatement).close();
verify(connection).close();
}
@Test
void updateWithGeneratedKeysAndKeyColumnNames() throws SQLException {
given(resultSetMetaData.getColumnCount()).willReturn(1);
@@ -386,4 +411,28 @@ class JdbcClientIndexedParameterTests {
verify(connection).close();
}
@Test
void batchUpdateWithGeneratedKeysAndKeyColumnNames() throws SQLException {
given(resultSetMetaData.getColumnCount()).willReturn(1);
given(resultSetMetaData.getColumnLabel(1)).willReturn("1");
given(resultSet.getMetaData()).willReturn(resultSetMetaData);
given(resultSet.next()).willReturn(true, false);
given(resultSet.getObject(1)).willReturn(11);
given(preparedStatement.executeUpdate()).willReturn(1);
given(preparedStatement.getGeneratedKeys()).willReturn(resultSet);
given(connection.prepareStatement(INSERT_GENERATE_KEYS, new String[] {"id"}))
.willReturn(preparedStatement);
KeyHolder generatedKeyHolder = new GeneratedKeyHolder();
int[] rowsAffected = client.sql(INSERT_GENERATE_KEYS).batch().param("rod").add().update(generatedKeyHolder, "id");
assertThat(rowsAffected).isEqualTo(new int[] { 1 });
assertThat(generatedKeyHolder.getKeyList()).hasSize(1);
assertThat(generatedKeyHolder.getKey()).isEqualTo(11);
verify(preparedStatement).setString(1, "rod");
verify(resultSet).close();
verify(preparedStatement).close();
verify(connection).close();
}
}
@@ -50,6 +50,7 @@ import static org.mockito.Mockito.verify;
/**
* @author Juergen Hoeller
* @author Yanming Zhou
* @since 6.1
*/
class JdbcClientNamedParameterTests {
@@ -429,6 +430,30 @@ class JdbcClientNamedParameterTests {
verify(connection).close();
}
@Test
void batchUpdateWithGeneratedKeys() throws SQLException {
given(resultSetMetaData.getColumnCount()).willReturn(1);
given(resultSetMetaData.getColumnLabel(1)).willReturn("1");
given(resultSet.getMetaData()).willReturn(resultSetMetaData);
given(resultSet.next()).willReturn(true, false);
given(resultSet.getObject(1)).willReturn(11);
given(preparedStatement.executeUpdate()).willReturn(1);
given(preparedStatement.getGeneratedKeys()).willReturn(resultSet);
given(connection.prepareStatement(INSERT_GENERATE_KEYS_PARSED, PreparedStatement.RETURN_GENERATED_KEYS))
.willReturn(preparedStatement);
KeyHolder generatedKeyHolder = new GeneratedKeyHolder();
int[] rowsAffected = client.sql(INSERT_GENERATE_KEYS).batch().param("name", "rod").add().update(generatedKeyHolder);
assertThat(rowsAffected).isEqualTo(new int[] { 1 });
assertThat(generatedKeyHolder.getKeyList()).hasSize(1);
assertThat(generatedKeyHolder.getKey()).isEqualTo(11);
verify(preparedStatement).setString(1, "rod");
verify(resultSet).close();
verify(preparedStatement).close();
verify(connection).close();
}
@Test
void updateWithGeneratedKeysAndKeyColumnNames() throws SQLException {
given(resultSetMetaData.getColumnCount()).willReturn(1);
@@ -453,4 +478,28 @@ class JdbcClientNamedParameterTests {
verify(connection).close();
}
@Test
void batchUpdateWithGeneratedKeysAndKeyColumnNames() throws SQLException {
given(resultSetMetaData.getColumnCount()).willReturn(1);
given(resultSetMetaData.getColumnLabel(1)).willReturn("1");
given(resultSet.getMetaData()).willReturn(resultSetMetaData);
given(resultSet.next()).willReturn(true, false);
given(resultSet.getObject(1)).willReturn(11);
given(preparedStatement.executeUpdate()).willReturn(1);
given(preparedStatement.getGeneratedKeys()).willReturn(resultSet);
given(connection.prepareStatement(INSERT_GENERATE_KEYS_PARSED, new String[] {"id"}))
.willReturn(preparedStatement);
KeyHolder generatedKeyHolder = new GeneratedKeyHolder();
int[] rowsAffected = client.sql(INSERT_GENERATE_KEYS).batch().param("name", "rod").add().update(generatedKeyHolder, "id");
assertThat(rowsAffected).isEqualTo(new int[] { 1 });
assertThat(generatedKeyHolder.getKeyList()).hasSize(1);
assertThat(generatedKeyHolder.getKey()).isEqualTo(11);
verify(preparedStatement).setString(1, "rod");
verify(resultSet).close();
verify(preparedStatement).close();
verify(connection).close();
}
}