diff --git a/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/EntityOperations.java b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/EntityOperations.java index a9b1d19d8..1ccd16023 100644 --- a/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/EntityOperations.java +++ b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/EntityOperations.java @@ -47,6 +47,7 @@ import com.mongodb.util.JSONParseException; * Common operations performed on an entity in the context of it's mapping metadata. * * @author Oliver Gierke + * @author Mark Paluch * @since 2.1 * @see MongoTemplate * @see ReactiveMongoTemplate @@ -70,11 +71,11 @@ class EntityOperations { Assert.notNull(entity, "Bean must not be null!"); if (entity instanceof String) { - return new SimpleEntity(parse(entity.toString())); + return new UnmappedEntity(parse(entity.toString())); } if (entity instanceof Map) { - return new SimpleEntity((Map) entity); + return new SimpleMappedEntity((Map) entity); } return MappedEntity.of(entity, context); @@ -94,11 +95,11 @@ class EntityOperations { Assert.notNull(conversionService, "ConversionService must not be null!"); if (entity instanceof String) { - return new SimpleEntity(parse(entity.toString())); + return new UnmappedEntity(parse(entity.toString())); } if (entity instanceof Map) { - return new SimpleEntity((Map) entity); + return new SimpleMappedEntity((Map) entity); } return AdaptibleMappedEntity.of(entity, context, conversionService); @@ -286,7 +287,7 @@ class EntityOperations { } @RequiredArgsConstructor - private static class SimpleEntity> implements AdaptibleEntity { + private static class UnmappedEntity> implements AdaptibleEntity { private final T map; @@ -388,6 +389,31 @@ class EntityOperations { } } + private static class SimpleMappedEntity> extends UnmappedEntity { + + public SimpleMappedEntity(T map) { + super(map); + } + + /* + * (non-Javadoc) + * @see org.springframework.data.mongodb.core.EntityOperations.PersistableSource#toMappedDocument(org.springframework.data.mongodb.core.convert.MongoWriter) + */ + @Override + @SuppressWarnings("unchecked") + public MappedDocument toMappedDocument(MongoWriter writer) { + + T bean = getBean(); + bean = (T) (bean instanceof Document // + ? (Document) bean // + : new Document(bean)); + Document document = new Document(); + writer.write(bean, document); + + return MappedDocument.of(document); + } + } + @RequiredArgsConstructor(access = AccessLevel.PROTECTED) private static class MappedEntity implements Entity { diff --git a/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/MongoTemplateTests.java b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/MongoTemplateTests.java index d2c869638..93dd66e08 100644 --- a/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/MongoTemplateTests.java +++ b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/MongoTemplateTests.java @@ -26,6 +26,8 @@ import static org.springframework.data.mongodb.core.query.Criteria.*; import static org.springframework.data.mongodb.core.query.Query.*; import static org.springframework.data.mongodb.core.query.Update.*; +import com.mongodb.BasicDBObject; +import com.mongodb.DBObject; import lombok.AllArgsConstructor; import lombok.Data; import lombok.EqualsAndHashCode; @@ -36,11 +38,13 @@ import lombok.experimental.Wither; import java.lang.reflect.InvocationTargetException; import java.math.BigDecimal; import java.math.BigInteger; +import java.time.Duration; import java.time.Instant; import java.util.*; import java.util.stream.Collectors; import java.util.stream.IntStream; +import org.bson.Document; import org.bson.types.ObjectId; import org.hamcrest.collection.IsMapContaining; import org.joda.time.DateTime; @@ -2061,6 +2065,18 @@ public class MongoTemplateTests { assertThat(result.get(0).field, is(value)); } + @Test // DATAMONGO-2028 + public void allowInsertOfDbObjectWithMappedTypes() { + + DBObject dbObject = new BasicDBObject("_id", "foo").append("duration", Duration.ofSeconds(100)); + template.insert(dbObject, "sample"); + List result = template.findAll(org.bson.Document.class, "sample"); + + assertThat(result.size(), is(1)); + assertThat(result.get(0).getString("_id"), is("foo")); + assertThat(result.get(0).getString("duration"), is("PT1M40S")); + } + @Test // DATAMONGO-816 public void shouldExecuteQueryShouldMapQueryBeforeQueryExecution() {