Refactoring to make receiver type safe

This commit is contained in:
Valentin Kipyatkov
2015-09-30 15:10:16 +03:00
parent c12520da7f
commit 2760b0bdb9
14 changed files with 154 additions and 118 deletions
@@ -54,7 +54,7 @@ public class ReferenceVariantsHelper(
expression: JetSimpleNameExpression, expression: JetSimpleNameExpression,
kindFilter: DescriptorKindFilter, kindFilter: DescriptorKindFilter,
nameFilter: (Name) -> Boolean, nameFilter: (Name) -> Boolean,
callTypeAndReceiver: CallTypeAndReceiver = CallTypeAndReceiver.detect(expression), callTypeAndReceiver: CallTypeAndReceiver<*, *> = CallTypeAndReceiver.detect(expression),
filterOutJavaGettersAndSetters: Boolean = false, filterOutJavaGettersAndSetters: Boolean = false,
useRuntimeReceiverType: Boolean = false useRuntimeReceiverType: Boolean = false
): Collection<DeclarationDescriptor> { ): Collection<DeclarationDescriptor> {
@@ -83,22 +83,31 @@ public class ReferenceVariantsHelper(
expression: JetSimpleNameExpression, expression: JetSimpleNameExpression,
kindFilter: DescriptorKindFilter, kindFilter: DescriptorKindFilter,
nameFilter: (Name) -> Boolean, nameFilter: (Name) -> Boolean,
callTypeAndReceiver: CallTypeAndReceiver, callTypeAndReceiver: CallTypeAndReceiver<*, *>,
useRuntimeReceiverType: Boolean useRuntimeReceiverType: Boolean
): Collection<DeclarationDescriptor> { ): Collection<DeclarationDescriptor> {
val (callType, receiverElement) = callTypeAndReceiver val receiverExpression: JetExpression?
when (callTypeAndReceiver) {
is CallTypeAndReceiver.IMPORT_DIRECTIVE -> {
return getVariantsForImportOrPackageDirective(callTypeAndReceiver.receiver, kindFilter, nameFilter)
}
if (callType == CallType.IMPORT_DIRECTIVE) { is CallTypeAndReceiver.PACKAGE_DIRECTIVE -> {
return getVariantsForImportOrPackageDirective(receiverElement as JetExpression?, kindFilter, nameFilter) val packageKindFilter = kindFilter restrictedToKinds DescriptorKindFilter.PACKAGES_MASK
} return getVariantsForImportOrPackageDirective(callTypeAndReceiver.receiver, packageKindFilter, nameFilter)
}
if (callType == CallType.PACKAGE_DIRECTIVE) { is CallTypeAndReceiver.TYPE -> {
val packageKindFilter = kindFilter restrictedToKinds DescriptorKindFilter.PACKAGES_MASK return getVariantsForUserType(callTypeAndReceiver.receiver, expression, kindFilter, nameFilter)
return getVariantsForImportOrPackageDirective(receiverElement as JetExpression?, packageKindFilter, nameFilter) }
}
if (expression.parent is JetUserType) { //TODO: special CallType? is CallTypeAndReceiver.CALLABLE_REFERENCE -> receiverExpression = null // handled below
return getVariantsForUserType(receiverElement as JetExpression?, expression, kindFilter, nameFilter)
is CallTypeAndReceiver.DEFAULT -> receiverExpression = null
is CallTypeAndReceiver.DOT -> receiverExpression = callTypeAndReceiver.receiver
is CallTypeAndReceiver.SAFE -> receiverExpression = callTypeAndReceiver.receiver
is CallTypeAndReceiver.INFIX -> receiverExpression = callTypeAndReceiver.receiver
is CallTypeAndReceiver.UNARY -> receiverExpression = null // can it happen at all?
} }
val resolutionScope = resolutionScope(expression) ?: return emptyList() val resolutionScope = resolutionScope(expression) ?: return emptyList()
@@ -110,27 +119,27 @@ public class ReferenceVariantsHelper(
smartCastManager.getSmartCastVariantsWithLessSpecificExcluded(it.value, context, containingDeclaration, dataFlowInfo) smartCastManager.getSmartCastVariantsWithLessSpecificExcluded(it.value, context, containingDeclaration, dataFlowInfo)
}.toSet() }.toSet()
if (callType == CallType.CALLABLE_REFERENCE) { if (callTypeAndReceiver is CallTypeAndReceiver.CALLABLE_REFERENCE) {
return getVariantsForCallableReference(receiverElement as JetTypeReference?, resolutionScope, implicitReceiverTypes, kindFilter, nameFilter) return getVariantsForCallableReference(callTypeAndReceiver.receiver, resolutionScope, implicitReceiverTypes, kindFilter, nameFilter)
} }
val callType = callTypeAndReceiver.callType
val descriptors = LinkedHashSet<DeclarationDescriptor>() val descriptors = LinkedHashSet<DeclarationDescriptor>()
if (receiverElement != null) { if (receiverExpression != null) {
receiverElement as JetExpression val qualifier = context[BindingContext.QUALIFIER, receiverExpression]
val qualifier = context[BindingContext.QUALIFIER, receiverElement]
if (qualifier != null) { if (qualifier != null) {
// It's impossible to add extension function for package or class (if it's companion object, expression type is not null) // It's impossible to add extension function for package or class (if it's companion object, expression type is not null)
qualifier.scope.getDescriptorsFiltered(kindFilter exclude DescriptorKindExclude.Extensions, nameFilter).filterTo(descriptors) { callType.canCall(it) } qualifier.scope.getDescriptorsFiltered(kindFilter exclude DescriptorKindExclude.Extensions, nameFilter).filterTo(descriptors) { callType.canCall(it) }
} }
val expressionType = if (useRuntimeReceiverType) val expressionType = if (useRuntimeReceiverType)
getQualifierRuntimeType(receiverElement) getQualifierRuntimeType(receiverExpression)
else else
context.getType(receiverElement) context.getType(receiverExpression)
if (expressionType != null && !expressionType.isError()) { if (expressionType != null && !expressionType.isError()) {
val receiverValue = ExpressionReceiver(receiverElement, expressionType) val receiverValue = ExpressionReceiver(receiverExpression, expressionType)
val explicitReceiverTypes = smartCastManager val explicitReceiverTypes = smartCastManager
.getSmartCastVariantsWithLessSpecificExcluded(receiverValue, context, containingDeclaration, dataFlowInfo) .getSmartCastVariantsWithLessSpecificExcluded(receiverValue, context, containingDeclaration, dataFlowInfo)
@@ -138,7 +147,7 @@ public class ReferenceVariantsHelper(
} }
} }
else { else {
descriptors.processAll(implicitReceiverTypes, implicitReceiverTypes, resolutionScope, CallType.NORMAL, kindFilter, nameFilter) descriptors.processAll(implicitReceiverTypes, implicitReceiverTypes, resolutionScope, callType, kindFilter, nameFilter)
// add non-instance members // add non-instance members
descriptors.addAll(resolutionScope.getDescriptorsFiltered(kindFilter exclude DescriptorKindExclude.Extensions, nameFilter)) descriptors.addAll(resolutionScope.getDescriptorsFiltered(kindFilter exclude DescriptorKindExclude.Extensions, nameFilter))
@@ -213,7 +222,7 @@ public class ReferenceVariantsHelper(
implicitReceiverTypes: Collection<JetType>, implicitReceiverTypes: Collection<JetType>,
receiverTypes: Collection<JetType>, receiverTypes: Collection<JetType>,
resolutionScope: JetScope, resolutionScope: JetScope,
callType: CallType, callType: CallType<*>,
kindFilter: DescriptorKindFilter, kindFilter: DescriptorKindFilter,
nameFilter: (Name) -> Boolean nameFilter: (Name) -> Boolean
) { ) {
@@ -225,7 +234,7 @@ public class ReferenceVariantsHelper(
private fun MutableSet<DeclarationDescriptor>.addMemberExtensions( private fun MutableSet<DeclarationDescriptor>.addMemberExtensions(
dispatchReceiverTypes: Collection<JetType>, dispatchReceiverTypes: Collection<JetType>,
extensionReceiverTypes: Collection<JetType>, extensionReceiverTypes: Collection<JetType>,
callType: CallType, callType: CallType<*>,
kindFilter: DescriptorKindFilter, kindFilter: DescriptorKindFilter,
nameFilter: (Name) -> Boolean nameFilter: (Name) -> Boolean
) { ) {
@@ -239,7 +248,7 @@ public class ReferenceVariantsHelper(
private fun MutableSet<DeclarationDescriptor>.addNonExtensionMembers( private fun MutableSet<DeclarationDescriptor>.addNonExtensionMembers(
receiverTypes: Collection<JetType>, receiverTypes: Collection<JetType>,
callType: CallType, callType: CallType<*>,
kindFilter: DescriptorKindFilter, kindFilter: DescriptorKindFilter,
nameFilter: (Name) -> Boolean, nameFilter: (Name) -> Boolean,
constructorsForInnerClassesOnly: Boolean constructorsForInnerClassesOnly: Boolean
@@ -251,7 +260,7 @@ public class ReferenceVariantsHelper(
private fun MutableSet<DeclarationDescriptor>.addNonExtensionCallablesAndConstructors( private fun MutableSet<DeclarationDescriptor>.addNonExtensionCallablesAndConstructors(
scope: JetScope, scope: JetScope,
callType: CallType, callType: CallType<*>,
kindFilter: DescriptorKindFilter, kindFilter: DescriptorKindFilter,
nameFilter: (Name) -> Boolean, nameFilter: (Name) -> Boolean,
constructorsForInnerClassesOnly: Boolean constructorsForInnerClassesOnly: Boolean
@@ -278,7 +287,7 @@ public class ReferenceVariantsHelper(
private fun MutableSet<DeclarationDescriptor>.addScopeAndSyntheticExtensions( private fun MutableSet<DeclarationDescriptor>.addScopeAndSyntheticExtensions(
resolutionScope: JetScope, resolutionScope: JetScope,
receiverTypes: Collection<JetType>, receiverTypes: Collection<JetType>,
callType: CallType, callType: CallType<*>,
kindFilter: DescriptorKindFilter, kindFilter: DescriptorKindFilter,
nameFilter: (Name) -> Boolean nameFilter: (Name) -> Boolean
) { ) {
@@ -26,85 +26,99 @@ import org.jetbrains.kotlin.psi.psiUtil.getReceiverExpression
import org.jetbrains.kotlin.psi.psiUtil.isImportDirectiveExpression import org.jetbrains.kotlin.psi.psiUtil.isImportDirectiveExpression
import org.jetbrains.kotlin.psi.psiUtil.isPackageDirectiveExpression import org.jetbrains.kotlin.psi.psiUtil.isPackageDirectiveExpression
public enum class CallType { public sealed class CallType<TReceiver : JetElement?> {
NORMAL, object DEFAULT : CallType<Nothing?>()
SAFE,
INFIX { object DOT : CallType<JetExpression>()
object SAFE : CallType<JetExpression>()
object INFIX : CallType<JetExpression>() {
override fun canCall(descriptor: DeclarationDescriptor) override fun canCall(descriptor: DeclarationDescriptor)
= descriptor is SimpleFunctionDescriptor && descriptor.getValueParameters().size() == 1 = descriptor is SimpleFunctionDescriptor && descriptor.getValueParameters().size() == 1
}, }
UNARY { object UNARY : CallType<JetExpression>() {
override fun canCall(descriptor: DeclarationDescriptor) override fun canCall(descriptor: DeclarationDescriptor)
= descriptor is SimpleFunctionDescriptor && descriptor.getValueParameters().size() == 0 = descriptor is SimpleFunctionDescriptor && descriptor.getValueParameters().size() == 0
}, }
CALLABLE_REFERENCE { object CALLABLE_REFERENCE : CallType<JetTypeReference?>() {
// currently callable references to locals and parameters are not supported // currently callable references to locals and parameters are not supported
override fun canCall(descriptor: DeclarationDescriptor) override fun canCall(descriptor: DeclarationDescriptor)
= descriptor is FunctionDescriptor || descriptor is PropertyDescriptor = descriptor is FunctionDescriptor || descriptor is PropertyDescriptor
}, }
//TODO: canCall //TODO: canCall
IMPORT_DIRECTIVE, object IMPORT_DIRECTIVE : CallType<JetExpression?>()
PACKAGE_DIRECTIVE object PACKAGE_DIRECTIVE : CallType<JetExpression?>()
; object TYPE : CallType<JetExpression?>()
public open fun canCall(descriptor: DeclarationDescriptor): Boolean = true public open fun canCall(descriptor: DeclarationDescriptor): Boolean = true
} }
public data class CallTypeAndReceiver( public sealed class CallTypeAndReceiver<TReceiver : JetElement?, TCallType : CallType<TReceiver>>(
val callType: CallType, val callType: TCallType,
val receiver: JetElement? val receiver: TReceiver
) { ) {
object DEFAULT : CallTypeAndReceiver<Nothing?, CallType.DEFAULT>(CallType.DEFAULT, null)
class DOT(receiver: JetExpression) : CallTypeAndReceiver<JetExpression, CallType.DOT>(CallType.DOT, receiver)
class SAFE(receiver: JetExpression) : CallTypeAndReceiver<JetExpression, CallType.SAFE>(CallType.SAFE, receiver)
class INFIX(receiver: JetExpression) : CallTypeAndReceiver<JetExpression, CallType.INFIX>(CallType.INFIX, receiver)
class UNARY(receiver: JetExpression) : CallTypeAndReceiver<JetExpression, CallType.UNARY>(CallType.UNARY, receiver)
class CALLABLE_REFERENCE(receiver: JetTypeReference?) : CallTypeAndReceiver<JetTypeReference?, CallType.CALLABLE_REFERENCE>(CallType.CALLABLE_REFERENCE, receiver)
class IMPORT_DIRECTIVE(receiver: JetExpression?) : CallTypeAndReceiver<JetExpression?, CallType.IMPORT_DIRECTIVE>(CallType.IMPORT_DIRECTIVE, receiver)
class PACKAGE_DIRECTIVE(receiver: JetExpression?) : CallTypeAndReceiver<JetExpression?, CallType.PACKAGE_DIRECTIVE>(CallType.PACKAGE_DIRECTIVE, receiver)
class TYPE(receiver: JetExpression?) : CallTypeAndReceiver<JetExpression?, CallType.TYPE>(CallType.TYPE, receiver)
companion object { companion object {
public fun detect(expression: JetSimpleNameExpression): CallTypeAndReceiver { public fun detect(expression: JetSimpleNameExpression): CallTypeAndReceiver<*, *> {
val parent = expression.parent val parent = expression.parent
if (parent is JetCallableReferenceExpression) { if (parent is JetCallableReferenceExpression) {
return CallTypeAndReceiver(CallType.CALLABLE_REFERENCE, parent.typeReference) return CallTypeAndReceiver.CALLABLE_REFERENCE(parent.typeReference)
} }
val receiverExpression = expression.getReceiverExpression() val receiverExpression = expression.getReceiverExpression()
if (expression.isImportDirectiveExpression()) { if (expression.isImportDirectiveExpression()) {
return CallTypeAndReceiver(CallType.IMPORT_DIRECTIVE, receiverExpression) return CallTypeAndReceiver.IMPORT_DIRECTIVE(receiverExpression)
} }
if (expression.isPackageDirectiveExpression()) { if (expression.isPackageDirectiveExpression()) {
return CallTypeAndReceiver(CallType.PACKAGE_DIRECTIVE, receiverExpression) return CallTypeAndReceiver.PACKAGE_DIRECTIVE(receiverExpression)
}
if (parent is JetUserType) {
return CallTypeAndReceiver.TYPE(receiverExpression)
} }
if (receiverExpression == null) { if (receiverExpression == null) {
return CallTypeAndReceiver(CallType.NORMAL, null) return CallTypeAndReceiver.DEFAULT
} }
val callType = when (parent) { return when (parent) {
is JetBinaryExpression -> CallType.INFIX is JetBinaryExpression -> CallTypeAndReceiver.INFIX(receiverExpression)
is JetCallExpression -> { is JetCallExpression -> {
if ((parent.getParent() as JetQualifiedExpression).getOperationSign() == JetTokens.SAFE_ACCESS) if ((parent.parent as JetQualifiedExpression).operationSign == JetTokens.SAFE_ACCESS)
CallType.SAFE CallTypeAndReceiver.SAFE(receiverExpression)
else else
CallType.NORMAL CallTypeAndReceiver.DOT(receiverExpression)
} }
is JetQualifiedExpression -> { is JetQualifiedExpression -> {
if (parent.getOperationSign() == JetTokens.SAFE_ACCESS) if (parent.operationSign == JetTokens.SAFE_ACCESS)
CallType.SAFE CallTypeAndReceiver.SAFE(receiverExpression)
else else
CallType.NORMAL CallTypeAndReceiver.DOT(receiverExpression)
} }
is JetUnaryExpression -> CallType.UNARY is JetUnaryExpression -> CallTypeAndReceiver.UNARY(receiverExpression)
is JetUserType -> CallType.NORMAL
else -> error("Unknown parent for expression with receiver: $parent") else -> error("Unknown parent for expression with receiver: $parent")
} }
return CallTypeAndReceiver(callType, receiverExpression)
} }
} }
} }
@@ -40,7 +40,7 @@ public class ShadowedDeclarationsFilter(
private val bindingContext: BindingContext, private val bindingContext: BindingContext,
private val resolutionFacade: ResolutionFacade, private val resolutionFacade: ResolutionFacade,
private val context: JetExpression, private val context: JetExpression,
callTypeAndReceiver: CallTypeAndReceiver callTypeAndReceiver: CallTypeAndReceiver<*, *>
) { ) {
private val psiFactory = JetPsiFactory(resolutionFacade.project) private val psiFactory = JetPsiFactory(resolutionFacade.project)
private val dummyExpressionFactory = DummyExpressionFactory(psiFactory) private val dummyExpressionFactory = DummyExpressionFactory(psiFactory)
@@ -18,7 +18,8 @@
package org.jetbrains.kotlin.idea.util package org.jetbrains.kotlin.idea.util
import org.jetbrains.kotlin.descriptors.* import org.jetbrains.kotlin.descriptors.CallableDescriptor
import org.jetbrains.kotlin.descriptors.DeclarationDescriptor
import org.jetbrains.kotlin.psi.JetPsiUtil import org.jetbrains.kotlin.psi.JetPsiUtil
import org.jetbrains.kotlin.psi.JetThisExpression import org.jetbrains.kotlin.psi.JetThisExpression
import org.jetbrains.kotlin.resolve.BindingContext import org.jetbrains.kotlin.resolve.BindingContext
@@ -37,7 +38,7 @@ public fun CallableDescriptor.substituteExtensionIfCallable(
receivers: Collection<ReceiverValue>, receivers: Collection<ReceiverValue>,
context: BindingContext, context: BindingContext,
dataFlowInfo: DataFlowInfo, dataFlowInfo: DataFlowInfo,
callType: CallType, callType: CallType<*>,
containingDeclarationOrModule: DeclarationDescriptor containingDeclarationOrModule: DeclarationDescriptor
): Collection<CallableDescriptor> { ): Collection<CallableDescriptor> {
val sequence = receivers.asSequence().flatMap { substituteExtensionIfCallable(it, callType, context, dataFlowInfo, containingDeclarationOrModule).asSequence() } val sequence = receivers.asSequence().flatMap { substituteExtensionIfCallable(it, callType, context, dataFlowInfo, containingDeclarationOrModule).asSequence() }
@@ -55,12 +56,12 @@ public fun CallableDescriptor.substituteExtensionIfCallableWithImplicitReceiver(
dataFlowInfo: DataFlowInfo dataFlowInfo: DataFlowInfo
): Collection<CallableDescriptor> { ): Collection<CallableDescriptor> {
val receiverValues = scope.getImplicitReceiversWithInstance().map { it.getValue() } val receiverValues = scope.getImplicitReceiversWithInstance().map { it.getValue() }
return substituteExtensionIfCallable(receiverValues, context, dataFlowInfo, CallType.NORMAL, scope.getContainingDeclaration()) return substituteExtensionIfCallable(receiverValues, context, dataFlowInfo, CallType.DEFAULT, scope.getContainingDeclaration())
} }
public fun CallableDescriptor.substituteExtensionIfCallable( public fun CallableDescriptor.substituteExtensionIfCallable(
receiver: ReceiverValue, receiver: ReceiverValue,
callType: CallType, callType: CallType<*>,
bindingContext: BindingContext, bindingContext: BindingContext,
dataFlowInfo: DataFlowInfo, dataFlowInfo: DataFlowInfo,
containingDeclarationOrModule: DeclarationDescriptor containingDeclarationOrModule: DeclarationDescriptor
@@ -73,7 +74,7 @@ public fun CallableDescriptor.substituteExtensionIfCallable(
public fun CallableDescriptor.substituteExtensionIfCallable( public fun CallableDescriptor.substituteExtensionIfCallable(
receiverTypes: Collection<JetType>, receiverTypes: Collection<JetType>,
callType: CallType callType: CallType<*>
): Collection<CallableDescriptor> { ): Collection<CallableDescriptor> {
if (!callType.canCall(this)) return listOf() if (!callType.canCall(this)) return listOf()
@@ -303,7 +303,7 @@ abstract class CompletionSession(protected val configuration: CompletionSessionC
val contextVariablesProvider = { val contextVariablesProvider = {
nameExpression?.let { nameExpression?.let {
referenceVariantsHelper.getReferenceVariants(it, DescriptorKindFilter.VARIABLES, { true }, CallTypeAndReceiver(CallType.NORMAL, null)) referenceVariantsHelper.getReferenceVariants(it, DescriptorKindFilter.VARIABLES, { true }, CallTypeAndReceiver.DEFAULT)
.map { it as VariableDescriptor } .map { it as VariableDescriptor }
} ?: emptyList() } ?: emptyList()
} }
@@ -314,23 +314,38 @@ abstract class CompletionSession(protected val configuration: CompletionSessionC
insertHandlerProvider, contextVariablesProvider) insertHandlerProvider, contextVariablesProvider)
} }
private fun detectCallTypeAndReceiverTypes(): Pair<CallType, Collection<JetType>> { private fun detectCallTypeAndReceiverTypes(): Pair<CallType<*>, Collection<JetType>> {
if (nameExpression == null) { if (nameExpression == null) {
return CallType.NORMAL to emptyList() return CallType.DEFAULT to emptyList()
} }
val (callType, receiverElement) = CallTypeAndReceiver.detect(nameExpression) val callTypeAndReceiver = CallTypeAndReceiver.detect(nameExpression)
if (callType == CallType.CALLABLE_REFERENCE && receiverElement != null) { val receiverExpression: JetExpression?
val type = bindingContext[BindingContext.TYPE, receiverElement as JetTypeReference] when (callTypeAndReceiver) {
return callType to type.singletonOrEmptyList() is CallTypeAndReceiver.CALLABLE_REFERENCE -> {
if (callTypeAndReceiver.receiver != null) {
val type = bindingContext[BindingContext.TYPE, callTypeAndReceiver.receiver]
return callTypeAndReceiver.callType to type.singletonOrEmptyList()
}
else {
receiverExpression = null
}
}
is CallTypeAndReceiver.DEFAULT -> receiverExpression = null
is CallTypeAndReceiver.DOT -> receiverExpression = callTypeAndReceiver.receiver
is CallTypeAndReceiver.SAFE -> receiverExpression = callTypeAndReceiver.receiver
is CallTypeAndReceiver.INFIX -> receiverExpression = callTypeAndReceiver.receiver
is CallTypeAndReceiver.UNARY -> receiverExpression = callTypeAndReceiver.receiver
is CallTypeAndReceiver.IMPORT_DIRECTIVE -> receiverExpression = callTypeAndReceiver.receiver
is CallTypeAndReceiver.PACKAGE_DIRECTIVE -> receiverExpression = callTypeAndReceiver.receiver
is CallTypeAndReceiver.TYPE -> receiverExpression = callTypeAndReceiver.receiver
} }
receiverElement as JetExpression? val receiverValues = if (receiverExpression != null) {
val expressionType = bindingContext.getType(receiverExpression)
val receiverValues = if (receiverElement != null) { expressionType?.let { listOf(ExpressionReceiver(receiverExpression, expressionType)) } ?: emptyList()
val expressionType = bindingContext.getType(receiverElement)
expressionType?.let { listOf(ExpressionReceiver(receiverElement, expressionType)) } ?: emptyList()
} }
else { else {
val resolutionScope = referenceVariantsHelper.resolutionScope(nameExpression) val resolutionScope = referenceVariantsHelper.resolutionScope(nameExpression)
@@ -350,10 +365,10 @@ abstract class CompletionSession(protected val configuration: CompletionSessionC
} }
} }
if (callType == CallType.SAFE) { if (callTypeAndReceiver is CallTypeAndReceiver.SAFE) {
receiverTypes = receiverTypes.map { it.makeNotNullable() } receiverTypes = receiverTypes.map { it.makeNotNullable() }
} }
return callType to receiverTypes return callTypeAndReceiver.callType to receiverTypes
} }
} }
@@ -27,7 +27,7 @@ import org.jetbrains.kotlin.types.JetType
import java.util.* import java.util.*
class InsertHandlerProvider( class InsertHandlerProvider(
private val callType: CallType, private val callType: CallType<*>,
expectedInfosCalculator: () -> Collection<ExpectedInfo> expectedInfosCalculator: () -> Collection<ExpectedInfo>
) { ) {
private val expectedInfos by lazy(LazyThreadSafetyMode.NONE) { expectedInfosCalculator() } private val expectedInfos by lazy(LazyThreadSafetyMode.NONE) { expectedInfosCalculator() }
@@ -100,7 +100,7 @@ class KDocNameCompletionSession(parameters: CompletionParameters,
val extensionReceiver = descriptor.getExtensionReceiverParameter() val extensionReceiver = descriptor.getExtensionReceiverParameter()
if (extensionReceiver != null) { if (extensionReceiver != null) {
val substituted = descriptor.substituteExtensionIfCallable(implicitReceivers, bindingContext, DataFlowInfo.EMPTY, val substituted = descriptor.substituteExtensionIfCallable(implicitReceivers, bindingContext, DataFlowInfo.EMPTY,
CallType.NORMAL, moduleDescriptor) CallType.DEFAULT, moduleDescriptor)
return !substituted.isEmpty() return !substituted.isEmpty()
} }
} }
@@ -77,7 +77,7 @@ object KeywordCompletion {
.withInsertHandler(if (keywordToken !in FUNCTION_KEYWORDS) .withInsertHandler(if (keywordToken !in FUNCTION_KEYWORDS)
KotlinKeywordInsertHandler KotlinKeywordInsertHandler
else else
KotlinFunctionInsertHandler(CallType.NORMAL, inputTypeArguments = false, inputValueArguments = false)) KotlinFunctionInsertHandler(CallType.DEFAULT, inputTypeArguments = false, inputValueArguments = false))
consumer(element) consumer(element)
} }
} }
@@ -45,7 +45,7 @@ import org.jetbrains.kotlin.types.typeUtil.isSubtypeOf
class LookupElementFactory( class LookupElementFactory(
private val resolutionFacade: ResolutionFacade, private val resolutionFacade: ResolutionFacade,
private val receiverTypes: Collection<JetType>, private val receiverTypes: Collection<JetType>,
private val callType: CallType, private val callType: CallType<*>,
private val isInStringTemplateAfterDollar: Boolean, private val isInStringTemplateAfterDollar: Boolean,
public val insertHandlerProvider: InsertHandlerProvider, public val insertHandlerProvider: InsertHandlerProvider,
contextVariablesProvider: () -> Collection<VariableDescriptor> contextVariablesProvider: () -> Collection<VariableDescriptor>
@@ -66,7 +66,7 @@ class LookupElementFactory(
result.add(lookupElement) result.add(lookupElement)
// add special item for function with one argument of function type with more than one parameter // add special item for function with one argument of function type with more than one parameter
if (descriptor is FunctionDescriptor && (callType == CallType.NORMAL || callType == CallType.SAFE)) { if (descriptor is FunctionDescriptor && (callType == CallType.DEFAULT || callType == CallType.DOT || callType == CallType.SAFE)) {
result.addSpecialFunctionCallElements(descriptor, useReceiverTypes) result.addSpecialFunctionCallElements(descriptor, useReceiverTypes)
} }
@@ -54,7 +54,7 @@ object PackageDirectiveCompletion {
val bindingContext = resolutionFacade.analyze(expression) val bindingContext = resolutionFacade.analyze(expression)
val variants = ReferenceVariantsHelper(bindingContext, resolutionFacade, { true }).getPackageReferenceVariants(expression, prefixMatcher.asNameFilter()) val variants = ReferenceVariantsHelper(bindingContext, resolutionFacade, { true }).getPackageReferenceVariants(expression, prefixMatcher.asNameFilter())
val lookupElementFactory = BasicLookupElementFactory(resolutionFacade.project, InsertHandlerProvider(callType = CallType.NORMAL/*TODO*/, expectedInfosCalculator = { emptyList() })) val lookupElementFactory = BasicLookupElementFactory(resolutionFacade.project, InsertHandlerProvider(callType = CallType.PACKAGE_DIRECTIVE, expectedInfosCalculator = { emptyList() }))
for (variant in variants) { for (variant in variants) {
val lookupElement = lookupElementFactory.createLookupElement(variant) val lookupElement = lookupElementFactory.createLookupElement(variant)
if (!lookupElement.getLookupString().contains(DUMMY_IDENTIFIER)) { if (!lookupElement.getLookupString().contains(DUMMY_IDENTIFIER)) {
@@ -35,7 +35,7 @@ import org.jetbrains.kotlin.types.JetType
class GenerateLambdaInfo(val lambdaType: JetType, val explicitParameters: Boolean) class GenerateLambdaInfo(val lambdaType: JetType, val explicitParameters: Boolean)
class KotlinFunctionInsertHandler( class KotlinFunctionInsertHandler(
val callType: CallType, val callType: CallType<*>,
val inputTypeArguments: Boolean, val inputTypeArguments: Boolean,
val inputValueArguments: Boolean, val inputValueArguments: Boolean,
val argumentText: String = "", val argumentText: String = "",
@@ -22,12 +22,10 @@ import com.intellij.codeInsight.completion.CompletionSorter
import com.intellij.psi.impl.source.tree.LeafPsiElement import com.intellij.psi.impl.source.tree.LeafPsiElement
import org.jetbrains.kotlin.descriptors.DeclarationDescriptor import org.jetbrains.kotlin.descriptors.DeclarationDescriptor
import org.jetbrains.kotlin.idea.completion.* import org.jetbrains.kotlin.idea.completion.*
import org.jetbrains.kotlin.idea.util.CallType
import org.jetbrains.kotlin.idea.util.CallTypeAndReceiver import org.jetbrains.kotlin.idea.util.CallTypeAndReceiver
import org.jetbrains.kotlin.load.java.descriptors.SamConstructorDescriptorKindExclude import org.jetbrains.kotlin.load.java.descriptors.SamConstructorDescriptorKindExclude
import org.jetbrains.kotlin.psi.FunctionLiteralArgument import org.jetbrains.kotlin.psi.FunctionLiteralArgument
import org.jetbrains.kotlin.psi.JetCodeFragment import org.jetbrains.kotlin.psi.JetCodeFragment
import org.jetbrains.kotlin.psi.JetExpression
import org.jetbrains.kotlin.psi.ValueArgumentName import org.jetbrains.kotlin.psi.ValueArgumentName
import org.jetbrains.kotlin.resolve.calls.callUtil.getCall import org.jetbrains.kotlin.resolve.calls.callUtil.getCall
import org.jetbrains.kotlin.resolve.calls.util.DelegatingCall import org.jetbrains.kotlin.resolve.calls.util.DelegatingCall
@@ -91,30 +89,28 @@ class SmartCompletionSession(configuration: CompletionSessionConfiguration, para
// special completion for outside parenthesis lambda argument // special completion for outside parenthesis lambda argument
private fun addFunctionLiteralArgumentCompletions() { private fun addFunctionLiteralArgumentCompletions() {
if (nameExpression != null) { if (nameExpression != null) {
val (callType, receiverElement) = CallTypeAndReceiver.detect(nameExpression) val callTypeAndReceiver = CallTypeAndReceiver.detect(nameExpression) as? CallTypeAndReceiver.INFIX ?: return
if (callType == CallType.INFIX) { val call = callTypeAndReceiver.receiver.getCall(bindingContext)
val call = (receiverElement as JetExpression).getCall(bindingContext) if (call != null && call.getFunctionLiteralArguments().isEmpty()) {
if (call != null && call.getFunctionLiteralArguments().isEmpty()) { val dummyArgument = object : FunctionLiteralArgument {
val dummyArgument = object : FunctionLiteralArgument { override fun getFunctionLiteral() = throw UnsupportedOperationException()
override fun getFunctionLiteral() = throw UnsupportedOperationException() override fun getArgumentExpression() = throw UnsupportedOperationException()
override fun getArgumentExpression() = throw UnsupportedOperationException() override fun getArgumentName(): ValueArgumentName? = null
override fun getArgumentName(): ValueArgumentName? = null override fun isNamed() = false
override fun isNamed() = false override fun asElement() = throw UnsupportedOperationException()
override fun asElement() = throw UnsupportedOperationException() override fun getSpreadElement(): LeafPsiElement? = null
override fun getSpreadElement(): LeafPsiElement? = null override fun isExternal() = false
override fun isExternal() = false
}
val dummyArguments = call.getValueArguments() + listOf(dummyArgument)
val dummyCall = object : DelegatingCall(call) {
override fun getValueArguments() = dummyArguments
override fun getFunctionLiteralArguments() = listOf(dummyArgument)
override fun getValueArgumentList() = throw UnsupportedOperationException()
}
val expectedInfos = ExpectedInfos(bindingContext, resolutionFacade)
.calculateForArgument(dummyCall, dummyArgument)
collector.addElements(LambdaItems.collect(expectedInfos))
} }
val dummyArguments = call.getValueArguments() + listOf(dummyArgument)
val dummyCall = object : DelegatingCall(call) {
override fun getValueArguments() = dummyArguments
override fun getFunctionLiteralArguments() = listOf(dummyArgument)
override fun getValueArgumentList() = throw UnsupportedOperationException()
}
val expectedInfos = ExpectedInfos(bindingContext, resolutionFacade)
.calculateForArgument(dummyCall, dummyArgument)
collector.addElements(LambdaItems.collect(expectedInfos))
} }
} }
} }
@@ -204,9 +204,9 @@ class TypeInstantiationItems(
} }
val baseInsertHandler = when (visibleConstructors.size()) { val baseInsertHandler = when (visibleConstructors.size()) {
0 -> KotlinFunctionInsertHandler(CallType.NORMAL, inputTypeArguments = false, inputValueArguments = false) 0 -> KotlinFunctionInsertHandler(CallType.DEFAULT, inputTypeArguments = false, inputValueArguments = false)
1 -> lookupElementFactory.insertHandlerProvider.insertHandler(visibleConstructors.single()) as KotlinFunctionInsertHandler 1 -> lookupElementFactory.insertHandlerProvider.insertHandler(visibleConstructors.single()) as KotlinFunctionInsertHandler
else -> KotlinFunctionInsertHandler(CallType.NORMAL, inputTypeArguments = false, inputValueArguments = true) else -> KotlinFunctionInsertHandler(CallType.DEFAULT, inputTypeArguments = false, inputValueArguments = true)
} }
insertHandler = object : InsertHandler<LookupElement> { insertHandler = object : InsertHandler<LookupElement> {
@@ -131,8 +131,9 @@ public class KotlinIndicesHelper(
constructor.getSupertypes().forEach { addTypeNames(it) } constructor.getSupertypes().forEach { addTypeNames(it) }
} }
private fun receiverValues(expression: JetSimpleNameExpression, bindingContext: BindingContext): Collection<Pair<ReceiverValue, CallType>> { private fun receiverValues(expression: JetSimpleNameExpression, bindingContext: BindingContext): Collection<Pair<ReceiverValue, CallType<*>>> {
val (callType, receiverElement) = CallTypeAndReceiver.detect(expression) val callTypeAndReceiver = CallTypeAndReceiver.detect(expression)
val receiverElement = callTypeAndReceiver.receiver
if (receiverElement != null) { if (receiverElement != null) {
if (receiverElement !is JetExpression) return emptyList() //TODO? if (receiverElement !is JetExpression) return emptyList() //TODO?
@@ -141,11 +142,11 @@ public class KotlinIndicesHelper(
val receiverValue = ExpressionReceiver(receiverElement, expressionType) val receiverValue = ExpressionReceiver(receiverElement, expressionType)
return listOf(receiverValue to callType) return listOf(receiverValue to callTypeAndReceiver.callType)
} }
else { else {
val resolutionScope = bindingContext[BindingContext.RESOLUTION_SCOPE, expression] ?: return emptyList() val resolutionScope = bindingContext[BindingContext.RESOLUTION_SCOPE, expression] ?: return emptyList()
return resolutionScope.getImplicitReceiversWithInstance().map { it.getValue() to callType } return resolutionScope.getImplicitReceiversWithInstance().map { it.getValue() to callTypeAndReceiver.callType }
} }
} }
@@ -154,7 +155,7 @@ public class KotlinIndicesHelper(
*/ */
private fun findSuitableExtensions( private fun findSuitableExtensions(
declarations: Sequence<JetCallableDeclaration>, declarations: Sequence<JetCallableDeclaration>,
receiverValues: Collection<Pair<ReceiverValue, CallType>>, receiverValues: Collection<Pair<ReceiverValue, CallType<*>>>,
dataFlowInfo: DataFlowInfo, dataFlowInfo: DataFlowInfo,
bindingContext: BindingContext bindingContext: BindingContext
): Collection<CallableDescriptor> { ): Collection<CallableDescriptor> {