From 4f78aacbf7e26ffb604d52cb084bb32bfeed9db9 Mon Sep 17 00:00:00 2001 From: Christoph Strobl Date: Wed, 15 Nov 2017 14:20:50 +0100 Subject: [PATCH] =?UTF-8?q?DATAMONGO-1824=20-=20Replace=20executeCommand?= =?UTF-8?q?=20for=20aggregations=20with=20MongoCollection.aggregate(?= =?UTF-8?q?=E2=80=A6)=20to=20support=20MongoDB=203.6.?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit We now use the driver native option for aggregate instead of a plain command execution. This is necessary as of MongoDB 3.6 the cursor option is required for aggregations changing the return and execution model. Still we maintain the raw values of AggregationResult as we used to do before. However, relying on executeCommand to change aggregation behavior in custom code will no longer work. Tested against MongoDB: 3.6.RC3, 3.4.9 and 3.2.6 Along the the way we opened up Aggregation itself to expose the AggregationOptions in use und deprecated the $pushAll option on Update since it has been removed in MongoDB 3.6. Original pull request: #515. --- .../data/mongodb/core/MongoTemplate.java | 70 +++++----- .../mongodb/core/aggregation/Aggregation.java | 10 ++ .../data/mongodb/core/query/Update.java | 5 +- .../data/mongodb/core/MongoTemplateTests.java | 30 +---- .../mongodb/core/MongoTemplateUnitTests.java | 31 ++--- .../core/aggregation/AggregationTests.java | 8 +- .../data/mongodb/test/util/MongoVersion.java | 49 +++++++ .../mongodb/test/util/MongoVersionRule.java | 125 ++++++++++++++---- 8 files changed, 223 insertions(+), 105 deletions(-) create mode 100644 spring-data-mongodb/src/test/java/org/springframework/data/mongodb/test/util/MongoVersion.java diff --git a/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/MongoTemplate.java b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/MongoTemplate.java index 30cec4d91..48331e10b 100644 --- a/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/MongoTemplate.java +++ b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/MongoTemplate.java @@ -24,20 +24,10 @@ import lombok.NonNull; import lombok.RequiredArgsConstructor; import java.io.IOException; -import java.util.ArrayList; -import java.util.Collection; -import java.util.Collections; -import java.util.HashMap; -import java.util.HashSet; -import java.util.Iterator; -import java.util.LinkedHashSet; -import java.util.List; -import java.util.Map; +import java.util.*; import java.util.Map.Entry; -import java.util.Optional; -import java.util.Scanner; -import java.util.Set; import java.util.concurrent.TimeUnit; +import java.util.stream.Collectors; import org.bson.Document; import org.bson.conversions.Bson; @@ -141,6 +131,7 @@ import com.mongodb.client.MapReduceIterable; import com.mongodb.client.MongoCollection; import com.mongodb.client.MongoCursor; import com.mongodb.client.MongoDatabase; +import com.mongodb.client.MongoIterable; import com.mongodb.client.model.CountOptions; import com.mongodb.client.model.CreateCollectionOptions; import com.mongodb.client.model.DeleteOptions; @@ -579,7 +570,7 @@ public class MongoTemplate implements MongoOperations, ApplicationContextAware, /* * (non-Javadoc) - * @see org.springframework.data.mongodb.core.ExecutableInsertOperation#getCollection(java.lang.String) + * @see org.springframework.data.mongodb.core.MongoOperations#getCollection(java.lang.String) */ public MongoCollection getCollection(final String collectionName) { @@ -1934,42 +1925,45 @@ public class MongoTemplate implements MongoOperations, ApplicationContextAware, Assert.notNull(outputType, "Output type must not be null!"); AggregationOperationContext rootContext = context == null ? Aggregation.DEFAULT_CONTEXT : context; - Document command = aggregation.toDocument(collectionName, rootContext); + + return doAggregate(aggregation, collectionName, outputType, rootContext); + } + + protected AggregationResults doAggregate(Aggregation aggregation, String collectionName, Class outputType, + AggregationOperationContext context) { + + Document command = aggregation.toDocument(collectionName, context); if (LOGGER.isDebugEnabled()) { LOGGER.debug("Executing aggregation: {}", serializeToJsonSafely(command)); } - Document commandResult = executeCommand(command, this.readPreference); + DocumentCallback callback = new UnwrapAndReadDocumentCallback<>(mongoConverter, outputType, collectionName); - return new AggregationResults(returnPotentiallyMappedResults(outputType, commandResult, collectionName), - commandResult); - } + if (aggregation.getOptions().isExplain()) { - /** - * Returns the potentially mapped results of the given {@code commandResult}. - * - * @param outputType - * @param commandResult - * @return - */ - private List returnPotentiallyMappedResults(Class outputType, Document commandResult, - String collectionName) { - - @SuppressWarnings("unchecked") - Iterable resultSet = (Iterable) commandResult.get("result"); - if (resultSet == null) { - return Collections.emptyList(); + Document commandResult = executeCommand(command); + return new AggregationResults<>(commandResult.get("results", new ArrayList(0)).stream() + .map(callback::doWith).collect(Collectors.toList()), commandResult); } - DocumentCallback callback = new UnwrapAndReadDocumentCallback(mongoConverter, outputType, collectionName); + return execute(collectionName, collection -> { - List mappedResults = new ArrayList(); - for (Document document : resultSet) { - mappedResults.add(callback.doWith(document)); - } + List rawResult = new ArrayList<>(); + + MongoIterable iterable = collection // + .aggregate(command.get("pipeline", new ArrayList(0)), Document.class) // + .collation(aggregation.getOptions().getCollation().map(Collation::toMongoCollation).orElse(null)).map(val -> { + + rawResult.add(val); + return callback.doWith(val); + }); + + return new AggregationResults<>(iterable.into(new ArrayList<>()), + new Document("results", rawResult).append("ok", 1.0D)); + + }); - return mappedResults; } protected CloseableIterator aggregateStream(Aggregation aggregation, String collectionName, diff --git a/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/aggregation/Aggregation.java b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/aggregation/Aggregation.java index dab1f69fe..2ec58c81c 100644 --- a/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/aggregation/Aggregation.java +++ b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/aggregation/Aggregation.java @@ -177,6 +177,16 @@ public class Aggregation { return aggregationOperations.indexOf(aggregationOperation) == aggregationOperations.size() - 1; } + /** + * Get the {@link AggregationOptions}. + * + * @return never {@literal null}. + * @since 2.0.2 + */ + public AggregationOptions getOptions() { + return options; + } + /** * A pointer to the previous {@link AggregationOperation}. * diff --git a/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/query/Update.java b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/query/Update.java index 1d517abbd..b42bbafac 100644 --- a/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/query/Update.java +++ b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/query/Update.java @@ -188,7 +188,8 @@ public class Update { /** * Update using the {@code $pushAll} update modifier.
- * Note: In mongodb 2.4 the usage of {@code $pushAll} has been deprecated in favor of {@code $push $each}. + * Note: In MongoDB 2.4 the usage of {@code $pushAll} has been deprecated in favor of {@code $push $each}. + * Important: As of MongoDB 3.6 {@code $pushAll} is not longer supported. Use {@code $push $each} instead. * {@link #push(String)}) returns a builder that can be used to populate the {@code $each} object. * * @param key @@ -196,7 +197,9 @@ public class Update { * @return * @see MongoDB Update operator: * $pushAll + * @deprecated as of MongoDB 2.4. Removed in MongoDB 3.6. Use {@link #push(String) $push $each} instead. */ + @Deprecated public Update pushAll(String key, Object[] values) { addMultiFieldOperation("$pushAll", key, Arrays.asList(values)); return this; 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 d0baae354..187fd613a 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 @@ -33,18 +33,7 @@ import lombok.NoArgsConstructor; import java.lang.reflect.InvocationTargetException; import java.math.BigDecimal; import java.math.BigInteger; -import java.util.ArrayList; -import java.util.Arrays; -import java.util.Calendar; -import java.util.Collections; -import java.util.Date; -import java.util.HashMap; -import java.util.HashSet; -import java.util.LinkedHashMap; -import java.util.List; -import java.util.Locale; -import java.util.Map; -import java.util.UUID; +import java.util.*; import org.bson.types.ObjectId; import org.hamcrest.collection.IsMapContaining; @@ -90,6 +79,8 @@ import org.springframework.data.mongodb.core.query.BasicQuery; import org.springframework.data.mongodb.core.query.Criteria; import org.springframework.data.mongodb.core.query.Query; import org.springframework.data.mongodb.core.query.Update; +import org.springframework.data.mongodb.test.util.MongoVersion; +import org.springframework.data.mongodb.test.util.MongoVersionRule; import org.springframework.data.mongodb.util.MongoClientVersion; import org.springframework.data.util.CloseableIterator; import org.springframework.test.annotation.DirtiesContext; @@ -137,9 +128,9 @@ public class MongoTemplateTests { ConfigurableApplicationContext context; MongoTemplate mappingTemplate; - org.springframework.data.util.Version mongoVersion; @Rule public ExpectedException thrown = ExpectedException.none(); + @Rule public MongoVersionRule mongoVersion = MongoVersionRule.any(); @Autowired public void setApplicationContext(ConfigurableApplicationContext context) { @@ -177,7 +168,6 @@ public class MongoTemplateTests { public void setUp() { cleanDb(); - queryMongoVersionIfNecessary(); this.mappingTemplate.setApplicationContext(context); } @@ -187,14 +177,6 @@ public class MongoTemplateTests { cleanDb(); } - private void queryMongoVersionIfNecessary() { - - if (mongoVersion == null) { - org.bson.Document result = template.executeCommand("{ buildInfo: 1 }"); - mongoVersion = org.springframework.data.util.Version.parse(result.get("version").toString()); - } - } - protected void cleanDb() { template.dropCollection(Person.class); template.dropCollection(PersonWithAList.class); @@ -2391,7 +2373,9 @@ public class MongoTemplateTests { assertThat(template.findOne(q, VersionedPerson.class), nullValue()); } - @Test // DATAMONGO-354 + @Test // DATAMONGO-354, DATAMONGO-1824 + @MongoVersion(until = "3.6") + @SuppressWarnings("deprecation") public void testUpdateShouldAllowMultiplePushAll() { DocumentWithMultipleCollections doc = new DocumentWithMultipleCollections(); diff --git a/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/MongoTemplateUnitTests.java b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/MongoTemplateUnitTests.java index f0e64d05c..84a2ecaac 100644 --- a/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/MongoTemplateUnitTests.java +++ b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/MongoTemplateUnitTests.java @@ -58,7 +58,6 @@ import org.springframework.data.domain.Sort; import org.springframework.data.geo.Point; import org.springframework.data.mongodb.MongoDbFactory; import org.springframework.data.mongodb.core.aggregation.Aggregation; -import org.springframework.data.mongodb.core.query.Collation; import org.springframework.data.mongodb.core.convert.DefaultDbRefResolver; import org.springframework.data.mongodb.core.convert.MappingMongoConverter; import org.springframework.data.mongodb.core.convert.MongoCustomConversions; @@ -71,6 +70,7 @@ import org.springframework.data.mongodb.core.mapping.event.BeforeConvertEvent; import org.springframework.data.mongodb.core.mapreduce.GroupBy; import org.springframework.data.mongodb.core.mapreduce.MapReduceOptions; import org.springframework.data.mongodb.core.query.BasicQuery; +import org.springframework.data.mongodb.core.query.Collation; import org.springframework.data.mongodb.core.query.Criteria; import org.springframework.data.mongodb.core.query.NearQuery; import org.springframework.data.mongodb.core.query.Query; @@ -81,6 +81,7 @@ import com.mongodb.DB; import com.mongodb.MongoClient; import com.mongodb.MongoException; import com.mongodb.ReadPreference; +import com.mongodb.client.AggregateIterable; import com.mongodb.client.FindIterable; import com.mongodb.client.MapReduceIterable; import com.mongodb.client.MongoCollection; @@ -111,6 +112,7 @@ public class MongoTemplateUnitTests extends MongoOperationsUnitTests { @Mock MongoCollection collection; @Mock MongoCursor cursor; @Mock FindIterable findIterable; + @Mock AggregateIterable aggregateIterable; @Mock MapReduceIterable mapReduceIterable; Document commandResultDocument = new Document(); @@ -130,6 +132,8 @@ public class MongoTemplateUnitTests extends MongoOperationsUnitTests { when(collection.find(Mockito.any(org.bson.Document.class))).thenReturn(findIterable); when(collection.mapReduce(Mockito.any(), Mockito.any())).thenReturn(mapReduceIterable); when(collection.count(any(Bson.class), any(CountOptions.class))).thenReturn(1L); + when(collection.aggregate(any(), any())).thenReturn(aggregateIterable); + when(collection.withReadPreference(any())).thenReturn(collection); when(findIterable.projection(Mockito.any())).thenReturn(findIterable); when(findIterable.sort(Mockito.any(org.bson.Document.class))).thenReturn(findIterable); when(findIterable.modifiers(Mockito.any(org.bson.Document.class))).thenReturn(findIterable); @@ -139,6 +143,9 @@ public class MongoTemplateUnitTests extends MongoOperationsUnitTests { when(mapReduceIterable.sort(Mockito.any())).thenReturn(mapReduceIterable); when(mapReduceIterable.iterator()).thenReturn(cursor); when(mapReduceIterable.filter(any())).thenReturn(mapReduceIterable); + when(aggregateIterable.collation(any())).thenReturn(aggregateIterable); + when(aggregateIterable.map(any())).thenReturn(aggregateIterable); + when(aggregateIterable.into(any())).thenReturn(Collections.emptyList()); this.mappingContext = new MongoMappingContext(); this.converter = new MappingMongoConverter(new DefaultDbRefResolver(factory), mappingContext); @@ -361,28 +368,22 @@ public class MongoTemplateUnitTests extends MongoOperationsUnitTests { assertThat(captor.getValue(), equalTo(new org.bson.Document("foo", 1))); } - @Test // DATAMONGO-1166 + @Test // DATAMONGO-1166, DATAMONGO-1824 public void aggregateShouldHonorReadPreferenceWhenSet() { - when(db.runCommand(Mockito.any(org.bson.Document.class), Mockito.any(ReadPreference.class), eq(Document.class))) - .thenReturn(mock(Document.class)); template.setReadPreference(ReadPreference.secondary()); template.aggregate(newAggregation(Aggregation.unwind("foo")), "collection-1", Wrapper.class); - verify(this.db, times(1)).runCommand(Mockito.any(org.bson.Document.class), eq(ReadPreference.secondary()), - eq(Document.class)); + verify(collection).withReadPreference(eq(ReadPreference.secondary())); } - @Test // DATAMONGO-1166 + @Test // DATAMONGO-1166, DATAMONGO-1824 public void aggregateShouldIgnoreReadPreferenceWhenNotSet() { - when(db.runCommand(Mockito.any(org.bson.Document.class), eq(org.bson.Document.class))) - .thenReturn(mock(Document.class)); - template.aggregate(newAggregation(Aggregation.unwind("foo")), "collection-1", Wrapper.class); - verify(this.db, times(1)).runCommand(Mockito.any(org.bson.Document.class), eq(org.bson.Document.class)); + verify(collection, never()).withReadPreference(any()); } @Test // DATAMONGO-1166 @@ -484,7 +485,6 @@ public class MongoTemplateUnitTests extends MongoOperationsUnitTests { when(output.iterator()).thenReturn(cursor); when(cursor.hasNext()).thenReturn(false); - when(collection.mapReduce(anyString(), anyString())).thenReturn(output); template.mapReduce("collection", "function(){}", "function(key,values){}", new MapReduceOptions().limit(1000), @@ -754,17 +754,14 @@ public class MongoTemplateUnitTests extends MongoOperationsUnitTests { assertThat(options.getValue().getCollation().getLocale(), is("fr")); } - @Test // DATAMONGO-1518 + @Test // DATAMONGO-1518, DATAMONGO-1824 public void aggregateShouldUseCollationWhenPresent() { Aggregation aggregation = newAggregation(project("id")) .withOptions(newAggregationOptions().collation(Collation.of("fr")).build()); template.aggregate(aggregation, AutogenerateableId.class, Document.class); - ArgumentCaptor cmd = ArgumentCaptor.forClass(Document.class); - verify(db).runCommand(cmd.capture(), Mockito.any(Class.class)); - - assertThat(cmd.getValue().get("collation", Document.class), equalTo(new Document("locale", "fr"))); + verify(aggregateIterable).collation(eq(com.mongodb.client.model.Collation.builder().locale("fr").build())); } @Test // DATAMONGO-1518 diff --git a/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/aggregation/AggregationTests.java b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/aggregation/AggregationTests.java index 11b39e7dd..a9729edaf 100644 --- a/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/aggregation/AggregationTests.java +++ b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/aggregation/AggregationTests.java @@ -126,6 +126,7 @@ public class AggregationTests { } private void cleanDb() { + mongoTemplate.dropCollection(INPUT_COLLECTION); mongoTemplate.dropCollection(Product.class); mongoTemplate.dropCollection(UserWithLikes.class); @@ -143,6 +144,7 @@ public class AggregationTests { mongoTemplate.dropCollection(Sales2.class); mongoTemplate.dropCollection(Employee.class); mongoTemplate.dropCollection(Art.class); + mongoTemplate.dropCollection("personQueryTemp"); } /** @@ -1357,6 +1359,7 @@ public class AggregationTests { Document rawResult = result.getRawResults(); assertThat(rawResult, is(notNullValue())); + assertThat(rawResult.containsKey("stages"), is(true)); } @@ -1565,7 +1568,7 @@ public class AggregationTests { assertThat(firstItem, isBsonObject().containing("linkedPerson.[0].firstname", "u1")); } - @Test // DATAMONGO-1418 + @Test // DATAMONGO-1418, DATAMONGO-1824 public void shouldCreateOutputCollection() { assumeTrue(mongoVersion.isGreaterThanOrEqualTo(TWO_DOT_SIX)); @@ -1579,7 +1582,8 @@ public class AggregationTests { out(tempOutCollection)); AggregationResults results = mongoTemplate.aggregate(agg, Document.class); - assertThat(results.getMappedResults(), is(empty())); + + assertThat(results.getMappedResults(), hasSize(2)); List list = mongoTemplate.findAll(Document.class, tempOutCollection); diff --git a/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/test/util/MongoVersion.java b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/test/util/MongoVersion.java new file mode 100644 index 000000000..e47938bbc --- /dev/null +++ b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/test/util/MongoVersion.java @@ -0,0 +1,49 @@ +/* + * Copyright 2017 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.data.mongodb.test.util; + +import java.lang.annotation.Documented; +import java.lang.annotation.ElementType; +import java.lang.annotation.Retention; +import java.lang.annotation.RetentionPolicy; +import java.lang.annotation.Target; + +/** + * {@link MongoVersion} allows specifying an version range of mongodb that is applicable for a specific test method. To + * be used along with {@link MongoVersionRule}. + * + * @author Christoph Strobl + * @since 2.0.2 + */ +@Retention(RetentionPolicy.RUNTIME) +@Target(ElementType.METHOD) +@Documented +public @interface MongoVersion { + + /** + * Inclusive lower bound of MongoDB server range. + * + * @return {@code 0.0.0} by default. + */ + String asOf() default "0.0.0"; + + /** + * Exclusive upper bound of MongoDB server range. + * + * @return {@code 9999.9999.9999} by default. + */ + String until() default "9999.9999.9999"; +} diff --git a/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/test/util/MongoVersionRule.java b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/test/util/MongoVersionRule.java index 51db7f26a..78c342e24 100644 --- a/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/test/util/MongoVersionRule.java +++ b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/test/util/MongoVersionRule.java @@ -15,9 +15,11 @@ */ package org.springframework.data.mongodb.test.util; +import java.util.concurrent.atomic.AtomicReference; + +import org.junit.AssumptionViolatedException; import org.junit.ClassRule; import org.junit.Rule; -import org.junit.internal.AssumptionViolatedException; import org.junit.rules.TestRule; import org.junit.runner.Description; import org.junit.runners.model.Statement; @@ -39,72 +41,147 @@ import com.mongodb.MongoClient; */ public class MongoVersionRule implements TestRule { - private String host = "localhost"; - private int port = 27017; + private static final Version ANY = new Version(9999, 9999, 9999); + private static final Version DEFAULT_HIGH = ANY; + private static final Version DEFAULT_LOW = new Version(0, 0, 0); + private final AtomicReference currentVersion = new AtomicReference<>(null); private final Version minVersion; private final Version maxVersion; - private Version currentVersion; + private String host = "localhost"; + private int port = 27017; public MongoVersionRule(Version min, Version max) { + this.minVersion = min; this.maxVersion = max; } public static MongoVersionRule any() { - return new MongoVersionRule(new Version(0, 0, 0), new Version(9999, 9999, 9999)); + return new MongoVersionRule(ANY, ANY); } public static MongoVersionRule atLeast(Version minVersion) { - return new MongoVersionRule(minVersion, new Version(9999, 9999, 9999)); + return new MongoVersionRule(minVersion, DEFAULT_HIGH); } public static MongoVersionRule atMost(Version maxVersion) { - return new MongoVersionRule(new Version(0, 0, 0), maxVersion); + return new MongoVersionRule(DEFAULT_LOW, maxVersion); } public MongoVersionRule withServerRunningAt(String host, int port) { + this.host = host; this.port = port; return this; } + /** + * @see Version#isGreaterThan(Version) + */ + public boolean isGreaterThan(Version version) { + return getCurrentVersion().isGreaterThan(version); + } + + /** + * @see Version#isGreaterThanOrEqualTo(Version) + */ + public boolean isGreaterThanOrEqualTo(Version version) { + return getCurrentVersion().isGreaterThanOrEqualTo(version); + } + + /** + * @see Version#is(Version) + */ + public boolean is(Version version) { + return getCurrentVersion().equals(version); + } + + /** + * @see Version#isLessThan(Version) + */ + public boolean isLessThan(Version version) { + return getCurrentVersion().isLessThan(version); + } + + /** + * @see Version#isLessThanOrEqualTo(Version) + */ + public boolean isLessThanOrEqualTo(Version version) { + return getCurrentVersion().isGreaterThanOrEqualTo(version); + } + @Override public Statement apply(final Statement base, Description description) { - initCurrentVersion(); return new Statement() { @Override public void evaluate() throws Throwable { - if (currentVersion != null) { - if (currentVersion.isLessThan(minVersion) || currentVersion.isGreaterThan(maxVersion)) { - throw new AssumptionViolatedException( - String.format("Expected mongodb server to be in range %s to %s but found %s", minVersion, maxVersion, - currentVersion)); + + if (!getCurrentVersion().equals(ANY)) { + + Version minVersion = MongoVersionRule.this.minVersion.equals(ANY) ? DEFAULT_LOW + : MongoVersionRule.this.minVersion; + Version maxVersion = MongoVersionRule.this.maxVersion.equals(ANY) ? DEFAULT_HIGH + : MongoVersionRule.this.maxVersion; + + if (MongoVersionRule.this.minVersion.equals(ANY) && MongoVersionRule.this.maxVersion.equals(ANY)) { + + MongoVersion version = description.getAnnotation(MongoVersion.class); + if (version != null) { + minVersion = Version.parse(version.asOf()); + maxVersion = Version.parse(version.until()); + } } + + validateVersion(minVersion, maxVersion); } + base.evaluate(); } }; } - private void initCurrentVersion() { + private void validateVersion(Version min, Version max) { - if (currentVersion == null) { - try { - MongoClient client; - client = new MongoClient(host, port); - DB db = client.getDB("test"); - CommandResult result = db.command(new BasicDBObject().append("buildInfo", 1)); - this.currentVersion = Version.parse(result.get("version").toString()); - } catch (Exception e) { - e.printStackTrace(); - } + if (getCurrentVersion().isLessThan(min) || getCurrentVersion().isGreaterThanOrEqualTo(max)) { + + throw new AssumptionViolatedException( + String.format("Expected MongoDB server to be in range (%s, %s] but found %s", min, max, currentVersion)); } } + private Version getCurrentVersion() { + + if (currentVersion.get() == null) { + currentVersion.compareAndSet(null, fetchCurrentVersion()); + } + + return currentVersion.get(); + } + + private Version fetchCurrentVersion() { + + try { + + MongoClient client; + client = new MongoClient(host, port); + DB db = client.getDB("test"); + CommandResult result = db.command(new BasicDBObject().append("buildInfo", 1)); + client.close(); + + return Version.parse(result.get("version").toString()); + } catch (Exception e) { + return ANY; + } + } + + @Override + public String toString() { + return getCurrentVersion().toString(); + } }