[evaluation-plugin] Context Collection for Unit Test Generation in Evaluation Plugin

Merge-request: IJ-MR-126952
Merged-by: Teodora Stojcheska <teodora.stojcheska@jetbrains.com>

GitOrigin-RevId: 7875780b0b5b15efc1e1cb37b01de7e7a323178e
This commit is contained in:
Alexey Kalina
2024-03-25 12:04:46 +00:00
committed by intellij-monorepo-bot
parent 1c447a08c2
commit 6706f8d13b
8 changed files with 239 additions and 0 deletions
@@ -0,0 +1,34 @@
package com.intellij.cce.metric
import com.intellij.cce.core.Session
import com.intellij.cce.metric.util.Sample
class EmptyContextSessionRatio(override val name: String = "EmptyContextSessionRatio") : Metric {
private val sample = Sample()
override val description: String = "EmptyContextSessionRatio"
override val showByDefault: Boolean = true
override val valueType = MetricValueType.DOUBLE
override val value: Double
get() = sample.mean()
override fun evaluate(sessions: List<Session>): Double {
val fileSample = Sample()
sessions
.flatMap { session -> session.lookups }
.forEach {
val context = it.additionalInfo["context"]
if (context is Map<*, *>) {
if (context["contents"] == "") {
sample.add(1.0)
fileSample.add(1.0)
}
else{
sample.add(0.0)
fileSample.add(0.0)
}
}
}
return fileSample.mean()
}
}
@@ -5,6 +5,7 @@
<completionEvaluationVisitor implementation="com.intellij.cce.visitor.JavaMultiLineEvaluationVisitor"/>
<setupSdkStep implementation="com.intellij.cce.evaluation.SetupJDKStep"/>
<lineCompletionVisitorFactory implementation="com.intellij.cce.visitor.JavaLineCompletionVisitorFactory"/>
<completionEvaluationVisitor implementation="com.intellij.cce.visitor.JavaTestGenerationVisitor"/>
<completionEvaluationVisitor implementation="com.intellij.cce.visitor.JavaCompletionContextEvaluationVisitor"/>
</extensions>
</idea-plugin>
@@ -0,0 +1,90 @@
package com.intellij.cce.visitor
import com.intellij.cce.core.*
import com.intellij.cce.visitor.exceptions.PsiConverterException
import com.intellij.codeInsight.TestFrameworks
import com.intellij.openapi.project.Project
import com.intellij.openapi.vfs.VirtualFile
import com.intellij.psi.*
import com.intellij.psi.impl.source.PsiJavaFileImpl
import com.intellij.psi.search.FilenameIndex
import com.intellij.psi.search.GlobalSearchScope
import com.intellij.psi.util.PsiTreeUtil
import com.intellij.psi.util.descendantsOfType
import java.nio.file.Paths
class JavaTestGenerationVisitor : EvaluationVisitor, JavaRecursiveElementVisitor() {
override val feature: String = "test-generation"
override val language: Language = Language.JAVA
private lateinit var testFile: PsiJavaFile
private var codeFragment: CodeFragment? = null
override fun getFile(): CodeFragment = codeFragment ?: throw PsiConverterException("Invoke 'accept' with visitor on PSI first")
override fun visitJavaFile(file: PsiJavaFile) {
codeFragment = CodeFragment(file.textOffset, file.textLength)
getTestFileForClass(file).let { if(it != null) {
testFile = it
super.visitJavaFile(file)
}
}
}
private fun getTestFileForClass(file: PsiJavaFile): PsiJavaFile? {
val scope = GlobalSearchScope.projectScope(file.project)
val files = FilenameIndex.getVirtualFilesByName(file.name.replace(".java", "Test.java"), scope)
if (files.isEmpty()) return null
return convertVirtualFileToJavaPsiFile(file.project, files.first())
}
private fun convertVirtualFileToJavaPsiFile(project: Project, file: VirtualFile): PsiJavaFile {
val psiManager = PsiManager.getInstance(project)
val psiFile = psiManager.findFile(file)
return psiFile as PsiJavaFileImpl
}
override fun visitMethod(method: PsiMethod) {
if (TestFrameworks.getInstance().isTestMethod(method)) {
return
}
val testMethod = findTestMethodInFile(testFile, method.name)
val testMethodCopy = testMethod?.copy() as? PsiMethod
if (testMethodCopy != null) {
removeComments(testMethodCopy)
val init: MutableMap<String, String>.() -> Unit = {
this["unitTest"] = testMethodCopy.text
this["testPath"] = method.project.basePath?.let { Paths.get(it).relativize(Paths.get(testMethod.containingFile.virtualFile.path)) }.toString()
this["testOffset"] = testMethod.textRange.startOffset.toString()
}
codeFragment?.addChild(CodeToken(method.text, method.textRange.startOffset, getMethodProperties(init)))
}
}
private fun removeComments(method: PsiMethod) {
for (comment in method.descendantsOfType<PsiComment>().toList()) {
if (comment.isValid) {
comment.delete()
}
}
}
private fun findTestMethodInFile(file: PsiJavaFile, methodName: String) : PsiMethod? {
for (fileClass: PsiClass in file.classes) {
for (child: PsiMethod in PsiTreeUtil.findChildrenOfType(fileClass, PsiMethod::class.java)){
if (TestFrameworks.getInstance().isTestMethod(child) && child.name.replace("Test", "") == methodName) {
return child
}
}
}
return null
}
private fun getMethodProperties(init: MutableMap<String, String>.() -> Unit): TokenProperties{
return SimpleTokenProperties.create(TypeProperty.METHOD, SymbolLocation.UNKNOWN, init)
}
}
@@ -43,6 +43,7 @@
<extensions defaultExtensionNs="com.intellij.cce">
<evaluableFeature implementation="com.intellij.cce.evaluable.rename.RenameFeature"/>
<evaluableFeature implementation="com.intellij.cce.evaluable.testGeneration.TestGenerationFeature"/>
<evaluableFeature implementation="com.intellij.cce.evaluable.completion.TokenCompletionFeature"/>
<evaluableFeature implementation="com.intellij.cce.actions.ContextCollectionFeature"/>
<suggestionsProvider implementation="com.intellij.cce.evaluable.completion.DefaultCompletionProvider"/>
@@ -0,0 +1,22 @@
// Copyright 2000-2023 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.intellij.cce.evaluable.testGeneration
import com.intellij.cce.actions.CallFeature
import com.intellij.cce.actions.MoveCaret
import com.intellij.cce.core.CodeFragment
import com.intellij.cce.core.CodeToken
import com.intellij.cce.processor.GenerateActionsProcessor
class TestGenerationActionsProcessor : GenerateActionsProcessor() {
override fun process(code: CodeFragment) {
for (token in code.getChildren()) {
processToken(token as CodeToken)
}
}
private fun processToken(token: CodeToken) {
addAction(MoveCaret(token.offset))
addAction(CallFeature(token.text, token.offset, token.properties))
}
}
@@ -0,0 +1,32 @@
package com.intellij.cce.evaluable.testGeneration
import com.google.gson.JsonObject
import com.google.gson.JsonSerializationContext
import com.intellij.cce.core.Language
import com.intellij.cce.evaluable.EvaluableFeatureBase
import com.intellij.cce.evaluable.StrategySerializer
import com.intellij.cce.evaluation.EvaluationStep
import com.intellij.cce.interpreter.FeatureInvoker
import com.intellij.cce.metric.EmptyContextSessionRatio
import com.intellij.cce.metric.Metric
import com.intellij.cce.metric.SessionsCountMetric
import com.intellij.cce.processor.GenerateActionsProcessor
import com.intellij.openapi.project.Project
import java.lang.reflect.Type
class TestGenerationFeature : EvaluableFeatureBase<TestGenerationStrategy>("test-generation") {
override fun getGenerateActionsProcessor(strategy: TestGenerationStrategy): GenerateActionsProcessor = TestGenerationActionsProcessor()
override fun getFeatureInvoker(project: Project, language: Language, strategy: TestGenerationStrategy): FeatureInvoker =
TestGenerationInvoker(project, language, strategy)
override fun getStrategySerializer(): StrategySerializer<TestGenerationStrategy> = object : StrategySerializer<TestGenerationStrategy> {
override fun deserialize(map: Map<String, Any>, language: String): TestGenerationStrategy = TestGenerationStrategy()
override fun serialize(src: TestGenerationStrategy, typeOfSrc: Type, context: JsonSerializationContext): JsonObject = JsonObject()
}
override fun getMetrics(): List<Metric> = listOf(SessionsCountMetric(), EmptyContextSessionRatio())
override fun getEvaluationSteps(language: Language, strategy: TestGenerationStrategy): List<EvaluationStep> = emptyList()
}
@@ -0,0 +1,47 @@
package com.intellij.cce.evaluable.testGeneration
import com.intellij.cce.core.*
import com.intellij.cce.evaluable.common.getEditorSafe
import com.intellij.cce.evaluation.SuggestionsProvider
import com.intellij.cce.interpreter.FeatureInvoker
import com.intellij.openapi.application.runInEdt
import com.intellij.openapi.application.runReadAction
import com.intellij.openapi.editor.Editor
import com.intellij.openapi.project.Project
import com.intellij.psi.PsiDocumentManager
class TestGenerationInvoker(private val project: Project,
private val language: Language,
private val strategy: TestGenerationStrategy) : FeatureInvoker {
override fun callFeature(expectedText: String, offset: Int, properties: TokenProperties): Session
{
val editor = runReadAction {
getEditorSafe(project)
}
runInEdt {
PsiDocumentManager.getInstance(project).commitDocument(editor.document)
}
val session = Session(offset, expectedText, expectedText.length, TokenProperties.UNKNOWN)
val lookup = getSuggestions(expectedText, editor, strategy.suggestionsProvider)
val res = mapOf(
"unitTest" to (properties as SimpleTokenProperties).additionalProperty("unitTest")!!,
"testPath" to properties.additionalProperty("testPath")!!,
"testOffset" to properties.additionalProperty("testOffset")!!
)
session.addLookup(Lookup.fromExpectedText("", "", lookup.suggestions, lookup.latency, null, false, additionalInfo = lookup.additionalInfo + res, comparator = this::comparator))
return session
}
protected fun getSuggestions(expectedLine: String, editor: Editor, suggestionsProviderName: String): com.intellij.cce.core.Lookup {
val lang = com.intellij.lang.Language.findLanguageByID(language.ideaLanguageId)
?: throw IllegalStateException("Can't find language \"${language.ideaLanguageId}\"")
val provider = SuggestionsProvider.find(project, suggestionsProviderName)
?: throw IllegalStateException("Can't find suggestions provider \"${suggestionsProviderName}\"")
return provider.getSuggestions(expectedLine, editor, lang, this::comparator)
}
override fun comparator(generated: String, expected: String): Boolean = expected == generated
}
@@ -0,0 +1,12 @@
package com.intellij.cce.evaluable.testGeneration
import com.intellij.cce.evaluable.EvaluationStrategy
import com.intellij.cce.filter.EvaluationFilter
class TestGenerationStrategy(val suggestionsProvider: String = DEFAULT_PROVIDER) : EvaluationStrategy {
override val filters: Map<String, EvaluationFilter> = emptyMap()
companion object {
private const val DEFAULT_PROVIDER: String = "LLM-test-gen"
}
}