diff --git a/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/MongoDatabaseUtils.java b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/MongoDatabaseUtils.java index 2c7d3903c..399286146 100644 --- a/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/MongoDatabaseUtils.java +++ b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/MongoDatabaseUtils.java @@ -79,7 +79,7 @@ public class MongoDatabaseUtils { * @param factory the {@link MongoDatabaseFactory} to get the {@link MongoDatabase} from. * @return the {@link MongoDatabase} that is potentially associated with a transactional {@link ClientSession}. */ - public static MongoDatabase getDatabase(String dbName, MongoDatabaseFactory factory) { + public static MongoDatabase getDatabase(@Nullable String dbName, MongoDatabaseFactory factory) { return doGetMongoDatabase(dbName, factory, SessionSynchronization.ON_ACTUAL_TRANSACTION); } @@ -94,7 +94,7 @@ public class MongoDatabaseUtils { * @param sessionSynchronization the synchronization to use. Must not be {@literal null}. * @return the {@link MongoDatabase} that is potentially associated with a transactional {@link ClientSession}. */ - public static MongoDatabase getDatabase(String dbName, MongoDatabaseFactory factory, + public static MongoDatabase getDatabase(@Nullable String dbName, MongoDatabaseFactory factory, SessionSynchronization sessionSynchronization) { return doGetMongoDatabase(dbName, factory, sessionSynchronization); } diff --git a/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/convert/DefaultDbRefResolver.java b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/convert/DefaultDbRefResolver.java index 45b9e6bb7..f0ee5ba41 100644 --- a/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/convert/DefaultDbRefResolver.java +++ b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/convert/DefaultDbRefResolver.java @@ -45,6 +45,7 @@ import org.springframework.dao.support.PersistenceExceptionTranslator; import org.springframework.data.mongodb.ClientSessionException; import org.springframework.data.mongodb.LazyLoadingException; import org.springframework.data.mongodb.MongoDatabaseFactory; +import org.springframework.data.mongodb.MongoDatabaseUtils; import org.springframework.data.mongodb.core.mapping.MongoPersistentProperty; import org.springframework.lang.Nullable; import org.springframework.objenesis.ObjenesisStd; @@ -497,7 +498,7 @@ public class DefaultDbRefResolver implements DbRefResolver { */ protected MongoCollection getCollection(DBRef dbref) { - return (StringUtils.hasText(dbref.getDatabaseName()) ? mongoDbFactory.getMongoDatabase(dbref.getDatabaseName()) - : mongoDbFactory.getMongoDatabase()).getCollection(dbref.getCollectionName(), Document.class); + return MongoDatabaseUtils.getDatabase(dbref.getDatabaseName(), mongoDbFactory) + .getCollection(dbref.getCollectionName(), Document.class); } } diff --git a/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/ClientSessionTests.java b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/ClientSessionTests.java index eac80d7c7..f0b660f2c 100644 --- a/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/ClientSessionTests.java +++ b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/ClientSessionTests.java @@ -28,6 +28,7 @@ import org.junit.jupiter.api.Test; import org.junit.jupiter.api.extension.ExtendWith; import org.springframework.data.annotation.Id; import org.springframework.data.mongodb.core.convert.MappingMongoConverter; +import org.springframework.data.mongodb.core.mapping.DBRef; import org.springframework.data.mongodb.core.query.Query; import org.springframework.data.mongodb.test.util.EnableIfMongoServerVersion; import org.springframework.data.mongodb.test.util.EnableIfReplicaSetAvailable; @@ -51,6 +52,7 @@ public class ClientSessionTests { private static final String DB_NAME = "client-session-tests"; private static final String COLLECTION_NAME = "test"; + private static final String REF_COLLECTION_NAME = "test-with-ref"; static @ReplSetClient MongoClient mongoClient; @@ -148,6 +150,31 @@ public class ClientSessionTests { assertThat(template.exists(query(where("id").is(saved.getId())), SomeDoc.class)).isFalse(); } + @Test // DATAMONGO-2490 + public void shouldBeAbleToReadDbRefDuringTransaction() { + + SomeDoc ref = new SomeDoc("ref-1", "da value"); + WithDbRef source = new WithDbRef("source-1", "da source", ref); + + ClientSession session = mongoClient.startSession(ClientSessionOptions.builder().causallyConsistent(true).build()); + + assertThat(session.getOperationTime()).isNull(); + + session.startTransaction(); + + WithDbRef saved = template.withSession(() -> session).execute(action -> { + + template.save(ref); + template.save(source); + + return template.findOne(query(where("id").is(source.id)), WithDbRef.class); + }); + + assertThat(saved.getSomeDocRef()).isEqualTo(ref); + + session.abortTransaction(); + } + @Data @AllArgsConstructor @org.springframework.data.mongodb.core.mapping.Document(COLLECTION_NAME) @@ -157,4 +184,14 @@ public class ClientSessionTests { String value; } + @Data + @AllArgsConstructor + @org.springframework.data.mongodb.core.mapping.Document(REF_COLLECTION_NAME) + static class WithDbRef { + + @Id String id; + String value; + @DBRef SomeDoc someDocRef; + } + } diff --git a/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/repository/PersonRepositoryTransactionalTests.java b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/repository/PersonRepositoryTransactionalTests.java index ad78b5969..e5c08cafc 100644 --- a/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/repository/PersonRepositoryTransactionalTests.java +++ b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/repository/PersonRepositoryTransactionalTests.java @@ -21,6 +21,7 @@ import static org.springframework.data.mongodb.test.util.MongoTestUtils.*; import java.util.Arrays; import java.util.Collections; import java.util.List; +import java.util.Optional; import java.util.Set; import java.util.concurrent.CopyOnWriteArrayList; @@ -180,6 +181,25 @@ public class PersonRepositoryTransactionalTests { assertAfterTransaction(hu).isNotPresent(); } + @Test // DATAMONGO-2490 + public void shouldBeAbleToReadDbRefDuringTransaction() { + + User rat = new User(); + rat.setUsername("rat"); + + template.save(rat); + + Person elene = new Person("Elene", "Cromwyll", 18); + elene.setCoworker(rat); + + repository.save(elene); + + Optional loaded = repository.findById(elene.getId()); + assertThat(loaded).isPresent(); + assertThat(loaded.get().getCoworker()).isNotNull(); + assertThat(loaded.get().getCoworker().getUsername()).isEqualTo(rat.getUsername()); + } + private AfterTransactionAssertion assertAfterTransaction(Person person) { AfterTransactionAssertion assertion = new AfterTransactionAssertion<>(new Persistable() {