mirror of
https://gitflic.ru/project/openide/openide.git
synced 2026-09-27 10:03:11 +07:00
ML in SE: Compute ML weight when creating FoundElementInfo (IDEA-294763)
GitOrigin-RevId: 48695b856e8b80427f606df5bb9bedc61a933a17
This commit is contained in:
committed by
intellij-monorepo-bot
parent
5b1ef53f46
commit
6c9af3fbe8
+1
-20
@@ -2,31 +2,12 @@
|
||||
package com.intellij.ide.actions.searcheverywhere;
|
||||
|
||||
public class FoundItemDescriptor<I> {
|
||||
private final static int MAX_WEIGHT = 10_000;
|
||||
private final static int ML_WEIGHT_DISABLED_VALUE = -1;
|
||||
|
||||
private final I item;
|
||||
private final int weight;
|
||||
private final int mlWeight;
|
||||
|
||||
public FoundItemDescriptor(I item, int weight) {
|
||||
this.item = item;
|
||||
this.weight = weight;
|
||||
mlWeight = ML_WEIGHT_DISABLED_VALUE; // not set
|
||||
}
|
||||
|
||||
public FoundItemDescriptor(I item, int weight, double mlWeight) {
|
||||
this.item = item;
|
||||
this.weight = weight;
|
||||
this.mlWeight = (int)(mlWeight * MAX_WEIGHT) * 100_000 + this.weight;
|
||||
}
|
||||
|
||||
public int getMlWeight() {
|
||||
return mlWeight;
|
||||
}
|
||||
|
||||
public boolean isMlWeightSet() {
|
||||
return mlWeight != ML_WEIGHT_DISABLED_VALUE;
|
||||
}
|
||||
|
||||
public I getItem() {
|
||||
@@ -34,6 +15,6 @@ public class FoundItemDescriptor<I> {
|
||||
}
|
||||
|
||||
public int getWeight() {
|
||||
return isMlWeightSet() ? mlWeight : weight;
|
||||
return weight;
|
||||
}
|
||||
}
|
||||
|
||||
+1
-22
@@ -109,15 +109,7 @@ public class ActionSearchEverywhereContributor implements WeightedSearchEverywhe
|
||||
return true;
|
||||
}
|
||||
|
||||
final FoundItemDescriptor<GotoActionModel.MatchedValue> descriptor;
|
||||
SearchEverywhereMlService mlService = SearchEverywhereMlService.getInstance();
|
||||
if (mlService != null) {
|
||||
descriptor = getMLWeightedItemDescriptor(mlService, element);
|
||||
}
|
||||
else {
|
||||
descriptor = new FoundItemDescriptor<>(element, element.getMatchingDegree());
|
||||
}
|
||||
|
||||
final var descriptor = new FoundItemDescriptor<GotoActionModel.MatchedValue>(element, element.getMatchingDegree());
|
||||
return consumer.process(descriptor);
|
||||
});
|
||||
}, progressIndicator);
|
||||
@@ -231,19 +223,6 @@ public class ActionSearchEverywhereContributor implements WeightedSearchEverywhe
|
||||
});
|
||||
}
|
||||
|
||||
private FoundItemDescriptor<GotoActionModel.MatchedValue> getMLWeightedItemDescriptor(@NotNull SearchEverywhereMlService service,
|
||||
@NotNull GotoActionModel.MatchedValue element) {
|
||||
double mlWeight = service.getMlWeight(this, element, element.getMatchingDegree());
|
||||
if (mlWeight > 0 && service.shouldOrderByMl()) {
|
||||
if (element.getType() == GotoActionModel.MatchedValueType.ABBREVIATION) {
|
||||
return new FoundItemDescriptor<>(element, element.getMatchingDegree(), 1.0);
|
||||
}
|
||||
|
||||
return new FoundItemDescriptor<>(element, element.getMatchingDegree(), mlWeight);
|
||||
}
|
||||
return new FoundItemDescriptor<>(element, element.getMatchingDegree());
|
||||
}
|
||||
|
||||
public static class Factory implements SearchEverywhereContributorFactory<GotoActionModel.MatchedValue> {
|
||||
@NotNull
|
||||
@Override
|
||||
|
||||
-9
@@ -107,15 +107,6 @@ public class FileSearchEverywhereContributor extends AbstractGotoSEContributor {
|
||||
return true;
|
||||
}
|
||||
|
||||
SearchEverywhereMlService mlService = SearchEverywhereMlService.getInstance();
|
||||
if (mlService != null) {
|
||||
double mlWeight = mlService.getMlWeight(this, element, degree);
|
||||
|
||||
if (mlWeight >= 0.0 && mlService.shouldOrderByMl()) {
|
||||
return consumer.process(new FoundItemDescriptor<>(element, degree, mlWeight));
|
||||
}
|
||||
}
|
||||
|
||||
return consumer.process(new FoundItemDescriptor<>(element, degree));
|
||||
}
|
||||
|
||||
|
||||
+17
-2
@@ -275,7 +275,14 @@ class GroupedResultsSearcher implements SESearcher {
|
||||
assert contributor == myExpandedContributor; // Only expanded contributor items allowed
|
||||
|
||||
Collection<SearchEverywhereFoundElementInfo> section = sections.get(contributor);
|
||||
SearchEverywhereFoundElementInfo newElementInfo = new SearchEverywhereFoundElementInfo(element, priority, contributor);
|
||||
final var mlService = SearchEverywhereMlService.getInstance();
|
||||
final SearchEverywhereFoundElementInfo newElementInfo;
|
||||
if (mlService == null) {
|
||||
newElementInfo = new SearchEverywhereFoundElementInfo(element, priority, contributor);
|
||||
}
|
||||
else {
|
||||
newElementInfo = mlService.createFoundElementInfo(contributor, element, priority);
|
||||
}
|
||||
|
||||
if (section.size() >= myNewLimit) {
|
||||
return false;
|
||||
@@ -347,7 +354,15 @@ class GroupedResultsSearcher implements SESearcher {
|
||||
|
||||
@Override
|
||||
public boolean addElement(Object element, SearchEverywhereContributor<?> contributor, int priority, ProgressIndicator indicator) throws InterruptedException {
|
||||
SearchEverywhereFoundElementInfo newElementInfo = new SearchEverywhereFoundElementInfo(element, priority, contributor);
|
||||
final var mlService = SearchEverywhereMlService.getInstance();
|
||||
SearchEverywhereFoundElementInfo newElementInfo;
|
||||
if (mlService == null) {
|
||||
newElementInfo = new SearchEverywhereFoundElementInfo(element, priority, contributor);
|
||||
}
|
||||
else {
|
||||
newElementInfo = mlService.createFoundElementInfo(contributor, element, priority);
|
||||
}
|
||||
|
||||
Condition condition = conditionsMap.get(contributor);
|
||||
Collection<SearchEverywhereFoundElementInfo> section = sections.get(contributor);
|
||||
int limit = sectionsLimits.get(contributor);
|
||||
|
||||
+9
-1
@@ -268,7 +268,15 @@ class MixedResultsSearcher implements SESearcher {
|
||||
}
|
||||
|
||||
public boolean addElement(Object element, SearchEverywhereContributor<?> contributor, int priority, ProgressIndicator indicator) throws InterruptedException {
|
||||
SearchEverywhereFoundElementInfo newElementInfo = new SearchEverywhereFoundElementInfo(element, priority, contributor);
|
||||
final var mlService = SearchEverywhereMlService.getInstance();
|
||||
final SearchEverywhereFoundElementInfo newElementInfo;
|
||||
if (mlService == null) {
|
||||
newElementInfo = new SearchEverywhereFoundElementInfo(element, priority, contributor);
|
||||
}
|
||||
else {
|
||||
newElementInfo = mlService.createFoundElementInfo(contributor, element, priority);
|
||||
}
|
||||
|
||||
Condition condition = conditionsMap.get(contributor);
|
||||
Collection<SearchEverywhereFoundElementInfo> section = mySections.get(contributor);
|
||||
int limit = sectionsLimits.get(contributor);
|
||||
|
||||
+2
-4
@@ -1,18 +1,16 @@
|
||||
// Copyright 2000-2021 JetBrains s.r.o. Use of this source code is governed by the Apache 2.0 license that can be found in the LICENSE file.
|
||||
package com.intellij.ide.actions.searcheverywhere;
|
||||
|
||||
import com.intellij.util.concurrency.annotations.RequiresEdt;
|
||||
import com.intellij.openapi.util.UserDataHolderBase;
|
||||
import org.jetbrains.annotations.ApiStatus;
|
||||
import org.jetbrains.annotations.NotNull;
|
||||
|
||||
import javax.swing.*;
|
||||
import java.util.Comparator;
|
||||
|
||||
/**
|
||||
* Class containing info about found elements
|
||||
*/
|
||||
@ApiStatus.Internal
|
||||
public class SearchEverywhereFoundElementInfo {
|
||||
public class SearchEverywhereFoundElementInfo extends UserDataHolderBase {
|
||||
public final int priority;
|
||||
public final Object element;
|
||||
public final SearchEverywhereContributor<?> contributor;
|
||||
|
||||
+6
@@ -5,6 +5,7 @@ import com.intellij.openapi.diagnostic.Logger
|
||||
import com.intellij.openapi.extensions.ExtensionPointName
|
||||
import com.intellij.openapi.project.Project
|
||||
import org.jetbrains.annotations.ApiStatus
|
||||
import org.jetbrains.annotations.Contract
|
||||
|
||||
@ApiStatus.Internal
|
||||
abstract class SearchEverywhereMlService {
|
||||
@@ -39,6 +40,11 @@ abstract class SearchEverywhereMlService {
|
||||
|
||||
abstract fun onSessionStarted(project: Project?, mixedListInfo: SearchEverywhereMixedListInfo)
|
||||
|
||||
@Contract("_, _, _ -> new")
|
||||
abstract fun createFoundElementInfo(contributor: SearchEverywhereContributor<*>,
|
||||
element: Any,
|
||||
priority: Int): SearchEverywhereFoundElementInfo
|
||||
|
||||
abstract fun onSearchRestart(project: Project?, tabId: String, reason: SearchRestartReason,
|
||||
keysTyped: Int, backspacesTyped: Int, searchQuery: String,
|
||||
previousElementsProvider: () -> List<SearchEverywhereFoundElementInfo>)
|
||||
|
||||
+12
-14
@@ -124,7 +124,7 @@ internal class SearchEverywhereMLStatisticsCollector : CounterUsagesCollector()
|
||||
data.addAll(additional)
|
||||
data.addAll(context.features)
|
||||
|
||||
val elementData = getElementsData(selectedElements, elements, elementIdProvider, selectedItems, project, state)
|
||||
val elementData = getElementsData(selectedElements, elements, elementIdProvider, selectedItems, project)
|
||||
data.add(IS_PROJECT_DISPOSED_KEY.with(elementData == null))
|
||||
if (elementData != null) {
|
||||
data.addAll(elementData)
|
||||
@@ -145,11 +145,10 @@ internal class SearchEverywhereMLStatisticsCollector : CounterUsagesCollector()
|
||||
elements: List<SearchEverywhereFoundElementInfo>,
|
||||
elementIdProvider: SearchEverywhereMlItemIdProvider,
|
||||
selectedItems: List<Any>,
|
||||
project: Project?,
|
||||
state: SearchEverywhereMlSearchState): List<EventPair<*>>? {
|
||||
project: Project?): List<EventPair<*>>? {
|
||||
return ArrayList<EventPair<*>>().apply {
|
||||
addAll(getSelectedElementsData(selectedElements, elements, elementIdProvider, selectedItems))
|
||||
addAll(getCollectedElementsData(elements, project, elementIdProvider, state) ?: return null)
|
||||
addAll(getCollectedElementsData(elements, project, elementIdProvider) ?: return null)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -182,8 +181,7 @@ internal class SearchEverywhereMLStatisticsCollector : CounterUsagesCollector()
|
||||
*/
|
||||
private fun getCollectedElementsData(elements: List<SearchEverywhereFoundElementInfo>,
|
||||
project: Project?,
|
||||
elementIdProvider: SearchEverywhereMlItemIdProvider,
|
||||
state: SearchEverywhereMlSearchState): ArrayList<EventPair<*>>? {
|
||||
elementIdProvider: SearchEverywhereMlItemIdProvider): ArrayList<EventPair<*>>? {
|
||||
val data = ArrayList<EventPair<*>>()
|
||||
val actionManager = ActionManager.getInstance()
|
||||
val value = elements.take(REPORTED_ITEMS_LIMIT).map {
|
||||
@@ -195,7 +193,7 @@ internal class SearchEverywhereMLStatisticsCollector : CounterUsagesCollector()
|
||||
CONTRIBUTOR_ID_KEY.with(it.contributor.searchProviderId)
|
||||
)
|
||||
|
||||
addElementFeatures(elementIdProvider.getId(it.element), it, state, result, actionManager)
|
||||
addElementFeatures(elementIdProvider.getId(it.element), it, result, actionManager)
|
||||
ObjectEventData(result)
|
||||
}
|
||||
data.add(COLLECTED_RESULTS_DATA_KEY.with(value))
|
||||
@@ -219,19 +217,19 @@ internal class SearchEverywhereMLStatisticsCollector : CounterUsagesCollector()
|
||||
|
||||
private fun addElementFeatures(elementId: Int?,
|
||||
elementInfo: SearchEverywhereFoundElementInfo,
|
||||
state: SearchEverywhereMlSearchState,
|
||||
result: MutableList<EventPair<*>>,
|
||||
actionManager: ActionManager) {
|
||||
val itemInfo = state.getElementFeatures(elementId, elementInfo.element, elementInfo.contributor, elementInfo.priority)
|
||||
if (itemInfo.features.isNotEmpty()) {
|
||||
result.add(FEATURES_DATA_KEY.with(ObjectEventData(itemInfo.features)))
|
||||
val mlFeatures = elementInfo.getUserData(SearchEverywhereMlSessionService.ML_FEATURES_KEY) ?: emptyList()
|
||||
if (mlFeatures.isNotEmpty()) {
|
||||
result.add(FEATURES_DATA_KEY.with(ObjectEventData(mlFeatures)))
|
||||
}
|
||||
|
||||
state.getMLWeightIfDefined(elementId)?.let { score ->
|
||||
result.add(ML_WEIGHT_KEY.with(roundDouble(score)))
|
||||
val mlWeight = elementInfo.getUserData(SearchEverywhereMlSessionService.ML_WEIGHT_KEY) ?: -1.0
|
||||
if (mlWeight >= 0.0) {
|
||||
result.add(ML_WEIGHT_KEY.with(roundDouble(mlWeight)))
|
||||
}
|
||||
|
||||
itemInfo.id?.let { result.add(ID_KEY.with(it)) }
|
||||
elementId?.let { result.add(ID_KEY.with(it)) }
|
||||
|
||||
doWhenIsActionWrapper(elementInfo.element) {
|
||||
val action = it.action
|
||||
|
||||
-34
@@ -18,9 +18,6 @@ internal class SearchEverywhereMlSearchState(
|
||||
private val modelProvider: SearchEverywhereModelProvider,
|
||||
private val providersCache: FeaturesProviderCache?
|
||||
) {
|
||||
private val cachedElementsInfo: MutableMap<Int, SearchEverywhereMLItemInfo> = hashMapOf()
|
||||
private val cachedMLWeight: MutableMap<Int, Double> = hashMapOf()
|
||||
|
||||
val searchStateFeatures = SearchEverywhereStateFeaturesProvider().getSearchStateFeatures(tabId, searchQuery)
|
||||
|
||||
private val model: SearchEverywhereRankingModel by lazy {
|
||||
@@ -32,18 +29,6 @@ internal class SearchEverywhereMlSearchState(
|
||||
element: Any,
|
||||
contributor: SearchEverywhereContributor<*>,
|
||||
priority: Int): SearchEverywhereMLItemInfo {
|
||||
if (elementId == null) {
|
||||
return computeElementFeatures(null, element, priority, contributor)
|
||||
}
|
||||
return cachedElementsInfo.computeIfAbsent(elementId) {
|
||||
return@computeIfAbsent computeElementFeatures(elementId, element, priority, contributor)
|
||||
}
|
||||
}
|
||||
|
||||
private fun computeElementFeatures(elementId: Int?,
|
||||
element: Any,
|
||||
priority: Int,
|
||||
contributor: SearchEverywhereContributor<*>): SearchEverywhereMLItemInfo {
|
||||
val features = arrayListOf<EventPair<*>>()
|
||||
val contributorId = contributor.searchProviderId
|
||||
SearchEverywhereElementFeaturesProvider.getFeatureProvidersForContributor(contributorId).forEach { provider ->
|
||||
@@ -53,31 +38,12 @@ internal class SearchEverywhereMlSearchState(
|
||||
return SearchEverywhereMLItemInfo(elementId, contributorId, features)
|
||||
}
|
||||
|
||||
@Synchronized
|
||||
fun getMLWeightIfDefined(elementId: Int?): Double? {
|
||||
return elementId?.let { cachedMLWeight[elementId] }
|
||||
}
|
||||
|
||||
@Synchronized
|
||||
fun getMLWeight(elementId: Int?,
|
||||
element: Any,
|
||||
contributor: SearchEverywhereContributor<*>,
|
||||
context: SearchEverywhereMLContextInfo,
|
||||
priority: Int): Double {
|
||||
if (elementId == null) {
|
||||
return computeMLWeight(context, null, element, contributor, priority)
|
||||
}
|
||||
|
||||
return cachedMLWeight.computeIfAbsent(elementId) {
|
||||
computeMLWeight(context, elementId, element, contributor, priority)
|
||||
}
|
||||
}
|
||||
|
||||
private fun computeMLWeight(context: SearchEverywhereMLContextInfo,
|
||||
elementId: Int?,
|
||||
element: Any,
|
||||
contributor: SearchEverywhereContributor<*>,
|
||||
priority: Int): Double {
|
||||
val features = ArrayList<EventPair<*>>()
|
||||
features.addAll(context.features)
|
||||
features.addAll(getElementFeatures(elementId, element, contributor, priority).features)
|
||||
|
||||
+44
@@ -3,13 +3,21 @@ package com.intellij.ide.actions.searcheverywhere.ml
|
||||
|
||||
import com.intellij.ide.actions.searcheverywhere.*
|
||||
import com.intellij.ide.actions.searcheverywhere.ml.settings.SearchEverywhereMlSettings
|
||||
import com.intellij.ide.util.gotoByName.GotoActionModel
|
||||
import com.intellij.internal.statistic.eventLog.events.EventPair
|
||||
import com.intellij.openapi.components.service
|
||||
import com.intellij.openapi.project.Project
|
||||
import com.intellij.openapi.util.Key
|
||||
import java.util.concurrent.atomic.AtomicInteger
|
||||
import java.util.concurrent.atomic.AtomicReference
|
||||
|
||||
internal class SearchEverywhereMlSessionService : SearchEverywhereMlService() {
|
||||
companion object {
|
||||
val ML_FEATURES_KEY = Key.create<List<EventPair<*>>>("se-ml-features")
|
||||
val ML_WEIGHT_KEY = Key.create<Double>("se-ml-weight")
|
||||
|
||||
private const val MAX_ELEMENT_WEIGHT = 10_000
|
||||
|
||||
internal const val RECORDER_CODE = "MLSE"
|
||||
|
||||
fun getService() = EP_NAME.findExtensionOrFail(SearchEverywhereMlSessionService::class.java).takeIf { it.isEnabled() }
|
||||
@@ -53,6 +61,42 @@ internal class SearchEverywhereMlSessionService : SearchEverywhereMlService() {
|
||||
}
|
||||
}
|
||||
|
||||
override fun createFoundElementInfo(contributor: SearchEverywhereContributor<*>,
|
||||
element: Any,
|
||||
priority: Int): SearchEverywhereFoundElementInfo {
|
||||
val foundElementInfoWithoutMl = SearchEverywhereFoundElementInfo(element, priority, contributor)
|
||||
|
||||
if (!isEnabled()) return foundElementInfoWithoutMl
|
||||
|
||||
val session = getCurrentSession() ?: return foundElementInfoWithoutMl
|
||||
val state = session.getCurrentSearchState() ?: return foundElementInfoWithoutMl
|
||||
|
||||
val elementId = session.itemIdProvider.getId(element)
|
||||
val mlFeatures = state.getElementFeatures(elementId, element, contributor, priority).features
|
||||
|
||||
// We can only compute the ML weight if we are in a tab which has a ML model
|
||||
val mlWeight = state.takeIf { SearchEverywhereTabWithMl.findById(it.tabId) != null }
|
||||
?.getMLWeight(elementId, element, contributor, session.cachedContextInfo, priority)
|
||||
|
||||
val foundElementInfo = if (shouldOrderByMl() && mlWeight != null) {
|
||||
val priorityWithMl = getPriorityWithMl(element, mlWeight, priority)
|
||||
SearchEverywhereFoundElementInfo(element, priorityWithMl, contributor)
|
||||
}
|
||||
else {
|
||||
foundElementInfoWithoutMl
|
||||
}
|
||||
|
||||
foundElementInfo.putUserData(ML_FEATURES_KEY, mlFeatures)
|
||||
foundElementInfo.putUserData(ML_WEIGHT_KEY, mlWeight)
|
||||
|
||||
return foundElementInfo
|
||||
}
|
||||
|
||||
private fun getPriorityWithMl(element: Any, mlWeight: Double, priority: Int): Int {
|
||||
val weight = if (element is GotoActionModel.MatchedValue && element.type == GotoActionModel.MatchedValueType.ABBREVIATION) 1.0 else mlWeight
|
||||
return (weight * MAX_ELEMENT_WEIGHT).toInt() * 100_000 + priority
|
||||
}
|
||||
|
||||
override fun onSearchRestart(project: Project?,
|
||||
tabId: String,
|
||||
reason: SearchRestartReason,
|
||||
|
||||
+7
-7
@@ -7,6 +7,8 @@ import com.intellij.ide.actions.searcheverywhere.SearchEverywhereUI
|
||||
import com.intellij.ide.actions.searcheverywhere.ml.SearchEverywhereMlSessionService
|
||||
import com.intellij.ide.actions.searcheverywhere.ml.SearchEverywhereTabWithMl
|
||||
import com.intellij.ide.actions.searcheverywhere.ml.features.SearchEverywhereContributorFeaturesProvider
|
||||
import com.intellij.ide.actions.searcheverywhere.ml.SearchEverywhereMlSessionService.Companion.ML_FEATURES_KEY
|
||||
import com.intellij.ide.actions.searcheverywhere.ml.SearchEverywhereMlSessionService.Companion.ML_WEIGHT_KEY
|
||||
import com.intellij.ide.scratch.ScratchFileCreationHelper
|
||||
import com.intellij.ide.scratch.ScratchFileService
|
||||
import com.intellij.ide.scratch.ScratchRootType
|
||||
@@ -57,12 +59,14 @@ class OpenFeaturesInScratchFileAction : AnAction() {
|
||||
val searchSession = mlSessionService.getCurrentSession()!!
|
||||
val state = searchSession.getCurrentSearchState()
|
||||
|
||||
val tabId = searchEverywhereUI.selectedTabID
|
||||
val features = searchEverywhereUI.foundElementsInfo.map { info ->
|
||||
val rankingWeight = info.priority
|
||||
val contributor = info.contributor.searchProviderId
|
||||
val elementName = StringUtil.notNullize(info.element.toString(), "undefined")
|
||||
val mlWeight = if (isTabWithMl(tabId)) mlSessionService.getMlWeight(info.contributor, info.element, rankingWeight) else null
|
||||
val mlWeight = info.getUserData(ML_WEIGHT_KEY)
|
||||
val mlFeatures: Map<String, Any> = info.getUserData(ML_FEATURES_KEY)
|
||||
?.associate { it.field.name to it.data as Any }
|
||||
?: emptyMap()
|
||||
|
||||
val elementId = searchSession.itemIdProvider.getId(info.element)
|
||||
return@map ElementFeatures(
|
||||
@@ -71,7 +75,7 @@ class OpenFeaturesInScratchFileAction : AnAction() {
|
||||
mlWeight,
|
||||
rankingWeight,
|
||||
contributor,
|
||||
state!!.getElementFeatures(elementId, info.element, info.contributor, rankingWeight).featuresAsMap().toSortedMap()
|
||||
mlFeatures.toSortedMap()
|
||||
)
|
||||
}
|
||||
|
||||
@@ -87,10 +91,6 @@ class OpenFeaturesInScratchFileAction : AnAction() {
|
||||
)
|
||||
}
|
||||
|
||||
private fun isTabWithMl(tabId: String): Boolean {
|
||||
return SearchEverywhereTabWithMl.findById(tabId) != null
|
||||
}
|
||||
|
||||
private fun createScratchFileContext(json: String) = ScratchFileCreationHelper.Context().apply {
|
||||
text = json
|
||||
fileExtension = JsonFileType.DEFAULT_EXTENSION
|
||||
|
||||
+4
-4
@@ -25,8 +25,10 @@ internal abstract class SearchEverywhereRankingModelTest
|
||||
protected fun performSearchFor(searchQuery: String, featuresProviderCache: FeaturesProviderCache? = null): RankingAssertion {
|
||||
VirtualFileManager.getInstance().syncRefresh()
|
||||
val rankedElements: List<FoundItemDescriptor<*>> = filterElements(searchQuery)
|
||||
.map { it.withMlWeight(getMlWeight(it, searchQuery, featuresProviderCache)) }
|
||||
.sortedByDescending { it.mlWeight }
|
||||
.associateWith { getMlWeight(it, searchQuery, featuresProviderCache) }
|
||||
.entries
|
||||
.sortedByDescending { it.value }
|
||||
.map { it.key }
|
||||
|
||||
assert(rankedElements.size > 1) { "Found ${rankedElements.size} which is unsuitable for ranking assertions" }
|
||||
|
||||
@@ -53,8 +55,6 @@ internal abstract class SearchEverywhereRankingModelTest
|
||||
}.fold(emptyMap()) { acc, value -> acc + value }
|
||||
}
|
||||
|
||||
private fun FoundItemDescriptor<*>.withMlWeight(mlWeight: Double) = FoundItemDescriptor(this.item, this.weight, mlWeight)
|
||||
|
||||
protected class RankingAssertion(private val results: List<FoundItemDescriptor<*>>) {
|
||||
fun thenAssertElement(element: FoundItemDescriptor<*>) = ElementAssertion(element)
|
||||
fun findElementAndAssert(predicate: (FoundItemDescriptor<*>) -> Boolean) = ElementAssertion(results.find(predicate)!!)
|
||||
|
||||
Reference in New Issue
Block a user