[ml-completion] fixed tests

GitOrigin-RevId: 3272533e32bf10c2bcf1ffcb455cdaf8262b651a
This commit is contained in:
Vadim Lomshakov
2021-04-09 20:02:37 +00:00
committed by intellij-monorepo-bot
parent 12dabb63a4
commit 8381dfdf9f
2 changed files with 21 additions and 8 deletions
@@ -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<String, Any>? {
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<String, Any>? {
weigh ?: return emptyMap()
return if (weigh is WeighingComparable<*, *>) weigh.weights else null
}
override fun toString(): String {
return weigh?.toString() ?: ""
}
@@ -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<RelevanceUtil>()
@@ -79,12 +79,13 @@ object RelevanceUtil {
}
private fun MutableMap<String, Any>.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
}
}