From 8fa9e3e1d5bc603cf90cf833e955c6f5e343a22d Mon Sep 17 00:00:00 2001 From: Mark Paluch Date: Tue, 26 Sep 2023 11:34:26 +0200 Subject: [PATCH] Polishing. Simplify type and interface arrangement. See #1601 Original pull request: #1617 --- .../jdbc/core/convert/AggregateReader.java | 75 ++++++++++--------- .../RowDocumentResultSetExtractor.java | 2 +- .../SingleQueryDataAccessStrategy.java | 13 ++-- ...JdbcAggregateTemplateIntegrationTests.java | 18 ++++- 4 files changed, 63 insertions(+), 45 deletions(-) diff --git a/spring-data-jdbc/src/main/java/org/springframework/data/jdbc/core/convert/AggregateReader.java b/spring-data-jdbc/src/main/java/org/springframework/data/jdbc/core/convert/AggregateReader.java index c765c2eb..40708947 100644 --- a/spring-data-jdbc/src/main/java/org/springframework/data/jdbc/core/convert/AggregateReader.java +++ b/spring-data-jdbc/src/main/java/org/springframework/data/jdbc/core/convert/AggregateReader.java @@ -37,10 +37,10 @@ import org.springframework.data.relational.core.sqlgeneration.SingleQuerySqlGene import org.springframework.data.relational.core.sqlgeneration.SqlGenerator; import org.springframework.data.relational.domain.RowDocument; import org.springframework.data.util.Streamable; +import org.springframework.jdbc.core.ResultSetExtractor; import org.springframework.jdbc.core.namedparam.MapSqlParameterSource; import org.springframework.jdbc.core.namedparam.NamedParameterJdbcOperations; import org.springframework.lang.Nullable; -import org.springframework.util.Assert; /** * Reads complete Aggregates from the database, by generating appropriate SQL using a {@link SingleQuerySqlGenerator} @@ -53,13 +53,14 @@ import org.springframework.util.Assert; * @author Mark Paluch * @since 3.2 */ -class AggregateReader { +class AggregateReader implements PathToColumnMapping { private final RelationalPersistentEntity aggregate; private final Table table; private final SqlGenerator sqlGenerator; private final JdbcConverter converter; private final NamedParameterJdbcOperations jdbcTemplate; + private final AliasFactory aliasFactory; private final RowDocumentResultSetExtractor extractor; AggregateReader(Dialect dialect, JdbcConverter converter, AliasFactory aliasFactory, @@ -70,8 +71,25 @@ class AggregateReader { this.jdbcTemplate = jdbcTemplate; this.table = Table.create(aggregate.getQualifiedTableName()); this.sqlGenerator = new SingleQuerySqlGenerator(converter.getMappingContext(), aliasFactory, dialect, aggregate); - this.extractor = new RowDocumentResultSetExtractor(converter.getMappingContext(), - createPathToColumnMapping(aliasFactory)); + this.aliasFactory = aliasFactory; + this.extractor = new RowDocumentResultSetExtractor(converter.getMappingContext(), this); + } + + @Override + public String column(AggregatePath path) { + + String alias = aliasFactory.getColumnAlias(path); + + if (alias == null) { + throw new IllegalStateException(String.format("Alias for '%s' must not be null", path)); + } + + return alias; + } + + @Override + public String keyColumn(AggregatePath path) { + return aliasFactory.getKeyAlias(path); } @Nullable @@ -84,30 +102,34 @@ class AggregateReader { @Nullable public T findOne(Query query) { - - MapSqlParameterSource parameterSource = new MapSqlParameterSource(); - Condition condition = createCondition(query, parameterSource); - - return jdbcTemplate.query(sqlGenerator.findAll(condition), parameterSource, this::extractZeroOrOne); - } - - public List findAll() { - return jdbcTemplate.query(sqlGenerator.findAll(), this::extractAll); + return doFind(query, this::extractZeroOrOne); } public List findAllById(Iterable ids) { Collection identifiers = ids instanceof Collection idl ? idl : Streamable.of(ids).toList(); - Query query = Query.query(Criteria.where(aggregate.getRequiredIdProperty().getName()).in(identifiers)).limit(1); + Query query = Query.query(Criteria.where(aggregate.getRequiredIdProperty().getName()).in(identifiers)); return findAll(query); } + @SuppressWarnings("ConstantConditions") + public List findAll() { + return jdbcTemplate.query(sqlGenerator.findAll(), this::extractAll); + } + public List findAll(Query query) { + return doFind(query, this::extractAll); + } + + @SuppressWarnings("ConstantConditions") + private R doFind(Query query, ResultSetExtractor extractor) { MapSqlParameterSource parameterSource = new MapSqlParameterSource(); Condition condition = createCondition(query, parameterSource); - return jdbcTemplate.query(sqlGenerator.findAll(condition), parameterSource, this::extractAll); + String sql = sqlGenerator.findAll(condition); + + return jdbcTemplate.query(sql, parameterSource, extractor); } @Nullable @@ -128,7 +150,7 @@ class AggregateReader { * * @param rs the {@link ResultSet} from which to extract the data. Must not be {(}@literal null}. * @return a {@code List} of aggregates, fully converted. - * @throws SQLException + * @throws SQLException on underlying JDBC errors. */ private List extractAll(ResultSet rs) throws SQLException { @@ -146,10 +168,10 @@ class AggregateReader { * {@link RowDocumentResultSetExtractor} and the {@link JdbcConverter}. When used as a method reference this conforms * to the {@link org.springframework.jdbc.core.ResultSetExtractor} contract. * - * @param @param rs the {@link ResultSet} from which to extract the data. Must not be {(}@literal null}. + * @param rs the {@link ResultSet} from which to extract the data. Must not be {(}@literal null}. * @return The single instance when the conversion results in exactly one instance. If the {@literal ResultSet} is * empty, null is returned. - * @throws SQLException + * @throws SQLException on underlying JDBC errors. * @throws IncorrectResultSizeDataAccessException when the conversion yields more than one instance. */ @Nullable @@ -167,21 +189,4 @@ class AggregateReader { return null; } - private PathToColumnMapping createPathToColumnMapping(AliasFactory aliasFactory) { - return new PathToColumnMapping() { - @Override - public String column(AggregatePath path) { - - String alias = aliasFactory.getColumnAlias(path); - Assert.notNull(alias, () -> "alias for >" + path + "< must not be null"); - return alias; - } - - @Override - public String keyColumn(AggregatePath path) { - return aliasFactory.getKeyAlias(path); - } - }; - } - } diff --git a/spring-data-jdbc/src/main/java/org/springframework/data/jdbc/core/convert/RowDocumentResultSetExtractor.java b/spring-data-jdbc/src/main/java/org/springframework/data/jdbc/core/convert/RowDocumentResultSetExtractor.java index cb3df05b..45b26405 100644 --- a/spring-data-jdbc/src/main/java/org/springframework/data/jdbc/core/convert/RowDocumentResultSetExtractor.java +++ b/spring-data-jdbc/src/main/java/org/springframework/data/jdbc/core/convert/RowDocumentResultSetExtractor.java @@ -159,7 +159,7 @@ class RowDocumentResultSetExtractor { */ private boolean hasNext; - RowDocumentIterator(RelationalPersistentEntity entity, ResultSet resultSet) throws SQLException { + RowDocumentIterator(RelationalPersistentEntity entity, ResultSet resultSet) { ResultSetAdapter adapter = ResultSetAdapter.INSTANCE; diff --git a/spring-data-jdbc/src/main/java/org/springframework/data/jdbc/core/convert/SingleQueryDataAccessStrategy.java b/spring-data-jdbc/src/main/java/org/springframework/data/jdbc/core/convert/SingleQueryDataAccessStrategy.java index a609619c..d5fc206e 100644 --- a/spring-data-jdbc/src/main/java/org/springframework/data/jdbc/core/convert/SingleQueryDataAccessStrategy.java +++ b/spring-data-jdbc/src/main/java/org/springframework/data/jdbc/core/convert/SingleQueryDataAccessStrategy.java @@ -16,6 +16,7 @@ package org.springframework.data.jdbc.core.convert; +import java.util.List; import java.util.Optional; import org.springframework.data.domain.Pageable; @@ -56,22 +57,22 @@ class SingleQueryDataAccessStrategy implements ReadingDataAccessStrategy { } @Override - public Iterable findAll(Class domainType) { + public List findAll(Class domainType) { return getReader(domainType).findAll(); } @Override - public Iterable findAllById(Iterable ids, Class domainType) { + public List findAllById(Iterable ids, Class domainType) { return getReader(domainType).findAllById(ids); } @Override - public Iterable findAll(Class domainType, Sort sort) { + public List findAll(Class domainType, Sort sort) { throw new UnsupportedOperationException(); } @Override - public Iterable findAll(Class domainType, Pageable pageable) { + public List findAll(Class domainType, Pageable pageable) { throw new UnsupportedOperationException(); } @@ -81,12 +82,12 @@ class SingleQueryDataAccessStrategy implements ReadingDataAccessStrategy { } @Override - public Iterable findAll(Query query, Class domainType) { + public List findAll(Query query, Class domainType) { return getReader(domainType).findAll(query); } @Override - public Iterable findAll(Query query, Class domainType, Pageable pageable) { + public List findAll(Query query, Class domainType, Pageable pageable) { throw new UnsupportedOperationException(); } diff --git a/spring-data-jdbc/src/test/java/org/springframework/data/jdbc/core/AbstractJdbcAggregateTemplateIntegrationTests.java b/spring-data-jdbc/src/test/java/org/springframework/data/jdbc/core/AbstractJdbcAggregateTemplateIntegrationTests.java index e76825b2..858a69a7 100644 --- a/spring-data-jdbc/src/test/java/org/springframework/data/jdbc/core/AbstractJdbcAggregateTemplateIntegrationTests.java +++ b/spring-data-jdbc/src/test/java/org/springframework/data/jdbc/core/AbstractJdbcAggregateTemplateIntegrationTests.java @@ -23,8 +23,16 @@ import static org.springframework.data.jdbc.testing.TestConfiguration.*; import static org.springframework.data.jdbc.testing.TestDatabaseFeatures.Feature.*; import java.time.LocalDateTime; -import java.util.*; import java.util.ArrayList; +import java.util.Collections; +import java.util.HashMap; +import java.util.HashSet; +import java.util.Iterator; +import java.util.List; +import java.util.Map; +import java.util.Objects; +import java.util.Optional; +import java.util.Set; import java.util.function.Function; import java.util.stream.IntStream; @@ -49,6 +57,7 @@ import org.springframework.data.jdbc.core.convert.DataAccessStrategy; import org.springframework.data.jdbc.core.convert.JdbcConverter; import org.springframework.data.jdbc.testing.EnabledOnFeature; import org.springframework.data.jdbc.testing.IntegrationTest; +import org.springframework.data.jdbc.testing.TestClass; import org.springframework.data.jdbc.testing.TestConfiguration; import org.springframework.data.jdbc.testing.TestDatabaseFeatures; import org.springframework.data.mapping.context.InvalidPersistentPropertyPath; @@ -63,6 +72,7 @@ import org.springframework.data.relational.core.query.CriteriaDefinition; import org.springframework.data.relational.core.query.Query; import org.springframework.jdbc.core.namedparam.NamedParameterJdbcOperations; import org.springframework.test.context.ActiveProfiles; +import org.springframework.test.context.ContextConfiguration; /** * Integration tests for {@link JdbcAggregateTemplate}. @@ -1927,8 +1937,8 @@ abstract class AbstractJdbcAggregateTemplateIntegrationTests { static class Config { @Bean - Class testClass() { - return JdbcAggregateTemplateIntegrationTests.class; + TestClass testClass() { + return TestClass.of(JdbcAggregateTemplateIntegrationTests.class); } @Bean @@ -1938,9 +1948,11 @@ abstract class AbstractJdbcAggregateTemplateIntegrationTests { } } + @ContextConfiguration(classes = Config.class) static class JdbcAggregateTemplateIntegrationTests extends AbstractJdbcAggregateTemplateIntegrationTests {} @ActiveProfiles(value = PROFILE_SINGLE_QUERY_LOADING) + @ContextConfiguration(classes = Config.class) static class JdbcAggregateTemplateSingleQueryLoadingIntegrationTests extends AbstractJdbcAggregateTemplateIntegrationTests {