JetPsiFactory: refactored method for creating argument + introduced ability to create arbitrary constructs by pattern

This commit is contained in:
Valentin Kipyatkov
2015-05-22 16:22:51 +03:00
parent 161630a449
commit 6b66e3b0e6
6 changed files with 61 additions and 39 deletions
@@ -23,13 +23,14 @@ import com.intellij.openapi.util.Key
import com.intellij.psi.PsiComment import com.intellij.psi.PsiComment
import com.intellij.psi.PsiElement import com.intellij.psi.PsiElement
import com.intellij.psi.PsiFileFactory import com.intellij.psi.PsiFileFactory
import com.intellij.psi.codeStyle.CodeStyleManager
import com.intellij.psi.util.PsiTreeUtil import com.intellij.psi.util.PsiTreeUtil
import com.intellij.util.LocalTimeCounter import com.intellij.util.LocalTimeCounter
import org.jetbrains.kotlin.analyzer.ModuleInfo import org.jetbrains.kotlin.analyzer.ModuleInfo
import org.jetbrains.kotlin.idea.JetFileType import org.jetbrains.kotlin.idea.JetFileType
import org.jetbrains.kotlin.lexer.JetKeywordToken import org.jetbrains.kotlin.lexer.JetKeywordToken
import org.jetbrains.kotlin.name.FqName import org.jetbrains.kotlin.name.FqName
import org.jetbrains.kotlin.name.Name
import org.jetbrains.kotlin.name.renderName
import org.jetbrains.kotlin.psi.JetPsiFactory.CallableBuilder.Target import org.jetbrains.kotlin.psi.JetPsiFactory.CallableBuilder.Target
import org.jetbrains.kotlin.resolve.ImportPath import org.jetbrains.kotlin.resolve.ImportPath
import java.io.PrintWriter import java.io.PrintWriter
@@ -187,12 +188,12 @@ public class JetPsiFactory(private val project: Project) {
return createDeclaration(text) return createDeclaration(text)
} }
public fun <T> createDeclaration(text: String): T { public fun <TDeclaration : JetDeclaration> createDeclaration(text: String): TDeclaration {
val file = createFile(text) val file = createFile(text)
val dcls = file.getDeclarations() val declarations = file.getDeclarations()
assert(dcls.size() == 1) { "${dcls.size()} declarations in $text" } assert(declarations.size() == 1) { "${declarations.size()} declarations in $text" }
[suppress("UNCHECKED_CAST")] @suppress("UNCHECKED_CAST")
val result = dcls.first() as T val result = declarations.first() as TDeclaration
return result return result
} }
@@ -361,13 +362,24 @@ public class JetPsiFactory(private val project: Project) {
createExpressionByPattern("if ($0) $1", condition, thenExpr)) as JetIfExpression createExpressionByPattern("if ($0) $1", condition, thenExpr)) as JetIfExpression
} }
public fun createArgumentWithName(name: String?, argumentExpression: JetExpression): JetValueArgument { public fun createArgument(expression: JetExpression, name: String? = null, isSpread: Boolean = false): JetValueArgument {
val argumentText = (if (name != null) "$name = " else "") + argumentExpression.getText() val argumentList = buildByPattern({ pattern, args -> createByPattern(pattern, *args) { createCallArguments(it) } }) {
return createCallArguments("($argumentText)").getArguments().first() appendFixedText("(")
}
public fun createArgument(argumentExpression: JetExpression): JetValueArgument { if (name != null) {
return createArgumentWithName(null, argumentExpression) appendFixedText(Name.identifier(name).renderName())
appendFixedText("=")
}
if (isSpread) {
appendFixedText("*")
}
appendExpression(expression)
appendFixedText(")")
}
return argumentList.getArguments().single()
} }
public fun createDelegatorToSuperCall(text: String): JetDelegatorToSuperCall { public fun createDelegatorToSuperCall(text: String): JetDelegatorToSuperCall {
@@ -30,13 +30,19 @@ import java.util.ArrayList
import java.util.HashMap import java.util.HashMap
import java.util.LinkedHashMap import java.util.LinkedHashMap
public fun JetPsiFactory.createExpressionByPattern(pattern: String, vararg args: Any): JetExpression { public fun JetPsiFactory.createExpressionByPattern(pattern: String, vararg args: Any): JetExpression
= createByPattern(pattern, *args) { createExpression(it) }
public fun <TDeclaration : JetDeclaration> JetPsiFactory.createDeclarationByPattern(pattern: String, vararg args: Any): TDeclaration
= createByPattern(pattern, *args) { createDeclaration<TDeclaration>(it) }
public fun <TElement : JetElement> createByPattern(pattern: String, vararg args: Any, factory: (String) -> TElement): TElement {
val (processedText, allPlaceholders) = processPattern(pattern, args) val (processedText, allPlaceholders) = processPattern(pattern, args)
var expression = createExpression(processedText.trim()) var resultElement = factory(processedText.trim())
val project = expression.getProject() val project = resultElement.getProject()
val start = expression.startOffset val start = resultElement.startOffset
val pointerManager = SmartPointerManager.getInstance(project) val pointerManager = SmartPointerManager.getInstance(project)
@@ -52,7 +58,7 @@ public fun JetPsiFactory.createExpressionByPattern(pattern: String, vararg args:
} }
for ((range, text) in placeholders) { for ((range, text) in placeholders) {
val token = expression.findElementAt(range.getStartOffset())!! val token = resultElement.findElementAt(range.getStartOffset())!!
for (element in token.parents()) { for (element in token.parents()) {
val elementRange = element.getTextRange().shiftRight(-start) val elementRange = element.getTextRange().shiftRight(-start)
if (elementRange == range && expectedElementType.isInstance(element)) { if (elementRange == range && expectedElementType.isInstance(element)) {
@@ -79,20 +85,20 @@ public fun JetPsiFactory.createExpressionByPattern(pattern: String, vararg args:
// reformat whole text except for String arguments (as they can contain user's formatting to be preserved) // reformat whole text except for String arguments (as they can contain user's formatting to be preserved)
if (stringPlaceholderRanges.none()) { if (stringPlaceholderRanges.none()) {
expression = codeStyleManager.reformat(expression, true) as JetExpression resultElement = codeStyleManager.reformat(resultElement, true) as TElement
} }
else { else {
var bound = expression.endOffset - 1 var bound = resultElement.endOffset - 1
for (range in stringPlaceholderRanges) { for (range in stringPlaceholderRanges) {
// we extend reformatting range by 1 to the right because otherwise some of spaces are not reformatted // we extend reformatting range by 1 to the right because otherwise some of spaces are not reformatted
expression = codeStyleManager.reformatRange(expression, range.getEndOffset() + start, bound + 1, true) as JetExpression resultElement = codeStyleManager.reformatRange(resultElement, range.getEndOffset() + start, bound + 1, true) as TElement
bound = range.getStartOffset() + start bound = range.getStartOffset() + start
} }
expression = codeStyleManager.reformatRange(expression, start, bound + 1, true) as JetExpression resultElement = codeStyleManager.reformatRange(resultElement, start, bound + 1, true) as TElement
} }
// do not reformat the whole expression in PostprocessReformattingAspect // do not reformat the whole expression in PostprocessReformattingAspect
CodeEditUtil.setNodeGeneratedRecursively(expression.getNode(), false) CodeEditUtil.setNodeGeneratedRecursively(resultElement.getNode(), false)
for ((pointer, n) in pointers) { for ((pointer, n) in pointers) {
var element = pointer.getElement()!! var element = pointer.getElement()!!
@@ -102,9 +108,9 @@ public fun JetPsiFactory.createExpressionByPattern(pattern: String, vararg args:
element.replace(args[n] as PsiElement) element.replace(args[n] as PsiElement)
} }
codeStyleManager.adjustLineIndent(expression.getContainingFile(), expression.getTextRange()) codeStyleManager.adjustLineIndent(resultElement.getContainingFile(), resultElement.getTextRange())
return expression return resultElement
} }
private data class Placeholder(val range: TextRange, val text: String) private data class Placeholder(val range: TextRange, val text: String)
@@ -182,22 +188,22 @@ private fun processPattern(pattern: String, args: Array<out Any>): PatternData {
return PatternData(text, ranges) return PatternData(text, ranges)
} }
public class ExpressionBuilder { public class BuilderByPattern<TElement> {
private val patternBuilder = StringBuilder() private val patternBuilder = StringBuilder()
private val arguments = ArrayList<Any>() private val arguments = ArrayList<Any>()
public fun appendFixedText(text: String): ExpressionBuilder { public fun appendFixedText(text: String): BuilderByPattern<TElement> {
patternBuilder.append(text) patternBuilder.append(text)
return this return this
} }
public fun appendNonFormattedText(text: String): ExpressionBuilder { public fun appendNonFormattedText(text: String): BuilderByPattern<TElement> {
patternBuilder.append("$" + arguments.size()) patternBuilder.append("$" + arguments.size())
arguments.add(text) arguments.add(text)
return this return this
} }
public fun appendExpression(expression: JetExpression?): ExpressionBuilder { public fun appendExpression(expression: JetExpression?): BuilderByPattern<TElement> {
if (expression != null) { if (expression != null) {
patternBuilder.append("$" + arguments.size()) patternBuilder.append("$" + arguments.size())
arguments.add(expression) arguments.add(expression)
@@ -205,7 +211,7 @@ public class ExpressionBuilder {
return this return this
} }
public fun appendTypeReference(typeRef: JetTypeReference?): ExpressionBuilder { public fun appendTypeReference(typeRef: JetTypeReference?): BuilderByPattern<TElement> {
if (typeRef != null) { if (typeRef != null) {
patternBuilder.append("$" + arguments.size()) patternBuilder.append("$" + arguments.size())
arguments.add(typeRef) arguments.add(typeRef)
@@ -213,13 +219,17 @@ public class ExpressionBuilder {
return this return this
} }
public fun createExpression(factory: JetPsiFactory): JetExpression { public fun create(factory: (String, Array<out Any>) -> TElement): TElement {
return factory.createExpressionByPattern(patternBuilder.toString(), *arguments.toArray()) return factory(patternBuilder.toString(), arguments.toArray())
} }
} }
public fun JetPsiFactory.buildExpression(build: ExpressionBuilder.() -> Unit): JetExpression { public fun JetPsiFactory.buildExpression(build: BuilderByPattern<JetExpression>.() -> Unit): JetExpression {
val builder = ExpressionBuilder() return buildByPattern({ pattern, args -> this.createExpressionByPattern(pattern, *args) }, build)
builder.build() }
return builder.createExpression(this)
public fun <TElement> buildByPattern(factory: (String, Array<out Any>) -> TElement, build: BuilderByPattern<TElement>.() -> Unit): TElement {
val builder = BuilderByPattern<TElement>()
builder.build()
return builder.create(factory)
} }
@@ -147,7 +147,7 @@ public class AddNameToArgumentFix extends JetIntentionAction<JetValueArgument> {
private static JetValueArgument getParsedArgumentWithName(@NotNull String name, @NotNull JetValueArgument argument) { private static JetValueArgument getParsedArgumentWithName(@NotNull String name, @NotNull JetValueArgument argument) {
JetExpression argumentExpression = argument.getArgumentExpression(); JetExpression argumentExpression = argument.getArgumentExpression();
assert argumentExpression != null : "Argument should be already parsed."; assert argumentExpression != null : "Argument should be already parsed.";
return JetPsiFactory(argument).createArgumentWithName(name, argumentExpression); return JetPsiFactory(argument).createArgument(argumentExpression, name, false);
} }
@NotNull @NotNull
@@ -369,7 +369,7 @@ public class JetFunctionCallUsage extends JetUsageInfo<JetCallElement> {
changeArgumentName(argumentNameExpression, parameterInfo); changeArgumentName(argumentNameExpression, parameterInfo);
//noinspection ConstantConditions //noinspection ConstantConditions
newArgument.replace(oldArgument instanceof JetFunctionLiteralArgument newArgument.replace(oldArgument instanceof JetFunctionLiteralArgument
? psiFactory.createArgument(oldArgument.getArgumentExpression()) ? psiFactory.createArgument(oldArgument.getArgumentExpression(), null, false)
: oldArgument.asElement()); : oldArgument.asElement());
} }
// TODO: process default arguments in the middle // TODO: process default arguments in the middle
@@ -48,7 +48,7 @@ fun JetFunctionLiteralArgument.moveInsideParenthesesAndReplaceWith(
val psiFactory = JetPsiFactory(getProject()) val psiFactory = JetPsiFactory(getProject())
val argument = if (newCallExpression.getValueArgumentsInParentheses().any { it.getArgumentName() != null }) { val argument = if (newCallExpression.getValueArgumentsInParentheses().any { it.getArgumentName() != null }) {
psiFactory.createArgumentWithName(functionLiteralArgumentName, replacement) psiFactory.createArgument(replacement, functionLiteralArgumentName)
} }
else { else {
psiFactory.createArgument(replacement) psiFactory.createArgument(replacement)
@@ -1,6 +1,6 @@
// IS_APPLICABLE: true // IS_APPLICABLE: true
fun foo() { fun foo() {
bar(2, { it * 3 }) bar(2, {it * 3})
} }
fun bar(a: Int, b: (Int) -> Int) { fun bar(a: Int, b: (Int) -> Int) {