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 6e3faac9b..6995da9a8 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 MongoDbFactory} 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, MongoDbFactory factory) {
+ public static MongoDatabase getDatabase(@Nullable String dbName, MongoDbFactory factory) {
return doGetMongoDatabase(dbName, factory, SessionSynchronization.ON_ACTUAL_TRANSACTION);
}
@@ -88,13 +88,13 @@ public class MongoDatabaseUtils {
*
* Registers a {@link MongoSessionSynchronization MongoDB specific transaction synchronization} within the current
* {@link Thread} if {@link TransactionSynchronizationManager#isSynchronizationActive() synchronization is active}.
- *
+ *
* @param dbName the name of the {@link MongoDatabase} to get.
* @param factory the {@link MongoDbFactory} to get the {@link MongoDatabase} from.
* @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, MongoDbFactory factory,
+ public static MongoDatabase getDatabase(@Nullable String dbName, MongoDbFactory 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 9370e7790..06e48b5dd 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.MongoDbFactory;
+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.getDb(dbref.getDatabaseName())
- : mongoDbFactory.getDb()).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 3cc48a5c8..d9688cc1d 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
@@ -31,6 +31,7 @@ import org.junit.Test;
import org.junit.rules.TestRule;
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.MongoTestUtils;
import org.springframework.data.mongodb.test.util.MongoVersion;
@@ -54,6 +55,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";
MongoTemplate template;
MongoClient client;
@@ -154,6 +156,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)
@@ -163,4 +190,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 d2b9b9c6e..c572d3782 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
@@ -20,6 +20,7 @@ import static org.springframework.data.mongodb.test.util.MongoTestUtils.*;
import java.util.Arrays;
import java.util.List;
+import java.util.Optional;
import java.util.concurrent.CopyOnWriteArrayList;
import org.bson.Document;
@@ -175,6 +176,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