JetCodeFragment: increment modification count after inserting import

This commit is contained in:
Natalia Ukhorskaya
2015-03-05 16:36:43 +03:00
parent c4f7bf6815
commit f6162dc726
11 changed files with 46 additions and 33 deletions
@@ -24,8 +24,9 @@ public class JetBlockCodeFragment(
project: Project, project: Project,
name: String, name: String,
text: CharSequence, text: CharSequence,
imports: String?,
context: PsiElement? context: PsiElement?
) : JetCodeFragment(project, name, text, JetNodeTypes.BLOCK_CODE_FRAGMENT, context) { ) : JetCodeFragment(project, name, text, imports, JetNodeTypes.BLOCK_CODE_FRAGMENT, context) {
override fun getContentElement() = findChildByClass(javaClass<JetBlockExpression>()) override fun getContentElement() = findChildByClass(javaClass<JetBlockExpression>())
?: throw IllegalStateException("Block expression should be parsed for BlockCodeFragment") ?: throw IllegalStateException("Block expression should be parsed for BlockCodeFragment")
@@ -26,24 +26,28 @@ import com.intellij.testFramework.LightVirtualFile
import org.jetbrains.kotlin.idea.JetFileType import org.jetbrains.kotlin.idea.JetFileType
import java.util.HashSet import java.util.HashSet
import com.intellij.openapi.util.Key import com.intellij.openapi.util.Key
import com.intellij.psi.impl.PsiModificationTrackerImpl
import com.intellij.psi.util.PsiModificationTracker
import org.jetbrains.kotlin.types.JetType import org.jetbrains.kotlin.types.JetType
import java.util.LinkedHashSet
public abstract class JetCodeFragment( public abstract class JetCodeFragment(
private val _project: Project, private val _project: Project,
name: String, name: String,
text: CharSequence, text: CharSequence,
imports: String?, // Should be separated by JetCodeFragment.IMPORT_SEPARATOR
elementType: IElementType, elementType: IElementType,
private val context: PsiElement? private val context: PsiElement?
): JetFile((PsiManager.getInstance(_project) as PsiManagerEx).getFileManager().createFileViewProvider(LightVirtualFile(name, JetFileType.INSTANCE, text), true), false), JavaCodeFragment { ): JetFile((PsiManager.getInstance(_project) as PsiManagerEx).getFileManager().createFileViewProvider(LightVirtualFile(name, JetFileType.INSTANCE, text), true), false), JavaCodeFragment {
private var viewProvider = super<JetFile>.getViewProvider() as SingleRootFileViewProvider private var viewProvider = super<JetFile>.getViewProvider() as SingleRootFileViewProvider
private var myImports = HashSet<String>(); private var myImports = LinkedHashSet<String>();
{ {
getViewProvider().forceCachedPsi(this) getViewProvider().forceCachedPsi(this)
init(TokenType.CODE_FRAGMENT, elementType) init(TokenType.CODE_FRAGMENT, elementType)
if (context != null) { if (context != null) {
addImportsFromString(getImportsForElement(context)) initImports(context, imports)
} }
} }
@@ -98,6 +102,11 @@ public abstract class JetCodeFragment(
override fun addImportsFromString(imports: String?) { override fun addImportsFromString(imports: String?) {
if (imports == null || imports.isEmpty()) return if (imports == null || imports.isEmpty()) return
// We should increment modification tracker after inserting import in code fragment to invalidate resolve caches.
// Without this modification references with new import won't be resolved without any modification in code fragment.
// Also shorten references won't work.
(PsiModificationTracker.SERVICE.getInstance(getProject()) as PsiModificationTrackerImpl).incOutOfCodeBlockModificationCounter()
myImports.addAll(imports.split(IMPORT_SEPARATOR)) myImports.addAll(imports.split(IMPORT_SEPARATOR))
} }
@@ -119,17 +128,27 @@ public abstract class JetCodeFragment(
return true return true
} }
private fun initImports(context: PsiElement, imports: String?) {
val containingFile = context.getContainingFile()
if (containingFile !is JetFile) return
val importListForContextElement = containingFile.getImportList()
if (importListForContextElement != null) {
myImports.addAll(importListForContextElement.getImports().map { it.getText() })
}
val packageName = containingFile.getPackageDirective()?.getFqName()?.asString()
if (packageName != null && packageName.isNotEmpty()) {
myImports.add("import $packageName.*")
}
if (imports != null && !imports.isEmpty()) {
myImports.addAll(imports.split(IMPORT_SEPARATOR))
}
}
class object { class object {
public val IMPORT_SEPARATOR: String = "," public val IMPORT_SEPARATOR: String = ","
public val RUNTIME_TYPE_EVALUATOR: Key<Function1<JetExpression, JetType?>> = Key.create("RUNTIME_TYPE_EVALUATOR") public val RUNTIME_TYPE_EVALUATOR: Key<Function1<JetExpression, JetType?>> = Key.create("RUNTIME_TYPE_EVALUATOR")
public fun getImportsForElement(elementAtCaret: PsiElement): String {
val containingFile = elementAtCaret.getContainingFile()
if (containingFile !is JetFile) return ""
return containingFile.getImportList()?.getImports()
?.map { it.getText() }
?.join(JetCodeFragment.IMPORT_SEPARATOR) ?: ""
}
} }
} }
@@ -24,8 +24,9 @@ public class JetExpressionCodeFragment(
project: Project, project: Project,
name: String, name: String,
text: CharSequence, text: CharSequence,
imports: String?,
context: PsiElement? context: PsiElement?
) : JetCodeFragment(project, name, text, JetNodeTypes.EXPRESSION_CODE_FRAGMENT, context) { ) : JetCodeFragment(project, name, text, imports, JetNodeTypes.EXPRESSION_CODE_FRAGMENT, context) {
override fun getContentElement() = findChildByClass(javaClass<JetExpression>()) override fun getContentElement() = findChildByClass(javaClass<JetExpression>())
} }
@@ -323,11 +323,11 @@ public class JetPsiFactory(private val project: Project) {
} }
public fun createExpressionCodeFragment(text: String, context: PsiElement?): JetExpressionCodeFragment { public fun createExpressionCodeFragment(text: String, context: PsiElement?): JetExpressionCodeFragment {
return JetExpressionCodeFragment(project, "fragment.kt", text, context) return JetExpressionCodeFragment(project, "fragment.kt", text, null, context)
} }
public fun createBlockCodeFragment(text: String, context: PsiElement?): JetBlockCodeFragment { public fun createBlockCodeFragment(text: String, context: PsiElement?): JetBlockCodeFragment {
return JetBlockCodeFragment(project, "fragment.kt", text, context) return JetBlockCodeFragment(project, "fragment.kt", text, null, context)
} }
public fun createReturn(text: String): JetReturnExpression { public fun createReturn(text: String): JetReturnExpression {
@@ -25,7 +25,7 @@ import org.jetbrains.kotlin.types.JetType;
public class JetTypeCodeFragment extends JetCodeFragment { public class JetTypeCodeFragment extends JetCodeFragment {
public JetTypeCodeFragment(Project project, String name, CharSequence text, PsiElement context) { public JetTypeCodeFragment(Project project, String name, CharSequence text, PsiElement context) {
super(project, name, text, JetNodeTypes.TYPE_CODE_FRAGMENT, context); super(project, name, text, null, JetNodeTypes.TYPE_CODE_FRAGMENT, context);
} }
@Nullable @Nullable
@@ -43,7 +43,7 @@ import org.jetbrains.kotlin.psi.JetArrayAccessExpression
class KotlinEditorTextProvider : EditorTextProvider { class KotlinEditorTextProvider : EditorTextProvider {
override fun getEditorText(elementAtCaret: PsiElement): TextWithImports? { override fun getEditorText(elementAtCaret: PsiElement): TextWithImports? {
val expression = findExpressionInner(elementAtCaret, true) val expression = findExpressionInner(elementAtCaret, true)
return TextWithImportsImpl(CodeFragmentKind.EXPRESSION, expression?.getText() ?: "", JetCodeFragment.getImportsForElement(elementAtCaret), JetFileType.INSTANCE) return TextWithImportsImpl(CodeFragmentKind.EXPRESSION, expression?.getText() ?: "", "", JetFileType.INSTANCE)
} }
override fun findExpression(elementAtCaret: PsiElement, allowMethodCalls: Boolean): Pair<PsiElement, TextRange>? { override fun findExpression(elementAtCaret: PsiElement, allowMethodCalls: Boolean): Pair<PsiElement, TextRange>? {
@@ -40,12 +40,11 @@ import com.intellij.openapi.progress.ProgressManager
class KotlinCodeFragmentFactory: CodeFragmentFactory() { class KotlinCodeFragmentFactory: CodeFragmentFactory() {
override fun createCodeFragment(item: TextWithImports, context: PsiElement?, project: Project): JavaCodeFragment { override fun createCodeFragment(item: TextWithImports, context: PsiElement?, project: Project): JavaCodeFragment {
val codeFragment = if (item.getKind() == CodeFragmentKind.EXPRESSION) { val codeFragment = if (item.getKind() == CodeFragmentKind.EXPRESSION) {
JetExpressionCodeFragment(project, "fragment.kt", item.getText(), getContextElement(context)) JetExpressionCodeFragment(project, "fragment.kt", item.getText(), item.getImports(), getContextElement(context))
} }
else { else {
JetBlockCodeFragment(project, "fragment.kt", item.getText(), getContextElement(context)) JetBlockCodeFragment(project, "fragment.kt", item.getText(), item.getImports(), getContextElement(context))
} }
codeFragment.addImportsFromString(item.getImports())
codeFragment.putCopyableUserData(JetCodeFragment.RUNTIME_TYPE_EVALUATOR, { codeFragment.putCopyableUserData(JetCodeFragment.RUNTIME_TYPE_EVALUATOR, {
(expression: JetExpression): JetType? -> (expression: JetExpression): JetType? ->
@@ -97,10 +97,6 @@ object KotlinEvaluationBuilder: EvaluatorBuilder {
throw EvaluateExceptionUtil.createEvaluateException("Couldn't evaluate kotlin expression in this context") throw EvaluateExceptionUtil.createEvaluateException("Couldn't evaluate kotlin expression in this context")
} }
val packageName = file.getPackageDirective()?.getFqName()?.asString()
if (packageName != null && packageName.isNotEmpty()) {
codeFragment.addImportsFromString("import $packageName.*")
}
return ExpressionEvaluatorImpl(KotlinEvaluator(codeFragment as JetCodeFragment, position)) return ExpressionEvaluatorImpl(KotlinEvaluator(codeFragment as JetCodeFragment, position))
} }
} }
@@ -179,7 +179,10 @@ public class ImportInsertHelperImpl(private val project: Project) : ImportInsert
return ImportDescriptorResult.FAIL return ImportDescriptorResult.FAIL
} }
val imports = file.getImportDirectives() val imports = if (file is JetCodeFragment)
file.importsAsImportList()?.getImports() ?: listOf()
else
file.getImportDirectives()
//TODO: is that correct? What if function is imported and we need to import class? //TODO: is that correct? What if function is imported and we need to import class?
if (imports.any { it.getImportedName() == name.asString() }) return ImportDescriptorResult.FAIL if (imports.any { it.getImportedName() == name.asString() }) return ImportDescriptorResult.FAIL
@@ -66,8 +66,6 @@ public abstract class AbstractCodeFragmentHighlightingTest : AbstractJetPsiCheck
.processImportReference(importDirective, scope, scope, null, BindingTraceContext(), LookupMode.EVERYTHING) .processImportReference(importDirective, scope, scope, null, BindingTraceContext(), LookupMode.EVERYTHING)
.singleOrNull() ?: error("Could not resolve descriptor to import: $it") .singleOrNull() ?: error("Could not resolve descriptor to import: $it")
ImportInsertHelper.getInstance(getProject()).importDescriptor(file, descriptor) ImportInsertHelper.getInstance(getProject()).importDescriptor(file, descriptor)
//TODO: it's a hack! we need to discuss it
(PsiModificationTracker.SERVICE.getInstance(getProject()) as PsiModificationTrackerImpl).incOutOfCodeBlockModificationCounter()
} }
} }
@@ -90,7 +88,7 @@ public abstract class AbstractCodeFragmentCompletionHandlerTest : AbstractComple
super.doTest(testPath) super.doTest(testPath)
val fragment = myFixture.getFile() as JetCodeFragment val fragment = myFixture.getFile() as JetCodeFragment
val importList = fragment.getImportList() val importList = fragment.importsAsImportList()
val fragmentAfterFile = File(testPath + ".after.imports") val fragmentAfterFile = File(testPath + ".after.imports")
if (importList != null && fragmentAfterFile.exists()) { if (importList != null && fragmentAfterFile.exists()) {
@@ -331,11 +331,7 @@ public abstract class AbstractKotlinEvaluateExpressionTest : KotlinDebuggerTestB
try { try {
val evaluator = val evaluator =
EvaluatorBuilderImpl.build(TextWithImportsImpl( EvaluatorBuilderImpl.build(TextWithImportsImpl(codeFragmentKind, text, "", JetFileType.INSTANCE),
codeFragmentKind,
text,
JetCodeFragment.getImportsForElement(contextElement),
JetFileType.INSTANCE),
contextElement, contextElement,
sourcePosition) sourcePosition)