From 25678f6154021ae03ea98025541ac3afe4c012bc Mon Sep 17 00:00:00 2001 From: Vincent Date: Fri, 21 Oct 2022 22:34:26 +0800 Subject: [PATCH] Finish part 3. Provide util methods to create score function and integrate to search request. --- .../es/repository/StudentEsRepository.java | 2 +- .../java/com/vincent/es/util/SearchInfo.java | 35 ++++- .../java/com/vincent/es/util/SearchUtils.java | 127 ++++++++++++++++-- src/test/java/com/vincent/es/SearchTests.java | 103 +++++++++++++- 4 files changed, 246 insertions(+), 21 deletions(-) diff --git a/src/main/java/com/vincent/es/repository/StudentEsRepository.java b/src/main/java/com/vincent/es/repository/StudentEsRepository.java index 3e8e2e5..a23dd11 100644 --- a/src/main/java/com/vincent/es/repository/StudentEsRepository.java +++ b/src/main/java/com/vincent/es/repository/StudentEsRepository.java @@ -127,7 +127,7 @@ public void deleteById(String id) { public List find(SearchInfo info) { var request = new SearchRequest.Builder() .index(indexName) - .query(info.getBoolQuery()._toQuery()) + .query(info.toQuery()) .sort(info.getSortOptions()) .from(info.getFrom()) .size(info.getSize()) diff --git a/src/main/java/com/vincent/es/util/SearchInfo.java b/src/main/java/com/vincent/es/util/SearchInfo.java index acf5257..5b44398 100644 --- a/src/main/java/com/vincent/es/util/SearchInfo.java +++ b/src/main/java/com/vincent/es/util/SearchInfo.java @@ -1,17 +1,23 @@ package com.vincent.es.util; import co.elastic.clients.elasticsearch._types.SortOptions; -import co.elastic.clients.elasticsearch._types.query_dsl.BoolQuery; -import co.elastic.clients.elasticsearch._types.query_dsl.Query; +import co.elastic.clients.elasticsearch._types.query_dsl.*; +import org.springframework.util.CollectionUtils; import java.util.List; public class SearchInfo { private BoolQuery boolQuery; - private List sortOptions; + private List functionScores = List.of(); + private List sortOptions = List.of(); private Integer from; private Integer size; + public SearchInfo() { + var matchAll = MatchAllQuery.of(b -> b)._toQuery(); + this.boolQuery = BoolQuery.of(b -> b.filter(matchAll)); + } + public static SearchInfo of(BoolQuery bool) { var info = new SearchInfo(); info.boolQuery = bool; @@ -32,6 +38,14 @@ public void setBoolQuery(BoolQuery boolQuery) { this.boolQuery = boolQuery; } + public List getFunctionScores() { + return functionScores; + } + + public void setFunctionScores(List functionScores) { + this.functionScores = functionScores; + } + public List getSortOptions() { return sortOptions; } @@ -55,4 +69,19 @@ public Integer getSize() { public void setSize(Integer size) { this.size = size; } + + public Query toQuery() { + if (CollectionUtils.isEmpty(functionScores)) { + return boolQuery._toQuery(); + } + + return new FunctionScoreQuery.Builder() + .query(boolQuery._toQuery()) + .functions(functionScores) + .scoreMode(FunctionScoreMode.Sum) + .boostMode(FunctionBoostMode.Replace) + .maxBoost(30.0) + .build() + ._toQuery(); + } } diff --git a/src/main/java/com/vincent/es/util/SearchUtils.java b/src/main/java/com/vincent/es/util/SearchUtils.java index bdc7069..ed8e58a 100644 --- a/src/main/java/com/vincent/es/util/SearchUtils.java +++ b/src/main/java/com/vincent/es/util/SearchUtils.java @@ -17,7 +17,7 @@ private SearchUtils() {} *
      *     {
      *         "term": {
-     *             "{field}": {value}
+     *             "{@param field}": {@param value}
      *         }
      *     }
      * 
@@ -30,7 +30,7 @@ public static Query createTermQuery(String field, Object value) { .value((int) value); // 此方法接受 long 型態 } else if (value instanceof String){ builder - .field(field + ".keyword") + .field(field) .value((String) value); } else { throw new UnsupportedOperationException("Please implement for other type additionally."); @@ -43,7 +43,10 @@ public static Query createTermQuery(String field, Object value) { *
      *     {
      *         "terms": {
-     *             "{field}": [{values}[0], {values}[1], ...]
+     *             "{@param field}": [
+     *                 {@param values[0]},
+     *                 {@param values[1]}, ...
+     *             ]
      *         }
      *     }
      * 
@@ -56,7 +59,6 @@ public static Query createTermsQuery(String field, Collection values) { fieldValueStream = values.stream() .map(value -> FieldValue.of(b -> b.longValue((int) value))); } else if (elem instanceof String) { - field = field + ".keyword"; fieldValueStream = values.stream() .map(value -> FieldValue.of(b -> b.stringValue((String) value))); } else { @@ -77,9 +79,9 @@ public static Query createTermsQuery(String field, Collection values) { *
      *     {
      *         "range": {
-     *             "{field}": {
-     *                 "gte": "{gte}",
-     *                 "lte": "{lte}"
+     *             "{@param field}": {
+     *                 "gte": "{@param gte}",
+     *                 "lte": "{@param lte}"
      *             }
      *         }
      *     }
@@ -119,10 +121,10 @@ public static Query createRangeQuery(String field, Date gte, Date lte) {
      *         "bool": {
      *             "should": [
      *                 {
-     *                     "match": { "{fields}[0]": "{searchText}" }
+     *                     "match": { "{@param fields[0]}": "{@param searchText}" }
      *                 },
      *                 {
-     *                     "match": { "{fields}[1]": "{searchText}" }
+     *                     "match": { "{@param fields[1]}": "{@param searchText}" }
      *                 }
      *             ]
      *         }
@@ -156,4 +158,109 @@ public static SortOptions createSortOption(String field, SortOrder order, SortMo
     public static SortOptions createSortOption(String field, SortOrder order) {
         return createSortOption(field, order, null);
     }
-}
+
+    /**
+     * 
+     *     {
+     *         "field_value_factor": {
+     *             "field": {@param field},
+     *             "factor": {@param factor},
+     *             "modifier": {@param modifier},
+     *             "missing": {@param missing}
+     *         }
+     *     }
+     * 
+ */ + public static FunctionScore createFieldValueFactor( + String field, Double factor, FieldValueFactorModifier modifier, Double missing) { + + return new FieldValueFactorScoreFunction.Builder() + .field(field) + .factor(factor) + .modifier(modifier) + .missing(missing) + .build() + ._toFunctionScore(); + } + + /** + *
+     *     {
+     *         "filter": {@param query},
+     *         "weight": {@param weight}
+     *     }
+     * 
+ */ + public static FunctionScore createConditionalWeightFunctionScore(Query query, Double weight) { + return new FunctionScore.Builder() + .filter(query) + .weight(weight) + .build(); + } + + /** + *
+     *     {
+     *         "field_value_factor": {@param function},
+     *         "weight": {@param weight}
+     *     }
+     * 
+ */ + public static FunctionScore createWeightedFieldValueFactor( + FieldValueFactorScoreFunction function, Double weight) { + + return new FunctionScore.Builder() + .fieldValueFactor(function) + .weight(weight) + .build(); + } + + /** + *
+     *     {
+     *         "gauss": {@param placement}
+     *     }
+     * 
+ */ + public static FunctionScore createGaussFunction(String field, DecayPlacement placement) { + var decayFunction = new DecayFunction.Builder() + .field(field) + .placement(placement) + .build(); + return new FunctionScore.Builder() + .gauss(decayFunction) + .build(); + } + + /** + *
+     *     {
+     *         "origin": {@param origin},
+     *         "offset": {@param offset},
+     *         "scale": {@param scale},
+     *         "decay": {@param decay}
+     *     }
+     * 
+ */ + public static DecayPlacement createDecayPlacement( + Number origin, Number offset, Number scale, Double decay) { + + return new DecayPlacement.Builder() + .origin(JsonData.of(origin)) + .offset(JsonData.of(offset)) + .scale(JsonData.of(scale)) + .decay(decay) + .build(); + } + + public static DecayPlacement createDecayPlacement( + String originExp, String offsetExp, String scaleExp, Double decay) { + + return new DecayPlacement.Builder() + .origin(JsonData.of(originExp)) + .offset(JsonData.of(offsetExp)) + .scale(JsonData.of(scaleExp)) + .decay(decay) + .build(); + } +} \ No newline at end of file diff --git a/src/test/java/com/vincent/es/SearchTests.java b/src/test/java/com/vincent/es/SearchTests.java index e31304a..1c2d747 100644 --- a/src/test/java/com/vincent/es/SearchTests.java +++ b/src/test/java/com/vincent/es/SearchTests.java @@ -3,6 +3,7 @@ import co.elastic.clients.elasticsearch._types.SortMode; import co.elastic.clients.elasticsearch._types.SortOrder; import co.elastic.clients.elasticsearch._types.query_dsl.BoolQuery; +import co.elastic.clients.elasticsearch._types.query_dsl.FieldValueFactorModifier; import co.elastic.clients.elasticsearch._types.query_dsl.MatchAllQuery; import com.vincent.es.entity.Student; import com.vincent.es.repository.StudentEsRepository; @@ -34,11 +35,12 @@ public class SearchTests { private StudentEsRepository repository; @Before - public void setup() throws IOException { + public void setup() throws IOException, InterruptedException { repository.init(); var documents = SampleData.get(); repository.insert(documents); + Thread.sleep(2000); } @Test @@ -138,17 +140,104 @@ public void testPaging() { assertDocumentIds(false, students, "101", "102"); } + @Test + public void testFunctionScore_FieldValueFactor() { + var fieldValueFactorScore = SearchUtils + .createFieldValueFactor("grade", 0.5, FieldValueFactorModifier.Square, 0.0); + + var searchInfo = new SearchInfo(); + searchInfo.setFunctionScores(List.of(fieldValueFactorScore)); + + var students = repository.find(searchInfo); + + // Dora (4.0) -> Mario (2.25) -> Vincent (1.0) -> Winnie (0.25) + assertDocumentIds(students, "101", "102", "103", "104"); + } + + @Test + public void testFunctionScore_ConditionalWeight() { + var departmentQuery = SearchUtils + .createTermQuery("departments.keyword", "財務金融"); + var departmentScore = SearchUtils + .createConditionalWeightFunctionScore(departmentQuery, 3.0); + + var courseQuery = SearchUtils + .createTermQuery("courses.name.keyword", "程式設計"); + var courseScore = SearchUtils + .createConditionalWeightFunctionScore(courseQuery, 1.5); + + var fieldValueFactorScore = SearchUtils + .createFieldValueFactor("grade", 1.0, FieldValueFactorModifier.None, 0.0); + var gradeScore = SearchUtils + .createWeightedFieldValueFactor(fieldValueFactorScore.fieldValueFactor(), 0.5); + + var searchInfo = new SearchInfo(); + searchInfo.setFunctionScores(List.of( + departmentScore, + courseScore, + gradeScore + )); + + var students = repository.find(searchInfo); + + // Vincent (5.5) -> Dora (5.0) -> Mario (1.5) -> Winnie (0.5) + assertDocumentIds(students, "103", "101", "102", "104"); + } + + @Test + public void testFunctionScore_DecayFunction_Number() { + var placement = SearchUtils + .createDecayPlacement(100, 15, 10, 0.5); + var decayFunctionScore = SearchUtils + .createGaussFunction("conductScore", placement); + + var searchInfo = new SearchInfo(); + searchInfo.setFunctionScores(List.of(decayFunctionScore)); + + var students = repository.find(searchInfo); + + // Vincent (1.0) -> Mario (0.9726) -> Dora (0.4322) -> Winnie (0.2570) + assertDocumentIds(students, "103", "102", "101", "104"); + } + + @Test + public void testFunctionScore_DecayFunction_Date() { + var placement = SearchUtils + .createDecayPlacement("now", "90d", "270d", 0.5); + var decayFunctionScore = SearchUtils + .createGaussFunction("englishIssuedDate", placement); + + var searchInfo = new SearchInfo(); + searchInfo.setFunctionScores(List.of(decayFunctionScore)); + + var students = repository.find(searchInfo); + + // Mario -> Winnie -> Dora -> Vincent + assertDocumentIds(students, "102", "104", "101", "103"); + } + private void assertDocumentIds(boolean ignoreOrder, List actualDocs, String... expectedIdArray) { + if (!ignoreOrder) { + assertDocumentIds(actualDocs, expectedIdArray); + return; + } + var expectedIds = List.of(expectedIdArray); var actualIds = actualDocs.stream() .map(Student::getId) .collect(Collectors.toList()); - if (ignoreOrder) { - assertTrue(expectedIds.containsAll(actualIds)); - assertTrue(actualIds.containsAll(expectedIds)); - } else { - assertEquals(expectedIds, actualIds); - } + assertTrue(expectedIds.containsAll(actualIds)); + assertTrue(actualIds.containsAll(expectedIds)); } + + private void assertDocumentIds(List actualDocs, String... expectedIdArray) { + var expectedIds = List.of(expectedIdArray); + var actualIds = actualDocs.stream() + .map(Student::getId) + .collect(Collectors.toList()); + + assertEquals(expectedIds, actualIds); + } + }