diff --git a/plugins/completion-ml-ranking/test/com/intellij/completion/ml/sorting/ChangeArrowsDrawingTest.kt b/plugins/completion-ml-ranking/test/com/intellij/completion/ml/sorting/ChangeArrowsDrawingTest.kt index da10809eb5d4..ae71d76f0c9f 100644 --- a/plugins/completion-ml-ranking/test/com/intellij/completion/ml/sorting/ChangeArrowsDrawingTest.kt +++ b/plugins/completion-ml-ranking/test/com/intellij/completion/ml/sorting/ChangeArrowsDrawingTest.kt @@ -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 { - return arrayOf(FloatFeature("position", 1000.0, false).createMapper(null)) - } - - override fun getRequiredFeatures(): List = listOf("position") - - override fun getUnknownFeatures(features: MutableCollection): List = 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) - } - } } } \ No newline at end of file diff --git a/plugins/completion-ml-ranking/test/com/intellij/completion/ml/sorting/MLSortingTestCase.kt b/plugins/completion-ml-ranking/test/com/intellij/completion/ml/sorting/MLSortingTestCase.kt new file mode 100644 index 000000000000..e60affec58d3 --- /dev/null +++ b/plugins/completion-ml-ranking/test/com/intellij/completion/ml/sorting/MLSortingTestCase.kt @@ -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 { + return arrayOf(FloatFeature("position", 1000.0, false).createMapper(null)) + } + + override fun getRequiredFeatures(): List = listOf("position") + + override fun getUnknownFeatures(features: MutableCollection): List = 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 + } +} \ No newline at end of file