diff --git a/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/MongoOperations.java b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/MongoOperations.java index b289d2b10..6b481de4f 100644 --- a/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/MongoOperations.java +++ b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/MongoOperations.java @@ -19,6 +19,8 @@ import java.util.Collection; import java.util.List; import java.util.Set; +import org.springframework.data.mongodb.core.aggregation.AggregationPipeline; +import org.springframework.data.mongodb.core.aggregation.AggregationResults; import org.springframework.data.mongodb.core.convert.MongoConverter; import org.springframework.data.mongodb.core.geo.GeoResult; import org.springframework.data.mongodb.core.geo.GeoResults; @@ -301,6 +303,16 @@ public interface MongoOperations { */ GroupByResults group(Criteria criteria, String inputCollectionName, GroupBy groupBy, Class entityClass); + /** + * Execute an aggregation operation. The raw results will be mapped to the given entity class. + * + * @param inputCollectionName the collection there the aggregation operation will read from. + * @param pipeline The pipeline holding the aggregation operations. + * @param entityClass The parameterized type of the returned list. + * @return The results of the aggregation operation. + */ + AggregationResults aggregate(String inputCollectionName, AggregationPipeline pipeline, Class entityClass); + /** * Execute a map-reduce operation. The map-reduce operation will be formed with an output type of INLINE * 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 3eeda7911..1782564da 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 @@ -15,8 +15,8 @@ */ package org.springframework.data.mongodb.core; -import static org.springframework.data.mongodb.core.query.Criteria.*; -import static org.springframework.data.mongodb.core.query.SerializationUtils.*; +import static org.springframework.data.mongodb.core.query.Criteria.where; +import static org.springframework.data.mongodb.core.query.SerializationUtils.serializeToJsonSafely; import java.io.IOException; import java.util.ArrayList; @@ -53,6 +53,8 @@ import org.springframework.data.mapping.context.MappingContext; import org.springframework.data.mapping.model.BeanWrapper; import org.springframework.data.mapping.model.MappingException; import org.springframework.data.mongodb.MongoDbFactory; +import org.springframework.data.mongodb.core.aggregation.AggregationPipeline; +import org.springframework.data.mongodb.core.aggregation.AggregationResults; import org.springframework.data.mongodb.core.convert.MappingMongoConverter; import org.springframework.data.mongodb.core.convert.MongoConverter; import org.springframework.data.mongodb.core.convert.MongoWriter; @@ -1208,6 +1210,31 @@ public class MongoTemplate implements MongoOperations, ApplicationContextAware { } + public AggregationResults aggregate(String inputCollectionName, AggregationPipeline pipeline, Class entityClass) { + Assert.notNull(inputCollectionName, "Collection name is missing"); + Assert.notNull(pipeline, "Aggregation pipeline is missing"); + Assert.notNull(entityClass, "Entity class is missing"); + + // prepare command + DBObject command = new BasicDBObject("aggregate", inputCollectionName ); + command.put( "pipeline", pipeline.getOperations() ); + + // execute command + CommandResult commandResult = executeCommand(command); + handleCommandError(commandResult, command); + + // map results + @SuppressWarnings("unchecked") + Iterable resultSet = (Iterable) commandResult.get("result"); + List mappedResults = new ArrayList(); + DbObjectCallback callback = new ReadDbObjectCallback(mongoConverter, entityClass); + for (DBObject dbObject : resultSet) { + mappedResults.add(callback.doWith(dbObject)); + } + + return new AggregationResults(mappedResults, commandResult); + } + protected String replaceWithResourceIfNecessary(String function) { String func = function; diff --git a/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/aggregation/AggregationPipeline.java b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/aggregation/AggregationPipeline.java new file mode 100644 index 000000000..cdaf46165 --- /dev/null +++ b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/aggregation/AggregationPipeline.java @@ -0,0 +1,178 @@ +/* + * Copyright 2011-2012 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.core.aggregation; + +import java.util.ArrayList; +import java.util.List; + +import org.springframework.data.domain.Sort; +import org.springframework.data.mongodb.core.query.Criteria; +import org.springframework.util.Assert; + +import com.mongodb.BasicDBObject; +import com.mongodb.DBObject; +import com.mongodb.util.JSON; +import com.mongodb.util.JSONParseException; + +/** + * Holds the operations of an aggregation pipeline. + * + * @author Tobias Trelle + */ +public class AggregationPipeline { + + private static final String OPERATOR_PREFIX = "$"; + + private List operations = new ArrayList(); + + /** + * Adds a projection operation to the pipeline. + * + * @param projection JSON string holding the projection. + * @return The pipeline. + */ + public AggregationPipeline project(String projection) { + return addDocumentOperation("project", projection); + } + + /** + * Adds a projection operation to the pipeline. + * + * @param projection Type safe projection object. + * @return The pipeline. + */ + public AggregationPipeline project(Projection projection) { + return addOperation("project", projection.toDBObject() ); + } + + /** + * Adds an unwind operation to the pipeline. + * + * @param field Name of the field to unwind (should be an array). + * @return The pipeline. + */ + public AggregationPipeline unwind(String field) { + Assert.notNull(field, "Missing field name"); + + if (!field.startsWith(OPERATOR_PREFIX)) { + field = OPERATOR_PREFIX + field; + } + + return addOperation("unwind", field); + } + + /** + * Adds a group operation to the pipeline. + * + * @param projection JSON string holding the group. + * @return The pipeline. + */ + public AggregationPipeline group(String group) { + return addDocumentOperation("group", group); + } + + /** + * Adds a sort operation to the pipeline. + * + * @param sort JSON string holding the sorting. + * @return The pipeline. + */ + public AggregationPipeline sort(String sort) { + return addDocumentOperation("sort", sort); + } + + /** + * Adds a sort operation to the pipeline. + * + * @param sort Type safe sort operation. + * @return The pipeline. + */ + public AggregationPipeline sort(Sort sort) { + Assert.notNull(sort); + + DBObject dbo = new BasicDBObject(); + + for (org.springframework.data.domain.Sort.Order order : sort) { + dbo.put(order.getProperty(), order.isAscending() ? 1 : -1); + } + return addOperation("sort", dbo); + } + + /** + * Adds a match operation to the pipeline that is basically a query on the collection.s + * + * @param projection JSON string holding the criteria. + * @return The pipeline. + */ + public AggregationPipeline match(String match) { + return addDocumentOperation("match", match); + } + + /** + * Adds a match operation to the pipeline that is basically a query on the collection.s + * + * @param criteria Type safe criteria to filter documents from the collection. + * @return The pipeline. + */ + public AggregationPipeline match(Criteria criteria) { + Assert.notNull(criteria); + + return addOperation("match", criteria.getCriteriaObject()); + } + + /** + * Adds an limit operation to the pipeline. + * + * @param n Number of document to consider. + * @return The pipeline. + */ + public AggregationPipeline limit(long n) { + return addOperation("limit", n); + } + + /** + * Adds an skip operation to the pipeline. + * + * @param n Number of documents to skip. + * @return The pipeline. + */ + public AggregationPipeline skip(long n) { + return addOperation("skip", n); + } + + public List getOperations() { + return operations; + } + + private AggregationPipeline addDocumentOperation(String opName, String operation) { + Assert.notNull(operation, "Missing " + opName); + return addOperation(opName, parseJson(operation)); + } + + private AggregationPipeline addOperation(String key, Object value) { + this.operations.add(new BasicDBObject(OPERATOR_PREFIX + key, value)); + return this; + } + + private DBObject parseJson(String json) { + try { + return (DBObject) JSON.parse(json); + } catch (JSONParseException e) { + throw new IllegalArgumentException("Not a valid JSON document: " + json, e); + } + } + +} \ No newline at end of file diff --git a/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/aggregation/AggregationResults.java b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/aggregation/AggregationResults.java new file mode 100644 index 000000000..c8b178584 --- /dev/null +++ b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/aggregation/AggregationResults.java @@ -0,0 +1,75 @@ +/* + * Copyright 2011-2012 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.core.aggregation; + +import java.util.ArrayList; +import java.util.Iterator; +import java.util.List; + +import org.springframework.util.Assert; + +import com.mongodb.DBObject; + +/** + * Collects the results of executing an aggregation operation. + * + * @author Tobias Trelle + * + * @param The class in which the results are mapped onto. + */ +public class AggregationResults implements Iterable { + + private final List mappedResults; + private final DBObject rawResults; + + private String serverUsed; + + public AggregationResults(List mappedResults, DBObject rawResults) { + Assert.notNull(mappedResults); + Assert.notNull(rawResults); + this.mappedResults = mappedResults; + this.rawResults = rawResults; + parseServerUsed(); + } + + public List getAggregationResult() { + List result = new ArrayList(); + Iterator it = iterator(); + + while (it.hasNext()) { + result.add(it.next()); + } + + return result; + } + + @Override + public Iterator iterator() { + return mappedResults.iterator(); + } + + public String getServerUsed() { + return serverUsed; + } + + private void parseServerUsed() { + Object object = rawResults.get("serverUsed"); + if (object instanceof String) { + serverUsed = (String) object; + } + } + +} diff --git a/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/aggregation/Projection.java b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/aggregation/Projection.java new file mode 100644 index 000000000..a596199bc --- /dev/null +++ b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/aggregation/Projection.java @@ -0,0 +1,156 @@ +/* + * Copyright 2011-2012 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.core.aggregation; + +import java.util.ArrayList; +import java.util.Collections; +import java.util.EmptyStackException; +import java.util.List; +import java.util.Stack; + +import org.springframework.dao.InvalidDataAccessApiUsageException; +import org.springframework.data.mongodb.core.query.Field; +import org.springframework.util.Assert; + +import com.mongodb.BasicDBObject; +import com.mongodb.DBObject; + +/** + * Projection of field to be used in an {@link AggregationPipeline}. + *

+ * A projection is similar to a {@link Field} inclusion/exclusion but more powerful. It can generate new fields, change + * values of given field etc. + * + * @author Tobias Trelle + */ +public class Projection { + + private static final String REFERENCE_PREFIX = "$"; + + private DBObject document = new BasicDBObject(); + + private DBObject rightHandExpression; + + /** Stack of key names. Size is 0 or 1. */ + private Stack reference = new Stack(); + + /** Create an empty projection. */ + public Projection() { + } + + /** + * This convenience constructor excludes the field _id and includes the given fields. + * + * @param includes Keys of field to include. + */ + public Projection(String... includes) { + Assert.notEmpty(includes); + exclude("_id"); + for (String key : includes) { + include(key); + } + } + + /** + * Excludes a given field. + * + * @param key The key of the field. + */ + public final void exclude(String key) { + Assert.notNull(key, "Missing key"); + document.put(key, 0); + } + + /** + * Includes a given field. + * + * @param key The key of the field. + */ + public final Projection include(String key) { + Assert.notNull(key, "Missing key"); + + safePop(); + reference.push(key); + + return this; + } + + /** + * Sets the key for a computed field. + * + */ + public final Projection as(String key) { + Assert.notNull(key, "Missing key"); + + try { + document.put(key, rightHandSide(safeReference(reference.pop())) ); + } catch (EmptyStackException e) { + throw new InvalidDataAccessApiUsageException("Invalid use of as()", e); + } + return this; + } + + public final Projection plus(Number n) { + return arithmeticOperation("add", n); + } + + public final Projection minus(Number n) { + return arithmeticOperation("substract", n); + } + + private Projection arithmeticOperation(String op, Number n) { + Assert.notNull(n, "Missing number"); + + rightHandExpression = createArrayObject(op, safeReference(reference.peek()), n); + + return this; + } + + private DBObject createArrayObject(String op, Object... items) { + List list = new ArrayList(); + Collections.addAll(list, items); + + return new BasicDBObject( safeReference(op), list ); + } + + private void safePop() { + if ( !reference.empty() ) { + document.put( reference.pop(), rightHandSide(1) ); + } + } + + private String safeReference(String key) { + Assert.notNull(key); + + if ( !key.startsWith(REFERENCE_PREFIX) ) { + return REFERENCE_PREFIX + key; + } else { + return key; + } + } + + private Object rightHandSide(Object defaultValue) { + Object value = rightHandExpression != null ? rightHandExpression : defaultValue; + rightHandExpression = null; + return value; + } + + DBObject toDBObject() { + safePop(); + return document; + } + +} diff --git a/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/aggregation/AggregationPipelineTests.java b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/aggregation/AggregationPipelineTests.java new file mode 100644 index 000000000..756d193a6 --- /dev/null +++ b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/aggregation/AggregationPipelineTests.java @@ -0,0 +1,137 @@ +package org.springframework.data.mongodb.core.aggregation; + +import static org.hamcrest.CoreMatchers.is; +import static org.hamcrest.CoreMatchers.notNullValue; +import static org.junit.Assert.assertThat; +import java.util.List; + +import org.junit.Before; +import org.junit.Test; +import org.springframework.data.domain.Sort; +import org.springframework.data.domain.Sort.Direction; +import org.springframework.data.mongodb.core.query.Criteria; + +import com.mongodb.DBObject; + +/** + * Tests of the {@link AggregationPipeline}. + * + * @author Tobias Trelle + */ +public class AggregationPipelineTests { + + /** Unit under test. */ + private AggregationPipeline pipeline; + + @Before public void setUp() { + pipeline = new AggregationPipeline(); + } + + @Test public void limitOperation() { + // given + pipeline.limit(42); + + // when + List rawPipeline = pipeline.getOperations(); + + // then + assertDBObject("$limit", 42L, rawPipeline); + } + + @Test public void skipOperation() { + // given + pipeline.skip(5); + + // when + List rawPipeline = pipeline.getOperations(); + + // then + assertDBObject("$skip", 5L, rawPipeline); + } + + @Test public void unwindOperation() { + // given + pipeline.unwind("$field"); + + // when + List rawPipeline = pipeline.getOperations(); + + // then + assertDBObject("$unwind", "$field", rawPipeline); + } + + @Test public void unwindOperationWithAddedPrefix() { + // given + pipeline.unwind("field"); + + // when + List rawPipeline = pipeline.getOperations(); + + // then + assertDBObject("$unwind", "$field", rawPipeline); + } + + + @Test public void matchOperation() { + // given + Criteria criteria = new Criteria("title").is("Doc 1"); + pipeline.match( criteria ); + + // when + List rawPipeline = pipeline.getOperations(); + + // then + assertOneDocument(rawPipeline); + DBObject match = rawPipeline.get(0); + DBObject criteriaDoc = (DBObject)match.get("$match"); + assertThat( criteriaDoc, notNullValue() ); + assertSingleDBObject( "title" , "Doc 1", criteriaDoc ); + } + + @Test public void sortOperation() { + // given + Sort sort = new Sort(new Sort.Order(Direction.ASC, "n")); + pipeline.sort( sort ); + + // when + List rawPipeline = pipeline.getOperations(); + + // then + assertOneDocument(rawPipeline); + DBObject sortDoc = rawPipeline.get(0); + DBObject orderDoc = (DBObject)sortDoc.get("$sort"); + assertThat( orderDoc, notNullValue() ); + assertSingleDBObject( "n" , 1, orderDoc ); + } + + @Test public void projectOperation() { + // given + Projection projection = new Projection("a"); + pipeline.project(projection); + + // when + List rawPipeline = pipeline.getOperations(); + + // then + assertOneDocument(rawPipeline); + DBObject projectionDoc = rawPipeline.get(0); + DBObject fields = (DBObject)projectionDoc.get("$project"); + assertThat( fields, notNullValue() ); + assertSingleDBObject( "a" , 1, fields ); + } + + private static void assertOneDocument(List result) { + assertThat( result, notNullValue() ); + assertThat( result.size(), is(1) ); + } + + private static void assertDBObject(String key, Object value, List result) { + assertOneDocument(result); + assertSingleDBObject( key, value, result.get(0) ); + } + + private static void assertSingleDBObject(String key, Object value, DBObject doc) { + assertThat( doc.get(key), is(value) ); + } + +} 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 new file mode 100644 index 000000000..f70ff7d98 --- /dev/null +++ b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/aggregation/AggregationTests.java @@ -0,0 +1,202 @@ +/* + * Copyright 2011-2012 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.core.aggregation; + +import static org.hamcrest.CoreMatchers.is; +import static org.hamcrest.CoreMatchers.notNullValue; +import static org.junit.Assert.assertThat; + +import java.util.ArrayList; +import java.util.List; + +import org.junit.After; +import org.junit.Before; +import org.junit.Test; +import org.junit.runner.RunWith; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.dao.InvalidDataAccessApiUsageException; +import org.springframework.data.domain.Sort; +import org.springframework.data.domain.Sort.Direction; +import org.springframework.data.mongodb.core.MongoTemplate; +import org.springframework.test.context.ContextConfiguration; +import org.springframework.test.context.junit4.SpringJUnit4ClassRunner; + +import com.mongodb.BasicDBObject; +import com.mongodb.DBCollection; +import com.mongodb.DBObject; + +/** + * Tests for {@link MongoTemplate#aggregate(String, AggregationPipeline, Class)}. + * + * @author Tobias Trelle + */ +@RunWith(SpringJUnit4ClassRunner.class) +@ContextConfiguration("classpath:infrastructure.xml") +public class AggregationTests { + + private static final String INPUT_COLLECTION = "aggregation_test_collection"; + + @Autowired + MongoTemplate mongoTemplate; + + @Before + public void setUp() { + cleanDb(); + } + + @After + public void cleanUp() { + cleanDb(); + } + + @Test(expected = IllegalArgumentException.class) + public void shouldHandleMissingInputCollection() { + mongoTemplate.aggregate(null, new AggregationPipeline(), TagCount.class); + } + + @Test(expected = IllegalArgumentException.class) + public void shouldHandleMissingAggregationPipeline() { + mongoTemplate.aggregate(INPUT_COLLECTION, null, TagCount.class); + } + + @Test(expected = IllegalArgumentException.class) + public void shouldHandleMissingEntityClass() { + mongoTemplate.aggregate(INPUT_COLLECTION, new AggregationPipeline(), null); + } + + @Test(expected = IllegalArgumentException.class) + public void shouldDetectIllegalJsonInOperation() { + // given + AggregationPipeline pipeline = new AggregationPipeline().project("{ foo bar"); + + // when + mongoTemplate.aggregate(INPUT_COLLECTION, pipeline, TagCount.class); + + // then: throw expected exception + } + + @Test + public void shouldAggregate() { + // given + createDocuments(); + AggregationPipeline pipeline = new AggregationPipeline() + .project("{_id:0,tags:1}}") + .unwind("tags") + .group("{_id:\"$tags\", n:{$sum:1}}") + .project("{tag: \"$_id\", n:1, _id:0}") + .sort( new Sort(new Sort.Order(Direction.DESC, "n")) ); + + // when + AggregationResults results = mongoTemplate.aggregate(INPUT_COLLECTION, pipeline, TagCount.class); + + // then + assertThat(results, notNullValue()); + assertThat(results.getServerUsed(), is("/127.0.0.1:27017")); + + List tagCount = results.getAggregationResult(); + assertThat(tagCount, notNullValue()); + assertThat(tagCount.size(), is(3)); + assertTagCount("spring", 3, tagCount.get(0)); + assertTagCount("mongodb", 2, tagCount.get(1)); + assertTagCount("nosql", 1, tagCount.get(2)); + } + + @Test(expected = InvalidDataAccessApiUsageException.class) + public void shouldDetectIllegalAggregationOperation() { + // given + createDocuments(); + AggregationPipeline pipeline = new AggregationPipeline().project("{$foobar:{_id:0,tags:1}}"); + + // when + mongoTemplate.aggregate(INPUT_COLLECTION, pipeline, TagCount.class); + + // then: throw expected exception + } + + @Test + public void shouldAggregateEmptyCollection() { + // given + AggregationPipeline pipeline = new AggregationPipeline() + .project("{_id:0,tags:1}}") + .unwind("$tags") + .group("{_id:\"$tags\", n:{$sum:1}}") + .project("{tag: \"$_id\", n:1, _id:0}") + .sort("{n:-1}"); + + // when + AggregationResults results = mongoTemplate.aggregate(INPUT_COLLECTION, pipeline, TagCount.class); + + // then + assertThat(results, notNullValue()); + assertThat(results.getServerUsed(), is("/127.0.0.1:27017")); + + List tagCount = results.getAggregationResult(); + assertThat(tagCount, notNullValue()); + assertThat(tagCount.size(), is(0)); + } + + @Test + public void shouldDetectResultMismatch() { + // given + createDocuments(); + AggregationPipeline pipeline = new AggregationPipeline() + .project("{_id:0,tags:1}}") + .unwind("$tags") + .group("{_id:\"$tags\", count:{$sum:1}}") + .limit(2); + + // when + AggregationResults results = mongoTemplate.aggregate(INPUT_COLLECTION, pipeline, TagCount.class); + + // then + assertThat(results, notNullValue()); + assertThat(results.getServerUsed(), is("/127.0.0.1:27017")); + + List tagCount = results.getAggregationResult(); + assertThat(tagCount, notNullValue()); + assertThat(tagCount.size(), is(2)); + assertTagCount(null, 0, tagCount.get(0)); + assertTagCount(null, 0, tagCount.get(1)); + } + + protected void cleanDb() { + mongoTemplate.dropCollection(INPUT_COLLECTION); + } + + private void createDocuments() { + DBCollection coll = mongoTemplate.getCollection(INPUT_COLLECTION); + coll.insert(createDocument("Doc1", "spring", "mongodb", "nosql")); + coll.insert(createDocument("Doc2", "spring", "mongodb")); + coll.insert(createDocument("Doc3", "spring")); + } + + private DBObject createDocument(String title, String... tags) { + DBObject doc = new BasicDBObject("title", title); + List tagList = new ArrayList(); + for (String tag : tags) { + tagList.add(tag); + } + doc.put("tags", tagList); + + return doc; + } + + private void assertTagCount(String tag, int n, TagCount tagCount) { + assertThat(tagCount.getTag(), is(tag)); + assertThat(tagCount.getN(), is(n)); + } + +} diff --git a/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/aggregation/ProjectionTests.java b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/aggregation/ProjectionTests.java new file mode 100644 index 000000000..25b537633 --- /dev/null +++ b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/aggregation/ProjectionTests.java @@ -0,0 +1,163 @@ +package org.springframework.data.mongodb.core.aggregation; + +import static org.hamcrest.CoreMatchers.is; +import static org.hamcrest.CoreMatchers.notNullValue; +import static org.junit.Assert.assertThat; + +import java.util.List; + +import org.junit.Before; +import org.junit.Test; +import org.springframework.dao.InvalidDataAccessApiUsageException; + +import com.mongodb.DBObject; + +/** + * Tests of {@link Projection}. + * + * @author Tobias Trelle + */ +public class ProjectionTests { + + /** Unit under test. */ + private Projection projection; + + @Before + public void setUp() { + projection = new Projection(); + } + + @Test + public void emptyProjection() { + // when + DBObject raw = projection.toDBObject(); + + // then + assertThat(raw, notNullValue()); + assertThat(raw.toMap().isEmpty(), is(true)); + } + + @Test(expected = IllegalArgumentException.class) + public void shouldDetectNullIncludesInConstructor() { + // when + new Projection((String[]) null); + // then: throw expected exception + } + + @Test + public void includesWithConstructor() { + // given + projection = new Projection("a", "b"); + + // when + DBObject raw = projection.toDBObject(); + + // then + assertThat(raw, notNullValue()); + assertThat(raw.toMap().size(), is(3)); + assertThat((Integer) raw.get("_id"), is(0)); + assertThat((Integer) raw.get("a"), is(1)); + assertThat((Integer) raw.get("b"), is(1)); + } + + @Test + public void include() { + // given + projection.include("a"); + + // when + DBObject raw = projection.toDBObject(); + + // then + assertSingleDBObject("a", 1, raw); + } + + @Test + public void exclude() { + // given + projection.exclude("a"); + + // when + DBObject raw = projection.toDBObject(); + + // then + assertSingleDBObject("a", 0, raw); + } + + @Test + public void includeAlias() { + // given + projection.include("a").as("b"); + + // when + DBObject raw = projection.toDBObject(); + + // then + assertSingleDBObject("b", "$a", raw); + } + + @Test(expected = InvalidDataAccessApiUsageException.class) + public void shouldDetectAliasWithoutInclude() { + // when + projection.as("b"); + // then: throw expected exception + } + + @Test(expected = InvalidDataAccessApiUsageException.class) + public void shouldDetectDuplicateAlias() { + // when + projection.include("a").as("b").as("c"); + // then: throw expected exception + } + + @Test + public void plus() { + // given + projection.include("a").plus(10); + + // when + DBObject raw = projection.toDBObject(); + + // then + assertNotNullDBObject(raw); + DBObject addition = (DBObject)raw.get("a"); + assertNotNullDBObject(addition); + @SuppressWarnings("unchecked") + List summands = (List)addition.get("$add"); + assertThat( summands, notNullValue() ); + assertThat( summands.size(), is(2) ); + assertThat( (String)summands.get(0), is("$a") ); + assertThat( (Integer)summands.get(1), is (10) ); + } + + @Test + public void plusWithAlias() { + // given + projection.include("a").plus(10).as("b"); + + // when + DBObject raw = projection.toDBObject(); + + // then + assertNotNullDBObject(raw); + DBObject addition = (DBObject)raw.get("b"); + assertNotNullDBObject(addition); + @SuppressWarnings("unchecked") + List summands = (List)addition.get("$add"); + assertThat( summands, notNullValue() ); + assertThat( summands.size(), is(2) ); + assertThat( (String)summands.get(0), is("$a") ); + assertThat( (Integer)summands.get(1), is (10) ); + } + + + private static void assertSingleDBObject(String key, Object value, DBObject doc) { + assertNotNullDBObject(doc); + assertThat(doc.get(key), is(value)); + } + + private static void assertNotNullDBObject(DBObject doc) { + assertThat(doc, notNullValue()); + assertThat(doc.toMap().size(), is(1)); + } +} diff --git a/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/aggregation/TagCount.java b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/aggregation/TagCount.java new file mode 100644 index 000000000..2bfdb2271 --- /dev/null +++ b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/aggregation/TagCount.java @@ -0,0 +1,30 @@ +package org.springframework.data.mongodb.core.aggregation; + +/** + * Simple value object holding the aggregation result. + * + * @author Tobias Trelle + */ +public class TagCount { + + private String tag; + + private int n; + + public String getTag() { + return tag; + } + + public void setTag(String tag) { + this.tag = tag; + } + + public int getN() { + return n; + } + + public void setN(int n) { + this.n = n; + } + +}