[IR] Rewrite transformer generator to write explicit return type

#KT-57812
This commit is contained in:
Ivan Kylchik
2023-04-28 14:56:54 +02:00
committed by Space Team
parent d75d182b26
commit 961d8a4905
3 changed files with 158 additions and 140 deletions
@@ -104,35 +104,35 @@ interface IrElementTransformer<in D> : IrElementVisitor<IrElement, D> {
return declaration return declaration
} }
override fun visitValueParameter(declaration: IrValueParameter, data: D) = override fun visitValueParameter(declaration: IrValueParameter, data: D): IrStatement =
visitDeclaration(declaration, data) visitDeclaration(declaration, data)
override fun visitClass(declaration: IrClass, data: D) = visitDeclaration(declaration, override fun visitClass(declaration: IrClass, data: D): IrStatement =
data)
override fun visitAnonymousInitializer(declaration: IrAnonymousInitializer, data: D) =
visitDeclaration(declaration, data) visitDeclaration(declaration, data)
override fun visitTypeParameter(declaration: IrTypeParameter, data: D) = override fun visitAnonymousInitializer(declaration: IrAnonymousInitializer, data: D):
IrStatement = visitDeclaration(declaration, data)
override fun visitTypeParameter(declaration: IrTypeParameter, data: D): IrStatement =
visitDeclaration(declaration, data) visitDeclaration(declaration, data)
override fun visitFunction(declaration: IrFunction, data: D) = override fun visitFunction(declaration: IrFunction, data: D): IrStatement =
visitDeclaration(declaration, data) visitDeclaration(declaration, data)
override fun visitConstructor(declaration: IrConstructor, data: D) = override fun visitConstructor(declaration: IrConstructor, data: D): IrStatement =
visitFunction(declaration, data) visitFunction(declaration, data)
override fun visitEnumEntry(declaration: IrEnumEntry, data: D) = override fun visitEnumEntry(declaration: IrEnumEntry, data: D): IrStatement =
visitDeclaration(declaration, data) visitDeclaration(declaration, data)
override fun visitErrorDeclaration(declaration: IrErrorDeclaration, data: D) = override fun visitErrorDeclaration(declaration: IrErrorDeclaration, data: D):
visitDeclaration(declaration, data) IrStatement = visitDeclaration(declaration, data)
override fun visitField(declaration: IrField, data: D) = visitDeclaration(declaration, override fun visitField(declaration: IrField, data: D): IrStatement =
data) visitDeclaration(declaration, data)
override fun visitLocalDelegatedProperty(declaration: IrLocalDelegatedProperty, override fun visitLocalDelegatedProperty(declaration: IrLocalDelegatedProperty,
data: D) = visitDeclaration(declaration, data) data: D): IrStatement = visitDeclaration(declaration, data)
override fun visitModuleFragment(declaration: IrModuleFragment, data: D): override fun visitModuleFragment(declaration: IrModuleFragment, data: D):
IrModuleFragment { IrModuleFragment {
@@ -140,22 +140,22 @@ interface IrElementTransformer<in D> : IrElementVisitor<IrElement, D> {
return declaration return declaration
} }
override fun visitProperty(declaration: IrProperty, data: D) = override fun visitProperty(declaration: IrProperty, data: D): IrStatement =
visitDeclaration(declaration, data) visitDeclaration(declaration, data)
override fun visitScript(declaration: IrScript, data: D) = override fun visitScript(declaration: IrScript, data: D): IrStatement =
visitDeclaration(declaration, data) visitDeclaration(declaration, data)
override fun visitSimpleFunction(declaration: IrSimpleFunction, data: D) = override fun visitSimpleFunction(declaration: IrSimpleFunction, data: D): IrStatement =
visitFunction(declaration, data) visitFunction(declaration, data)
override fun visitTypeAlias(declaration: IrTypeAlias, data: D) = override fun visitTypeAlias(declaration: IrTypeAlias, data: D): IrStatement =
visitDeclaration(declaration, data) visitDeclaration(declaration, data)
override fun visitVariable(declaration: IrVariable, data: D) = override fun visitVariable(declaration: IrVariable, data: D): IrStatement =
visitDeclaration(declaration, data) visitDeclaration(declaration, data)
override fun visitPackageFragment(declaration: IrPackageFragment, data: D) = override fun visitPackageFragment(declaration: IrPackageFragment, data: D): IrElement =
visitElement(declaration, data) visitElement(declaration, data)
override fun visitExternalPackageFragment(declaration: IrExternalPackageFragment, override fun visitExternalPackageFragment(declaration: IrExternalPackageFragment,
@@ -179,13 +179,13 @@ interface IrElementTransformer<in D> : IrElementVisitor<IrElement, D> {
return body return body
} }
override fun visitExpressionBody(body: IrExpressionBody, data: D) = visitBody(body, override fun visitExpressionBody(body: IrExpressionBody, data: D): IrBody =
data) visitBody(body, data)
override fun visitBlockBody(body: IrBlockBody, data: D) = visitBody(body, data) override fun visitBlockBody(body: IrBlockBody, data: D): IrBody = visitBody(body, data)
override fun visitDeclarationReference(expression: IrDeclarationReference, data: D) = override fun visitDeclarationReference(expression: IrDeclarationReference, data: D):
visitExpression(expression, data) IrExpression = visitExpression(expression, data)
override fun visitMemberAccess(expression: IrMemberAccessExpression<*>, data: D): override fun visitMemberAccess(expression: IrMemberAccessExpression<*>, data: D):
IrElement = visitDeclarationReference(expression, data) IrElement = visitDeclarationReference(expression, data)
@@ -196,57 +196,60 @@ interface IrElementTransformer<in D> : IrElementVisitor<IrElement, D> {
override fun visitConstructorCall(expression: IrConstructorCall, data: D): IrElement = override fun visitConstructorCall(expression: IrConstructorCall, data: D): IrElement =
visitFunctionAccess(expression, data) visitFunctionAccess(expression, data)
override fun visitSingletonReference(expression: IrGetSingletonValue, data: D) = override fun visitSingletonReference(expression: IrGetSingletonValue, data: D):
visitDeclarationReference(expression, data) IrExpression = visitDeclarationReference(expression, data)
override fun visitGetObjectValue(expression: IrGetObjectValue, data: D) = override fun visitGetObjectValue(expression: IrGetObjectValue, data: D): IrExpression =
visitSingletonReference(expression, data) visitSingletonReference(expression, data)
override fun visitGetEnumValue(expression: IrGetEnumValue, data: D) = override fun visitGetEnumValue(expression: IrGetEnumValue, data: D): IrExpression =
visitSingletonReference(expression, data) visitSingletonReference(expression, data)
override fun visitRawFunctionReference(expression: IrRawFunctionReference, data: D) = override fun visitRawFunctionReference(expression: IrRawFunctionReference, data: D):
visitDeclarationReference(expression, data) IrExpression = visitDeclarationReference(expression, data)
override fun visitContainerExpression(expression: IrContainerExpression, data: D) = override fun visitContainerExpression(expression: IrContainerExpression, data: D):
visitExpression(expression, data) IrExpression = visitExpression(expression, data)
override fun visitBlock(expression: IrBlock, data: D) = override fun visitBlock(expression: IrBlock, data: D): IrExpression =
visitContainerExpression(expression, data) visitContainerExpression(expression, data)
override fun visitComposite(expression: IrComposite, data: D) = override fun visitComposite(expression: IrComposite, data: D): IrExpression =
visitContainerExpression(expression, data) visitContainerExpression(expression, data)
override fun visitSyntheticBody(body: IrSyntheticBody, data: D) = visitBody(body, data) override fun visitSyntheticBody(body: IrSyntheticBody, data: D): IrBody =
visitBody(body, data)
override fun visitBreakContinue(jump: IrBreakContinue, data: D) = visitExpression(jump, override fun visitBreakContinue(jump: IrBreakContinue, data: D): IrExpression =
data) visitExpression(jump, data)
override fun visitBreak(jump: IrBreak, data: D) = visitBreakContinue(jump, data) override fun visitBreak(jump: IrBreak, data: D): IrExpression =
visitBreakContinue(jump, data)
override fun visitContinue(jump: IrContinue, data: D) = visitBreakContinue(jump, data) override fun visitContinue(jump: IrContinue, data: D): IrExpression =
visitBreakContinue(jump, data)
override fun visitCall(expression: IrCall, data: D) = visitFunctionAccess(expression, override fun visitCall(expression: IrCall, data: D): IrElement =
data) visitFunctionAccess(expression, data)
override fun visitCallableReference(expression: IrCallableReference<*>, data: D) = override fun visitCallableReference(expression: IrCallableReference<*>, data: D):
visitMemberAccess(expression, data) IrElement = visitMemberAccess(expression, data)
override fun visitFunctionReference(expression: IrFunctionReference, data: D) = override fun visitFunctionReference(expression: IrFunctionReference, data: D):
visitCallableReference(expression, data) IrElement = visitCallableReference(expression, data)
override fun visitPropertyReference(expression: IrPropertyReference, data: D) = override fun visitPropertyReference(expression: IrPropertyReference, data: D):
visitCallableReference(expression, data) IrElement = visitCallableReference(expression, data)
override override
fun visitLocalDelegatedPropertyReference(expression: IrLocalDelegatedPropertyReference, fun visitLocalDelegatedPropertyReference(expression: IrLocalDelegatedPropertyReference,
data: D) = visitCallableReference(expression, data) data: D): IrElement = visitCallableReference(expression, data)
override fun visitClassReference(expression: IrClassReference, data: D) = override fun visitClassReference(expression: IrClassReference, data: D): IrExpression =
visitDeclarationReference(expression, data) visitDeclarationReference(expression, data)
override fun visitConst(expression: IrConst<*>, data: D) = visitExpression(expression, override fun visitConst(expression: IrConst<*>, data: D): IrExpression =
data) visitExpression(expression, data)
override fun visitConstantValue(expression: IrConstantValue, data: D): override fun visitConstantValue(expression: IrConstantValue, data: D):
IrConstantValue { IrConstantValue {
@@ -254,103 +257,107 @@ interface IrElementTransformer<in D> : IrElementVisitor<IrElement, D> {
return expression return expression
} }
override fun visitConstantPrimitive(expression: IrConstantPrimitive, data: D) = override fun visitConstantPrimitive(expression: IrConstantPrimitive, data: D):
visitConstantValue(expression, data) IrConstantValue = visitConstantValue(expression, data)
override fun visitConstantObject(expression: IrConstantObject, data: D) = override fun visitConstantObject(expression: IrConstantObject, data: D):
visitConstantValue(expression, data) IrConstantValue = visitConstantValue(expression, data)
override fun visitConstantArray(expression: IrConstantArray, data: D) = override fun visitConstantArray(expression: IrConstantArray, data: D): IrConstantValue
visitConstantValue(expression, data) = visitConstantValue(expression, data)
override fun visitDelegatingConstructorCall(expression: IrDelegatingConstructorCall, override fun visitDelegatingConstructorCall(expression: IrDelegatingConstructorCall,
data: D) = visitFunctionAccess(expression, data) data: D): IrElement = visitFunctionAccess(expression, data)
override fun visitDynamicExpression(expression: IrDynamicExpression, data: D) = override fun visitDynamicExpression(expression: IrDynamicExpression, data: D):
visitExpression(expression, data) IrExpression = visitExpression(expression, data)
override fun visitDynamicOperatorExpression(expression: IrDynamicOperatorExpression, override fun visitDynamicOperatorExpression(expression: IrDynamicOperatorExpression,
data: D) = visitDynamicExpression(expression, data) data: D): IrExpression = visitDynamicExpression(expression, data)
override fun visitDynamicMemberExpression(expression: IrDynamicMemberExpression, override fun visitDynamicMemberExpression(expression: IrDynamicMemberExpression,
data: D) = visitDynamicExpression(expression, data) data: D): IrExpression = visitDynamicExpression(expression, data)
override fun visitEnumConstructorCall(expression: IrEnumConstructorCall, data: D) = override fun visitEnumConstructorCall(expression: IrEnumConstructorCall, data: D):
visitFunctionAccess(expression, data) IrElement = visitFunctionAccess(expression, data)
override fun visitErrorExpression(expression: IrErrorExpression, data: D) = override fun visitErrorExpression(expression: IrErrorExpression, data: D): IrExpression
visitExpression(expression, data) = visitExpression(expression, data)
override fun visitErrorCallExpression(expression: IrErrorCallExpression, data: D) = override fun visitErrorCallExpression(expression: IrErrorCallExpression, data: D):
visitErrorExpression(expression, data) IrExpression = visitErrorExpression(expression, data)
override fun visitFieldAccess(expression: IrFieldAccessExpression, data: D) = override fun visitFieldAccess(expression: IrFieldAccessExpression, data: D):
visitDeclarationReference(expression, data) IrExpression = visitDeclarationReference(expression, data)
override fun visitGetField(expression: IrGetField, data: D) = override fun visitGetField(expression: IrGetField, data: D): IrExpression =
visitFieldAccess(expression, data) visitFieldAccess(expression, data)
override fun visitSetField(expression: IrSetField, data: D) = override fun visitSetField(expression: IrSetField, data: D): IrExpression =
visitFieldAccess(expression, data) visitFieldAccess(expression, data)
override fun visitFunctionExpression(expression: IrFunctionExpression, data: D): override fun visitFunctionExpression(expression: IrFunctionExpression, data: D):
IrElement = visitExpression(expression, data) IrElement = visitExpression(expression, data)
override fun visitGetClass(expression: IrGetClass, data: D) = override fun visitGetClass(expression: IrGetClass, data: D): IrExpression =
visitExpression(expression, data) visitExpression(expression, data)
override fun visitInstanceInitializerCall(expression: IrInstanceInitializerCall, override fun visitInstanceInitializerCall(expression: IrInstanceInitializerCall,
data: D) = visitExpression(expression, data) data: D): IrExpression = visitExpression(expression, data)
override fun visitLoop(loop: IrLoop, data: D) = visitExpression(loop, data) override fun visitLoop(loop: IrLoop, data: D): IrExpression = visitExpression(loop,
override fun visitWhileLoop(loop: IrWhileLoop, data: D) = visitLoop(loop, data)
override fun visitDoWhileLoop(loop: IrDoWhileLoop, data: D) = visitLoop(loop, data)
override fun visitReturn(expression: IrReturn, data: D) = visitExpression(expression,
data) data)
override fun visitStringConcatenation(expression: IrStringConcatenation, data: D) = override fun visitWhileLoop(loop: IrWhileLoop, data: D): IrExpression = visitLoop(loop,
visitExpression(expression, data)
override fun visitSuspensionPoint(expression: IrSuspensionPoint, data: D) =
visitExpression(expression, data)
override fun visitSuspendableExpression(expression: IrSuspendableExpression, data: D) =
visitExpression(expression, data)
override fun visitThrow(expression: IrThrow, data: D) = visitExpression(expression,
data) data)
override fun visitTry(aTry: IrTry, data: D) = visitExpression(aTry, data) override fun visitDoWhileLoop(loop: IrDoWhileLoop, data: D): IrExpression =
visitLoop(loop, data)
override fun visitReturn(expression: IrReturn, data: D): IrExpression =
visitExpression(expression, data)
override fun visitStringConcatenation(expression: IrStringConcatenation, data: D):
IrExpression = visitExpression(expression, data)
override fun visitSuspensionPoint(expression: IrSuspensionPoint, data: D): IrExpression
= visitExpression(expression, data)
override fun visitSuspendableExpression(expression: IrSuspendableExpression, data: D):
IrExpression = visitExpression(expression, data)
override fun visitThrow(expression: IrThrow, data: D): IrExpression =
visitExpression(expression, data)
override fun visitTry(aTry: IrTry, data: D): IrExpression = visitExpression(aTry, data)
override fun visitCatch(aCatch: IrCatch, data: D): IrCatch { override fun visitCatch(aCatch: IrCatch, data: D): IrCatch {
aCatch.transformChildren(this, data) aCatch.transformChildren(this, data)
return aCatch return aCatch
} }
override fun visitTypeOperator(expression: IrTypeOperatorCall, data: D) = override fun visitTypeOperator(expression: IrTypeOperatorCall, data: D): IrExpression =
visitExpression(expression, data) visitExpression(expression, data)
override fun visitValueAccess(expression: IrValueAccessExpression, data: D) = override fun visitValueAccess(expression: IrValueAccessExpression, data: D):
visitDeclarationReference(expression, data) IrExpression = visitDeclarationReference(expression, data)
override fun visitGetValue(expression: IrGetValue, data: D) = override fun visitGetValue(expression: IrGetValue, data: D): IrExpression =
visitValueAccess(expression, data) visitValueAccess(expression, data)
override fun visitSetValue(expression: IrSetValue, data: D) = override fun visitSetValue(expression: IrSetValue, data: D): IrExpression =
visitValueAccess(expression, data) visitValueAccess(expression, data)
override fun visitVararg(expression: IrVararg, data: D) = visitExpression(expression, override fun visitVararg(expression: IrVararg, data: D): IrExpression =
data) visitExpression(expression, data)
override fun visitSpreadElement(spread: IrSpreadElement, data: D): IrSpreadElement { override fun visitSpreadElement(spread: IrSpreadElement, data: D): IrSpreadElement {
spread.transformChildren(this, data) spread.transformChildren(this, data)
return spread return spread
} }
override fun visitWhen(expression: IrWhen, data: D) = visitExpression(expression, data) override fun visitWhen(expression: IrWhen, data: D): IrExpression =
visitExpression(expression, data)
override fun visitBranch(branch: IrBranch, data: D): IrBranch { override fun visitBranch(branch: IrBranch, data: D): IrBranch {
branch.transformChildren(this, data) branch.transformChildren(this, data)
@@ -9,6 +9,7 @@
package org.jetbrains.kotlin.ir.visitors package org.jetbrains.kotlin.ir.visitors
import org.jetbrains.kotlin.ir.IrElement import org.jetbrains.kotlin.ir.IrElement
import org.jetbrains.kotlin.ir.IrStatement
import org.jetbrains.kotlin.ir.declarations.IrClass import org.jetbrains.kotlin.ir.declarations.IrClass
import org.jetbrains.kotlin.ir.declarations.IrField import org.jetbrains.kotlin.ir.declarations.IrField
import org.jetbrains.kotlin.ir.declarations.IrFunction import org.jetbrains.kotlin.ir.declarations.IrFunction
@@ -20,6 +21,7 @@ import org.jetbrains.kotlin.ir.declarations.IrValueParameter
import org.jetbrains.kotlin.ir.declarations.IrVariable import org.jetbrains.kotlin.ir.declarations.IrVariable
import org.jetbrains.kotlin.ir.expressions.IrClassReference import org.jetbrains.kotlin.ir.expressions.IrClassReference
import org.jetbrains.kotlin.ir.expressions.IrConstantObject import org.jetbrains.kotlin.ir.expressions.IrConstantObject
import org.jetbrains.kotlin.ir.expressions.IrConstantValue
import org.jetbrains.kotlin.ir.expressions.IrExpression import org.jetbrains.kotlin.ir.expressions.IrExpression
import org.jetbrains.kotlin.ir.expressions.IrMemberAccessExpression import org.jetbrains.kotlin.ir.expressions.IrMemberAccessExpression
import org.jetbrains.kotlin.ir.expressions.IrTypeOperatorCall import org.jetbrains.kotlin.ir.expressions.IrTypeOperatorCall
@@ -33,92 +35,93 @@ interface IrTypeTransformerVoid<in D> : IrElementTransformer<D> {
data: D, data: D,
): Type ): Type
override fun visitValueParameter(declaration: IrValueParameter, data: D) = run { override fun visitValueParameter(declaration: IrValueParameter, data: D): IrStatement {
declaration.varargElementType = transformType(declaration, declaration.varargElementType, declaration.varargElementType = transformType(declaration, declaration.varargElementType,
data) data)
declaration.type = transformType(declaration, declaration.type, data) declaration.type = transformType(declaration, declaration.type, data)
return@run super.visitValueParameter(declaration, data) return super.visitValueParameter(declaration, data)
} }
override fun visitClass(declaration: IrClass, data: D) = run { override fun visitClass(declaration: IrClass, data: D): IrStatement {
declaration.valueClassRepresentation?.mapUnderlyingType { declaration.valueClassRepresentation?.mapUnderlyingType {
transformType(declaration, it, data) transformType(declaration, it, data)
} }
declaration.superTypes = declaration.superTypes.map { transformType(declaration, it, data) } declaration.superTypes = declaration.superTypes.map { transformType(declaration, it, data) }
return@run super.visitClass(declaration, data) return super.visitClass(declaration, data)
} }
override fun visitTypeParameter(declaration: IrTypeParameter, data: D) = run { override fun visitTypeParameter(declaration: IrTypeParameter, data: D): IrStatement {
declaration.superTypes = declaration.superTypes.map { transformType(declaration, it, data) } declaration.superTypes = declaration.superTypes.map { transformType(declaration, it, data) }
return@run super.visitTypeParameter(declaration, data) return super.visitTypeParameter(declaration, data)
} }
override fun visitFunction(declaration: IrFunction, data: D) = run { override fun visitFunction(declaration: IrFunction, data: D): IrStatement {
declaration.returnType = transformType(declaration, declaration.returnType, data) declaration.returnType = transformType(declaration, declaration.returnType, data)
return@run super.visitFunction(declaration, data) return super.visitFunction(declaration, data)
} }
override fun visitField(declaration: IrField, data: D) = run { override fun visitField(declaration: IrField, data: D): IrStatement {
declaration.type = transformType(declaration, declaration.type, data) declaration.type = transformType(declaration, declaration.type, data)
return@run super.visitField(declaration, data) return super.visitField(declaration, data)
} }
override fun visitLocalDelegatedProperty(declaration: IrLocalDelegatedProperty, override fun visitLocalDelegatedProperty(declaration: IrLocalDelegatedProperty,
data: D) = run { data: D): IrStatement {
declaration.type = transformType(declaration, declaration.type, data) declaration.type = transformType(declaration, declaration.type, data)
return@run super.visitLocalDelegatedProperty(declaration, data) return super.visitLocalDelegatedProperty(declaration, data)
} }
override fun visitScript(declaration: IrScript, data: D) = run { override fun visitScript(declaration: IrScript, data: D): IrStatement {
declaration.baseClass = transformType(declaration, declaration.baseClass, data) declaration.baseClass = transformType(declaration, declaration.baseClass, data)
return@run super.visitScript(declaration, data) return super.visitScript(declaration, data)
} }
override fun visitTypeAlias(declaration: IrTypeAlias, data: D) = run { override fun visitTypeAlias(declaration: IrTypeAlias, data: D): IrStatement {
declaration.expandedType = transformType(declaration, declaration.expandedType, data) declaration.expandedType = transformType(declaration, declaration.expandedType, data)
return@run super.visitTypeAlias(declaration, data) return super.visitTypeAlias(declaration, data)
} }
override fun visitVariable(declaration: IrVariable, data: D) = run { override fun visitVariable(declaration: IrVariable, data: D): IrStatement {
declaration.type = transformType(declaration, declaration.type, data) declaration.type = transformType(declaration, declaration.type, data)
return@run super.visitVariable(declaration, data) return super.visitVariable(declaration, data)
} }
override fun visitExpression(expression: IrExpression, data: D) = run { override fun visitExpression(expression: IrExpression, data: D): IrExpression {
expression.type = transformType(expression, expression.type, data) expression.type = transformType(expression, expression.type, data)
return@run super.visitExpression(expression, data) return super.visitExpression(expression, data)
} }
override fun visitMemberAccess(expression: IrMemberAccessExpression<*>, data: D) = override fun visitMemberAccess(expression: IrMemberAccessExpression<*>, data: D):
run { IrElement {
(0 until expression.typeArgumentsCount).forEach { (0 until expression.typeArgumentsCount).forEach {
expression.getTypeArgument(it)?.let { type -> expression.getTypeArgument(it)?.let { type ->
expression.putTypeArgument(it, transformType(expression, type, data)) expression.putTypeArgument(it, transformType(expression, type, data))
} }
} }
return@run super.visitMemberAccess(expression, data) return super.visitMemberAccess(expression, data)
} }
override fun visitClassReference(expression: IrClassReference, data: D) = run { override fun visitClassReference(expression: IrClassReference, data: D): IrExpression {
expression.classType = transformType(expression, expression.classType, data) expression.classType = transformType(expression, expression.classType, data)
return@run super.visitClassReference(expression, data) return super.visitClassReference(expression, data)
} }
override fun visitConstantObject(expression: IrConstantObject, data: D) = run { override fun visitConstantObject(expression: IrConstantObject, data: D):
IrConstantValue {
for (i in 0 until expression.typeArguments.size) { for (i in 0 until expression.typeArguments.size) {
expression.typeArguments[i] = transformType(expression, expression.typeArguments[i], expression.typeArguments[i] = transformType(expression, expression.typeArguments[i],
data) data)
} }
return@run super.visitConstantObject(expression, data) return super.visitConstantObject(expression, data)
} }
override fun visitTypeOperator(expression: IrTypeOperatorCall, data: D) = run { override fun visitTypeOperator(expression: IrTypeOperatorCall, data: D): IrExpression {
expression.typeOperand = transformType(expression, expression.typeOperand, data) expression.typeOperand = transformType(expression, expression.typeOperand, data)
return@run super.visitTypeOperator(expression, data) return super.visitTypeOperator(expression, data)
} }
override fun visitVararg(expression: IrVararg, data: D) = run { override fun visitVararg(expression: IrVararg, data: D): IrExpression {
expression.varargElementType = transformType(expression, expression.varargElementType, data) expression.varargElementType = transformType(expression, expression.varargElementType, data)
return@run super.visitVararg(expression, data) return super.visitVararg(expression, data)
} }
} }
@@ -93,19 +93,18 @@ fun printTransformer(generationPath: File, model: Model): GeneratedFile {
} }
for (element in model.elements) { for (element in model.elements) {
val returnType = element.getTransformExplicitType()
if (element.transformByChildren) { if (element.transformByChildren) {
addFunction(buildVisitFun(element).apply { addFunction(buildVisitFun(element).apply {
addStatement("${element.visitorParam}.transformChildren(this, data)") addStatement("${element.visitorParam}.transformChildren(this, data)")
addStatement("return ${element.visitorParam}") addStatement("return ${element.visitorParam}")
returns((element.transformerReturnType ?: element).toPoetStarParameterized()) returns(returnType.toPoetStarParameterized())
}.build()) }.build())
} else { } else {
element.visitorParent?.let { parent -> element.visitorParent?.let { parent ->
addFunction(buildVisitFun(element).apply { addFunction(buildVisitFun(element).apply {
addStatement("return ${parent.element.visitFunName}(${element.visitorParam}, data)") addStatement("return ${parent.element.visitFunName}(${element.visitorParam}, data)")
element.transformerReturnType?.let { returns(returnType.toPoetStarParameterized())
returns(it.toPoetStarParameterized())
}
}.build()) }.build())
} }
} }
@@ -179,10 +178,10 @@ fun printTypeVisitor(generationPath: File, model: Model): GeneratedFile {
val irTypeFields = element.getFieldsWithIrTypeType() val irTypeFields = element.getFieldsWithIrTypeType()
if (irTypeFields.isEmpty()) continue if (irTypeFields.isEmpty()) continue
val returnType = element.getTransformExplicitType()
element.visitorParent?.let { _ -> element.visitorParent?.let { _ ->
addFunction(buildVisitFun(element).apply { addFunction(buildVisitFun(element).apply {
// Note: using `run` here to infer return type automatically returns(returnType.toPoetStarParameterized())
beginControlFlow("return run")
val visitorParam = element.visitorParam val visitorParam = element.visitorParam
when (element.name) { when (element.name) {
@@ -201,8 +200,7 @@ fun printTypeVisitor(generationPath: File, model: Model): GeneratedFile {
} }
else -> irTypeFields.forEach { addVisitTypeStatement(element, it) } else -> irTypeFields.forEach { addVisitTypeStatement(element, it) }
} }
addStatement("return@run super.${element.visitFunName}($visitorParam, data)") addStatement("return super.${element.visitFunName}($visitorParam, data)")
endControlFlow()
}.build()) }.build())
} }
} }
@@ -210,3 +208,13 @@ fun printTypeVisitor(generationPath: File, model: Model): GeneratedFile {
return printTypeCommon(generationPath, typeTransformerVoidTypeName.packageName, visitorType) return printTypeCommon(generationPath, typeTransformerVoidTypeName.packageName, visitorType)
} }
private fun Element.getTransformExplicitType(): Element {
return generateSequence(this) { it.visitorParent?.element }
.firstNotNullOfOrNull {
when {
it.transformByChildren -> it.transformerReturnType ?: it
else -> it.transformerReturnType
}
} ?: this
}