[java] Fix retrieving failed call in editable test diff

Also rewrites part of literal discovery to UAST. #IDEA-314369 Fixed

GitOrigin-RevId: cf87f54ba6c487f74dc0215e81a065bb2e013fdb
This commit is contained in:
Bart van Helvert
2023-03-09 17:32:37 +00:00
committed by intellij-monorepo-bot
parent 42a397fc09
commit 7b735a0270
5 changed files with 153 additions and 124 deletions
@@ -1,55 +1,8 @@
// 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.execution.testframework
import com.intellij.psi.*
import com.intellij.psi.util.PsiTreeUtil
import com.intellij.psi.util.parentOfType
import com.intellij.refactoring.suggested.startOffset
import com.intellij.util.asSafely
import com.siyeh.ig.testFrameworks.AssertHint
import org.jetbrains.uast.UMethod
import org.jetbrains.uast.UParameter
import com.intellij.psi.PsiElement
class JavaTestDiffProvider : JvmTestDiffProvider() {
override fun failedCall(file: PsiFile, startOffset: Int, endOffset: Int, method: UMethod?): PsiElement? {
val failedCalls = findCallsInRange(file, startOffset, endOffset)
if (failedCalls.isEmpty()) return null
if (failedCalls.size == 1) return failedCalls.first()
if (method == null) return null
return failedCalls.firstOrNull { it.resolveMethod()?.isEquivalentTo(method.sourcePsi) == true }
}
private fun findCallsInRange(file: PsiFile, startOffset: Int, endOffset: Int): List<PsiMethodCallExpression> {
val element = file.findElementAt(startOffset)
val codeBlock = PsiTreeUtil.getParentOfType(element, PsiCodeBlock::class.java)
return PsiTreeUtil.findChildrenOfAnyType(codeBlock, false, PsiMethodCallExpression::class.java)
.filter { it.startOffset in startOffset..endOffset }
}
override fun getExpected(call: PsiElement, param: UParameter?): PsiElement? {
if (call !is PsiMethodCallExpression) return null
val expr = if (param == null) {
val assertHint = AssertHint.createAssertEqualsHint(call) ?: return null
if (assertHint.actual.type != PsiType.getJavaLangString(call.manager, call.resolveScope)) return null
if (assertHint.expected.type != PsiType.getJavaLangString(call.manager, call.resolveScope)) return null
assertHint.expected
} else {
val srcParam = param.sourcePsi?.asSafely<PsiParameter>()
val paramList = srcParam?.parentOfType<PsiParameterList>()
val argIndex = paramList?.parameters?.indexOf<PsiElement>(srcParam)
if (argIndex != null && argIndex != -1) call.argumentList.expressions.getOrNull(argIndex) else null
}
if (expr is PsiLiteralExpression) return expr
// disabled for now
//if (expr is PsiPolyadicExpression && expr.operands.all { it is PsiLiteralExpression }) return expr
if (expr is PsiReference) {
val resolved = expr.resolve()
if (resolved is PsiParameter) return resolved
if (resolved is PsiLocalVariable || resolved is PsiField) {
return resolved.asSafely<PsiVariable>()?.initializer.asSafely<PsiLiteralExpression>()
}
return null
}
return null
}
override fun getStringLiteral(expected: PsiElement) = expected
}
@@ -6,43 +6,94 @@ import com.intellij.execution.filters.ExceptionLineParserFactory
import com.intellij.execution.testframework.actions.TestDiffProvider
import com.intellij.openapi.fileEditor.FileDocumentManager
import com.intellij.openapi.project.Project
import com.intellij.openapi.roots.ProjectFileIndex
import com.intellij.psi.PsiElement
import com.intellij.psi.PsiFile
import com.intellij.psi.PsiType
import com.intellij.psi.search.GlobalSearchScope
import com.intellij.refactoring.suggested.startOffset
import com.intellij.util.asSafely
import com.siyeh.ig.testFrameworks.UAssertHint
import org.jetbrains.uast.*
abstract class JvmTestDiffProvider : TestDiffProvider {
final override fun findExpected(project: Project, stackTrace: String): PsiElement? {
var expectedParam: UParameter? = null
var enclosingMethod: UMethod? = null
var (searchStacktrace, expectedParam) = findExpectedEntryPoint(project, stackTrace) ?: return null
val lineParser = ExceptionLineParserFactory.getInstance().create(ExceptionInfoCache(project, GlobalSearchScope.allScope(project)))
stackTrace.lineSequence().forEach { line ->
lineParser.execute(line, line.length) ?: return@forEach
searchStacktrace.lineSequence().forEach { line ->
lineParser.execute(line, line.length) ?: return@findExpected null
val file = lineParser.file ?: return@findExpected null
val diffProvider = TestDiffProvider.TEST_DIFF_PROVIDER_LANGUAGE_EXTENSION.forLanguage(file.language).asSafely<JvmTestDiffProvider>()
if (diffProvider == null) return@findExpected null
if (!ProjectFileIndex.getInstance(project).isInSourceContent(file.virtualFile)) {
if (expectedParam == null) return@forEach // keep looking, test class wasn't found yet
else return@findExpected null // return null when tracking expected param
val diffProvider = TestDiffProvider.TEST_DIFF_PROVIDER_LANGUAGE_EXTENSION
.forLanguage(file.language).asSafely<JvmTestDiffProvider>() ?: return@findExpected null
val containingMethod = expectedParam.getContainingUMethod() ?: return@findExpected null
val failedCall = findFailedCall(file, lineParser.info.lineNumber, expectedParam.getContainingUMethod()) ?: return@findExpected null
val expectedArg = failedCall.getArgumentForParameter(containingMethod.uastParameters.indexOf(expectedParam)) ?: return@findExpected null
if (expectedArg.isInjectionHost()) {
return diffProvider.getStringLiteral(expectedArg.sourcePsi ?: return@forEach)
}
if (expectedArg is UReferenceExpression) {
val resolved = expectedArg.resolveToUElement()
if (resolved is UVariable && resolved.uastInitializer.isInjectionHost()) {
return diffProvider.getStringLiteral(resolved.uastInitializer?.sourcePsi ?: return@forEach)
}
if (resolved is UParameter) {
expectedParam = resolved
return@forEach
}
}
val virtualFile = file.virtualFile ?: return@findExpected null
val document = FileDocumentManager.getInstance().getDocument(virtualFile) ?: return@findExpected null
val lineNumber = lineParser.info.lineNumber
if (lineNumber < 1 || lineNumber > document.lineCount) return@findExpected null
val startOffset = document.getLineStartOffset(lineNumber - 1)
val endOffset = document.getLineEndOffset(lineNumber - 1)
val failedCall = diffProvider.failedCall(file, startOffset, endOffset, enclosingMethod) ?: return@findExpected null
val expected = diffProvider.getExpected(failedCall, expectedParam) ?: return@findExpected null
enclosingMethod = failedCall.toUElement()?.getParentOfType<UMethod>(true)
expectedParam = expected.toUElementOfType<UParameter>()
if (expectedParam == null) return expected
}
return null
}
abstract fun failedCall(file: PsiFile, startOffset: Int, endOffset: Int, method: UMethod?): PsiElement?
abstract fun getStringLiteral(expected: PsiElement): PsiElement?
abstract fun getExpected(call: PsiElement, param: UParameter?): PsiElement?
private data class ExpectedEntryPoint(val stackTrace: String, val param: UParameter)
private fun findExpectedEntryPoint(project: Project, stackTrace: String): ExpectedEntryPoint? {
val lineParser = ExceptionLineParserFactory.getInstance().create(ExceptionInfoCache(project, GlobalSearchScope.allScope(project)))
stackTrace.lineSequence().forEach { line ->
lineParser.execute(line, line.length) ?: return@forEach
val file = lineParser.file ?: return@findExpectedEntryPoint null
val failedCall = findFailedCall(file, lineParser.info.lineNumber, null) ?: return@forEach
val entryParam = findExpectedEntryPointParam(failedCall) ?: return@forEach
return ExpectedEntryPoint(line + stackTrace.substringAfter(line), entryParam)
}
return null
}
private fun findExpectedEntryPointParam(call: UCallExpression): UParameter? {
val assertHint = UAssertHint.createAssertEqualsHint(call) ?: return null
val srcCall = call.sourcePsi ?: return null
val stringType = PsiType.getJavaLangString(srcCall.manager, srcCall.resolveScope)
if (assertHint.expected.getExpressionType() != stringType || assertHint.actual.getExpressionType() != stringType) return null
val method = call.resolveToUElement()?.asSafely<UMethod>() ?: return null
if (method.name != "assertEquals") return null
return method.uastParameters.firstOrNull()
}
private fun findFailedCall(file: PsiFile, lineNumber: Int, resolvedMethod: UMethod?): UCallExpression? {
val virtualFile = file.virtualFile ?: return null
val document = FileDocumentManager.getInstance().getDocument(virtualFile) ?: return null
if (lineNumber < 1 || lineNumber > document.lineCount) return null
val startOffset = document.getLineStartOffset(lineNumber - 1)
val endOffset = document.getLineEndOffset(lineNumber - 1)
val candidateCalls = getCallElementsInRange(file, startOffset, endOffset) ?: return null
return if (candidateCalls.size != 1) {
candidateCalls.firstOrNull { call ->
call.resolveToUElement().asSafely<UMethod>()?.sourcePsi?.isEquivalentTo(resolvedMethod?.sourcePsi) == true
}
} else candidateCalls.first()
}
private fun getCallElementsInRange(file: PsiFile, startOffset: Int, endOffset: Int): List<UCallExpression>? {
val startElement = file.findElementAt(startOffset) ?: return null
val searchStartOffset = startElement.startOffset
val calls = mutableListOf<UCallExpression>()
var curElement: PsiElement? = startElement
while (curElement != null && curElement.startOffset in searchStartOffset..endOffset) {
val callExpression = curElement.toUElement().getUCallExpression(searchLimit = 2)
if (callExpression != null) calls.add(callExpression)
curElement = curElement.nextSibling
}
return calls
}
}
@@ -6,8 +6,6 @@ import org.intellij.lang.annotations.Language
@Suppress("AssertBetweenInconvertibleTypes", "NewClassNamingConvention", "SameParameterValue")
class JavaTestDiffUpdateTest : JvmTestDiffUpdateTest() {
private val fileExt = "java"
@Suppress("SameParameterValue")
private fun checkHasNoDiff(
@Language("Java") before: String,
@@ -97,6 +95,42 @@ class JavaTestDiffUpdateTest : JvmTestDiffUpdateTest() {
""".trimIndent())
}
fun `test accept string literal diff with actual call`() {
checkAcceptFullDiff("""
import org.junit.Assert;
import org.junit.Test;
public class MyJUnitTest {
@Test
public void testFoo() {
Assert.assertEquals("expected", getActual(getActual(getActual("actual"))));
}
private static String getActual(String str) {
return str;
}
}
""".trimIndent(), """
import org.junit.Assert;
import org.junit.Test;
public class MyJUnitTest {
@Test
public void testFoo() {
Assert.assertEquals("actual", getActual(getActual(getActual("actual"))));
}
private static String getActual(String str) {
return str;
}
}
""".trimIndent(), "MyJUnitTest", "testFoo", "expected", "actual", """
at org.junit.Assert.assertEquals(Assert.java:117)
at org.junit.Assert.assertEquals(Assert.java:146)
at MyJUnitTest.testFoo(MyJUnitTest.java:7)
""".trimIndent())
}
fun `test accept string literal diff with carriage return and line feed in expected`() {
checkAcceptFullDiff("""
import org.junit.Assert;
@@ -3,57 +3,10 @@ package org.jetbrains.kotlin.idea.testIntegration
import com.intellij.execution.testframework.JvmTestDiffProvider
import com.intellij.psi.PsiElement
import com.intellij.psi.PsiFile
import com.intellij.psi.PsiType
import com.intellij.psi.util.parentOfType
import com.intellij.util.asSafely
import com.siyeh.ig.testFrameworks.UAssertHint
import org.jetbrains.kotlin.backend.jvm.ir.psiElement
import org.jetbrains.kotlin.idea.caches.resolve.resolveToCall
import org.jetbrains.kotlin.idea.util.findElementsOfClassInRange
import org.jetbrains.kotlin.psi.*
import org.jetbrains.uast.UCallExpression
import org.jetbrains.uast.UMethod
import org.jetbrains.uast.UParameter
import org.jetbrains.uast.toUElementOfType
import org.jetbrains.kotlin.psi.KtStringTemplateEntry
class KotlinTestDiffProvider : JvmTestDiffProvider() {
override fun failedCall(file: PsiFile, startOffset: Int, endOffset: Int, method: UMethod?): PsiElement? {
val failedCalls = findElementsOfClassInRange(file, startOffset, endOffset, KtCallExpression::class.java)
.map { it as KtCallExpression }
if (failedCalls.isEmpty()) return null
if (failedCalls.size == 1) return failedCalls.first()
if (method == null) return null
return failedCalls.firstOrNull { it.resolveToCall()?.resultingDescriptor?.psiElement?.isEquivalentTo(method.sourcePsi) == true }
}
override fun getExpected(call: PsiElement, param: UParameter?): PsiElement? {
if (call !is KtCallExpression) return null
val expr = if (param == null) {
val uCallElement = call.toUElementOfType<UCallExpression>() ?: return null
val assertHint = UAssertHint.createAssertEqualsHint(uCallElement) ?: return null
if (assertHint.expected.getExpressionType() != PsiType.getJavaLangString(call.manager, call.resolveScope)) return null
if (assertHint.actual.getExpressionType() != PsiType.getJavaLangString(call.manager, call.resolveScope)) return null
assertHint.expected.sourcePsi ?: return null
} else {
val argument = call.valueArguments.firstOrNull {it.getArgumentName()?.asName?.asString() == param.name } ?: let {
val srcParam = param.sourcePsi?.asSafely<KtParameter>()
val paramList = srcParam?.parentOfType<KtParameterList>()
val argIndex = paramList?.parameters?.indexOf<PsiElement>(srcParam)
if (argIndex != null && argIndex != -1) call.valueArguments.getOrNull(argIndex) else null
}
argument?.getArgumentExpression()
}
if (expr is KtStringTemplateExpression && expr.entries.size == 1) return expr
if (expr is KtStringTemplateEntry) return expr.parent
if (expr is KtNameReferenceExpression) {
val resolved = expr.reference?.resolve()
if (resolved is KtVariableDeclaration) {
val initializer = resolved.initializer
if (initializer is KtStringTemplateExpression) return initializer
}
return resolved
}
return null
override fun getStringLiteral(expected: PsiElement): PsiElement? {
return if (expected is KtStringTemplateEntry) expected.parent else null
}
}
@@ -68,6 +68,44 @@ class KotlinTestDiffUpdateTest : JvmTestDiffUpdateTest() {
)
}
fun `test accept string literal diff with actual call`() {
checkAcceptFullDiff(
"""
import org.junit.Assert
import org.junit.Test
class MyJUnitTest {
@Test
fun testFoo() {
Assert.assertEquals("expected", getActual(getActual(getActual("actual"))))
}
private fun getActual(str: String): String {
return str
}
}
""".trimIndent(), """
import org.junit.Assert
import org.junit.Test
class MyJUnitTest {
@Test
fun testFoo() {
Assert.assertEquals("actual", getActual(getActual(getActual("actual"))))
}
private fun getActual(str: String): String {
return str
}
}
""".trimIndent(), "MyJUnitTest", "testFoo", "expected", "actual", """
at org.junit.Assert.assertEquals(Assert.java:117)
at org.junit.Assert.assertEquals(Assert.java:146)
at MyJUnitTest.testFoo(MyJUnitTest.kt:7)
""".trimIndent()
)
}
fun `test accept diff is not available when expected is not a string literal`() {
checkHasNoDiff(
"""