Support unwrapping parameters properly in Kotlin KTIJ-11687

- Top level call unwrapping
- Non-nested call unwrapping
- Delete receiver when unwrapping qualified call
- Support multi argument calls
- Allow caret on method name, paren, comma or whitespace
- Prefer parameter unwrapping over removal for top level function calls
- Fix strikethrough preview

closes https://github.com/JetBrains/intellij-community/pull/1996

GitOrigin-RevId: ae278ace872176ba46c8cac3e8df30b6428dbb0b
This commit is contained in:
Christian Vonrüti
2022-07-20 15:25:37 +00:00
committed by intellij-monorepo-bot
parent f73da4c0ea
commit d77149caca
49 changed files with 695 additions and 18 deletions
@@ -19,4 +19,7 @@ fun <K, V> merge(vararg maps: Map<K, V>?): Map<K, V> {
}
}
return result
}
}
@Suppress("UNCHECKED_CAST")
inline fun <reified T> Sequence<*>.takeWhileIsInstance(): Sequence<T> = takeWhile { it is T } as Sequence<T>
@@ -2,26 +2,190 @@
package org.jetbrains.kotlin.idea.codeInsight.unwrap
import com.intellij.psi.PsiComment
import com.intellij.psi.PsiElement
import org.jetbrains.kotlin.psi.KtCallExpression
import org.jetbrains.kotlin.psi.KtValueArgument
import com.intellij.psi.PsiWhiteSpace
import com.intellij.psi.util.elementType
import org.jetbrains.kotlin.idea.base.resources.KotlinBundle
import org.jetbrains.kotlin.idea.refactoring.getExpressionShortText
import org.jetbrains.kotlin.idea.util.isComma
import org.jetbrains.kotlin.lexer.KtTokens
import org.jetbrains.kotlin.psi.*
import org.jetbrains.kotlin.psi.psiUtil.getNextSiblingIgnoringWhitespaceAndComments
import org.jetbrains.kotlin.psi.psiUtil.getPrevSiblingIgnoringWhitespaceAndComments
import org.jetbrains.kotlin.psi.psiUtil.getStrictParentOfType
import org.jetbrains.kotlin.psi.psiUtil.parentsWithSelf
import org.jetbrains.kotlin.util.takeWhileIsInstance
import org.jetbrains.kotlin.utils.addToStdlib.ifTrue
import org.jetbrains.kotlin.utils.addToStdlib.safeAs
class KotlinFunctionParameterUnwrapper(key: String) : KotlinUnwrapRemoveBase(key) {
class KotlinFunctionParameterUnwrapper(val key: String) : KotlinUnwrapRemoveBase(key) {
override fun isApplicableTo(element: PsiElement): Boolean {
if (element !is KtCallExpression) return false
if (element.valueArguments.size != 1) return false
val argument = element.parent as? KtValueArgument ?: return false
if (argument.getStrictParentOfType<KtCallExpression>() == null) return false
val argumentToUnwrap = argumentToUnwrap(element) ?: return false
deletionTarget(argumentToUnwrap) ?: return false
return true
}
private fun argumentToUnwrap(element: PsiElement?): KtValueArgument? {
if (element == null) return null
nearbySingleArg(element)?.let {
return it
}
when (element) {
is KtValueArgument -> return element
else -> {
if (element.parent !is KtValueArgumentList) return null
if (element.isComma || element.elementType == KtTokens.RPAR || element is PsiWhiteSpace || element is PsiComment) {
return findAdjacentValueArgumentInsideParen(element)
} else return null
}
}
}
/**
* When outside a function call's parenthesized arguments but the function has a single argument, choose that
* ```
* ┌ LPAR
* │
* v
* methodCall(arguments)
* ^^^^^^^^^^
* │
* └ Identifier
* ```
*
*/
private fun nearbySingleArg(element: PsiElement): KtValueArgument? {
// a.b.c.d(123)
// ^^^^^^^
// walk up chained dot expressions until we find the one that has a call as rhs
// we use this complicated rather than just have isApplicableTo apply to a KtQualifiedExpression to get priority over
// "Remove" unwrap item
fun caretOnQualifier(qualifiedExpression: KtQualifiedExpression): KtCallExpression? =
qualifiedExpression.parentsWithSelf.takeWhileIsInstance<KtQualifiedExpression>().last().selectorExpression as? KtCallExpression
val callExpression: KtCallExpression = when (element.elementType) {
KtTokens.LPAR -> {
element.parent?.safeAs<KtValueArgumentList>()
?.parent?.safeAs<KtCallExpression>() ?: return null
}
KtTokens.DOT, KtTokens.SAFE_ACCESS -> {
element.parent?.safeAs<KtQualifiedExpression>()
?.let { caretOnQualifier(it) } ?: return null
}
KtTokens.IDENTIFIER -> {
val referenceExpression = element.parent?.safeAs<KtReferenceExpression>()
when (val parent = referenceExpression?.parent) {
is KtCallExpression -> parent
is KtQualifiedExpression -> {
caretOnQualifier(parent)
}
else -> null
} ?: return null
// we could just handle this on
}
else -> return null
}
val ktValueArgumentList = callExpression.valueArgumentList
if (ktValueArgumentList?.arguments.isNullOrEmpty()) {
// if there's no parenthesized arguments, unwrap trailing lambda if any
return callExpression.lambdaArguments.singleOrNull()
}
// there are parenthesized arguments, only consider if there's exactly one
return ktValueArgumentList?.arguments?.singleOrNull()
}
/**
* When inside a function call's parenthesized arguments.
*
* * for Comma and RPAR, we scan towards the left
* * for whitespace, we scan leftwards and rightwards
*
* ```
* ┌ COMMA
* │ ┌ RPAR
* │ │
* v v
* methodCall(1, 2 , 3 )
* ^ ^ ^
* │ │ │
* Whitespace ┴───┴─┘
* ```
*
*/
private fun findAdjacentValueArgumentInsideParen(element: PsiElement): KtValueArgument? {
if (element.parent !is KtValueArgumentList) {
return null
}
return when {
element.elementType == KtTokens.RPAR -> {
// caret after last argument before closing paren, argument is the last element
val valueArgumentList = element.parent as? KtValueArgumentList ?: return null
if (valueArgumentList.parent !is KtCallExpression) return null
return valueArgumentList.arguments.lastOrNull()
}
element.isComma -> {
// caret before comma, choose previous argument
val previous = element.getPrevSiblingIgnoringWhitespaceAndComments()
isCallArgument(previous).ifTrue { previous as KtValueArgument }
}
element is PsiWhiteSpace || element is PsiComment -> {
// caret before blank or in comment, argument could be towards the left or the right
val previous = element.getPrevSiblingIgnoringWhitespaceAndComments()
if (isCallArgument(previous)) return previous as KtValueArgument
val next = element.getNextSiblingIgnoringWhitespaceAndComments()
if (isCallArgument(next)) return next as KtValueArgument
null
}
else -> null
}
}
private fun isCallArgument(element: PsiElement?): Boolean {
if (element is KtLambdaArgument) return false
if (element !is KtValueArgument) return false
val argumentList = element.parent as KtValueArgumentList
if (argumentList.parent !is KtCallExpression) return false
return true
}
override fun doUnwrap(element: PsiElement?, context: Context?) {
val function = element as? KtCallExpression ?: return
val argument = element.valueArguments.firstOrNull()?.getArgumentExpression() ?: return
context?.extractFromExpression(argument, function)
context?.delete(function)
val valueArgument = argumentToUnwrap(element) ?: return
val deletionTarget = deletionTarget(valueArgument) ?: return
val argument = valueArgument.getArgumentExpression() ?: return
context?.extractFromExpression(argument, deletionTarget)
context?.delete(deletionTarget)
}
private fun deletionTarget(valueArgument: KtValueArgument): KtElement? {
var function: KtElement = valueArgument.getStrictParentOfType<KtCallExpression>() ?: return null
val parent = function.parent
if (parent is KtQualifiedExpression) {
function = parent
}
return function
}
override fun collectAffectedElements(e: PsiElement, toExtract: MutableList<PsiElement>): PsiElement {
super.collectAffectedElements(e, toExtract)
return argumentToUnwrap(e)?.let { deletionTarget(it) } ?: e
}
override fun getDescription(e: PsiElement): String {
// necessary because the base implementation expects to be called for KtElement
// but because we support the caret to be on LPAR/COMMA and other tokens this doesn't work
val target = argumentToUnwrap(e) ?: error("Description asked for a non applicable unwrapper")
return KotlinBundle.message(key, getExpressionShortText(target))
}
}
@@ -303,19 +303,129 @@ public abstract class UnwrapRemoveTestGenerated extends AbstractUnwrapRemoveTest
KotlinTestUtils.runTest(this::doTestFunctionParameterUnwrapper, this, testDataFilePath);
}
@TestMetadata("caretBeforeComma.kt")
public void testCaretBeforeComma() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/caretBeforeComma.kt");
}
@TestMetadata("caretInComment1.kt")
public void testCaretInComment1() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/caretInComment1.kt");
}
@TestMetadata("caretInComment2.kt")
public void testCaretInComment2() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/caretInComment2.kt");
}
@TestMetadata("caretInLambda.kt")
public void testCaretInLambda() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/caretInLambda.kt");
}
@TestMetadata("caretInLambda2.kt")
public void testCaretInLambda2() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/caretInLambda2.kt");
}
@TestMetadata("caretInLambda3.kt")
public void testCaretInLambda3() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/caretInLambda3.kt");
}
@TestMetadata("caretInQualifier1.kt")
public void testCaretInQualifier1() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/caretInQualifier1.kt");
}
@TestMetadata("caretInQualifier2.kt")
public void testCaretInQualifier2() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/caretInQualifier2.kt");
}
@TestMetadata("caretInQualifier3.kt")
public void testCaretInQualifier3() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/caretInQualifier3.kt");
}
@TestMetadata("caretInQualifier4.kt")
public void testCaretInQualifier4() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/caretInQualifier4.kt");
}
@TestMetadata("caretInQualifier5.kt")
public void testCaretInQualifier5() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/caretInQualifier5.kt");
}
@TestMetadata("caretInQualifier6.kt")
public void testCaretInQualifier6() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/caretInQualifier6.kt");
}
@TestMetadata("caretInQualifier7.kt")
public void testCaretInQualifier7() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/caretInQualifier7.kt");
}
@TestMetadata("functionHasMultiParam.kt")
public void testFunctionHasMultiParam() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/functionHasMultiParam.kt");
}
@TestMetadata("functionHasMultiParam1.kt")
public void testFunctionHasMultiParam1() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/functionHasMultiParam1.kt");
}
@TestMetadata("functionHasMultiParam2.kt")
public void testFunctionHasMultiParam2() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/functionHasMultiParam2.kt");
}
@TestMetadata("functionHasSingleParam.kt")
public void testFunctionHasSingleParam() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/functionHasSingleParam.kt");
}
@TestMetadata("functionWithReceiver.kt")
public void testFunctionWithReceiver() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/functionWithReceiver.kt");
@TestMetadata("nestedFunctionWithReceiver.kt")
public void testNestedFunctionWithReceiver() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/nestedFunctionWithReceiver.kt");
}
@TestMetadata("nestedSingleParamCall.kt")
public void testNestedSingleParamCall() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/nestedSingleParamCall.kt");
}
@TestMetadata("nestedSingleParamCall.option1.kt")
public void testNestedSingleParamCall_option1() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/nestedSingleParamCall.option1.kt");
}
@TestMetadata("topLevelMultiParamCall_beforeIdentifier.kt")
public void testTopLevelMultiParamCall_beforeIdentifier() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/topLevelMultiParamCall_beforeIdentifier.kt");
}
@TestMetadata("topLevelMultiParamCall_beforeLeftParen.kt")
public void testTopLevelMultiParamCall_beforeLeftParen() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/topLevelMultiParamCall_beforeLeftParen.kt");
}
@TestMetadata("topLevelMultiParamCall_beforeRightParen.kt")
public void testTopLevelMultiParamCall_beforeRightParen() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/topLevelMultiParamCall_beforeRightParen.kt");
}
@TestMetadata("topLevelMultiParamCall_onIdentifier.kt")
public void testTopLevelMultiParamCall_onIdentifier() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/topLevelMultiParamCall_onIdentifier.kt");
}
@TestMetadata("topLevelSingleParamCall_onIdentifier.kt")
public void testTopLevelSingleParamCall_onIdentifier() throws Exception {
runTest("testData/codeInsight/unwrapAndRemove/unwrapFunctionParameter/topLevelSingleParamCall_onIdentifier.kt");
}
}
}
@@ -0,0 +1,7 @@
// OPTION: 0
fun test() {
val i = 1
foo(1 <caret>, 2)
}
fun foo(i: Int, i2: Int) {}
@@ -0,0 +1,7 @@
// OPTION: 0
fun test() {
val i = 1
1<caret>
}
fun foo(i: Int, i2: Int) {}
@@ -0,0 +1,7 @@
// OPTION: 0
fun test() {
val i = 1
foo(/* <caret> */ 1, 2)
}
fun foo(i: Int, i2: Int) {}
@@ -0,0 +1,7 @@
// OPTION: 0
fun test() {
val i = 1
1
}
fun foo(i: Int, i2: Int) {}
@@ -0,0 +1,7 @@
// OPTION: 0
fun test() {
val i = 1
foo(1 /* <caret> */, 2)
}
fun foo(i: Int, i2: Int) {}
@@ -0,0 +1,7 @@
// OPTION: 0
fun test() {
val i = 1
1
}
fun foo(i: Int, i2: Int) {}
@@ -0,0 +1,8 @@
// OPTION: 2
fun test() {
foo {
<caret>val x = 1
}
}
fun foo(body: () -> Unit) {}
@@ -0,0 +1,8 @@
// OPTION: 2
fun test() {
{
val x = 1
}<caret>
}
fun foo(body: () -> Unit) {}
@@ -0,0 +1,8 @@
// OPTION: 2
fun test() {
foo(1) {
<caret>val x = 1
}
}
fun foo(a: Int, body: () -> Unit) {}
@@ -0,0 +1,8 @@
// OPTION: 2
fun test() {
{
val x = 1
}<caret>
}
fun foo(a: Int, body: () -> Unit) {}
@@ -0,0 +1,8 @@
// OPTION: 2
fun test() {
foo(1, body = {
<caret>val x = 1
})
}
fun foo(a: Int, body: () -> Unit) {}
@@ -0,0 +1,8 @@
// OPTION: 2
fun test() {
{
val x = 1
}<caret>
}
fun foo(a: Int, body: () -> Unit) {}
@@ -0,0 +1,10 @@
// OPTION: 0
fun test() {
val i = 1
val test = Test()
te<caret>st.qux(i)
}
class Test {
fun qux(i: Int) = 1
}
@@ -0,0 +1,10 @@
// OPTION: 0
fun test() {
val i = 1
val test = Test()
i
}
class Test {
fun qux(i: Int) = 1
}
@@ -0,0 +1,12 @@
// OPTION: 0
fun test() {
val i = 1
val test = Test()
test.te<caret>st.qux(i)
}
class Test {
val test: Test
get() = Test()
fun qux(i: Int) = 1
}
@@ -0,0 +1,12 @@
// OPTION: 0
fun test() {
val i = 1
val test = Test()
i
}
class Test {
val test: Test
get() = Test()
fun qux(i: Int) = 1
}
@@ -0,0 +1,12 @@
// OPTION: 0
fun test() {
val i = 1
val test = Test()
te<caret>st.test.qux(i)
}
class Test {
val test: Test
get() = Test()
fun qux(i: Int) = 1
}
@@ -0,0 +1,12 @@
// OPTION: 0
fun test() {
val i = 1
val test = Test()
i
}
class Test {
val test: Test
get() = Test()
fun qux(i: Int) = 1
}
@@ -0,0 +1,12 @@
// OPTION: 0
fun test() {
val i = 1
val test = Test()
test<caret>.test.qux(i)
}
class Test {
val test: Test
get() = Test()
fun qux(i: Int) = 1
}
@@ -0,0 +1,12 @@
// OPTION: 0
fun test() {
val i = 1
val test = Test()
i
}
class Test {
val test: Test
get() = Test()
fun qux(i: Int) = 1
}
@@ -0,0 +1,12 @@
// OPTION: 0
fun test() {
val i = 1
val test = Test()
test?<caret>.test.qux(i)
}
class Test {
val test: Test
get() = Test()
fun qux(i: Int) = 1
}
@@ -0,0 +1,12 @@
// OPTION: 0
fun test() {
val i = 1
val test = Test()
i
}
class Test {
val test: Test
get() = Test()
fun qux(i: Int) = 1
}
@@ -0,0 +1,12 @@
// OPTION: 0
fun test() {
val i = 1
val test = Test()
test<caret>.test.test.qux(i)
}
class Test {
val test: Test
get() = Test()
fun qux(i: Int) = 1
}
@@ -0,0 +1,12 @@
// OPTION: 0
fun test() {
val i = 1
val test = Test()
i
}
class Test {
val test: Test
get() = Test()
fun qux(i: Int) = 1
}
@@ -0,0 +1,11 @@
// OPTION: 0
fun test() {
val test = Test()
test<caret>.test.test.qux { val i = 1 }
}
class Test {
val test: Test
get() = Test()
fun qux(body: () -> Unit) = 1
}
@@ -0,0 +1,11 @@
// OPTION: 0
fun test() {
val test = Test()
{ val i = 1 }
}
class Test {
val test: Test
get() = Test()
fun qux(body: () -> Unit) = 1
}
@@ -1,4 +1,4 @@
// IS_APPLICABLE: false
// OPTION: 0
fun test() {
val i = 1
foo(baz<caret>(i, 2))
@@ -0,0 +1,9 @@
// OPTION: 0
fun test() {
val i = 1
baz(i, 2)<caret>
}
fun foo(i: Int) {}
fun baz(i: Int, j: Int) = 1
@@ -0,0 +1,9 @@
// OPTION: 0
fun test() {
val i = 1
foo(baz(<caret>i, 2))
}
fun foo(i: Int) {}
fun baz(i: Int, j: Int) = 1
@@ -0,0 +1,9 @@
// OPTION: 0
fun test() {
val i = 1
foo(i<caret>)
}
fun foo(i: Int) {}
fun baz(i: Int, j: Int) = 1
@@ -0,0 +1,9 @@
// OPTION: 0
fun test() {
val i = 1
foo(123, 456<caret>)
}
fun foo(i: Int) {}
fun baz(i: Int, j: Int) = 1
@@ -0,0 +1,9 @@
// OPTION: 0
fun test() {
val i = 1
456<caret>
}
fun foo(i: Int) {}
fun baz(i: Int, j: Int) = 1
@@ -1,7 +1,7 @@
// OPTION: 0
fun test() {
val i = 1
foo(bar<caret>(1))
foo(bar(<caret>1))
}
fun foo(i: Int) {}
@@ -0,0 +1,12 @@
// OPTION: 0
fun test() {
val i = 1
val test = Test()
foo(i<caret>)
}
fun foo(i: Int) {}
class Test {
fun qux(i: Int) = 1
}
@@ -0,0 +1,9 @@
// OPTION: 0
fun test() {
val i = 1
foo(bar<caret>(1))
}
fun foo(i: Int) {}
fun bar(i: Int) = 1
@@ -0,0 +1,9 @@
// OPTION: 0
fun test() {
val i = 1
foo(1<caret>)
}
fun foo(i: Int) {}
fun bar(i: Int) = 1
@@ -0,0 +1,9 @@
// OPTION: 1
fun test() {
val i = 1
foo(bar<caret>(1))
}
fun foo(i: Int) {}
fun bar(i: Int) = 1
@@ -0,0 +1,9 @@
// OPTION: 1
fun test() {
val i = 1
bar(1)<caret>
}
fun foo(i: Int) {}
fun bar(i: Int) = 1
@@ -0,0 +1,8 @@
// IS_APPLICABLE: false
fun test() {
val i = 1
<caret>foo(123, 456)
}
fun foo(i: Int) {}
@@ -0,0 +1,8 @@
// IS_APPLICABLE: false
fun test() {
val i = 1
foo<caret>(123, 456)
}
fun foo(i: Int, j: Int) {}
@@ -0,0 +1,8 @@
// OPTION: 0
fun test() {
val i = 1
foo(123, 456<caret>)
}
fun foo(i: Int, j: Int) {}
@@ -0,0 +1,8 @@
// OPTION: 0
fun test() {
val i = 1
456<caret>
}
fun foo(i: Int, j: Int) {}
@@ -0,0 +1,9 @@
// IS_APPLICABLE: false
fun test() {
val i = 1
fo<caret>o(123, 456)
}
fun foo(i: Int) {}
fun baz(i: Int, j: Int) = 1
@@ -0,0 +1,7 @@
// OPTION: 0
fun test() {
val i = 1
fo<caret>o(123)
}
fun foo(i: Int) {}
@@ -0,0 +1,7 @@
// OPTION: 0
fun test() {
val i = 1
123<caret>
}
fun foo(i: Int) {}