From 6c9af3fbe84995beb13be2ddcd59644b289fe09f Mon Sep 17 00:00:00 2001 From: Adam Malek Date: Wed, 1 Jun 2022 20:32:47 +0200 Subject: [PATCH] ML in SE: Compute ML weight when creating FoundElementInfo (IDEA-294763) GitOrigin-RevId: 48695b856e8b80427f606df5bb9bedc61a933a17 --- .../searcheverywhere/FoundItemDescriptor.java | 21 +-------- .../ActionSearchEverywhereContributor.java | 23 +--------- .../FileSearchEverywhereContributor.java | 9 ---- .../GroupedResultsSearcher.java | 19 +++++++- .../MixedResultsSearcher.java | 10 ++++- .../SearchEverywhereFoundElementInfo.java | 6 +-- .../SearchEverywhereMlService.kt | 6 +++ .../SearchEverywhereMLStatisticsCollector.kt | 26 +++++------ .../ml/SearchEverywhereMlSearchState.kt | 34 -------------- .../ml/SearchEverywhereMlSessionService.kt | 44 +++++++++++++++++++ .../OpenFeaturesInScratchFileAction.kt | 14 +++--- .../ml/SearchEverywhereRankingModelTest.kt | 8 ++-- 12 files changed, 103 insertions(+), 117 deletions(-) diff --git a/platform/lang-api/src/com/intellij/ide/actions/searcheverywhere/FoundItemDescriptor.java b/platform/lang-api/src/com/intellij/ide/actions/searcheverywhere/FoundItemDescriptor.java index b4b839d430ec..4bb5c2bd7d60 100644 --- a/platform/lang-api/src/com/intellij/ide/actions/searcheverywhere/FoundItemDescriptor.java +++ b/platform/lang-api/src/com/intellij/ide/actions/searcheverywhere/FoundItemDescriptor.java @@ -2,31 +2,12 @@ package com.intellij.ide.actions.searcheverywhere; public class FoundItemDescriptor { - 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 { } public int getWeight() { - return isMlWeightSet() ? mlWeight : weight; + return weight; } } diff --git a/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/ActionSearchEverywhereContributor.java b/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/ActionSearchEverywhereContributor.java index fbeaf125f2e2..6851321f1107 100644 --- a/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/ActionSearchEverywhereContributor.java +++ b/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/ActionSearchEverywhereContributor.java @@ -109,15 +109,7 @@ public class ActionSearchEverywhereContributor implements WeightedSearchEverywhe return true; } - final FoundItemDescriptor descriptor; - SearchEverywhereMlService mlService = SearchEverywhereMlService.getInstance(); - if (mlService != null) { - descriptor = getMLWeightedItemDescriptor(mlService, element); - } - else { - descriptor = new FoundItemDescriptor<>(element, element.getMatchingDegree()); - } - + final var descriptor = new FoundItemDescriptor(element, element.getMatchingDegree()); return consumer.process(descriptor); }); }, progressIndicator); @@ -231,19 +223,6 @@ public class ActionSearchEverywhereContributor implements WeightedSearchEverywhe }); } - private FoundItemDescriptor 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 { @NotNull @Override diff --git a/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/FileSearchEverywhereContributor.java b/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/FileSearchEverywhereContributor.java index 991c07906fb4..2d2489f2e2df 100644 --- a/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/FileSearchEverywhereContributor.java +++ b/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/FileSearchEverywhereContributor.java @@ -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)); } diff --git a/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/GroupedResultsSearcher.java b/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/GroupedResultsSearcher.java index ee081e879da9..bf4c3cdd9d1e 100644 --- a/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/GroupedResultsSearcher.java +++ b/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/GroupedResultsSearcher.java @@ -275,7 +275,14 @@ class GroupedResultsSearcher implements SESearcher { assert contributor == myExpandedContributor; // Only expanded contributor items allowed Collection 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 section = sections.get(contributor); int limit = sectionsLimits.get(contributor); diff --git a/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/MixedResultsSearcher.java b/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/MixedResultsSearcher.java index ff7f34904b8f..595cec4d4cca 100644 --- a/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/MixedResultsSearcher.java +++ b/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/MixedResultsSearcher.java @@ -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 section = mySections.get(contributor); int limit = sectionsLimits.get(contributor); diff --git a/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/SearchEverywhereFoundElementInfo.java b/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/SearchEverywhereFoundElementInfo.java index 5f0e0dacc395..07540593e8d9 100644 --- a/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/SearchEverywhereFoundElementInfo.java +++ b/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/SearchEverywhereFoundElementInfo.java @@ -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; diff --git a/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/SearchEverywhereMlService.kt b/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/SearchEverywhereMlService.kt index 61fa9bfc2be7..fdd3aedc9467 100644 --- a/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/SearchEverywhereMlService.kt +++ b/platform/lang-impl/src/com/intellij/ide/actions/searcheverywhere/SearchEverywhereMlService.kt @@ -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) diff --git a/plugins/search-everywhere-ml/src/com/intellij/ide/actions/searcheverywhere/ml/SearchEverywhereMLStatisticsCollector.kt b/plugins/search-everywhere-ml/src/com/intellij/ide/actions/searcheverywhere/ml/SearchEverywhereMLStatisticsCollector.kt index 2810d6a51ecc..84ad958d0716 100644 --- a/plugins/search-everywhere-ml/src/com/intellij/ide/actions/searcheverywhere/ml/SearchEverywhereMLStatisticsCollector.kt +++ b/plugins/search-everywhere-ml/src/com/intellij/ide/actions/searcheverywhere/ml/SearchEverywhereMLStatisticsCollector.kt @@ -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, elementIdProvider: SearchEverywhereMlItemIdProvider, selectedItems: List, - project: Project?, - state: SearchEverywhereMlSearchState): List>? { + project: Project?): List>? { return ArrayList>().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, project: Project?, - elementIdProvider: SearchEverywhereMlItemIdProvider, - state: SearchEverywhereMlSearchState): ArrayList>? { + elementIdProvider: SearchEverywhereMlItemIdProvider): ArrayList>? { val data = ArrayList>() 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>, 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 diff --git a/plugins/search-everywhere-ml/src/com/intellij/ide/actions/searcheverywhere/ml/SearchEverywhereMlSearchState.kt b/plugins/search-everywhere-ml/src/com/intellij/ide/actions/searcheverywhere/ml/SearchEverywhereMlSearchState.kt index 75f10c95916a..6856c0b4ff1e 100644 --- a/plugins/search-everywhere-ml/src/com/intellij/ide/actions/searcheverywhere/ml/SearchEverywhereMlSearchState.kt +++ b/plugins/search-everywhere-ml/src/com/intellij/ide/actions/searcheverywhere/ml/SearchEverywhereMlSearchState.kt @@ -18,9 +18,6 @@ internal class SearchEverywhereMlSearchState( private val modelProvider: SearchEverywhereModelProvider, private val providersCache: FeaturesProviderCache? ) { - private val cachedElementsInfo: MutableMap = hashMapOf() - private val cachedMLWeight: MutableMap = 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>() 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>() features.addAll(context.features) features.addAll(getElementFeatures(elementId, element, contributor, priority).features) diff --git a/plugins/search-everywhere-ml/src/com/intellij/ide/actions/searcheverywhere/ml/SearchEverywhereMlSessionService.kt b/plugins/search-everywhere-ml/src/com/intellij/ide/actions/searcheverywhere/ml/SearchEverywhereMlSessionService.kt index 357f0c243849..dda7a717f0a3 100644 --- a/plugins/search-everywhere-ml/src/com/intellij/ide/actions/searcheverywhere/ml/SearchEverywhereMlSessionService.kt +++ b/plugins/search-everywhere-ml/src/com/intellij/ide/actions/searcheverywhere/ml/SearchEverywhereMlSessionService.kt @@ -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>>("se-ml-features") + val ML_WEIGHT_KEY = Key.create("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, diff --git a/plugins/search-everywhere-ml/src/com/intellij/ide/actions/searcheverywhere/ml/actions/OpenFeaturesInScratchFileAction.kt b/plugins/search-everywhere-ml/src/com/intellij/ide/actions/searcheverywhere/ml/actions/OpenFeaturesInScratchFileAction.kt index d1dd4ed6d9e3..922cc1f59583 100644 --- a/plugins/search-everywhere-ml/src/com/intellij/ide/actions/searcheverywhere/ml/actions/OpenFeaturesInScratchFileAction.kt +++ b/plugins/search-everywhere-ml/src/com/intellij/ide/actions/searcheverywhere/ml/actions/OpenFeaturesInScratchFileAction.kt @@ -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 = 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 diff --git a/plugins/search-everywhere-ml/test/com/intellij/ide/actions/searcheverywhere/ml/SearchEverywhereRankingModelTest.kt b/plugins/search-everywhere-ml/test/com/intellij/ide/actions/searcheverywhere/ml/SearchEverywhereRankingModelTest.kt index 8b8f7611c750..76daeaab9d31 100644 --- a/plugins/search-everywhere-ml/test/com/intellij/ide/actions/searcheverywhere/ml/SearchEverywhereRankingModelTest.kt +++ b/plugins/search-everywhere-ml/test/com/intellij/ide/actions/searcheverywhere/ml/SearchEverywhereRankingModelTest.kt @@ -25,8 +25,10 @@ internal abstract class SearchEverywhereRankingModelTest protected fun performSearchFor(searchQuery: String, featuresProviderCache: FeaturesProviderCache? = null): RankingAssertion { VirtualFileManager.getInstance().syncRefresh() val rankedElements: List> = 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>) { fun thenAssertElement(element: FoundItemDescriptor<*>) = ElementAssertion(element) fun findElementAndAssert(predicate: (FoundItemDescriptor<*>) -> Boolean) = ElementAssertion(results.find(predicate)!!)