diff --git a/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/query/TextQuery.java b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/query/TextQuery.java index bccbbc72e..8b8a0ab02 100644 --- a/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/query/TextQuery.java +++ b/spring-data-mongodb/src/main/java/org/springframework/data/mongodb/core/query/TextQuery.java @@ -175,32 +175,44 @@ public class TextQuery extends Query { public Document getSortObject() { if (this.sortByScore) { - if (sortByScoreIndex == 0) { - Document sort = new Document(); - sort.put(getScoreFieldName(), META_TEXT_SCORE); - sort.putAll(super.getSortObject()); - return sort; - } - return fitInSortByScoreAtPosition(super.getSortObject()); + + int sortByScoreIndex = this.sortByScoreIndex; + + return sortByScoreIndex != 0 + ? sortByScoreAtPosition(super.getSortObject(), sortByScoreIndex) + : sortByScoreAtPositionZero(); } return super.getSortObject(); } - private Document fitInSortByScoreAtPosition(Document source) { + private Document sortByScoreAtPositionZero() { + + Document sort = new Document(); + + sort.put(getScoreFieldName(), META_TEXT_SCORE); + sort.putAll(super.getSortObject()); + + return sort; + } + + private Document sortByScoreAtPosition(Document source, int sortByScoreIndex) { Document target = new Document(); - int i = 0; + int index = 0; + for (Entry entry : source.entrySet()) { - if (i == sortByScoreIndex) { + if (index == sortByScoreIndex) { target.put(getScoreFieldName(), META_TEXT_SCORE); } target.put(entry.getKey(), entry.getValue()); - i++; + index++; } - if (i == sortByScoreIndex) { + + if (index == sortByScoreIndex) { target.put(getScoreFieldName(), META_TEXT_SCORE); } + return target; } diff --git a/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/query/TextQueryUnitTests.java b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/query/TextQueryUnitTests.java index ffdeaf39b..228e7efca 100644 --- a/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/query/TextQueryUnitTests.java +++ b/spring-data-mongodb/src/test/java/org/springframework/data/mongodb/core/query/TextQueryUnitTests.java @@ -15,9 +15,7 @@ */ package org.springframework.data.mongodb.core.query; -import static org.springframework.data.mongodb.test.util.Assertions.*; - -import java.util.Map.Entry; +import static org.springframework.data.mongodb.test.util.Assertions.assertThat; import org.junit.jupiter.api.Test; import org.springframework.data.domain.Sort; @@ -103,20 +101,20 @@ public class TextQueryUnitTests { query.sortByScore(); query.with(Sort.by(Direction.DESC, "two")); - assertThat(query.getSortObject().entrySet().stream().map(Entry::getKey)).containsExactly("one", "score", "two"); + assertThat(query.getSortObject().keySet().stream()).containsExactly("one", "score", "two"); query = new TextQuery(QUERY); query.with(Sort.by(Direction.DESC, "one")); query.sortByScore(); - assertThat(query.getSortObject().entrySet().stream().map(Entry::getKey)).containsExactly("one", "score"); + assertThat(query.getSortObject().keySet().stream()).containsExactly("one", "score"); query = new TextQuery(QUERY); query.sortByScore(); query.with(Sort.by(Direction.DESC, "one")); query.with(Sort.by(Direction.DESC, "two")); - assertThat(query.getSortObject().entrySet().stream().map(Entry::getKey)).containsExactly("score", "one", "two"); + assertThat(query.getSortObject().keySet().stream()).containsExactly("score", "one", "two"); } }