Implement proper bulk loads in the transaction manager (#3233)

Previously we just looped over and loaded each entity one by one. It's
better to try to limit it to fewer queries (ideally one). Fortunately,
Hibernate supports bulk load of entities with both compound and simple
IDs.
This commit is contained in:
gbrodman
2026-09-21 19:19:40 +00:00
committed by GitHub
parent 7b26018490
commit 901e6ab789
3 changed files with 131 additions and 40 deletions
@@ -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<? extends VKey<? extends T>> keys) {
checkArgumentNotNull(keys, "keys must be specified");
assertInTransaction();
return StreamSupport.stream(keys.spliterator(), false)
// Accept duplicate keys.
.distinct()
.map(
key ->
new SimpleEntry<VKey<? extends T>, 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<Class<? extends T>, VKey<? extends T>> keysByObjectType =
Multimaps.index(Streams.stream(keys).distinct().collect(toImmutableList()), VKey::getKind);
ImmutableMap.Builder<VKey<? extends T>, T> builder = new ImmutableMap.Builder<>();
for (Class<? extends T> objectClass : keysByObjectType.keySet()) {
ImmutableList<VKey<? extends T>> singleObjectTypeKeys = keysByObjectType.get(objectClass);
ImmutableList<Serializable> ids =
singleObjectTypeKeys.stream().map(VKey::getKey).collect(toImmutableList());
// Note: Hibernate batches SQL queries for us if necessary under the hood
List<? extends T> 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 <T> ImmutableList<T> loadByEntitiesIfPresent(Iterable<T> entities) {
return Streams.stream(entities)
.filter(this::exists)
.map(this::loadByEntity)
.collect(toImmutableList());
checkArgumentNotNull(entities, "entities must be specified");
assertInTransaction();
ImmutableList<VKey<T>> 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> 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<VKey<T>> optionalKey = getKeyFromEntity(entity);
return loadByKey(
optionalKey.orElseThrow(
() -> new NoSuchElementException(String.format("No entity found %s", entity))));
}
@Override
public <T> ImmutableList<T> loadByEntities(Iterable<T> entities) {
return Streams.stream(entities).map(this::loadByEntity).collect(toImmutableList());
checkArgumentNotNull(entities, "entities must be specified");
assertInTransaction();
ImmutableList.Builder<VKey<T>> 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 <T> Optional<VKey<T>> 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<T>) VKey.create(entity.getClass(), s));
}
/**
* A SQL Sequence based ID allocator that generates an ID from a monotonically increasing {@link
* AtomicLong}
@@ -608,41 +608,103 @@ class JpaTransactionManagerImplTest {
}
@Test
void loadByKeys_succeeds() {
void loadByKeysIfPresent_mixedEntityTypes_succeeds() {
persistResource(theEntity);
persistResource(compoundIdEntity);
tm().transact(
() -> {
ImmutableMap<VKey<? extends TestEntity>, TestEntity> results =
tm().loadByKeysIfPresent(ImmutableList.of(theEntityKey));
assertThat(results).containsExactly(theEntityKey, theEntity);
ImmutableMap<VKey<? extends ImmutableObject>, 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<VKey<? extends ImmutableObject>, 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<TestEntity> results =
ImmutableList<ImmutableObject> 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<TestEntity> results = tm().loadByEntities(ImmutableList.of(theEntity));
assertThat(results).containsExactly(theEntity);
assertDetachedFromEntityManager(results.get(0));
ImmutableList<ImmutableObject> 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;
@@ -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);
}