diff --git a/platform/lang-api/src/com/intellij/codeInsight/completion/ml/GenericFeatureProvider.kt b/platform/lang-api/src/com/intellij/codeInsight/completion/ml/GenericFeatureProvider.kt deleted file mode 100644 index 853a2e6329de..000000000000 --- a/platform/lang-api/src/com/intellij/codeInsight/completion/ml/GenericFeatureProvider.kt +++ /dev/null @@ -1,37 +0,0 @@ -// Copyright 2000-2022 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. -package com.intellij.codeInsight.completion.ml - -import com.intellij.lang.Language -import com.intellij.lang.LanguageExtension -import org.jetbrains.annotations.ApiStatus -import kotlin.system.measureTimeMillis - -@ApiStatus.Internal -interface GenericFeatureProvider { - val name: String - fun knownFeatures(): Set - fun calculateFeatures(context: MLContext): Map - fun isApplicable(context: MLContext): Boolean - - data class FeaturesCalculationResult(val features: Map, val performance: Map) - - companion object { - private val EP_NAME = LanguageExtension("com.intellij.completion.ml.genericFeatures") - - fun forLanguage(language: Language): List = EP_NAME.allForLanguageOrAny(language) - - fun calculateFeatures(context: MLContext, features: Set): FeaturesCalculationResult { - val featureValues = mutableMapOf() - val performance = mutableMapOf() - for (provider in forLanguage(context.position.language)) { - if (provider.knownFeatures().any { it in features }) { - val time = measureTimeMillis { - featureValues.putAll(provider.calculateFeatures(context)) - } - performance[provider.name] = time - } - } - return FeaturesCalculationResult(featureValues, performance) - } - } -} diff --git a/platform/lang-api/src/com/intellij/codeInsight/completion/ml/MLFeature.kt b/platform/lang-api/src/com/intellij/codeInsight/completion/ml/MLFeature.kt new file mode 100644 index 000000000000..db54fb7a095b --- /dev/null +++ b/platform/lang-api/src/com/intellij/codeInsight/completion/ml/MLFeature.kt @@ -0,0 +1,11 @@ +// Copyright 2000-2022 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. +package com.intellij.codeInsight.completion.ml + +interface MLFeature { + val name: String + fun compute(): MLFeatureValue +} + +data class SimpleMLFeature(override val name: String, private val value: MLFeatureValue) : MLFeature { + override fun compute(): MLFeatureValue = value +} diff --git a/platform/lang-api/src/com/intellij/codeInsight/completion/ml/MLFeatureProvider.kt b/platform/lang-api/src/com/intellij/codeInsight/completion/ml/MLFeatureProvider.kt new file mode 100644 index 000000000000..911a4a73a98d --- /dev/null +++ b/platform/lang-api/src/com/intellij/codeInsight/completion/ml/MLFeatureProvider.kt @@ -0,0 +1,20 @@ +// Copyright 2000-2022 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license. +package com.intellij.codeInsight.completion.ml + +import com.intellij.lang.LanguageExtension +import org.jetbrains.annotations.ApiStatus + +@ApiStatus.Internal +interface MLFeatureProvider { + val name: String + fun knownFeatures(): Set + fun calculateFeatures(context: MLContext): List + fun isApplicable(context: MLContext): Boolean + + companion object { + private val EP_NAME = LanguageExtension("com.intellij.completion.ml.features") + + fun applicableProviders(context: MLContext): List = + EP_NAME.allForLanguageOrAny(context.position.language).filter { it.isApplicable(context) } + } +} diff --git a/platform/lang-api/src/com/intellij/codeInsight/completion/ml/MLFeatureValue.kt b/platform/lang-api/src/com/intellij/codeInsight/completion/ml/MLFeatureValue.kt index 793cdba3b0fe..f50f7ac6d186 100644 --- a/platform/lang-api/src/com/intellij/codeInsight/completion/ml/MLFeatureValue.kt +++ b/platform/lang-api/src/com/intellij/codeInsight/completion/ml/MLFeatureValue.kt @@ -1,4 +1,4 @@ -// 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. +// Copyright 2000-2022 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.codeInsight.completion.ml import org.jetbrains.annotations.ApiStatus @@ -57,15 +57,32 @@ sealed class MLFeatureValue { @JvmStatic fun > className(value: T, useSimpleName: Boolean = true): MLFeatureValue = ClassNameValue(value, useSimpleName) + @JvmStatic + fun string(value: String): MLFeatureValue = StringValue(value) + + @JvmStatic + fun list(value: List<*>): MLFeatureValue = ListValue(value) + @JvmStatic fun version(value: String): MLFeatureValue = VersionValue(value) } abstract val value: Any + abstract val isSafe: Boolean - data class BinaryValue internal constructor(override val value: Boolean) : MLFeatureValue() - data class FloatValue internal constructor(override val value: Double) : MLFeatureValue() - data class CategoricalValue internal constructor(override val value: String) : MLFeatureValue() - data class ClassNameValue internal constructor(override val value: Class<*>, val useSimpleName: Boolean) : MLFeatureValue() - data class VersionValue internal constructor(override val value: String) : MLFeatureValue() + abstract class SafeValue : MLFeatureValue() { + override val isSafe: Boolean = true + } + + abstract class UnSafeValue : MLFeatureValue() { + override val isSafe: Boolean = false + } + + data class BinaryValue internal constructor(override val value: Boolean) : SafeValue() + data class FloatValue internal constructor(override val value: Double) : SafeValue() + data class CategoricalValue internal constructor(override val value: String) : SafeValue() + data class ClassNameValue internal constructor(override val value: Class<*>, val useSimpleName: Boolean) : UnSafeValue() + data class VersionValue internal constructor(override val value: String) : UnSafeValue() + data class StringValue internal constructor(override val value: String) : UnSafeValue() + data class ListValue internal constructor(override val value: List<*>) : UnSafeValue() } diff --git a/platform/platform-resources/src/META-INF/CompletionExtensionPoints.xml b/platform/platform-resources/src/META-INF/CompletionExtensionPoints.xml index 96b6d2f6f5e4..3d94100a844c 100644 --- a/platform/platform-resources/src/META-INF/CompletionExtensionPoints.xml +++ b/platform/platform-resources/src/META-INF/CompletionExtensionPoints.xml @@ -15,8 +15,8 @@ - - + + diff --git a/plugins/completion-ml-ranking/src/com/intellij/completion/ml/features/MLFeaturesUtil.kt b/plugins/completion-ml-ranking/src/com/intellij/completion/ml/features/MLFeaturesUtil.kt index 0fe2e1f5f0ff..42af051dfbd1 100644 --- a/plugins/completion-ml-ranking/src/com/intellij/completion/ml/features/MLFeaturesUtil.kt +++ b/plugins/completion-ml-ranking/src/com/intellij/completion/ml/features/MLFeaturesUtil.kt @@ -14,6 +14,7 @@ internal object MLFeaturesUtil { is MLFeatureValue.CategoricalValue -> featureValue.value is MLFeatureValue.ClassNameValue -> getClassNameSafe(featureValue) is MLFeatureValue.VersionValue -> getVersionSafe(featureValue) + else -> if (featureValue.isSafe) featureValue.value else throw IllegalArgumentException("Feature value $featureValue is unsafe") } }