mirror of
https://gitflic.ru/project/openide/openide.git
synced 2026-09-27 10:03:11 +07:00
[ml-ranking] Extract test case for running completion with Model-based reordering
GitOrigin-RevId: ea325981d9f0b3a6d420cc8f5dbfcb48e14d2ed0
This commit is contained in:
committed by
intellij-monorepo-bot
parent
cf2cc5cbf8
commit
102fb895e1
+15
-85
@@ -1,41 +1,23 @@
|
||||
// Copyright 2000-2020 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.completion.ml.sorting
|
||||
|
||||
import com.intellij.codeInsight.completion.LightFixtureCompletionTestCase
|
||||
import com.intellij.codeInsight.lookup.LookupManagerListener
|
||||
import com.intellij.codeInsight.lookup.impl.LookupImpl
|
||||
import com.intellij.completion.ml.experiment.ExperimentInfo
|
||||
import com.intellij.completion.ml.experiment.ExperimentStatus
|
||||
import com.intellij.completion.ml.ranker.ExperimentModelProvider
|
||||
import com.intellij.completion.ml.settings.CompletionMLRankingSettings
|
||||
import com.intellij.completion.ml.storage.MutableLookupStorage
|
||||
import com.intellij.completion.ml.tracker.LookupTracker
|
||||
import com.intellij.completion.ml.tracker.setupCompletionContext
|
||||
import com.intellij.internal.ml.DecisionFunction
|
||||
import com.intellij.internal.ml.FeatureMapper
|
||||
import com.intellij.internal.ml.FloatFeature
|
||||
import com.intellij.lang.Language
|
||||
import com.intellij.openapi.Disposable
|
||||
import com.intellij.openapi.application.ApplicationManager
|
||||
import com.intellij.openapi.util.Disposer
|
||||
import com.intellij.testFramework.replaceService
|
||||
import junit.framework.TestCase
|
||||
|
||||
class ChangeArrowsDrawingTest : LightFixtureCompletionTestCase() {
|
||||
class ChangeArrowsDrawingTest : MLSortingTestCase() {
|
||||
private val arrowChecker = ArrowPresenceChecker()
|
||||
|
||||
override fun customizeSettings(settings: MLRankingSettingsState): MLRankingSettingsState {
|
||||
return settings.withRankingEnabled(true).withDiffEnabled(true)
|
||||
}
|
||||
|
||||
override fun setUp() {
|
||||
super.setUp()
|
||||
val settings = CompletionMLRankingSettings.getInstance()
|
||||
val settingsStateBefore = MLRankingSettingsState.build("Java", settings)
|
||||
with(settings) {
|
||||
isRankingEnabled = true
|
||||
isShowDiffEnabled = true
|
||||
setLanguageEnabled("Java", true)
|
||||
}
|
||||
|
||||
RankingSupport.enableInTests(testRootDisposable)
|
||||
ApplicationManager.getApplication().replaceService(ExperimentStatus::class.java, TestArrowsExperimentStatus(), testRootDisposable)
|
||||
Disposer.register(testRootDisposable, Disposable { settingsStateBefore.restore(CompletionMLRankingSettings.getInstance()) })
|
||||
project.messageBus.connect(testRootDisposable).subscribe(
|
||||
LookupManagerListener.TOPIC,
|
||||
ItemsDecoratorInitializer()
|
||||
@@ -44,34 +26,30 @@ class ChangeArrowsDrawingTest : LightFixtureCompletionTestCase() {
|
||||
LookupManagerListener.TOPIC,
|
||||
arrowChecker
|
||||
)
|
||||
ExperimentModelProvider.registerProvider(TestInverseRankingModelProvider(), testRootDisposable)
|
||||
}
|
||||
|
||||
fun testArrowsAreShowingAndUpdated() {
|
||||
withRanker(Ranker.Inverse)
|
||||
setupCompletionContext(myFixture)
|
||||
complete()
|
||||
val invokedCountBeforeTyping = arrowChecker.invokedCount
|
||||
type('r')
|
||||
arrowChecker.assertArrowsAvailable()
|
||||
TestCase.assertTrue(arrowChecker.invokedCount > 1)
|
||||
TestCase.assertTrue(arrowChecker.arrowsFound)
|
||||
TestCase.assertTrue(invokedCountBeforeTyping < arrowChecker.invokedCount)
|
||||
}
|
||||
|
||||
private class TestArrowsExperimentStatus : ExperimentStatus {
|
||||
companion object {
|
||||
const val VERSION = 1
|
||||
}
|
||||
fun testNoArrowsWithDefaultRanker() {
|
||||
setupCompletionContext(myFixture)
|
||||
complete()
|
||||
|
||||
override fun forLanguage(language: Language): ExperimentInfo =
|
||||
ExperimentInfo(true, VERSION, true, true, true)
|
||||
|
||||
override fun disable() = Unit
|
||||
|
||||
override fun isDisabled(): Boolean = false
|
||||
TestCase.assertTrue(arrowChecker.invokedCount > 1)
|
||||
TestCase.assertFalse(arrowChecker.arrowsFound)
|
||||
}
|
||||
|
||||
private class ArrowPresenceChecker : LookupTracker() {
|
||||
var invokedCount: Int = 0
|
||||
private var arrowsFound: Boolean = false
|
||||
var arrowsFound: Boolean = false
|
||||
|
||||
override fun lookupCreated(lookup: LookupImpl, storage: MutableLookupStorage) {
|
||||
lookup.addPresentationCustomizer { _, presentation ->
|
||||
@@ -81,53 +59,5 @@ class ChangeArrowsDrawingTest : LightFixtureCompletionTestCase() {
|
||||
presentation
|
||||
}
|
||||
}
|
||||
|
||||
fun assertArrowsAvailable() {
|
||||
TestCase.assertTrue(invokedCount > 1)
|
||||
TestCase.assertTrue(arrowsFound)
|
||||
}
|
||||
}
|
||||
|
||||
private class TestInverseRankingModelProvider : ExperimentModelProvider {
|
||||
override fun experimentGroupNumber(): Int = TestArrowsExperimentStatus.VERSION
|
||||
|
||||
override fun getModel(): DecisionFunction {
|
||||
return object : DecisionFunction {
|
||||
override fun getFeaturesOrder(): Array<FeatureMapper> {
|
||||
return arrayOf(FloatFeature("position", 1000.0, false).createMapper(null))
|
||||
}
|
||||
|
||||
override fun getRequiredFeatures(): List<String> = listOf("position")
|
||||
|
||||
override fun getUnknownFeatures(features: MutableCollection<String>): List<String> = emptyList()
|
||||
|
||||
override fun version(): String? = null
|
||||
|
||||
override fun predict(features: DoubleArray): Double = features[0]
|
||||
}
|
||||
}
|
||||
|
||||
override fun getDisplayNameInSettings(): String = "Test provider"
|
||||
|
||||
override fun isLanguageSupported(language: Language): Boolean = true
|
||||
}
|
||||
|
||||
private data class MLRankingSettingsState(val language: String,
|
||||
val rankingEnabled: Boolean,
|
||||
val languageEnabled: Boolean,
|
||||
val diffEnabled: Boolean) {
|
||||
companion object {
|
||||
fun build(language: String, settings: CompletionMLRankingSettings): MLRankingSettingsState {
|
||||
return MLRankingSettingsState(language, settings.isRankingEnabled, settings.isLanguageEnabled(language), settings.isShowDiffEnabled)
|
||||
}
|
||||
}
|
||||
|
||||
fun restore(settings: CompletionMLRankingSettings) {
|
||||
with(settings) {
|
||||
isRankingEnabled = rankingEnabled
|
||||
isShowDiffEnabled = diffEnabled
|
||||
setLanguageEnabled(language, languageEnabled)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
+124
@@ -0,0 +1,124 @@
|
||||
package com.intellij.completion.ml.sorting
|
||||
|
||||
import com.intellij.codeInsight.completion.LightFixtureCompletionTestCase
|
||||
import com.intellij.completion.ml.experiment.ExperimentInfo
|
||||
import com.intellij.completion.ml.experiment.ExperimentStatus
|
||||
import com.intellij.completion.ml.ranker.ExperimentModelProvider
|
||||
import com.intellij.completion.ml.settings.CompletionMLRankingSettings
|
||||
import com.intellij.internal.ml.DecisionFunction
|
||||
import com.intellij.internal.ml.FeatureMapper
|
||||
import com.intellij.internal.ml.FloatFeature
|
||||
import com.intellij.lang.Language
|
||||
import com.intellij.openapi.Disposable
|
||||
import com.intellij.openapi.application.ApplicationManager
|
||||
import com.intellij.openapi.util.Disposer
|
||||
import com.intellij.testFramework.replaceService
|
||||
|
||||
abstract class MLSortingTestCase : LightFixtureCompletionTestCase() {
|
||||
private var ranker: Ranker = Ranker.Default
|
||||
|
||||
protected abstract fun customizeSettings(settings: MLRankingSettingsState): MLRankingSettingsState
|
||||
|
||||
override fun setUp() {
|
||||
super.setUp()
|
||||
|
||||
setUpRankingSettings()
|
||||
setUpRankingModel()
|
||||
}
|
||||
|
||||
protected fun withRanker(value: Ranker) {
|
||||
ranker = value
|
||||
}
|
||||
|
||||
private fun setUpRankingSettings() {
|
||||
val settings = CompletionMLRankingSettings.getInstance()
|
||||
val settingsStateBefore = MLRankingSettingsState.build("Java", settings)
|
||||
val actualSettingsState = customizeSettings(settingsStateBefore)
|
||||
|
||||
actualSettingsState.setStateTo(settings)
|
||||
|
||||
Disposer.register(testRootDisposable, Disposable { settingsStateBefore.setStateTo(CompletionMLRankingSettings.getInstance()) })
|
||||
|
||||
RankingSupport.enableInTests(testRootDisposable)
|
||||
|
||||
ApplicationManager.getApplication().replaceService(ExperimentStatus::class.java, TestExperimentStatus(actualSettingsState),
|
||||
testRootDisposable)
|
||||
}
|
||||
|
||||
private fun setUpRankingModel() {
|
||||
ExperimentModelProvider.registerProvider(TestModelProvider(::ranker), testRootDisposable)
|
||||
}
|
||||
|
||||
protected sealed interface Ranker {
|
||||
fun score(position: Double): Double
|
||||
|
||||
object Default : Ranker {
|
||||
override fun score(position: Double): Double = -position
|
||||
}
|
||||
|
||||
object Inverse : Ranker {
|
||||
override fun score(position: Double): Double = position
|
||||
}
|
||||
}
|
||||
|
||||
private class TestExperimentStatus(private val settings: MLRankingSettingsState) : ExperimentStatus {
|
||||
companion object {
|
||||
const val VERSION = 1
|
||||
}
|
||||
|
||||
override fun forLanguage(language: Language): ExperimentInfo =
|
||||
ExperimentInfo(true, VERSION, settings.rankingEnabled, settings.diffEnabled, true)
|
||||
|
||||
override fun disable() = Unit
|
||||
|
||||
override fun isDisabled(): Boolean = false
|
||||
}
|
||||
|
||||
protected class MLRankingSettingsState private constructor(private val language: String,
|
||||
val rankingEnabled: Boolean,
|
||||
val diffEnabled: Boolean) {
|
||||
companion object {
|
||||
fun build(language: String, settings: CompletionMLRankingSettings): MLRankingSettingsState {
|
||||
return MLRankingSettingsState(language,
|
||||
settings.isRankingEnabled && settings.isLanguageEnabled(language),
|
||||
settings.isShowDiffEnabled)
|
||||
}
|
||||
}
|
||||
|
||||
fun withRankingEnabled(value: Boolean): MLRankingSettingsState = MLRankingSettingsState(language, value, diffEnabled)
|
||||
|
||||
fun withDiffEnabled(value: Boolean): MLRankingSettingsState = MLRankingSettingsState(language, rankingEnabled, value)
|
||||
|
||||
fun setStateTo(settings: CompletionMLRankingSettings) {
|
||||
with(settings) {
|
||||
isRankingEnabled = rankingEnabled
|
||||
isShowDiffEnabled = diffEnabled
|
||||
setLanguageEnabled(language, rankingEnabled)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private class TestModelProvider(private val rankerSupplier: () -> Ranker) : ExperimentModelProvider {
|
||||
override fun experimentGroupNumber(): Int = TestExperimentStatus.VERSION
|
||||
|
||||
override fun getModel(): DecisionFunction {
|
||||
return object : DecisionFunction {
|
||||
override fun getFeaturesOrder(): Array<FeatureMapper> {
|
||||
return arrayOf(FloatFeature("position", 1000.0, false).createMapper(null))
|
||||
}
|
||||
|
||||
override fun getRequiredFeatures(): List<String> = listOf("position")
|
||||
|
||||
override fun getUnknownFeatures(features: MutableCollection<String>): List<String> = emptyList()
|
||||
|
||||
override fun version(): String? = null
|
||||
|
||||
override fun predict(features: DoubleArray): Double = rankerSupplier.invoke().score(features[0])
|
||||
}
|
||||
}
|
||||
|
||||
override fun getDisplayNameInSettings(): String = "Test provider"
|
||||
|
||||
override fun isLanguageSupported(language: Language): Boolean = true
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user