diff --git a/spring-jdbc/src/main/java/org/springframework/jdbc/core/metadata/TableMetaDataContext.java b/spring-jdbc/src/main/java/org/springframework/jdbc/core/metadata/TableMetaDataContext.java index 7c741ef9fb0..b4480a22519 100644 --- a/spring-jdbc/src/main/java/org/springframework/jdbc/core/metadata/TableMetaDataContext.java +++ b/spring-jdbc/src/main/java/org/springframework/jdbc/core/metadata/TableMetaDataContext.java @@ -205,13 +205,23 @@ public class TableMetaDataContext { if (generatedKeyNames.length > 0) { this.generatedKeyColumnsUsed = true; } - if (!declaredColumns.isEmpty()) { - return new ArrayList<>(declaredColumns); - } Set keys = CollectionUtils.newLinkedHashSet(generatedKeyNames.length); for (String key : generatedKeyNames) { keys.add(key.toUpperCase(Locale.ROOT)); } + if (!declaredColumns.isEmpty()) { + List overlapping = new ArrayList<>(); + for (String column : declaredColumns) { + if (keys.contains(column.toUpperCase(Locale.ROOT))) { + overlapping.add(column); + } + } + if (!overlapping.isEmpty()) { + throw new InvalidDataAccessApiUsageException( + "Declared columns " + overlapping + " must not overlap with generated key columns"); + } + return new ArrayList<>(declaredColumns); + } List columns = new ArrayList<>(); for (TableParameterMetaData meta : obtainMetaDataProvider().getTableParameterMetaData()) { if (!keys.contains(meta.getParameterName().toUpperCase(Locale.ROOT))) { diff --git a/spring-jdbc/src/test/java/org/springframework/jdbc/core/simple/TableMetaDataContextTests.java b/spring-jdbc/src/test/java/org/springframework/jdbc/core/simple/TableMetaDataContextTests.java index f2b06097ae7..d61dde4424e 100644 --- a/spring-jdbc/src/test/java/org/springframework/jdbc/core/simple/TableMetaDataContextTests.java +++ b/spring-jdbc/src/test/java/org/springframework/jdbc/core/simple/TableMetaDataContextTests.java @@ -29,11 +29,13 @@ import javax.sql.DataSource; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; +import org.springframework.dao.InvalidDataAccessApiUsageException; import org.springframework.jdbc.core.SqlParameterValue; import org.springframework.jdbc.core.metadata.TableMetaDataContext; import org.springframework.jdbc.core.namedparam.MapSqlParameterSource; 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.Mockito.atLeastOnce; import static org.mockito.Mockito.mock; @@ -153,4 +155,62 @@ class TableMetaDataContextTests { verify(columnsResultSet).close(); } + @Test + void overlappingDeclaredAndGeneratedKeyColumnsAreRejected() throws Exception { + initializeTwoColumnCustomersTable(); + context.setTableName("customers"); + + assertThatExceptionOfType(InvalidDataAccessApiUsageException.class) + .isThrownBy(() -> context.processMetaData( + dataSource, List.of("id", "name"), new String[] { "id" })) + .withMessage("Declared columns [id] must not overlap with generated key columns"); + } + + @Test + void overlappingDeclaredAndGeneratedKeyColumnsAreRejectedRegardlessOfCase() throws Exception { + initializeTwoColumnCustomersTable(); + context.setTableName("customers"); + + assertThatExceptionOfType(InvalidDataAccessApiUsageException.class) + .isThrownBy(() -> context.processMetaData( + dataSource, List.of("ID", "name"), new String[] { "id" })) + .withMessage("Declared columns [ID] must not overlap with generated key columns"); + } + + @Test + void declaredColumnsWithoutOverlapAreUsedAsIs() throws Exception { + initializeTwoColumnCustomersTable(); + MapSqlParameterSource map = new MapSqlParameterSource(); + map.addValue("name", "Sven"); + String[] keyCols = new String[] { "id" }; + context.setTableName("customers"); + context.processMetaData(dataSource, List.of("name"), keyCols); + List values = context.matchInParameterValuesWithInsertColumns(map); + String insertString = context.createInsertString(keyCols); + + assertThat(insertString).isEqualTo("INSERT INTO customers (name) VALUES(?)"); + assertThat(values).containsExactly("Sven"); + } + + private void initializeTwoColumnCustomersTable() throws Exception { + ResultSet metaDataResultSet = mock(); + given(metaDataResultSet.next()).willReturn(true, false); + given(metaDataResultSet.getString("TABLE_SCHEM")).willReturn("me"); + given(metaDataResultSet.getString("TABLE_NAME")).willReturn("customers"); + given(metaDataResultSet.getString("TABLE_TYPE")).willReturn("TABLE"); + + ResultSet columnsResultSet = mock(); + given(columnsResultSet.next()).willReturn(true, true, false); + given(columnsResultSet.getString("COLUMN_NAME")).willReturn("id", "name"); + given(columnsResultSet.getInt("DATA_TYPE")).willReturn(Types.INTEGER, Types.VARCHAR); + given(columnsResultSet.getBoolean("NULLABLE")).willReturn(false, true); + + given(databaseMetaData.getDatabaseProductName()).willReturn("MyDB"); + given(databaseMetaData.getDatabaseProductVersion()).willReturn("1.0"); + given(databaseMetaData.getUserName()).willReturn("me"); + given(databaseMetaData.storesLowerCaseIdentifiers()).willReturn(true); + given(databaseMetaData.getTables(null, null, "customers", null)).willReturn(metaDataResultSet); + given(databaseMetaData.getColumns(null, "me", "customers", null)).willReturn(columnsResultSet); + } + }