From 8381dfdf9f5801eca2403502dac1fe3a1147dcf3 Mon Sep 17 00:00:00 2001 From: Vadim Lomshakov Date: Fri, 9 Apr 2021 15:35:35 +0300 Subject: [PATCH] [ml-completion] fixed tests GitOrigin-RevId: 3272533e32bf10c2bcf1ffcb455cdaf8262b651a --- .../completion/ml/MLWeigherUtil.kt | 20 +++++++++++++++---- .../completion/ml/util/RelevanceUtil.kt | 9 +++++---- 2 files changed, 21 insertions(+), 8 deletions(-) diff --git a/platform/lang-impl/src/com/intellij/codeInsight/completion/ml/MLWeigherUtil.kt b/platform/lang-impl/src/com/intellij/codeInsight/completion/ml/MLWeigherUtil.kt index de3feea1b66c..1999070271fd 100644 --- a/platform/lang-impl/src/com/intellij/codeInsight/completion/ml/MLWeigherUtil.kt +++ b/platform/lang-impl/src/com/intellij/codeInsight/completion/ml/MLWeigherUtil.kt @@ -6,10 +6,7 @@ import com.intellij.codeInsight.completion.CompletionService import com.intellij.codeInsight.completion.CompletionSorter import com.intellij.codeInsight.lookup.LookupElement import com.intellij.codeInsight.lookup.LookupElementWeigher -import com.intellij.psi.ForceableComparable -import com.intellij.psi.Weigher -import com.intellij.psi.WeigherExtensionPoint -import com.intellij.psi.WeighingService +import com.intellij.psi.* import org.jetbrains.annotations.ApiStatus @ApiStatus.Internal @@ -48,6 +45,15 @@ object MLWeigherUtil { return mlWeigher.weigh(element, location) } } + + @ApiStatus.Internal + fun extractWeightsOrNull(comparable: Any): Map? { + return when (comparable) { + is DummyWeigherComparableDelegate -> comparable.getWeights() + is WeighingComparable<*, *> -> comparable.weights + else -> null + } + } } private class DummyWeigherComparableDelegate(private val weigh: Comparable<*>?) @@ -67,6 +73,12 @@ private class DummyWeigherComparableDelegate(private val weigh: Comparable<*>?) return 0 } + fun getWeights(): Map? { + weigh ?: return emptyMap() + + return if (weigh is WeighingComparable<*, *>) weigh.weights else null + } + override fun toString(): String { return weigh?.toString() ?: "" } diff --git a/plugins/completion-ml-ranking/src/com/intellij/completion/ml/util/RelevanceUtil.kt b/plugins/completion-ml-ranking/src/com/intellij/completion/ml/util/RelevanceUtil.kt index a137df9e42d4..da39565dbe1e 100644 --- a/plugins/completion-ml-ranking/src/com/intellij/completion/ml/util/RelevanceUtil.kt +++ b/plugins/completion-ml-ranking/src/com/intellij/completion/ml/util/RelevanceUtil.kt @@ -1,5 +1,6 @@ package com.intellij.completion.ml.util +import com.intellij.codeInsight.completion.ml.MLWeigherUtil import com.intellij.completion.ml.features.MLCompletionWeigher import com.intellij.completion.ml.sorting.FeatureUtils import com.intellij.internal.statistic.utils.PluginType @@ -7,7 +8,6 @@ import com.intellij.internal.statistic.utils.getPluginInfo import com.intellij.openapi.diagnostic.logger import com.intellij.openapi.util.Pair import com.intellij.openapi.util.text.StringUtil -import com.intellij.psi.WeighingComparable object RelevanceUtil { private val LOG = logger() @@ -79,12 +79,13 @@ object RelevanceUtil { } private fun MutableMap.addProximityValues(prefix: String, proximity: Any) { - if (proximity !is WeighingComparable<*, *>) { - LOG.error("Unexpected value type of `$prefix`: ${proximity.javaClass.simpleName}") + val weights = MLWeigherUtil.extractWeightsOrNull(proximity) + if (weights == null) { + LOG.error("Unexpected comparable type for `$prefix` weigher: ${proximity.javaClass.simpleName}") return } - for ((name, weight) in proximity.weights) { + for ((name, weight) in weights) { this["${prefix}_$name"] = weight } }