[ml-ranking] Extract test case for running completion with Model-based reordering

GitOrigin-RevId: ea325981d9f0b3a6d420cc8f5dbfcb48e14d2ed0
This commit is contained in:
Vitaliy.Bibaev
2022-08-24 15:43:35 +00:00
committed by intellij-monorepo-bot
parent cf2cc5cbf8
commit 102fb895e1
2 changed files with 139 additions and 85 deletions
@@ -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)
}
}
}
}
@@ -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
}
}