K1: Support referencing class context receivers in a form of this@Name

This commit is contained in:
Denis.Zharkov
2022-03-24 13:35:40 +03:00
committed by teamcity
parent c33f06b9e4
commit 4349060db1
10 changed files with 83 additions and 43 deletions
@@ -123,6 +123,7 @@ public interface BindingContext {
WritableSlice<KtSuperExpression, Boolean> SUPER_EXPRESSION_FROM_ANY_MIGRATION = Slices.createSimpleSlice(); WritableSlice<KtSuperExpression, Boolean> SUPER_EXPRESSION_FROM_ANY_MIGRATION = Slices.createSimpleSlice();
WritableSlice<KtReferenceExpression, DeclarationDescriptor> REFERENCE_TARGET = new BasicWritableSlice<>(DO_NOTHING); WritableSlice<KtReferenceExpression, DeclarationDescriptor> REFERENCE_TARGET = new BasicWritableSlice<>(DO_NOTHING);
WritableSlice<KtReferenceExpression, ReceiverParameterDescriptor> THIS_REFERENCE_TARGET = new BasicWritableSlice<>(DO_NOTHING);
// if 'A' really means 'A.Companion' then this slice stores class descriptor for A, REFERENCE_TARGET stores descriptor Companion in this case // if 'A' really means 'A.Companion' then this slice stores class descriptor for A, REFERENCE_TARGET stores descriptor Companion in this case
WritableSlice<KtReferenceExpression, ClassifierDescriptorWithTypeParameters> SHORT_REFERENCE_TO_COMPANION_OBJECT = WritableSlice<KtReferenceExpression, ClassifierDescriptorWithTypeParameters> SHORT_REFERENCE_TO_COMPANION_OBJECT =
new BasicWritableSlice<>(DO_NOTHING); new BasicWritableSlice<>(DO_NOTHING);
@@ -5,6 +5,7 @@
package org.jetbrains.kotlin.resolve.lazy.descriptors; package org.jetbrains.kotlin.resolve.lazy.descriptors;
import com.google.common.collect.HashMultimap;
import com.intellij.psi.PsiElement; import com.intellij.psi.PsiElement;
import com.intellij.psi.PsiNameIdentifierOwner; import com.intellij.psi.PsiNameIdentifierOwner;
import kotlin.Pair; import kotlin.Pair;
@@ -305,7 +306,8 @@ public class LazyClassDescriptor extends ClassDescriptorBase implements ClassDes
if (classOrObject == null) { if (classOrObject == null) {
return CollectionsKt.emptyList(); return CollectionsKt.emptyList();
} }
return classOrObject.getContextReceivers().stream() List<KtContextReceiver> contextReceivers = classOrObject.getContextReceivers();
List<ReceiverParameterDescriptor> contextReceiverDescriptors = contextReceivers.stream()
.map(KtContextReceiver::typeReference) .map(KtContextReceiver::typeReference)
.filter(Objects::nonNull) .filter(Objects::nonNull)
.map(typeReference -> { .map(typeReference -> {
@@ -316,7 +318,19 @@ public class LazyClassDescriptor extends ClassDescriptorBase implements ClassDes
kotlinType, kotlinType,
Annotations.Companion.getEMPTY() Annotations.Companion.getEMPTY()
); );
}).collect(Collectors.toList()); }).collect(Collectors.toList());
if (c.getLanguageVersionSettings().supportsFeature(LanguageFeature.ContextReceivers)) {
HashMultimap<String, ReceiverParameterDescriptor> labelNameToReceiverMap = HashMultimap.create();
for (int i = 0; i < contextReceivers.size(); i++) {
labelNameToReceiverMap.put(contextReceivers.get(i).name(), contextReceiverDescriptors.get(i));
}
c.getTrace().record(BindingContext.DESCRIPTOR_TO_CONTEXT_RECEIVER_MAP, this, labelNameToReceiverMap);
}
return contextReceiverDescriptors;
}); });
} }
@@ -404,7 +404,9 @@ public class BasicExpressionTypingVisitor extends ExpressionTypingVisitor {
context.trace.report(NO_THIS.on(expression)); context.trace.report(NO_THIS.on(expression));
break; break;
case SUCCESS: case SUCCESS:
result = resolutionResult.getReceiverParameterDescriptor().getType(); ReceiverParameterDescriptor descriptor = resolutionResult.getReceiverParameterDescriptor();
context.trace.record(THIS_REFERENCE_TARGET, expression.getInstanceReference(), descriptor);
result = descriptor.getType();
context.trace.recordType(expression.getInstanceReference(), result); context.trace.recordType(expression.getInstanceReference(), result);
break; break;
} }
@@ -25,8 +25,11 @@ import org.jetbrains.kotlin.name.Name
import org.jetbrains.kotlin.psi.* import org.jetbrains.kotlin.psi.*
import org.jetbrains.kotlin.psi.psiUtil.checkReservedYield import org.jetbrains.kotlin.psi.psiUtil.checkReservedYield
import org.jetbrains.kotlin.psi.psiUtil.parents import org.jetbrains.kotlin.psi.psiUtil.parents
import org.jetbrains.kotlin.resolve.*
import org.jetbrains.kotlin.resolve.BindingContext.* import org.jetbrains.kotlin.resolve.BindingContext.*
import org.jetbrains.kotlin.resolve.BindingContextUtils
import org.jetbrains.kotlin.resolve.BindingTrace
import org.jetbrains.kotlin.resolve.DescriptorResolver
import org.jetbrains.kotlin.resolve.DescriptorToSourceUtils
import org.jetbrains.kotlin.resolve.calls.context.ResolutionContext import org.jetbrains.kotlin.resolve.calls.context.ResolutionContext
import org.jetbrains.kotlin.resolve.scopes.utils.getDeclarationsByLabel import org.jetbrains.kotlin.resolve.scopes.utils.getDeclarationsByLabel
import org.jetbrains.kotlin.utils.addIfNotNull import org.jetbrains.kotlin.utils.addIfNotNull
@@ -64,6 +67,12 @@ object LabelResolver {
is KtFunctionLiteral -> return getLabelNamesIfAny(element.parent, false) is KtFunctionLiteral -> return getLabelNamesIfAny(element.parent, false)
is KtLambdaExpression -> result.addIfNotNull(getLabelForFunctionalExpression(element)) is KtLambdaExpression -> result.addIfNotNull(getLabelForFunctionalExpression(element))
} }
if (element is KtClass) {
element.contextReceivers
.mapNotNullTo(result) { it.name()?.let { s -> Name.identifier(s) } }
}
val functionOrProperty = when (element) { val functionOrProperty = when (element) {
is KtNamedFunction -> { is KtNamedFunction -> {
result.addIfNotNull(element.nameAsName ?: getLabelForFunctionalExpression(element)) result.addIfNotNull(element.nameAsName ?: getLabelForFunctionalExpression(element))
@@ -210,14 +219,15 @@ object LabelResolver {
trace.record(LABEL_TARGET, targetLabelExpression, it) trace.record(LABEL_TARGET, targetLabelExpression, it)
} }
val declarationDescriptor = trace.bindingContext[DECLARATION_TO_DESCRIPTOR, element] val declarationDescriptor = trace.bindingContext[DECLARATION_TO_DESCRIPTOR, element]
if (declarationDescriptor is FunctionDescriptor) { if (declarationDescriptor is FunctionDescriptor || declarationDescriptor is ClassDescriptor) {
val labelNameToReceiverMap = trace.bindingContext[ val labelNameToReceiverMap = trace.bindingContext[
DESCRIPTOR_TO_CONTEXT_RECEIVER_MAP, DESCRIPTOR_TO_CONTEXT_RECEIVER_MAP,
if (declarationDescriptor is PropertyAccessorDescriptor) declarationDescriptor.correspondingProperty else declarationDescriptor if (declarationDescriptor is PropertyAccessorDescriptor) declarationDescriptor.correspondingProperty else declarationDescriptor
] ]
val thisReceivers = labelNameToReceiverMap?.get(labelName.identifier) val thisReceivers = labelNameToReceiverMap?.get(labelName.identifier)
val thisReceiver = when { val thisReceiver = when {
thisReceivers.isNullOrEmpty() -> declarationDescriptor.extensionReceiverParameter thisReceivers.isNullOrEmpty() ->
(declarationDescriptor as? FunctionDescriptor)?.extensionReceiverParameter
thisReceivers.size == 1 -> thisReceivers.single() thisReceivers.size == 1 -> thisReceivers.single()
else -> { else -> {
BindingContextUtils.reportAmbiguousLabel(trace, targetLabelExpression, declarationsByLabel) BindingContextUtils.reportAmbiguousLabel(trace, targetLabelExpression, declarationsByLabel)
@@ -294,4 +304,4 @@ object LabelResolver {
} }
} }
} }
} }
@@ -121,20 +121,7 @@ fun StatementGenerator.generateReceiver(defaultStartOffset: Int, defaultEndOffse
context.symbolTable.referenceValueParameter(receiverClassDescriptor.thisAsReceiverParameter) context.symbolTable.referenceValueParameter(receiverClassDescriptor.thisAsReceiverParameter)
) )
} }
is ContextClassReceiver -> { is ContextClassReceiver -> loadContextReceiver(receiver, defaultStartOffset, defaultEndOffset)
val receiverClassDescriptor = receiver.classDescriptor
val thisAsReceiverParameter = receiverClassDescriptor.thisAsReceiverParameter
val thisReceiver = IrGetValueImpl(
defaultStartOffset, defaultEndOffset,
thisAsReceiverParameter.type.toIrType(),
context.symbolTable.referenceValue(thisAsReceiverParameter)
)
IrGetFieldImpl(
defaultStartOffset, defaultEndOffset,
context.additionalDescriptorStorage.getSyntheticField(receiver).symbol,
irReceiverType, thisReceiver
)
}
is ThisClassReceiver -> is ThisClassReceiver ->
generateThisOrSuperReceiver(receiver, receiver.classDescriptor) generateThisOrSuperReceiver(receiver, receiver.classDescriptor)
is SuperCallReceiverValue -> is SuperCallReceiverValue ->
@@ -161,6 +148,26 @@ fun StatementGenerator.generateReceiver(defaultStartOffset: Int, defaultEndOffse
} }
} }
fun StatementGenerator.loadContextReceiver(
receiver: ContextClassReceiver,
defaultStartOffset: Int, defaultEndOffset: Int,
): IrGetFieldImpl {
val receiverClassDescriptor = receiver.classDescriptor
val thisAsReceiverParameter = receiverClassDescriptor.thisAsReceiverParameter
val thisReceiver = IrGetValueImpl(
defaultStartOffset, defaultEndOffset,
thisAsReceiverParameter.type.toIrType(),
context.symbolTable.referenceValue(thisAsReceiverParameter)
)
return IrGetFieldImpl(
defaultStartOffset, defaultEndOffset,
context.additionalDescriptorStorage.getSyntheticField(receiver).symbol,
receiver.type.toIrType(), thisReceiver
)
}
fun StatementGenerator.generateSingletonReference( fun StatementGenerator.generateSingletonReference(
descriptor: ClassDescriptor, descriptor: ClassDescriptor,
startOffset: Int, startOffset: Int,
@@ -17,10 +17,7 @@
package org.jetbrains.kotlin.psi2ir.generators package org.jetbrains.kotlin.psi2ir.generators
import org.jetbrains.kotlin.backend.common.BackendException import org.jetbrains.kotlin.backend.common.BackendException
import org.jetbrains.kotlin.descriptors.CallableDescriptor import org.jetbrains.kotlin.descriptors.*
import org.jetbrains.kotlin.descriptors.ClassDescriptor
import org.jetbrains.kotlin.descriptors.DeclarationDescriptor
import org.jetbrains.kotlin.descriptors.VariableDescriptorWithAccessors
import org.jetbrains.kotlin.ir.IrStatement import org.jetbrains.kotlin.ir.IrStatement
import org.jetbrains.kotlin.ir.builders.Scope import org.jetbrains.kotlin.ir.builders.Scope
import org.jetbrains.kotlin.ir.declarations.IrDeclaration import org.jetbrains.kotlin.ir.declarations.IrDeclaration
@@ -46,6 +43,7 @@ import org.jetbrains.kotlin.resolve.calls.model.VariableAsFunctionResolvedCall
import org.jetbrains.kotlin.resolve.calls.tasks.isDynamic import org.jetbrains.kotlin.resolve.calls.tasks.isDynamic
import org.jetbrains.kotlin.resolve.constants.CompileTimeConstant import org.jetbrains.kotlin.resolve.constants.CompileTimeConstant
import org.jetbrains.kotlin.resolve.constants.evaluate.ConstantExpressionEvaluator import org.jetbrains.kotlin.resolve.constants.evaluate.ConstantExpressionEvaluator
import org.jetbrains.kotlin.resolve.scopes.receivers.ContextClassReceiver
import org.jetbrains.kotlin.types.KotlinType import org.jetbrains.kotlin.types.KotlinType
import org.jetbrains.kotlin.types.expressions.ExpressionTypingUtils import org.jetbrains.kotlin.types.expressions.ExpressionTypingUtils
import org.jetbrains.kotlin.util.OperatorNameConventions import org.jetbrains.kotlin.util.OperatorNameConventions
@@ -415,23 +413,26 @@ class StatementGenerator(
override fun visitThisExpression(expression: KtThisExpression, data: Nothing?): IrExpression { override fun visitThisExpression(expression: KtThisExpression, data: Nothing?): IrExpression {
val referenceTarget = getOrFail(BindingContext.REFERENCE_TARGET, expression.instanceReference) { "No reference target for this" } val referenceTarget = getOrFail(BindingContext.REFERENCE_TARGET, expression.instanceReference) { "No reference target for this" }
val receiverParameter =
getOrFail<KtReferenceExpression, ReceiverParameterDescriptor>(
BindingContext.THIS_REFERENCE_TARGET, expression.instanceReference
) { "No reference target for this" }
val startOffset = expression.startOffsetSkippingComments val startOffset = expression.startOffsetSkippingComments
val endOffset = expression.endOffset val endOffset = expression.endOffset
return when (referenceTarget) { return when (referenceTarget) {
is ClassDescriptor -> is ClassDescriptor ->
generateThisReceiver(startOffset, endOffset, referenceTarget.thisAsReceiverParameter.type, referenceTarget) when (receiverParameter.value) {
is ContextClassReceiver -> loadContextReceiver(receiverParameter.value as ContextClassReceiver, startOffset, endOffset)
else -> generateThisReceiver(
startOffset, endOffset, referenceTarget.thisAsReceiverParameter.type, referenceTarget
)
}
is CallableDescriptor -> { is CallableDescriptor -> {
val resolvedCall = getResolvedCall(expression) val receiverType = receiverParameter.type.toIrType()
val receivers = listOfNotNull(referenceTarget.extensionReceiverParameter) + referenceTarget.contextReceiverParameters
val receiver = receivers.find {
it == resolvedCall?.candidateDescriptor
} ?: referenceTarget.extensionReceiverParameter ?: error("No receiver: $referenceTarget")
val receiverType = receiver.type.toIrType()
IrGetValueImpl( IrGetValueImpl(
startOffset, endOffset, startOffset, endOffset,
receiverType, receiverType,
context.symbolTable.referenceValueParameter(receiver) context.symbolTable.referenceValueParameter(receiverParameter)
) )
} }
@@ -513,4 +514,4 @@ abstract class StatementGeneratorExtension(val statementGenerator: StatementGene
fun KtExpression.genStmt() = statementGenerator.generateStatement(this) fun KtExpression.genStmt() = statementGenerator.generateStatement(this)
fun KotlinType.toIrType() = with(statementGenerator) { toIrType() } fun KotlinType.toIrType() = with(statementGenerator) { toIrType() }
fun translateType(kotlinType: KotlinType) = kotlinType.toIrType() fun translateType(kotlinType: KotlinType) = kotlinType.toIrType()
} }
@@ -22,11 +22,16 @@ class Foo {
fun four(dummy: Any?) = this@Int fun four(dummy: Any?) = this@Int
} }
context(Int)
class Bar {
fun five() = this@Int
}
// MODULE: main(lib) // MODULE: main(lib)
// FILE: B.kt // FILE: B.kt
fun box(): String { fun box(): String {
return with(1) { return with(1) {
if (a.one(null) + a.two + a.Foo().three + a.Foo().four(null) == 4) "OK" else "fail" if (a.one(null) + a.two + a.Foo().three + a.Foo().four(null) + a.Bar().five() == 5) "OK" else "fail"
} }
} }
@@ -15,7 +15,7 @@ class B : A() {
inner class C { inner class C {
fun g() { fun g() {
super@B.f() super@B.f()
<!DEBUG_INFO_MISSING_UNRESOLVED!>super<!><!UNRESOLVED_REFERENCE!>@Context<!>.<!DEBUG_INFO_MISSING_UNRESOLVED!>h<!>() <!SUPERCLASS_NOT_ACCESSIBLE_FROM_INTERFACE!>super@Context<!>.<!UNRESOLVED_REFERENCE!>h<!>()
} }
} }
} }
@@ -6,7 +6,7 @@ class A {
} }
context(A) class B { context(A) class B {
val prop = x + this<!UNRESOLVED_REFERENCE!>@A<!>.<!DEBUG_INFO_MISSING_UNRESOLVED!>x<!> val prop = x + this@A.x
fun f() = x + this<!UNRESOLVED_REFERENCE!>@A<!>.<!DEBUG_INFO_MISSING_UNRESOLVED!>x<!> fun f() = x + this@A.x
} }
@@ -10,9 +10,9 @@ public final class A {
context(A) public final class B { context(A) public final class B {
public constructor B() public constructor B()
public final val prop: [Error type: Not found recorded type for x + this@A.x] public final val prop: kotlin.Int = 2
public open override /*1*/ /*fake_override*/ fun equals(/*0*/ other: kotlin.Any?): kotlin.Boolean public open override /*1*/ /*fake_override*/ fun equals(/*0*/ other: kotlin.Any?): kotlin.Boolean
public final fun f(): [Error type: Return type for function cannot be resolved] public final fun f(): kotlin.Int
public open override /*1*/ /*fake_override*/ fun hashCode(): kotlin.Int public open override /*1*/ /*fake_override*/ fun hashCode(): kotlin.Int
public open override /*1*/ /*fake_override*/ fun toString(): kotlin.String public open override /*1*/ /*fake_override*/ fun toString(): kotlin.String
} }