diff --git a/core/src/main/java/google/registry/persistence/transaction/JpaTransactionManagerImpl.java b/core/src/main/java/google/registry/persistence/transaction/JpaTransactionManagerImpl.java index fe975a1cc..2895380e0 100644 --- a/core/src/main/java/google/registry/persistence/transaction/JpaTransactionManagerImpl.java +++ b/core/src/main/java/google/registry/persistence/transaction/JpaTransactionManagerImpl.java @@ -17,19 +17,19 @@ package google.registry.persistence.transaction; import static com.google.common.base.Preconditions.checkArgument; import static com.google.common.base.Throwables.throwIfUnchecked; import static com.google.common.collect.ImmutableList.toImmutableList; -import static com.google.common.collect.ImmutableMap.toImmutableMap; import static com.google.common.collect.ImmutableSet.toImmutableSet; import static google.registry.config.RegistryConfig.getHibernateAllowNestedTransactions; import static google.registry.persistence.transaction.DatabaseException.throwIfSqlException; import static google.registry.util.PreconditionsUtils.checkArgumentNotNull; -import static java.util.AbstractMap.SimpleEntry; import static java.util.stream.Collectors.joining; import com.google.common.annotations.VisibleForTesting; import com.google.common.collect.ImmutableCollection; import com.google.common.collect.ImmutableList; +import com.google.common.collect.ImmutableListMultimap; import com.google.common.collect.ImmutableMap; import com.google.common.collect.ImmutableSet; +import com.google.common.collect.Multimaps; import com.google.common.collect.Streams; import com.google.common.flogger.FluentLogger; import com.google.common.flogger.StackSize; @@ -78,7 +78,6 @@ import java.util.function.Supplier; import java.util.function.UnaryOperator; import java.util.regex.Pattern; import java.util.stream.Stream; -import java.util.stream.StreamSupport; import javax.annotation.Nullable; import org.hibernate.Session; import org.hibernate.SessionFactory; @@ -473,23 +472,41 @@ public class JpaTransactionManagerImpl implements JpaTransactionManager { Iterable> keys) { checkArgumentNotNull(keys, "keys must be specified"); assertInTransaction(); - return StreamSupport.stream(keys.spliterator(), false) - // Accept duplicate keys. - .distinct() - .map( - key -> - new SimpleEntry, T>( - key, detach(getEntityManager().find(key.getKind(), key.getKey())))) - .filter(entry -> entry.getValue() != null) - .collect(toImmutableMap(Map.Entry::getKey, Map.Entry::getValue)); + // Group keys by entity type; T may be a common superclass with keys pointing to different + // concrete entity tables (e.g. EppResource, Domain, and Host). Session::findMultiple requires a + // single concrete entity class per call. + ImmutableListMultimap, VKey> keysByObjectType = + Multimaps.index(Streams.stream(keys).distinct().collect(toImmutableList()), VKey::getKind); + ImmutableMap.Builder, T> builder = new ImmutableMap.Builder<>(); + for (Class objectClass : keysByObjectType.keySet()) { + ImmutableList> singleObjectTypeKeys = keysByObjectType.get(objectClass); + ImmutableList ids = + singleObjectTypeKeys.stream().map(VKey::getKey).collect(toImmutableList()); + // Note: Hibernate batches SQL queries for us if necessary under the hood + List entities = + getEntityManager().unwrap(Session.class).findMultiple(objectClass, ids); + // Session::findMultiple keeps the entities in the same order with null values for missing + // keys. As a result, we can zip the keys+values as a map, ignoring null values. + for (int i = 0; i < ids.size(); i++) { + T entity = entities.get(i); + if (entity != null) { + builder.put(singleObjectTypeKeys.get(i), detach(entity)); + } + } + } + return builder.build(); } @Override public ImmutableList loadByEntitiesIfPresent(Iterable entities) { - return Streams.stream(entities) - .filter(this::exists) - .map(this::loadByEntity) - .collect(toImmutableList()); + checkArgumentNotNull(entities, "entities must be specified"); + assertInTransaction(); + ImmutableList> keys = + Streams.stream(entities) + .map(this::getKeyFromEntity) + .flatMap(Optional::stream) + .collect(toImmutableList()); + return loadByKeysIfPresent(keys).values().asList(); } @Override @@ -521,20 +538,24 @@ public class JpaTransactionManagerImpl implements JpaTransactionManager { public T loadByEntity(T entity) { checkArgumentNotNull(entity, "entity must be specified"); assertInTransaction(); - @SuppressWarnings("unchecked") - T returnValue = - (T) - loadByKey( - VKey.create( - entity.getClass(), - // Casting to Serializable is safe according to JPA (JSR 338 sec. 2.4). - (Serializable) emf.getPersistenceUnitUtil().getIdentifier(entity))); - return returnValue; + Optional> optionalKey = getKeyFromEntity(entity); + return loadByKey( + optionalKey.orElseThrow( + () -> new NoSuchElementException(String.format("No entity found %s", entity)))); } @Override public ImmutableList loadByEntities(Iterable entities) { - return Streams.stream(entities).map(this::loadByEntity).collect(toImmutableList()); + checkArgumentNotNull(entities, "entities must be specified"); + assertInTransaction(); + ImmutableList.Builder> keyBuilder = new ImmutableList.Builder<>(); + for (T entity : entities) { + keyBuilder.add( + getKeyFromEntity(entity) + .orElseThrow( + () -> new NoSuchElementException(String.format("No entity found %s", entity)))); + } + return loadByKeys(keyBuilder.build()).values().asList(); } @Override @@ -644,6 +665,14 @@ public class JpaTransactionManagerImpl implements JpaTransactionManager { return emf.getMetamodel().entity(clazz); } + @SuppressWarnings("unchecked") + private Optional> getKeyFromEntity(T entity) { + checkArgumentNotNull(entity, "entity must be specified"); + // Casting to Serializable is safe according to JPA (JSR 338 sec. 2.4). + Serializable key = (Serializable) emf.getPersistenceUnitUtil().getIdentifier(entity); + return Optional.ofNullable(key).map(s -> (VKey) VKey.create(entity.getClass(), s)); + } + /** * A SQL Sequence based ID allocator that generates an ID from a monotonically increasing {@link * AtomicLong} diff --git a/core/src/test/java/google/registry/persistence/transaction/JpaTransactionManagerImplTest.java b/core/src/test/java/google/registry/persistence/transaction/JpaTransactionManagerImplTest.java index 6f4012afc..580fe2882 100644 --- a/core/src/test/java/google/registry/persistence/transaction/JpaTransactionManagerImplTest.java +++ b/core/src/test/java/google/registry/persistence/transaction/JpaTransactionManagerImplTest.java @@ -608,41 +608,103 @@ class JpaTransactionManagerImplTest { } @Test - void loadByKeys_succeeds() { + void loadByKeysIfPresent_mixedEntityTypes_succeeds() { persistResource(theEntity); + persistResource(compoundIdEntity); tm().transact( () -> { - ImmutableMap, TestEntity> results = - tm().loadByKeysIfPresent(ImmutableList.of(theEntityKey)); - assertThat(results).containsExactly(theEntityKey, theEntity); + ImmutableMap, ImmutableObject> results = + tm().loadByKeysIfPresent( + ImmutableList.of( + theEntityKey, + compoundIdEntityKey, + VKey.create(TestEntity.class, "does-not-exist"))); + + assertThat(results) + .containsExactly(theEntityKey, theEntity, compoundIdEntityKey, compoundIdEntity); assertDetachedFromEntityManager(results.get(theEntityKey)); + assertDetachedFromEntityManager(results.get(compoundIdEntityKey)); }); } + @Test + void loadByKeys_succeeds() { + persistResource(theEntity); + persistResource(compoundIdEntity); + tm().transact( + () -> { + ImmutableMap, ImmutableObject> results = + tm().loadByKeys(ImmutableList.of(theEntityKey, compoundIdEntityKey)); + assertThat(results) + .containsExactly(theEntityKey, theEntity, compoundIdEntityKey, compoundIdEntity); + assertDetachedFromEntityManager(results.get(theEntityKey)); + assertDetachedFromEntityManager(results.get(compoundIdEntityKey)); + }); + } + + @Test + void loadByKeys_missingKey_throws() { + persistResource(theEntity); + assertThat( + assertThrows( + NoSuchElementException.class, + () -> + tm().transact( + () -> + tm().loadByKeys( + ImmutableList.of( + theEntityKey, + VKey.create(TestEntity.class, "does-not-exist")))))) + .hasMessageThat() + .contains("does-not-exist"); + } + @Test void loadByEntitiesIfPresent_succeeds() { persistResource(theEntity); + persistResource(compoundIdEntity); tm().transact( () -> { - ImmutableList results = + ImmutableList results = tm().loadByEntitiesIfPresent( - ImmutableList.of(theEntity, new TestEntity("does-not-exist", "bar"))); - assertThat(results).containsExactly(theEntity); - assertDetachedFromEntityManager(results.get(0)); + ImmutableList.of( + theEntity, + compoundIdEntity, + new TestEntity("does-not-exist", "bar"))); + assertThat(results).containsExactly(theEntity, compoundIdEntity); + results.forEach(DatabaseHelper::assertDetachedFromEntityManager); }); } @Test void loadByEntities_succeeds() { persistResource(theEntity); + persistResource(compoundIdEntity); tm().transact( () -> { - ImmutableList results = tm().loadByEntities(ImmutableList.of(theEntity)); - assertThat(results).containsExactly(theEntity); - assertDetachedFromEntityManager(results.get(0)); + ImmutableList results = + tm().loadByEntities(ImmutableList.of(theEntity, compoundIdEntity)); + assertThat(results).containsExactly(theEntity, compoundIdEntity); + results.forEach(DatabaseHelper::assertDetachedFromEntityManager); }); } + @Test + void loadByEntities_missingEntity_throws() { + persistResource(theEntity); + assertThat( + assertThrows( + NoSuchElementException.class, + () -> + tm().transact( + () -> + tm().loadByEntities( + ImmutableList.of( + theEntity, new TestEntity("does-not-exist", "bar")))))) + .hasMessageThat() + .contains("does-not-exist"); + } + @Test void loadAll_succeeds() { persistResources(moreEntities); @@ -911,7 +973,7 @@ class JpaTransactionManagerImplTest { } } - private static class CompoundId implements Serializable { + private static class CompoundId extends ImmutableObject implements Serializable { String name; int age; @@ -959,7 +1021,7 @@ class JpaTransactionManagerImplTest { } } - private static class NamedCompoundId implements Serializable { + private static class NamedCompoundId extends ImmutableObject implements Serializable { String nameField; int ageField; diff --git a/core/src/test/java/google/registry/testing/DatabaseHelper.java b/core/src/test/java/google/registry/testing/DatabaseHelper.java index d82688456..f61785091 100644 --- a/core/src/test/java/google/registry/testing/DatabaseHelper.java +++ b/core/src/test/java/google/registry/testing/DatabaseHelper.java @@ -912,7 +912,7 @@ public final class DatabaseHelper { .that(resource) .isNotInstanceOf(Buildable.Builder.class); } - tm().transact(() -> resources.forEach(e -> tm().put(e))); + tm().transact(() -> tm().putAll(ImmutableList.copyOf(resources))); maybeAdvanceClock(); return loadByEntitiesIfPresent(resources); }