From fb9a4f363da81b28ce539acc24ffa6052d905dcb Mon Sep 17 00:00:00 2001 From: Jens Schauder Date: Thu, 22 Dec 2022 15:57:42 +0100 Subject: [PATCH] Polishing. Original pull request #1396 See #1395 --- .../data/jdbc/core/JdbcAggregateTemplate.java | 28 ++++++++++++------- ...JdbcAggregateTemplateIntegrationTests.java | 1 + 2 files changed, 19 insertions(+), 10 deletions(-) diff --git a/spring-data-jdbc/src/main/java/org/springframework/data/jdbc/core/JdbcAggregateTemplate.java b/spring-data-jdbc/src/main/java/org/springframework/data/jdbc/core/JdbcAggregateTemplate.java index 02455011..c2368be8 100644 --- a/spring-data-jdbc/src/main/java/org/springframework/data/jdbc/core/JdbcAggregateTemplate.java +++ b/spring-data-jdbc/src/main/java/org/springframework/data/jdbc/core/JdbcAggregateTemplate.java @@ -170,6 +170,7 @@ public class JdbcAggregateTemplate implements JdbcAggregateOperations { @Override public Iterable saveAll(Iterable instances) { + Assert.notNull(instances, "Aggregate instances must not be null"); Assert.isTrue(instances.iterator().hasNext(), "Aggregate instances must not be empty"); List> entityAndChangeCreators = new ArrayList<>(); @@ -191,19 +192,22 @@ public class JdbcAggregateTemplate implements JdbcAggregateOperations { Assert.notNull(instance, "Aggregate instance must not be null"); - return performSave(new EntityAndChangeCreator<>( - instance, entity -> createInsertChange(prepareVersionForInsert(entity)))); + return performSave( + new EntityAndChangeCreator<>(instance, entity -> createInsertChange(prepareVersionForInsert(entity)))); } @Override public Iterable insertAll(Iterable instances) { + Assert.notNull(instances, "Aggregate instances must not be null"); Assert.isTrue(instances.iterator().hasNext(), "Aggregate instances must not be empty"); List> entityAndChangeCreators = new ArrayList<>(); for (T instance : instances) { - entityAndChangeCreators.add(new EntityAndChangeCreator<>( - instance, entity -> createInsertChange(prepareVersionForInsert(entity)))); + + Function> changeCreator = entity -> createInsertChange(prepareVersionForInsert(entity)); + EntityAndChangeCreator entityChange = new EntityAndChangeCreator<>(instance, changeCreator); + entityAndChangeCreators.add(entityChange); } return performSaveAll(entityAndChangeCreators); } @@ -220,19 +224,22 @@ public class JdbcAggregateTemplate implements JdbcAggregateOperations { Assert.notNull(instance, "Aggregate instance must not be null"); - return performSave(new EntityAndChangeCreator<>( - instance, entity -> createUpdateChange(prepareVersionForUpdate(entity)))); + return performSave( + new EntityAndChangeCreator<>(instance, entity -> createUpdateChange(prepareVersionForUpdate(entity)))); } @Override public Iterable updateAll(Iterable instances) { + Assert.notNull(instances, "Aggregate instances must not be null"); Assert.isTrue(instances.iterator().hasNext(), "Aggregate instances must not be empty"); List> entityAndChangeCreators = new ArrayList<>(); for (T instance : instances) { - entityAndChangeCreators.add(new EntityAndChangeCreator<>( - instance, entity -> createUpdateChange(prepareVersionForUpdate(entity)))); + + Function> changeCreator = entity -> createUpdateChange(prepareVersionForUpdate(entity)); + EntityAndChangeCreator entityChange = new EntityAndChangeCreator<>(instance, changeCreator); + entityAndChangeCreators.add(entityChange); } return performSaveAll(entityAndChangeCreators); } @@ -393,6 +400,7 @@ public class JdbcAggregateTemplate implements JdbcAggregateOperations { Map> groupedByType = new HashMap<>(); for (T instance : instances) { + Class type = instance.getClass(); final List list = groupedByType.computeIfAbsent(type, __ -> new ArrayList<>()); list.add(instance); @@ -474,13 +482,13 @@ public class JdbcAggregateTemplate implements JdbcAggregateOperations { } private List performSaveAll(Iterable> instances) { + BatchingAggregateChange> batchingAggregateChange = null; for (EntityAndChangeCreator instance : instances) { if (batchingAggregateChange == null) { // noinspection unchecked - batchingAggregateChange = BatchingAggregateChange.forSave( - (Class) ClassUtils.getUserClass(instance.entity)); + batchingAggregateChange = BatchingAggregateChange.forSave((Class) ClassUtils.getUserClass(instance.entity)); } batchingAggregateChange.add(beforeExecute(instance)); } diff --git a/spring-data-jdbc/src/test/java/org/springframework/data/jdbc/core/JdbcAggregateTemplateIntegrationTests.java b/spring-data-jdbc/src/test/java/org/springframework/data/jdbc/core/JdbcAggregateTemplateIntegrationTests.java index 4d92ffd3..e80c7e27 100644 --- a/spring-data-jdbc/src/test/java/org/springframework/data/jdbc/core/JdbcAggregateTemplateIntegrationTests.java +++ b/spring-data-jdbc/src/test/java/org/springframework/data/jdbc/core/JdbcAggregateTemplateIntegrationTests.java @@ -390,6 +390,7 @@ class JdbcAggregateTemplateIntegrationTests { AggregateWithImmutableVersion aggregate1 = new AggregateWithImmutableVersion(null, null); AggregateWithImmutableVersion aggregate2 = new AggregateWithImmutableVersion(null, null); AggregateWithImmutableVersion aggregate3 = new AggregateWithImmutableVersion(null, null); + Iterator savedAggregatesIterator = template .insertAll(List.of(aggregate1, aggregate2, aggregate3)).iterator(); assertThat(template.count(AggregateWithImmutableVersion.class)).isEqualTo(3);