diff --git a/fleet/compiler-plugins/rpc/.gitignore b/fleet/compiler-plugins/rpc/.gitignore new file mode 100644 index 000000000000..01d8fb504537 --- /dev/null +++ b/fleet/compiler-plugins/rpc/.gitignore @@ -0,0 +1,4 @@ +out/ +build/ +.idea/ +.gradle/ \ No newline at end of file diff --git a/fleet/compiler-plugins/rpc/build.gradle.kts b/fleet/compiler-plugins/rpc/build.gradle.kts new file mode 100644 index 000000000000..e9eca7ff1d60 --- /dev/null +++ b/fleet/compiler-plugins/rpc/build.gradle.kts @@ -0,0 +1,89 @@ +plugins { + kotlin("jvm") + `maven-publish` + `java-gradle-plugin` +} + +val pluginVersion = run { + val pluginVersionOverride = System.getProperty("fleet.plugin.version.override") + if (pluginVersionOverride != null) { + return@run pluginVersionOverride + } + val versionRegex = Regex("const val VERSION = \"(.+)\"") + project.file("src/main/kotlin/com/jetbrains/fleet/rpc/plugin/gradle/RpcGradlePlugin.kt").readLines().firstNotNullOfOrNull { line -> + versionRegex.find(line)?.let { match -> + match.groupValues[1] + } + } ?: error("Cannot find `const val VERSION = \"...\"` in RpcGradlePlugin.kt") +} + +// the compiler plugin will be used together with this Kotlin compiler +val KOTLIN_VERSION = "2.2.0" + +// the compiler plugin will be built with these Kotlin LV/APIV +val KOTLIN_LANGUAGE_VERSION = "2.1" +val KOTLIN_API_VERSION = "2.1" + +group = "com.jetbrains.fleet" +version = pluginVersion + +repositories { + mavenCentral() + if ("SNAPSHOT" in KOTLIN_VERSION || "dev" in KOTLIN_VERSION) { + mavenLocal() + } +} + +dependencies { + compileOnly("org.jetbrains.kotlin:kotlin-compiler-embeddable:$KOTLIN_VERSION") + implementation("org.jetbrains.kotlin:kotlin-gradle-plugin-api:$KOTLIN_VERSION") + implementation("org.jetbrains:annotations:24.1.0") +} + +java { + withSourcesJar() + targetCompatibility = JavaVersion.VERSION_17 +} + +kotlin { + compilerOptions { + jvmTarget.set(org.jetbrains.kotlin.gradle.dsl.JvmTarget.JVM_17) + apiVersion.set(org.jetbrains.kotlin.gradle.dsl.KotlinVersion.fromVersion(KOTLIN_LANGUAGE_VERSION)) + languageVersion.set(org.jetbrains.kotlin.gradle.dsl.KotlinVersion.fromVersion(KOTLIN_API_VERSION)) + // FREE COMPILER ARGS PLACEHOLDER PLEASE DO NOT CHANGE IT BEFORE DISCUSS IN #ij-monorepo-kotlin + } +} + +gradlePlugin { + plugins { + create("rpcGradle") { + id = "rpc" + implementationClass = "com.jetbrains.fleet.rpc.plugin.gradle.RpcGradlePlugin" + displayName = "RPC compiler plugin" + description = "Use this plugin for modules, which use RPC" + } + } +} + +publishing { + repositories { + maven { + url = uri("https://packages.jetbrains.team/maven/p/ij/intellij-dependencies") + credentials { + username = prop("publishingUsername") + password = prop("publishingPassword") + } + } + } + publications { + create("rpcPlugin") { + groupId = "com.jetbrains.fleet" + artifactId = "rpc-compiler-plugin" + version = pluginVersion + from(components["java"]) + } + } +} + +fun prop(name: String): String? = + extra.properties[name] as? String diff --git a/fleet/compiler-plugins/rpc/gradle/wrapper/gradle-wrapper.jar b/fleet/compiler-plugins/rpc/gradle/wrapper/gradle-wrapper.jar new file mode 100644 index 000000000000..9bbc975c742b Binary files /dev/null and b/fleet/compiler-plugins/rpc/gradle/wrapper/gradle-wrapper.jar differ diff --git a/fleet/compiler-plugins/rpc/gradle/wrapper/gradle-wrapper.properties b/fleet/compiler-plugins/rpc/gradle/wrapper/gradle-wrapper.properties new file mode 100644 index 000000000000..37f853b1c84d --- /dev/null +++ b/fleet/compiler-plugins/rpc/gradle/wrapper/gradle-wrapper.properties @@ -0,0 +1,7 @@ +distributionBase=GRADLE_USER_HOME +distributionPath=wrapper/dists +distributionUrl=https\://services.gradle.org/distributions/gradle-8.13-bin.zip +networkTimeout=10000 +validateDistributionUrl=true +zipStoreBase=GRADLE_USER_HOME +zipStorePath=wrapper/dists diff --git a/fleet/compiler-plugins/rpc/gradlew b/fleet/compiler-plugins/rpc/gradlew new file mode 100755 index 000000000000..faf93008b77e --- /dev/null +++ b/fleet/compiler-plugins/rpc/gradlew @@ -0,0 +1,251 @@ +#!/bin/sh + +# +# Copyright © 2015-2021 the original authors. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# SPDX-License-Identifier: Apache-2.0 +# + +############################################################################## +# +# Gradle start up script for POSIX generated by Gradle. +# +# Important for running: +# +# (1) You need a POSIX-compliant shell to run this script. If your /bin/sh is +# noncompliant, but you have some other compliant shell such as ksh or +# bash, then to run this script, type that shell name before the whole +# command line, like: +# +# ksh Gradle +# +# Busybox and similar reduced shells will NOT work, because this script +# requires all of these POSIX shell features: +# * functions; +# * expansions «$var», «${var}», «${var:-default}», «${var+SET}», +# «${var#prefix}», «${var%suffix}», and «$( cmd )»; +# * compound commands having a testable exit status, especially «case»; +# * various built-in commands including «command», «set», and «ulimit». +# +# Important for patching: +# +# (2) This script targets any POSIX shell, so it avoids extensions provided +# by Bash, Ksh, etc; in particular arrays are avoided. +# +# The "traditional" practice of packing multiple parameters into a +# space-separated string is a well documented source of bugs and security +# problems, so this is (mostly) avoided, by progressively accumulating +# options in "$@", and eventually passing that to Java. +# +# Where the inherited environment variables (DEFAULT_JVM_OPTS, JAVA_OPTS, +# and GRADLE_OPTS) rely on word-splitting, this is performed explicitly; +# see the in-line comments for details. +# +# There are tweaks for specific operating systems such as AIX, CygWin, +# Darwin, MinGW, and NonStop. +# +# (3) This script is generated from the Groovy template +# https://github.com/gradle/gradle/blob/HEAD/platforms/jvm/plugins-application/src/main/resources/org/gradle/api/internal/plugins/unixStartScript.txt +# within the Gradle project. +# +# You can find Gradle at https://github.com/gradle/gradle/. +# +############################################################################## + +# Attempt to set APP_HOME + +# Resolve links: $0 may be a link +app_path=$0 + +# Need this for daisy-chained symlinks. +while + APP_HOME=${app_path%"${app_path##*/}"} # leaves a trailing /; empty if no leading path + [ -h "$app_path" ] +do + ls=$( ls -ld "$app_path" ) + link=${ls#*' -> '} + case $link in #( + /*) app_path=$link ;; #( + *) app_path=$APP_HOME$link ;; + esac +done + +# This is normally unused +# shellcheck disable=SC2034 +APP_BASE_NAME=${0##*/} +# Discard cd standard output in case $CDPATH is set (https://github.com/gradle/gradle/issues/25036) +APP_HOME=$( cd -P "${APP_HOME:-./}" > /dev/null && printf '%s\n' "$PWD" ) || exit + +# Use the maximum available, or set MAX_FD != -1 to use that value. +MAX_FD=maximum + +warn () { + echo "$*" +} >&2 + +die () { + echo + echo "$*" + echo + exit 1 +} >&2 + +# OS specific support (must be 'true' or 'false'). +cygwin=false +msys=false +darwin=false +nonstop=false +case "$( uname )" in #( + CYGWIN* ) cygwin=true ;; #( + Darwin* ) darwin=true ;; #( + MSYS* | MINGW* ) msys=true ;; #( + NONSTOP* ) nonstop=true ;; +esac + +CLASSPATH=$APP_HOME/gradle/wrapper/gradle-wrapper.jar + + +# Determine the Java command to use to start the JVM. +if [ -n "$JAVA_HOME" ] ; then + if [ -x "$JAVA_HOME/jre/sh/java" ] ; then + # IBM's JDK on AIX uses strange locations for the executables + JAVACMD=$JAVA_HOME/jre/sh/java + else + JAVACMD=$JAVA_HOME/bin/java + fi + if [ ! -x "$JAVACMD" ] ; then + die "ERROR: JAVA_HOME is set to an invalid directory: $JAVA_HOME + +Please set the JAVA_HOME variable in your environment to match the +location of your Java installation." + fi +else + JAVACMD=java + if ! command -v java >/dev/null 2>&1 + then + die "ERROR: JAVA_HOME is not set and no 'java' command could be found in your PATH. + +Please set the JAVA_HOME variable in your environment to match the +location of your Java installation." + fi +fi + +# Increase the maximum file descriptors if we can. +if ! "$cygwin" && ! "$darwin" && ! "$nonstop" ; then + case $MAX_FD in #( + max*) + # In POSIX sh, ulimit -H is undefined. That's why the result is checked to see if it worked. + # shellcheck disable=SC2039,SC3045 + MAX_FD=$( ulimit -H -n ) || + warn "Could not query maximum file descriptor limit" + esac + case $MAX_FD in #( + '' | soft) :;; #( + *) + # In POSIX sh, ulimit -n is undefined. That's why the result is checked to see if it worked. + # shellcheck disable=SC2039,SC3045 + ulimit -n "$MAX_FD" || + warn "Could not set maximum file descriptor limit to $MAX_FD" + esac +fi + +# Collect all arguments for the java command, stacking in reverse order: +# * args from the command line +# * the main class name +# * -classpath +# * -D...appname settings +# * --module-path (only if needed) +# * DEFAULT_JVM_OPTS, JAVA_OPTS, and GRADLE_OPTS environment variables. + +# For Cygwin or MSYS, switch paths to Windows format before running java +if "$cygwin" || "$msys" ; then + APP_HOME=$( cygpath --path --mixed "$APP_HOME" ) + CLASSPATH=$( cygpath --path --mixed "$CLASSPATH" ) + + JAVACMD=$( cygpath --unix "$JAVACMD" ) + + # Now convert the arguments - kludge to limit ourselves to /bin/sh + for arg do + if + case $arg in #( + -*) false ;; # don't mess with options #( + /?*) t=${arg#/} t=/${t%%/*} # looks like a POSIX filepath + [ -e "$t" ] ;; #( + *) false ;; + esac + then + arg=$( cygpath --path --ignore --mixed "$arg" ) + fi + # Roll the args list around exactly as many times as the number of + # args, so each arg winds up back in the position where it started, but + # possibly modified. + # + # NB: a `for` loop captures its iteration list before it begins, so + # changing the positional parameters here affects neither the number of + # iterations, nor the values presented in `arg`. + shift # remove old arg + set -- "$@" "$arg" # push replacement arg + done +fi + + +# Add default JVM options here. You can also use JAVA_OPTS and GRADLE_OPTS to pass JVM options to this script. +DEFAULT_JVM_OPTS='"-Xmx64m" "-Xms64m"' + +# Collect all arguments for the java command: +# * DEFAULT_JVM_OPTS, JAVA_OPTS, and optsEnvironmentVar are not allowed to contain shell fragments, +# and any embedded shellness will be escaped. +# * For example: A user cannot expect ${Hostname} to be expanded, as it is an environment variable and will be +# treated as '${Hostname}' itself on the command line. + +set -- \ + "-Dorg.gradle.appname=$APP_BASE_NAME" \ + -classpath "$CLASSPATH" \ + org.gradle.wrapper.GradleWrapperMain \ + "$@" + +# Stop when "xargs" is not available. +if ! command -v xargs >/dev/null 2>&1 +then + die "xargs is not available" +fi + +# Use "xargs" to parse quoted args. +# +# With -n1 it outputs one arg per line, with the quotes and backslashes removed. +# +# In Bash we could simply go: +# +# readarray ARGS < <( xargs -n1 <<<"$var" ) && +# set -- "${ARGS[@]}" "$@" +# +# but POSIX shell has neither arrays nor command substitution, so instead we +# post-process each arg (as a line of input to sed) to backslash-escape any +# character that might be a shell metacharacter, then use eval to reverse +# that process (while maintaining the separation between arguments), and wrap +# the whole thing up as a single "set" statement. +# +# This will of course break if any of these variables contains a newline or +# an unmatched quote. +# + +eval "set -- $( + printf '%s\n' "$DEFAULT_JVM_OPTS $JAVA_OPTS $GRADLE_OPTS" | + xargs -n1 | + sed ' s~[^-[:alnum:]+,./:=@_]~\\&~g; ' | + tr '\n' ' ' + )" '"$@"' + +exec "$JAVACMD" "$@" diff --git a/fleet/compiler-plugins/rpc/gradlew.bat b/fleet/compiler-plugins/rpc/gradlew.bat new file mode 100644 index 000000000000..9d21a21834d5 --- /dev/null +++ b/fleet/compiler-plugins/rpc/gradlew.bat @@ -0,0 +1,94 @@ +@rem +@rem Copyright 2015 the original author or authors. +@rem +@rem Licensed under the Apache License, Version 2.0 (the "License"); +@rem you may not use this file except in compliance with the License. +@rem You may obtain a copy of the License at +@rem +@rem https://www.apache.org/licenses/LICENSE-2.0 +@rem +@rem Unless required by applicable law or agreed to in writing, software +@rem distributed under the License is distributed on an "AS IS" BASIS, +@rem WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +@rem See the License for the specific language governing permissions and +@rem limitations under the License. +@rem +@rem SPDX-License-Identifier: Apache-2.0 +@rem + +@if "%DEBUG%"=="" @echo off +@rem ########################################################################## +@rem +@rem Gradle startup script for Windows +@rem +@rem ########################################################################## + +@rem Set local scope for the variables with windows NT shell +if "%OS%"=="Windows_NT" setlocal + +set DIRNAME=%~dp0 +if "%DIRNAME%"=="" set DIRNAME=. +@rem This is normally unused +set APP_BASE_NAME=%~n0 +set APP_HOME=%DIRNAME% + +@rem Resolve any "." and ".." in APP_HOME to make it shorter. +for %%i in ("%APP_HOME%") do set APP_HOME=%%~fi + +@rem Add default JVM options here. You can also use JAVA_OPTS and GRADLE_OPTS to pass JVM options to this script. +set DEFAULT_JVM_OPTS="-Xmx64m" "-Xms64m" + +@rem Find java.exe +if defined JAVA_HOME goto findJavaFromJavaHome + +set JAVA_EXE=java.exe +%JAVA_EXE% -version >NUL 2>&1 +if %ERRORLEVEL% equ 0 goto execute + +echo. 1>&2 +echo ERROR: JAVA_HOME is not set and no 'java' command could be found in your PATH. 1>&2 +echo. 1>&2 +echo Please set the JAVA_HOME variable in your environment to match the 1>&2 +echo location of your Java installation. 1>&2 + +goto fail + +:findJavaFromJavaHome +set JAVA_HOME=%JAVA_HOME:"=% +set JAVA_EXE=%JAVA_HOME%/bin/java.exe + +if exist "%JAVA_EXE%" goto execute + +echo. 1>&2 +echo ERROR: JAVA_HOME is set to an invalid directory: %JAVA_HOME% 1>&2 +echo. 1>&2 +echo Please set the JAVA_HOME variable in your environment to match the 1>&2 +echo location of your Java installation. 1>&2 + +goto fail + +:execute +@rem Setup the command line + +set CLASSPATH=%APP_HOME%\gradle\wrapper\gradle-wrapper.jar + + +@rem Execute Gradle +"%JAVA_EXE%" %DEFAULT_JVM_OPTS% %JAVA_OPTS% %GRADLE_OPTS% "-Dorg.gradle.appname=%APP_BASE_NAME%" -classpath "%CLASSPATH%" org.gradle.wrapper.GradleWrapperMain %* + +:end +@rem End local scope for the variables with windows NT shell +if %ERRORLEVEL% equ 0 goto mainEnd + +:fail +rem Set variable GRADLE_EXIT_CONSOLE if you need the _script_ return code instead of +rem the _cmd.exe /c_ return code! +set EXIT_CODE=%ERRORLEVEL% +if %EXIT_CODE% equ 0 set EXIT_CODE=1 +if not ""=="%GRADLE_EXIT_CONSOLE%" exit %EXIT_CODE% +exit /b %EXIT_CODE% + +:mainEnd +if "%OS%"=="Windows_NT" endlocal + +:omega diff --git a/fleet/compiler-plugins/rpc/settings.gradle.kts b/fleet/compiler-plugins/rpc/settings.gradle.kts new file mode 100644 index 000000000000..c7dec56ee7b7 --- /dev/null +++ b/fleet/compiler-plugins/rpc/settings.gradle.kts @@ -0,0 +1,26 @@ +rootProject.name = "rpc-compiler-plugin" + +pluginManagement { + // the compiler plugin will be built by this Kotlin compiler + val KOTLIN_VERSION = "2.2.0" + + plugins { + kotlin("jvm") version KOTLIN_VERSION + } + resolutionStrategy { + eachPlugin { + repositories { + gradlePluginPortal() + if ("SNAPSHOT" in KOTLIN_VERSION || "dev" in KOTLIN_VERSION) { + mavenLocal() + } + } + } + } + repositories { + gradlePluginPortal() + if ("SNAPSHOT" in KOTLIN_VERSION || "dev" in KOTLIN_VERSION) { + mavenLocal() + } + } +} diff --git a/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/Plugin.kt b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/Plugin.kt new file mode 100644 index 000000000000..8c23b2625d2c --- /dev/null +++ b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/Plugin.kt @@ -0,0 +1,27 @@ +package com.jetbrains.fleet.rpc.plugin + +import org.jetbrains.kotlin.backend.common.extensions.IrGenerationExtension +import org.jetbrains.kotlin.cli.common.CLIConfigurationKeys +import org.jetbrains.kotlin.cli.common.messages.MessageCollector +import org.jetbrains.kotlin.compiler.plugin.AbstractCliOption +import org.jetbrains.kotlin.compiler.plugin.CommandLineProcessor +import org.jetbrains.kotlin.compiler.plugin.CompilerPluginRegistrar +import org.jetbrains.kotlin.compiler.plugin.ExperimentalCompilerApi +import org.jetbrains.kotlin.config.CompilerConfiguration +import com.jetbrains.fleet.rpc.plugin.ir.ServiceGenerationExtension +import org.jetbrains.kotlin.config.CommonConfigurationKeys + +@OptIn(ExperimentalCompilerApi::class) +class RpcCommandLineProcessor : CommandLineProcessor { + override val pluginId: String = "rpc" + override val pluginOptions: Collection = emptyList() +} + +@OptIn(ExperimentalCompilerApi::class) +class RpcComponentRegistrar : CompilerPluginRegistrar() { + override val supportsK2: Boolean = true + override fun ExtensionStorage.registerExtensions(configuration: CompilerConfiguration) { + val messageCollector = configuration.get(CommonConfigurationKeys.MESSAGE_COLLECTOR_KEY, MessageCollector.NONE) + IrGenerationExtension.registerExtension(ServiceGenerationExtension(messageCollector)) + } +} \ No newline at end of file diff --git a/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/gradle/RpcGradlePlugin.kt b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/gradle/RpcGradlePlugin.kt new file mode 100644 index 000000000000..12c8b0581ae7 --- /dev/null +++ b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/gradle/RpcGradlePlugin.kt @@ -0,0 +1,25 @@ +package com.jetbrains.fleet.rpc.plugin.gradle + +import org.gradle.api.provider.Provider +import org.jetbrains.kotlin.gradle.plugin.KotlinCompilation +import org.jetbrains.kotlin.gradle.plugin.KotlinCompilerPluginSupportPlugin +import org.jetbrains.kotlin.gradle.plugin.SubpluginArtifact +import org.jetbrains.kotlin.gradle.plugin.SubpluginOption + +private const val VERSION = "2.2.0-0.2" + +class RpcGradlePlugin : KotlinCompilerPluginSupportPlugin { + override fun isApplicable(kotlinCompilation: KotlinCompilation<*>): Boolean = true + + override fun getCompilerPluginId(): String = "rpc" + + override fun getPluginArtifact(): SubpluginArtifact = SubpluginArtifact( + groupId = "com.jetbrains.fleet", + artifactId = "rpc-compiler-plugin", + version = VERSION + ) + + override fun applyToCompilation(kotlinCompilation: KotlinCompilation<*>): Provider> { + return kotlinCompilation.target.project.provider { emptyList() } + } +} \ No newline at end of file diff --git a/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/ClientStubBuilder.kt b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/ClientStubBuilder.kt new file mode 100644 index 000000000000..0d3aef61a9f8 --- /dev/null +++ b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/ClientStubBuilder.kt @@ -0,0 +1,255 @@ +package com.jetbrains.fleet.rpc.plugin.ir + +import com.jetbrains.fleet.rpc.plugin.ir.util.* +import org.jetbrains.kotlin.backend.common.extensions.IrPluginContext +import org.jetbrains.kotlin.backend.jvm.codegen.AnnotationCodegen.Companion.annotationClass +import org.jetbrains.kotlin.descriptors.DescriptorVisibilities +import org.jetbrains.kotlin.ir.builders.declarations.addDefaultGetter +import org.jetbrains.kotlin.ir.builders.declarations.addProperty +import org.jetbrains.kotlin.ir.builders.declarations.addValueParameter +import org.jetbrains.kotlin.ir.builders.declarations.buildField +import org.jetbrains.kotlin.ir.builders.irGet +import org.jetbrains.kotlin.ir.builders.irReturn +import org.jetbrains.kotlin.ir.builders.irString +import org.jetbrains.kotlin.ir.declarations.* +import org.jetbrains.kotlin.ir.expressions.IrCall +import org.jetbrains.kotlin.ir.expressions.IrConstructorCall +import org.jetbrains.kotlin.ir.expressions.IrExpression +import org.jetbrains.kotlin.ir.expressions.IrStatementOrigin +import org.jetbrains.kotlin.ir.expressions.impl.IrCallImpl +import org.jetbrains.kotlin.ir.expressions.impl.fromSymbolOwner +import org.jetbrains.kotlin.ir.types.* +import org.jetbrains.kotlin.ir.util.* +import org.jetbrains.kotlin.name.CallableId +import org.jetbrains.kotlin.name.ClassId +import org.jetbrains.kotlin.name.FqName +import org.jetbrains.kotlin.name.Name + +private val flowPackage = FqName.fromSegments(listOf("kotlinx", "coroutines", "flow")) +private val flowClassId = ClassId(packageFqName = flowPackage, topLevelName = "Flow".name) +private val flowCollectorClassId = ClassId(packageFqName = flowPackage, topLevelName = "FlowCollector".name) +private val flowFnId = CallableId(packageName = flowPackage, callableName = "flow".name) +private val flowCollectFnId = CallableId(classId = flowClassId, callableName = "collect".name) + +private val asyncPackage = FqName.fromSegments(listOf("fleet", "util", "async")) +private val resourceFnId = CallableId(packageName = asyncPackage, callableName = "resource".name) +private val resourceUseFnId = CallableId(classId = ClassId(asyncPackage, "Resource".name), callableName = "use".name) +private val consumedClassId = ClassId(asyncPackage, "Consumed".name) + +private val coroutineScopeClassId = ClassId(FqName.fromSegments(listOf("kotlinx", "coroutines")), "CoroutineScope".name) + +fun buildClientStub(interfaceClass: IrClass, rpcMethods: List, context: CompilerPluginContext): IrClass { + val pluginContext = context.pluginContext + + // Resolve necessary refs + val flowCollectorClass = pluginContext.referenceClass(flowCollectorClassId)!! + val flowFunction = pluginContext.referenceFunctions(flowFnId).first() + val collectFunction = pluginContext.referenceFunctions(flowCollectFnId).first() + val resourceFunction = pluginContext.referenceFunctions(resourceFnId).first() + val useFunction = pluginContext.referenceFunctions(resourceUseFnId).first() + val consumed = pluginContext.referenceClass(consumedClassId)!! + val coroutineScope = pluginContext.referenceClass(coroutineScopeClassId)!! + + val generatedClass = pluginContext.buildClassBase { + name = Name.identifier(interfaceClass.name.identifier + "ClientStub") + visibility = interfaceClass.visibility + }.also { + it.parent = interfaceClass.parent + it.superTypes = listOfNotNull(interfaceClass.defaultType) + it.annotations += interfaceClass.extractApiStatusAnnotations() + + val container = interfaceClass.parent as? IrClass ?: interfaceClass.file + container.addChild(it) + } + + val invocationHandler = generatedClass.addInvocationHandlerToPrimaryConstructor(pluginContext) + + for (rpcMethod in rpcMethods) { + val originalFunction = rpcMethod.symbol.owner + pluginContext.addFunctionOverride(originalFunction, generatedClass) { generatedFunction -> + generatedFunction.annotations += originalFunction.extractApiStatusAnnotations() + generatedFunction.generateBody(context) { + val irThis = irGet(generatedFunction.parameters[0]) + + // Obtain the 'invoke' function from the 'invocationHandler' + val invokeFunction = invocationHandler.backingField!!.type.getClass()!!.functions.first { it.name.asString() == "invoke" } + + val invocationHandlerReceiver = irGetValue(invocationHandler, irThis) // this.invocationHandler + val invocationHandlerCall = irCall(invokeFunction.symbol).also { call -> // this.invocationHandler(methodName, args) + call.arguments[0] = invocationHandlerReceiver + call.arguments[1] = irString(originalFunction.name.identifier) // methodName + + val args = createArrayOfExpression( + pluginContext.irBuiltIns.anyNType, + // All parameters except the first (dispatch receiver) + generatedFunction.parameters.drop(1).map { irGet(it) }, + pluginContext + ) + call.arguments[2] = args // args + } + + when { + originalFunction.isNonSuspendFlowFunction() -> { + val flowElementType = when (val arg = (originalFunction.returnType as IrSimpleType).arguments.first()) { + is IrStarProjection -> pluginContext.irBuiltIns.anyNType + is IrTypeProjection -> arg.type + } + + // flow { it -> invocationHandler(...).collect(it) } + val wrappedCall = irCall(flowFunction).also { flowCall -> + val flowLambda = pluginContext.irLambda( + originalFunction, + parent = generatedFunction, + pluginContext.irBuiltIns.suspendFunctionN(1).typeWith( + flowCollectorClass.typeWith(flowElementType), + pluginContext.irBuiltIns.unitType + ) + ) { flowLambda -> + flowLambda.generateBody(context) { + +irReturn( + irCall(collectFunction).also { collectCall -> + collectCall.arguments[0] = invocationHandlerCall + + // We're getting a flow collector in the lambda, just gotta pass it + collectCall.arguments[1] = irGet(flowLambda.parameters.first()) + } + ) + } + } + + flowCall.arguments[0] = flowLambda + } + + +irReturn(wrappedCall) + } + originalFunction.isNonSuspendResourceFunction() -> { + val resourceElementType = + ((originalFunction.returnType as IrSimpleType) + .arguments + .first() + as IrTypeProjection + ).type + + /* + * resource { consume -> remoteCall().use { it -> consume(it) } } + */ + val wrappedCall = irCall(resourceFunction).also { flowCall -> + val consumerType = pluginContext.irBuiltIns.suspendFunctionN(1).typeWith( + resourceElementType, + consumed.defaultType, + ) + val resourceLambda = pluginContext.irLambda( + originalFunction, + parent = generatedFunction, + + // suspend CoroutineScope.(consumer: Consumer) -> Consumed + pluginContext.irBuiltIns.suspendFunctionN(2).typeWith( + coroutineScope.defaultType, + consumerType, + consumed.defaultType + ) + ) { resourceLambda -> + resourceLambda.generateBody(context) { + +irReturn( + // use { it -> consume(it) } + irCall(useFunction).also { useCall -> + useCall.arguments[0] = invocationHandlerCall + useCall.arguments[1] = pluginContext.irLambda( + originalFunction, + parent = resourceLambda, + // suspend CoroutineScope.(T) -> U + pluginContext.irBuiltIns.suspendFunctionN(2).typeWith( + coroutineScope.defaultType, + resourceElementType, + consumed.defaultType, + ) + ) { useLambda -> + useLambda.generateBody(context) { + +irReturn(irCall(consumerType.classOrFail.getSimpleFunction("invoke")!!).also { collectCall -> + // Call the resource consumer (second parameter: coroutine scope is first) + collectCall.arguments[0] = irGet(resourceLambda.parameters[1]) + + // ... with our result + collectCall.arguments[1] = irGet(useLambda.parameters[1]) + }) + } + } + } + ) + } + } + + flowCall.arguments[0] = resourceLambda + } + + +irReturn(wrappedCall) + } + else -> { + +irReturn(invocationHandlerCall) + } + } + } + } + } + + return generatedClass +} + +private fun irGetValue(irProperty: IrProperty, receiver: IrExpression?): IrCall = + IrCallImpl.fromSymbolOwner( + startOffset = SYNTHETIC_OFFSET, + endOffset = SYNTHETIC_OFFSET, + type = irProperty.backingField!!.type, + symbol = irProperty.getter!!.symbol, + origin = IrStatementOrigin.GET_PROPERTY, + ).apply { + dispatchReceiver = receiver + } + +private fun IrClass.addInvocationHandlerToPrimaryConstructor(pluginContext: IrPluginContext): IrProperty { + val primaryConstructor = this.primaryConstructor!! + + val invocationHandlerType = pluginContext.getProxyType() + val invocationHandler = primaryConstructor.addValueParameter("invocationHandler".name, invocationHandlerType) + + return addIrValProperty( + valName = invocationHandler.name, + valType = invocationHandlerType, + valInitializer = irGet(invocationHandler), + pluginContext = pluginContext, + ) +} + +internal fun IrAnnotationContainer.extractApiStatusAnnotations(): List { + return annotations.filter { it.annotationClass.kotlinFqName.startsWith(FqName.fromSegments(listOf("org", "jetbrains", "annotations", "ApiStatus"))) } +} + +private fun IrClass.addIrValProperty( + valName: Name, + valType: IrType, + valInitializer: IrExpression, + pluginContext: IrPluginContext, +): IrProperty { + val irClass = this + + val irField = pluginContext.irFactory.buildField { + name = valName + type = valType + origin = IrDeclarationOrigin.PROPERTY_BACKING_FIELD + visibility = DescriptorVisibilities.PRIVATE + }.apply { + parent = irClass + initializer = pluginContext.irFactory.createExpressionBody(valInitializer) + } + + val irProperty = irClass.addProperty { + name = valName + isVar = false + }.apply { + parent = irClass + backingField = irField + addDefaultGetter(irClass, pluginContext.irBuiltIns) + } + + return irProperty +} diff --git a/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/GenerateClasses.kt b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/GenerateClasses.kt new file mode 100644 index 000000000000..72a6654a92b9 --- /dev/null +++ b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/GenerateClasses.kt @@ -0,0 +1,39 @@ +package com.jetbrains.fleet.rpc.plugin.ir + +import com.jetbrains.fleet.rpc.plugin.ir.util.collectRpcMethods +import org.jetbrains.kotlin.backend.common.IrElementTransformerVoidWithContext +import org.jetbrains.kotlin.ir.IrStatement +import org.jetbrains.kotlin.ir.declarations.IrClass +import org.jetbrains.kotlin.ir.util.isInterface +import com.jetbrains.fleet.rpc.plugin.ir.util.hasRpcAnnotation +import com.jetbrains.fleet.rpc.plugin.ir.util.isNonSuspendFlowFunction +import com.jetbrains.fleet.rpc.plugin.ir.util.isNonSuspendResourceFunction +import org.jetbrains.kotlin.backend.common.getCompilerMessageLocation +import org.jetbrains.kotlin.cli.common.messages.CompilerMessageSeverity +import org.jetbrains.kotlin.ir.util.file + + +class GenerateClasses(private val context: CompilerPluginContext) : IrElementTransformerVoidWithContext() { + override fun visitClassNew(declaration: IrClass): IrStatement { + if (declaration.isInterface && declaration.hasRpcAnnotation()) { + val rpcMethods = collectRpcMethods(declaration).filter { + if (it.isSuspend || it.isNonSuspendFlowFunction() || it.isNonSuspendResourceFunction()) { + true + } + else { + // Forbid unsupported methods + context.messageCollector.report( + CompilerMessageSeverity.ERROR, + "${declaration.name} should either be suspend, return Resource or return Flow<*> to be supported in @Rpc interface", + it.getCompilerMessageLocation(declaration.file) + ) + false + } + } + val clientStubClass = buildClientStub(declaration, rpcMethods, context) + buildRemoteApiDescriptorImpl(declaration, clientStubClass, rpcMethods, context) + } + + return super.visitClassNew(declaration) + } +} diff --git a/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/RemoteApiDescriptorImplBuilder.kt b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/RemoteApiDescriptorImplBuilder.kt new file mode 100644 index 000000000000..3a1986dd16a7 --- /dev/null +++ b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/RemoteApiDescriptorImplBuilder.kt @@ -0,0 +1,294 @@ +package com.jetbrains.fleet.rpc.plugin.ir + +import com.jetbrains.fleet.rpc.plugin.ir.remoteKind.REMOTE_KIND_FQN +import com.jetbrains.fleet.rpc.plugin.ir.remoteKind.toRemoteKind +import com.jetbrains.fleet.rpc.plugin.ir.util.* +import org.jetbrains.kotlin.backend.common.extensions.IrPluginContext +import org.jetbrains.kotlin.descriptors.ClassKind +import org.jetbrains.kotlin.ir.IrBuiltIns +import org.jetbrains.kotlin.ir.UNDEFINED_OFFSET +import org.jetbrains.kotlin.ir.builders.* +import org.jetbrains.kotlin.ir.builders.declarations.addValueParameter +import org.jetbrains.kotlin.ir.builders.declarations.buildValueParameter +import org.jetbrains.kotlin.ir.declarations.IrClass +import org.jetbrains.kotlin.ir.declarations.IrParameterKind +import org.jetbrains.kotlin.ir.declarations.IrSimpleFunction +import org.jetbrains.kotlin.ir.declarations.IrVariable +import org.jetbrains.kotlin.ir.expressions.* +import org.jetbrains.kotlin.ir.expressions.impl.IrConstImpl +import org.jetbrains.kotlin.ir.types.IrSimpleType +import org.jetbrains.kotlin.ir.types.defaultType +import org.jetbrains.kotlin.ir.types.impl.IrSimpleTypeImpl +import org.jetbrains.kotlin.ir.util.* +import org.jetbrains.kotlin.name.CallableId +import org.jetbrains.kotlin.name.ClassId +import org.jetbrains.kotlin.name.FqName +import kotlin.collections.plus + +fun buildRemoteApiDescriptorImpl( + interfaceClass: IrClass, + clientStub: IrClass, + rpcMethods: List, + context: CompilerPluginContext +): IrClass { + val generatedClass = context.pluginContext.buildClassBase { + name = interfaceClass.name.getRemoteApiDescriptorImplName() + visibility = interfaceClass.visibility + kind = ClassKind.OBJECT + }.also { + it.parent = interfaceClass.parent + val newTypeArgument = IrSimpleTypeImpl( + classifier = interfaceClass.symbol, + hasQuestionMark = false, + arguments = emptyList(), + annotations = emptyList() + ) + val parametrizedClientType: IrSimpleType = IrSimpleTypeImpl( + classifier = context.remoteApiDescriptorClass.symbol, + hasQuestionMark = false, + arguments = listOf(newTypeArgument), + annotations = emptyList() + ) + it.annotations += interfaceClass.extractApiStatusAnnotations() + it.superTypes = listOf(parametrizedClientType) + val container = interfaceClass.parent as? IrClass ?: interfaceClass.file + container.addChild(it) + } + + overrideGetSignatureFunction(rpcMethods, generatedClass, interfaceClass, context) + overrideCallFunction(rpcMethods, generatedClass, interfaceClass, context) + overrideClientStubFunction(generatedClass, interfaceClass, clientStub, context) + overrideGetApiFqnFunction(generatedClass, interfaceClass, context) + + return generatedClass +} + +private fun overrideGetSignatureFunction(rpcMethods: List, + generatedClass: IrClass, + interfaceClass: IrClass, + context: CompilerPluginContext) { // fun getSignature(methodName: String): RpcSignature + overrideFunctionWithName("getSignature", generatedClass, context) { function -> + function.generateBody(context) { + val methodName = irTemporary(irGet(function.parameters[1])) + whenBlockForMethods(methodName, rpcMethods, interfaceClass, context) { remoteApiFunction -> + val errorMessage = "${interfaceClass.kotlinFqName}/${remoteApiFunction.name}" + val parameters = createArrayOfExpression( + arrayElementType = context.parameterDescriptorClass.defaultType, + arrayElements = remoteApiFunction.parameters.filter { it.kind == IrParameterKind.Regular }.map { parameter -> + irCallConstructor( + callee = context.parameterDescriptorConstructor, + typeArguments = listOf(context.pluginContext.irBuiltIns.stringType, context.remoteKindClass.defaultType) + ).apply { + arguments[0] = irString(parameter.name.identifier) // parameterName + arguments[1] = toRemoteKind(parameter.type, context, errorMessage) // parameterKind + } + }, + pluginContext = context.pluginContext + ) + val rpcSignature = irCallConstructor( + callee = context.rpcSignatureConstructor, + typeArguments = listOf(context.pluginContext.irBuiltIns.stringType, parameters.type, context.remoteKindClass.defaultType) + ).apply { + arguments[0] = irString(remoteApiFunction.name.identifier) // methodName + arguments[1] = parameters // parameters + arguments[2] = toRemoteKind(remoteApiFunction.returnType, context, errorMessage) // returnType + } + irReturn(rpcSignature) + } + } + } +} + +private val CompilerPluginContext.rpcSignatureConstructor + get() = cache.remember { + val rpcSignatureClass = pluginContext.referenceClass(ClassId.topLevel(RPC_FQN.child("RpcSignature".name)))!! + rpcSignatureClass.constructors.single() + } + +private val CompilerPluginContext.parameterDescriptorClass + get() = cache.remember { + pluginContext.referenceClass(ClassId.topLevel(RPC_FQN.child("ParameterDescriptor".name)))!! + } + +private val CompilerPluginContext.parameterDescriptorConstructor + get() = cache.remember { + parameterDescriptorClass.constructors.single() + } + +private val CompilerPluginContext.remoteKindClass + get() = cache.remember { + pluginContext.referenceClass(ClassId.topLevel(REMOTE_KIND_FQN))!! + } + +private fun overrideClientStubFunction(generatedClass: IrClass, + interfaceClass: IrClass, + clientStub: IrClass, + context: CompilerPluginContext) { // fun clientStub(proxy: Proxy): T + overrideFunctionWithName("clientStub", generatedClass, context) { function -> + function.returnType = interfaceClass.defaultType + + function.generateBody(context) { + val clientStubConstructor = clientStub.primaryConstructor?.symbol!! + + val proxyType = context.pluginContext.getProxyType() + val clientStubInstance = irCallConstructor(clientStubConstructor, listOf(proxyType)).apply { + arguments[0] = irGet(function.parameters[1]) + } + +irReturn(clientStubInstance) + } + } +} + +private fun overrideGetApiFqnFunction(generatedClass: IrClass, + interfaceClass: IrClass, + context: CompilerPluginContext) { // fun getApiFqn(): String + overrideFunctionWithName("getApiFqn", generatedClass, context) { function -> + function.generateBody(context) { + +irReturn(irString(interfaceClass.kotlinFqName.asString())) + } + } +} + +private fun overrideCallFunction(rpcMethods: List, + generatedClass: IrClass, + interfaceClass: IrClass, + context: CompilerPluginContext) { // suspend fun call(impl: T, methodName: String, args: Array): Any? + val oldFunction = checkNotNull(context.remoteApiDescriptorClass.getSimpleFunction("call")).owner + + context.pluginContext.addFunctionOverride(oldFunction, generatedClass) { function -> + // change first regular parameter type from T to actual interface type (interfaceClass.defaultType) + function.parameters = function.parameters.mapIndexed { index, parameter -> + if (index == 1) { + buildValueParameter(function) { + name = parameter.name + updateFrom(parameter) + type = interfaceClass.defaultType + } + } else { + parameter + } + } + + function.generateBody(context) { + val methodName = irTemporary(irGet(function.parameters[2])) + whenBlockForMethods(methodName, rpcMethods, interfaceClass, context) { remoteApiFunction -> + val call = callRemoteApiFunction(remoteApiFunction, irGet(function.parameters[1]), irGet(function.parameters[3]), context.pluginContext) + irReturn(call) + } + } + } +} + +private fun overrideFunctionWithName(name: String, + generatedClass: IrClass, + context: CompilerPluginContext, + block: (IrSimpleFunction) -> Unit) { + val function = checkNotNull(context.remoteApiDescriptorClass.getSimpleFunction(name)).owner + context.pluginContext.addFunctionOverride(function, generatedClass, block) +} + +private fun IrBlockBodyBuilder.whenBlockForMethods(methodName: IrVariable, + rpcMethods: List, + interfaceClass: IrClass, + context: CompilerPluginContext, + block: IrBlockBodyBuilder.(IrSimpleFunction) -> IrExpression) { + +irBlock( + origin = IrStatementOrigin.WHEN + ) { + +irWhen( + context.pluginContext.irBuiltIns.unitType, + rpcMethods.map { whenBranch(it, methodName, context, block) } + irElseBranch( + generateErrorFunctionCall( + irString("${interfaceClass.kotlinFqName} does not have method ").plusString(irGet(methodName), context), + context + ) + ) + ) + } +} + +private fun IrConst.plusString(otherString: IrExpression, context: CompilerPluginContext): IrExpression { + return irCall(stringPlusFunction(context)).apply { + arguments[0] = this@plusString + arguments[1] = otherString + } +} + +private fun stringPlusFunction(context: CompilerPluginContext) = context.cache.remember { + context.pluginContext.referenceFunctions(CallableId( + classId = ClassId(FqName.fromSegments(listOf("kotlin")), "String".name), + callableName = "plus".name + )).first().owner.symbol +} + +private fun generateErrorFunctionCall(argument: IrExpression, context: CompilerPluginContext): IrExpression { + return irCall(errorFunction(context)).apply { + arguments[0] = argument + } +} + +private fun errorFunction(context: CompilerPluginContext) = context.cache.remember { + context.pluginContext.referenceFunctions(CallableId( + packageName = FqName.fromSegments(listOf("kotlin")), + className = null, + callableName = "error".name + )).first().owner.symbol +} + +private fun IrBlockBodyBuilder.whenBranch( + remoteApiFunction: IrSimpleFunction, + methodName: IrVariable, + context: CompilerPluginContext, + block: IrBlockBodyBuilder.(IrSimpleFunction) -> IrExpression, +): IrBranch = + irBranch( + irEquals( + irGet(methodName), + context.pluginContext.irBuiltIns.irString(remoteApiFunction.name.identifier), + IrStatementOrigin.EQEQ + ), + block(remoteApiFunction) + ) + +private fun IrBuiltIns.irString(value: String): IrConst = IrConstImpl( + UNDEFINED_OFFSET, + UNDEFINED_OFFSET, + stringType, + IrConstKind.String, + value +) + +private fun IrBlockBodyBuilder.callRemoteApiFunction(remoteApiFunction: IrSimpleFunction, + dispatchReceiverExpression: IrExpression, + argsExpression: IrExpression, + pluginContext: IrPluginContext): IrExpression { + val getFunction = pluginContext.irBuiltIns.arrayClass.owner.declarations + .filterIsInstance() + .single { it.name.asString() == "get" } + + return irCall( + remoteApiFunction.symbol, + ).also { call -> + call.arguments[0] = dispatchReceiverExpression + + // Parameters without dispatcher + val valueParameters = remoteApiFunction.parameters.drop(1) + val args = irTemporary(argsExpression) // vararg with args + + for (index in valueParameters.indices) { + val argument = irCall(getFunction.symbol).also { getCall -> + getCall.arguments[0] = irGet(args) + getCall.arguments[1] = irInt(index) + } + call.arguments[index + 1] = irCastIfNeeded(argument, valueParameters[index].type) + } + } +} + +private val REMOTE_API_DESCRIPTOR_INTERFACE_FQN = RPC_FQN.child("RemoteApiDescriptor".name) + +private val CompilerPluginContext.remoteApiDescriptorClass + get() = cache.remember { + pluginContext.referenceClass(ClassId.topLevel(REMOTE_API_DESCRIPTOR_INTERFACE_FQN))?.owner!! + } + diff --git a/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/RemoteApiDescriptorTransform.kt b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/RemoteApiDescriptorTransform.kt new file mode 100644 index 000000000000..5ccf2f5d96db --- /dev/null +++ b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/RemoteApiDescriptorTransform.kt @@ -0,0 +1,54 @@ +package com.jetbrains.fleet.rpc.plugin.ir + +import org.jetbrains.kotlin.backend.common.IrElementTransformerVoidWithContext +import org.jetbrains.kotlin.descriptors.ClassKind +import org.jetbrains.kotlin.ir.builders.declarations.buildClass +import org.jetbrains.kotlin.ir.declarations.impl.IrFactoryImpl +import org.jetbrains.kotlin.ir.expressions.IrCall +import org.jetbrains.kotlin.ir.expressions.IrExpression +import org.jetbrains.kotlin.ir.types.IrType +import org.jetbrains.kotlin.ir.util.* +import org.jetbrains.kotlin.name.Name +import com.jetbrains.fleet.rpc.plugin.ir.util.RPC_FQN +import com.jetbrains.fleet.rpc.plugin.ir.util.name +import org.jetbrains.kotlin.ir.UNDEFINED_OFFSET +import org.jetbrains.kotlin.ir.expressions.impl.IrGetObjectValueImpl +import org.jetbrains.kotlin.ir.types.classOrNull + +class RemoteApiDescriptorTransform : IrElementTransformerVoidWithContext() { + override fun visitCall(expression: IrCall): IrExpression { + if (expression.symbol.owner.kotlinFqName == REMOTE_API_DESCRIPTOR_FUNCTION_FQN) { + val descriptorInstance = getRemoteApiDescriptorInstance(checkNotNull(expression.typeArguments[0])) { + currentDeclarationParent?.dumpKotlinLike() + } + expression.arguments[0] = descriptorInstance + } + return super.visitCall(expression) + } +} + +fun getRemoteApiDescriptorInstance(type: IrType, debugInfo: () -> String?): IrExpression { + val irClass = requireNotNull(type.classOrNull) { + "Error in RPC compiler plugin: can't get class by type ${type.dumpKotlinLike()}, error occured in\n${debugInfo()}" + }.owner + + val implClass = + IrFactoryImpl.buildClass { + name = irClass.name.getRemoteApiDescriptorImplName() + kind = ClassKind.OBJECT + }.apply { + parent = irClass.parent + createThisReceiverParameter() + } + + return IrGetObjectValueImpl( + startOffset = UNDEFINED_OFFSET, + endOffset = UNDEFINED_OFFSET, + type = implClass.defaultType, + symbol = implClass.symbol + ) +} + +fun Name.getRemoteApiDescriptorImplName() = Name.identifier(identifier + "RemoteApiDescriptorImpl") + +val REMOTE_API_DESCRIPTOR_FUNCTION_FQN = RPC_FQN.child("remoteApiDescriptor".name) diff --git a/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/ServiceGenerationExtension.kt b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/ServiceGenerationExtension.kt new file mode 100644 index 000000000000..ea3976893112 --- /dev/null +++ b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/ServiceGenerationExtension.kt @@ -0,0 +1,20 @@ +package com.jetbrains.fleet.rpc.plugin.ir + +import com.jetbrains.fleet.rpc.plugin.ir.util.Cache +import org.jetbrains.kotlin.backend.common.extensions.IrGenerationExtension +import org.jetbrains.kotlin.backend.common.extensions.IrPluginContext +import org.jetbrains.kotlin.cli.common.messages.MessageCollector +import org.jetbrains.kotlin.ir.declarations.IrModuleFragment + +class ServiceGenerationExtension(private val messageCollector: MessageCollector) : IrGenerationExtension { + override fun generate(moduleFragment: IrModuleFragment, pluginContext: IrPluginContext) { + val context = CompilerPluginContext(moduleFragment, pluginContext, messageCollector, Cache()) + moduleFragment.accept(GenerateClasses(context), null) + moduleFragment.accept(RemoteApiDescriptorTransform(), null) + } +} + +data class CompilerPluginContext(val moduleFragment: IrModuleFragment, + val pluginContext: IrPluginContext, + val messageCollector: MessageCollector, + val cache: Cache) \ No newline at end of file diff --git a/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/remoteKind/DefaultSerializers.kt b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/remoteKind/DefaultSerializers.kt new file mode 100644 index 000000000000..d33b0d7309cc --- /dev/null +++ b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/remoteKind/DefaultSerializers.kt @@ -0,0 +1,165 @@ +package com.jetbrains.fleet.rpc.plugin.ir.remoteKind + +import org.jetbrains.kotlin.name.CallableId +import org.jetbrains.kotlin.name.FqName +import com.jetbrains.fleet.rpc.plugin.ir.CompilerPluginContext +import com.jetbrains.fleet.rpc.plugin.ir.util.name + +internal fun CompilerPluginContext.getSpecialSerializer(fqName: FqName?) = cache.remember { + val listSerializer = pluginContext.referenceFunctions(CallableId( + packageName = FqName.fromSegments(listOf("kotlinx", "serialization", "builtins")), + className = null, + callableName = "ListSerializer".name + )).first() + + val setSerializer = pluginContext.referenceFunctions(CallableId( + packageName = FqName.fromSegments(listOf("kotlinx", "serialization", "builtins")), + className = null, + callableName = "SetSerializer".name + )).first() + + val mapSerializer = pluginContext.referenceFunctions(CallableId( + packageName = FqName.fromSegments(listOf("kotlinx", "serialization", "builtins")), + className = null, + callableName = "MapSerializer".name + )).first() + + val pairSerializer = pluginContext.referenceFunctions(CallableId( + packageName = FqName.fromSegments(listOf("kotlinx", "serialization", "builtins")), + className = null, + callableName = "PairSerializer".name + )).first() + + val mapEntrySerializer = pluginContext.referenceFunctions(CallableId( + packageName = FqName.fromSegments(listOf("kotlinx", "serialization", "builtins")), + className = null, + callableName = "MapEntrySerializer".name + )).first() + + val tripleSerializer = pluginContext.referenceFunctions(CallableId( + packageName = FqName.fromSegments(listOf("kotlinx", "serialization", "builtins")), + className = null, + callableName = "TripleSerializer".name + )).first() + + val charArraySerializer = pluginContext.referenceFunctions(CallableId( + packageName = FqName.fromSegments(listOf("kotlinx", "serialization", "builtins")), + className = null, + callableName = "CharArraySerializer".name + )).first() + + val byteArraySerializer = pluginContext.referenceFunctions(CallableId( + packageName = FqName.fromSegments(listOf("kotlinx", "serialization", "builtins")), + className = null, + callableName = "ByteArraySerializer".name + )).first() + + val uByteArraySerializer = pluginContext.referenceFunctions(CallableId( + packageName = FqName.fromSegments(listOf("kotlinx", "serialization", "builtins")), + className = null, + callableName = "UByteArraySerializer".name + )).first() + + val shortArraySerializer = pluginContext.referenceFunctions(CallableId( + packageName = FqName.fromSegments(listOf("kotlinx", "serialization", "builtins")), + className = null, + callableName = "ShortArraySerializer".name + )).first() + + val uShortArraySerializer = pluginContext.referenceFunctions(CallableId( + packageName = FqName.fromSegments(listOf("kotlinx", "serialization", "builtins")), + className = null, + callableName = "UShortArraySerializer".name + )).first() + + val intArraySerializer = pluginContext.referenceFunctions(CallableId( + packageName = FqName.fromSegments(listOf("kotlinx", "serialization", "builtins")), + className = null, + callableName = "IntArraySerializer".name + )).first() + + val uIntArraySerializer = pluginContext.referenceFunctions(CallableId( + packageName = FqName.fromSegments(listOf("kotlinx", "serialization", "builtins")), + className = null, + callableName = "UIntArraySerializer".name + )).first() + + val longArraySerializer = pluginContext.referenceFunctions(CallableId( + packageName = FqName.fromSegments(listOf("kotlinx", "serialization", "builtins")), + className = null, + callableName = "LongArraySerializer".name + )).first() + + val uLongArraySerializer = pluginContext.referenceFunctions(CallableId( + packageName = FqName.fromSegments(listOf("kotlinx", "serialization", "builtins")), + className = null, + callableName = "ULongArraySerializer".name + )).first() + + val floatArraySerializer = pluginContext.referenceFunctions(CallableId( + packageName = FqName.fromSegments(listOf("kotlinx", "serialization", "builtins")), + className = null, + callableName = "FloatArraySerializer".name + )).first() + + val doubleArraySerializer = pluginContext.referenceFunctions(CallableId( + packageName = FqName.fromSegments(listOf("kotlinx", "serialization", "builtins")), + className = null, + callableName = "DoubleArraySerializer".name + )).first() + + val booleanArraySerializer = pluginContext.referenceFunctions(CallableId( + packageName = FqName.fromSegments(listOf("kotlinx", "serialization", "builtins")), + className = null, + callableName = "BooleanArraySerializer".name + )).first() + + val arraySerializer = pluginContext.referenceFunctions(CallableId( + packageName = FqName.fromSegments(listOf("kotlinx", "serialization", "builtins")), + className = null, + callableName = "ArraySerializer".name + )).first { it.owner.parameters.size == 1 } + + mapOf( + LIST_FQN to listSerializer, + SET_FQN to setSerializer, + MAP_FQN to mapSerializer, + PAIR_FQN to pairSerializer, + MAP_ENTRY_FQN to mapEntrySerializer, + TRIPLE_FQN to tripleSerializer, + CHAR_ARRAY_FQN to charArraySerializer, + BYTE_ARRAY_FQN to byteArraySerializer, + U_BYTE_ARRAY_FQN to uByteArraySerializer, + SHORT_ARRAY_FQN to shortArraySerializer, + U_SHORT_ARRAY_FQN to uShortArraySerializer, + INT_ARRAY_FQN to intArraySerializer, + U_INT_ARRAY_FQN to uIntArraySerializer, + LONG_ARRAY_FQN to longArraySerializer, + U_LONG_ARRAY_FQN to uLongArraySerializer, + FLOAT_ARRAY_FQN to floatArraySerializer, + DOUBLE_ARRAY_FQN to doubleArraySerializer, + BOOLEAN_ARRAY_FQN to booleanArraySerializer, + ARRAY_FQN to arraySerializer, + ) +}[fqName] + + +private val LIST_FQN = FqName.fromSegments(listOf("kotlin", "collections", "List")) +private val SET_FQN = FqName.fromSegments(listOf("kotlin", "collections", "Set")) +private val MAP_FQN = FqName.fromSegments(listOf("kotlin", "collections", "Map")) +private val PAIR_FQN = FqName.fromSegments(listOf("kotlin", "Pair")) +private val MAP_ENTRY_FQN = FqName.fromSegments(listOf("kotlin", "collections", "Map", "Entry")) +private val TRIPLE_FQN = FqName.fromSegments(listOf("kotlin", "Triple")) +private val CHAR_ARRAY_FQN = FqName.fromSegments(listOf("kotlin", "CharArray")) +private val BYTE_ARRAY_FQN = FqName.fromSegments(listOf("kotlin", "ByteArray")) +private val U_BYTE_ARRAY_FQN = FqName.fromSegments(listOf("kotlin", "UByteArray")) +private val SHORT_ARRAY_FQN = FqName.fromSegments(listOf("kotlin", "ShortArray")) +private val U_SHORT_ARRAY_FQN = FqName.fromSegments(listOf("kotlin", "UShortArray")) +private val INT_ARRAY_FQN = FqName.fromSegments(listOf("kotlin", "IntArray")) +private val U_INT_ARRAY_FQN = FqName.fromSegments(listOf("kotlin", "UIntArray")) +private val LONG_ARRAY_FQN = FqName.fromSegments(listOf("kotlin", "LongArray")) +private val U_LONG_ARRAY_FQN = FqName.fromSegments(listOf("kotlin", "ULongArray")) +private val FLOAT_ARRAY_FQN = FqName.fromSegments(listOf("kotlin", "FloatArray")) +private val DOUBLE_ARRAY_FQN = FqName.fromSegments(listOf("kotlin", "DoubleArray")) +private val BOOLEAN_ARRAY_FQN = FqName.fromSegments(listOf("kotlin", "BooleanArray")) +private val ARRAY_FQN = FqName.fromSegments(listOf("kotlin", "Array")) diff --git a/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/remoteKind/RemoteKind.kt b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/remoteKind/RemoteKind.kt new file mode 100644 index 000000000000..d6ca5c98ebd5 --- /dev/null +++ b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/remoteKind/RemoteKind.kt @@ -0,0 +1,222 @@ +package com.jetbrains.fleet.rpc.plugin.ir.remoteKind + +import com.jetbrains.fleet.rpc.plugin.ir.CompilerPluginContext +import com.jetbrains.fleet.rpc.plugin.ir.getRemoteApiDescriptorInstance +import com.jetbrains.fleet.rpc.plugin.ir.util.RPC_FQN +import com.jetbrains.fleet.rpc.plugin.ir.util.name +import org.jetbrains.kotlin.backend.wasm.ir2wasm.allSuperInterfaces +import org.jetbrains.kotlin.cli.common.messages.CompilerMessageSeverity +import org.jetbrains.kotlin.ir.builders.* +import org.jetbrains.kotlin.ir.declarations.IrClass +import org.jetbrains.kotlin.ir.declarations.IrParameterKind +import org.jetbrains.kotlin.ir.expressions.IrExpression +import org.jetbrains.kotlin.ir.expressions.IrFunctionAccessExpression +import org.jetbrains.kotlin.ir.types.IrSimpleType +import org.jetbrains.kotlin.ir.types.IrType +import org.jetbrains.kotlin.ir.types.IrTypeProjection +import org.jetbrains.kotlin.ir.types.classFqName +import org.jetbrains.kotlin.ir.types.classOrFail +import org.jetbrains.kotlin.ir.types.classOrNull +import org.jetbrains.kotlin.ir.types.typeOrFail +import org.jetbrains.kotlin.ir.util.* +import org.jetbrains.kotlin.name.CallableId +import org.jetbrains.kotlin.name.ClassId +import org.jetbrains.kotlin.name.FqName + +fun IrBuilderWithScope.toRemoteKind(irType: IrType, context: CompilerPluginContext, debugInfo: String): IrExpression { + return handleSpecialTypes(irType, context, debugInfo) + ?: handleRemoteObject(irType, context) + ?: handleResource(irType, context) + ?: run { // RemoteKind.Data + val serializer = generateSerializerCall(irType, context, debugInfo) + + irCallConstructor(context.remoteKindDataConstructor, listOf(serializer.type)).apply { + arguments[0] = serializer + } + } +} + +private fun IrBuilderWithScope.handleResource(irType: IrType, context: CompilerPluginContext): IrExpression? { + return if (irType.classFqName == RESOURCE_FQN) { + // Get the class of the type argument + val serviceType = ((irType as? IrSimpleType)?.arguments?.singleOrNull() as? IrTypeProjection)?.type + val irClass = serviceType?.classOrNull?.owner + + if (irClass != null && irClass.allSuperInterfaces().any { it.kotlinFqName == REMOTE_RESOURCE_FQN }) { + val remoteApiDescriptor = getRemoteApiDescriptorInstance(serviceType) { + "Resource generation" + } + irCallConstructor(context.remoteKindResourceConstructor, listOf(remoteApiDescriptor.type)).apply { + arguments[0] = remoteApiDescriptor + } + } + else { + null + } + } + else { + null + } + +} + +private fun IrBuilderWithScope.handleRemoteObject(irType: IrType, context: CompilerPluginContext): IrExpression? { + val irClass = irType.classOrFail.owner + + return if (irClass.allSuperInterfaces().any { it.kotlinFqName == REMOTE_OBJECT_FQN }) { + val remoteApiDescriptor = getRemoteApiDescriptorInstance(irType) { + "RemoteObject generation" + } + irCallConstructor(context.remoteKindRemoteObjectConstructor, listOf(remoteApiDescriptor.type)).apply { + arguments[0] = remoteApiDescriptor + } + } + else { + null + } +} + +private fun IrBuilderWithScope.generateSerializerCall(irType: IrType, + context: CompilerPluginContext, + debugInfo: String, + isTypeArgument: Boolean = false): IrExpression { + val irClass = irType.classOrFail.owner + + // fallback to the class itself, for example Unit serializer() is defined on Unit since it has no Companion + val companionOrClass = irClass.companionObject() ?: irClass + + val serializerCall = + // check if the class is one of the special ones that don't have serializer on Companion (List/Set/etc) + context.getSpecialSerializer(irClass.kotlinFqName)?.let { irCall(it) } + ?: findBuiltinSerializerExtensionFunctions(companionOrClass, context) + // check that class has @Serializable annotation before trying to find serializer() method on it + ?: irClass.getAnnotation(FqName.fromSegments(listOf("kotlinx", "serialization", "Serializable")))?.let { + findSerializerAsMethod(companionOrClass) + } + + val serializer = serializerCall?.apply { + val valueParameters = symbol.owner.parameters.filter { it.kind == IrParameterKind.Regular } + + (irType as IrSimpleType).arguments.zip(valueParameters).forEachIndexed { index, (typeArgument, parameter) -> + arguments[parameter.indexInParameters] = generateSerializerCall(typeArgument.typeOrFail, context, debugInfo, isTypeArgument = true) + typeArguments[index] = typeArgument.typeOrFail + } + } + + return if (serializer == null) { + val message = "Couldn't find serializer function on companion object for ${irClass.dumpKotlinLike()}\nAdditional info:\n$debugInfo" + if (isTypeArgument) { + context.messageCollector.report( + CompilerMessageSeverity.WARNING, + message + ) + + irCallConstructor(context.throwingSerializerConstructor, listOf(this.context.irBuiltIns.stringType)).apply { + val errorString = "ThrowingSerializer was used for serialization/deserialization. " + + "This could happen if there is a @Serializable class with non-serializable type argument in rpc interface.\n" + + "Additional info:\n$debugInfo" + arguments[0] = irString(errorString) + } + } else { + error(message) + } + } else if (irType.isNullable()) { + irCall(context.nullableSerializerProperty.getter!!).apply { + insertExtensionReceiver(serializer) + } + } + else { + serializer + } +} + +// find .serializer() as an extension function +private fun IrBuilderWithScope.findBuiltinSerializerExtensionFunctions( + companionOrClass: IrClass, + context: CompilerPluginContext, +): IrFunctionAccessExpression? { + val serializerFunction = + context.pluginContext.referenceFunctions(CallableId( + packageName = FqName.fromSegments(listOf("kotlinx", "serialization", "builtins")), + className = null, + callableName = "serializer".name + )).singleOrNullOrThrow { fn -> + fn.owner.parameters.firstOrNull { it.kind == IrParameterKind.ExtensionReceiver }?.type == companionOrClass.defaultType + } + + return serializerFunction?.let { + irCall(serializerFunction).apply { + insertExtensionReceiver(irGetObject(companionOrClass.symbol)) + } + } +} + +// find .serializer() as a method of irClass +private fun IrBuilderWithScope.findSerializerAsMethod(irClass: IrClass): IrFunctionAccessExpression? { + return irClass.functions.asIterable().singleOrNullOrThrow { function -> + function.name.identifierOrNullIfSpecial == "serializer" && + function.parameters.singleOrNull { it.kind == IrParameterKind.Regular }?.isVararg != true + }?.let { + irCall(it).apply { + dispatchReceiver = irGetObject(irClass.symbol) + } + } +} + +private val CompilerPluginContext.nullableSerializerProperty + get() = cache.remember { + pluginContext.referenceProperties(CallableId( + packageName = FqName.fromSegments(listOf("kotlinx", "serialization", "builtins")), + className = null, + callableName = "nullable".name + )).first().owner + } + +private val CompilerPluginContext.remoteKindDataConstructor + get() = cache.remember { + pluginContext.referenceClass( + ClassId(RPC_FQN, FqName.fromSegments(listOf("RemoteKind", "Data")), false) + )!!.constructors.first() + } + +private val CompilerPluginContext.remoteKindRemoteObjectConstructor + get() = cache.remember { + pluginContext.referenceClass( + ClassId(RPC_FQN, FqName.fromSegments(listOf("RemoteKind", "RemoteObject")), false) + )!!.constructors.first() + } + +private val CompilerPluginContext.remoteKindResourceConstructor + get() = cache.remember { + pluginContext.referenceClass( + ClassId(RPC_FQN, FqName.fromSegments(listOf("RemoteKind", "Resource")), false) + )!!.constructors.first() + } + +private val CompilerPluginContext.throwingSerializerConstructor + get() = pluginContext.referenceClass( + ClassId.topLevel(RPC_FQN.child("core".name).child("ThrowingSerializer".name)) + )!!.constructors.first() + +/** + * Same idea as [kotlin.collections.singleOrNull] but will throw if the collection contains more than one element. + * */ +private inline fun Iterable.singleOrNullOrThrow(p: (T) -> Boolean = { true }): T? { + var single: T? = null + var found = false + for (element in this) { + if (p(element)) { + if (found) { + throw IllegalArgumentException("Collection contains more than one matching element: $single, $element") + } + single = element + found = true + } + } + return single +} + +val REMOTE_KIND_FQN = RPC_FQN.child("RemoteKind".name) + +private val REMOTE_OBJECT_FQN = FqName.fromSegments(listOf("fleet", "rpc", "core", "RemoteObject")) +private val REMOTE_RESOURCE_FQN = FqName.fromSegments(listOf("fleet", "rpc", "core", "RemoteResource")) diff --git a/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/remoteKind/SpecialSerializers.kt b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/remoteKind/SpecialSerializers.kt new file mode 100644 index 000000000000..8b7d6d6e0124 --- /dev/null +++ b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/remoteKind/SpecialSerializers.kt @@ -0,0 +1,57 @@ +package com.jetbrains.fleet.rpc.plugin.ir.remoteKind + +import org.jetbrains.kotlin.ir.builders.IrBuilderWithScope +import org.jetbrains.kotlin.ir.builders.irCallConstructor +import org.jetbrains.kotlin.ir.expressions.IrExpression +import org.jetbrains.kotlin.ir.util.constructors +import org.jetbrains.kotlin.name.ClassId +import org.jetbrains.kotlin.name.FqName +import com.jetbrains.fleet.rpc.plugin.ir.CompilerPluginContext +import com.jetbrains.fleet.rpc.plugin.ir.util.RPC_FQN +import com.jetbrains.fleet.rpc.plugin.ir.util.name +import org.jetbrains.kotlin.ir.builders.irBoolean +import org.jetbrains.kotlin.ir.types.* +import org.jetbrains.kotlin.ir.util.isNullable + +internal fun IrBuilderWithScope.handleSpecialTypes(irType: IrType, context: CompilerPluginContext, debugInfo: String): IrExpression? { + return context.fqnToRemoteKind(irType.classFqName)?.let { constructor -> + val elementKindType = (irType as IrSimpleType).arguments.first().typeOrFail + val remoteKind = toRemoteKind(elementKindType, context, debugInfo) + irCallConstructor(constructor, listOf(remoteKind.type)).apply { + arguments[0] = remoteKind + arguments[1] = irBoolean(irType.isNullable()) + } + } +} + +private fun CompilerPluginContext.fqnToRemoteKind(fqName: FqName?) = cache.remember { + val remoteKindFlowConstructor = pluginContext.referenceClass( + ClassId(RPC_FQN, FqName.fromSegments(listOf("RemoteKind", "Flow")), false) + )!!.constructors.first() + + val remoteKindReceiveChannelConstructor = pluginContext.referenceClass( + ClassId(RPC_FQN, FqName.fromSegments(listOf("RemoteKind", "ReceiveChannel")), false) + )!!.constructors.first() + + val remoteKindSendChannelConstructor = pluginContext.referenceClass( + ClassId(RPC_FQN, FqName.fromSegments(listOf("RemoteKind", "SendChannel")), false) + )!!.constructors.first() + + val remoteKindDeferredConstructor = pluginContext.referenceClass( + ClassId(RPC_FQN, FqName.fromSegments(listOf("RemoteKind", "Deferred")), false) + )!!.constructors.first() + + mapOf( + FLOW_FQN to remoteKindFlowConstructor, + RECEIVE_CHANNEL_FQN to remoteKindReceiveChannelConstructor, + SEND_CHANNEL_FQN to remoteKindSendChannelConstructor, + DEFERRED_FQN to remoteKindDeferredConstructor, + ) +}[fqName] + + +internal val FLOW_FQN = FqName.fromSegments(listOf("kotlinx", "coroutines", "flow", "Flow")) +internal val RESOURCE_FQN = FqName.fromSegments(listOf("fleet", "util", "async", "Resource")) +private val RECEIVE_CHANNEL_FQN = FqName.fromSegments(listOf("kotlinx", "coroutines", "channels", "ReceiveChannel")) +private val SEND_CHANNEL_FQN = FqName.fromSegments(listOf("kotlinx", "coroutines", "channels", "SendChannel")) +private val DEFERRED_FQN = FqName.fromSegments(listOf("kotlinx", "coroutines", "Deferred")) diff --git a/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/util/Cache.kt b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/util/Cache.kt new file mode 100644 index 000000000000..5813b85edb2f --- /dev/null +++ b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/util/Cache.kt @@ -0,0 +1,7 @@ +package com.jetbrains.fleet.rpc.plugin.ir.util + +class Cache { + private val cache = mutableMapOf() + + fun remember(block: () -> T): T = cache.getOrPut(block) { block() } as T +} \ No newline at end of file diff --git a/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/util/CollectRpcMethods.kt b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/util/CollectRpcMethods.kt new file mode 100644 index 000000000000..41a6932e7c3f --- /dev/null +++ b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/util/CollectRpcMethods.kt @@ -0,0 +1,48 @@ +package com.jetbrains.fleet.rpc.plugin.ir.util + +import com.jetbrains.fleet.rpc.plugin.ir.remoteKind.FLOW_FQN +import com.jetbrains.fleet.rpc.plugin.ir.remoteKind.RESOURCE_FQN +import org.jetbrains.kotlin.backend.wasm.ir2wasm.allSuperInterfaces +import org.jetbrains.kotlin.ir.declarations.IrClass +import org.jetbrains.kotlin.ir.declarations.IrSimpleFunction +import org.jetbrains.kotlin.ir.types.* +import org.jetbrains.kotlin.ir.util.superTypes + +fun collectRpcMethods(irClass: IrClass): List { + return irClass.allSuperInterfaces().flatMap { superType -> + getRpcMethods(superType) + } +} + +/** + * /not suspend/ fun(): Flow<*> + */ +fun IrSimpleFunction.isNonSuspendFlowFunction(): Boolean { + return !isSuspend && returnType.isClassType(FLOW_FQN.toUnsafe(), false) +} + +/** + * /not suspend/ fun(): Resource + */ +fun IrSimpleFunction.isNonSuspendResourceFunction(): Boolean { + return when { + isSuspend -> false + !returnType.isClassType(RESOURCE_FQN.toUnsafe(), false) -> false + else -> { + when (val arg = (returnType as IrSimpleType).arguments.first()) { + is IrStarProjection -> false + is IrTypeProjection -> { + arg.type.superTypes().any { it.classFqName == REMOTE_RESOURCE_FQN } + } + } + } + } +} + +private val FUN_NAMES_TO_SKIP = listOf("equals", "hashCode", "toString", "dispatch", "signature") + +private fun getRpcMethods(rpcClass: IrClass): List { + return rpcClass.declarations.filterIsInstance().filter { declaration -> + !declaration.isFakeOverride && declaration.overriddenSymbols.isEmpty() && declaration.name.identifier !in FUN_NAMES_TO_SKIP + } +} diff --git a/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/util/IrUtil.kt b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/util/IrUtil.kt new file mode 100644 index 000000000000..3694bd3e09ce --- /dev/null +++ b/fleet/compiler-plugins/rpc/src/main/kotlin/com/jetbrains/fleet/rpc/plugin/ir/util/IrUtil.kt @@ -0,0 +1,250 @@ +package com.jetbrains.fleet.rpc.plugin.ir.util + +import com.jetbrains.fleet.rpc.plugin.ir.CompilerPluginContext +import org.jetbrains.kotlin.backend.common.extensions.IrPluginContext +import org.jetbrains.kotlin.backend.common.ir.addDispatchReceiver +import org.jetbrains.kotlin.backend.common.lower.DeclarationIrBuilder +import org.jetbrains.kotlin.descriptors.ClassKind +import org.jetbrains.kotlin.descriptors.DescriptorVisibilities +import org.jetbrains.kotlin.descriptors.Modality +import org.jetbrains.kotlin.ir.UNDEFINED_OFFSET +import org.jetbrains.kotlin.ir.builders.IrBlockBodyBuilder +import org.jetbrains.kotlin.ir.builders.IrBuilderWithScope +import org.jetbrains.kotlin.ir.builders.declarations.* +import org.jetbrains.kotlin.ir.builders.irBlockBody +import org.jetbrains.kotlin.ir.builders.irCall +import org.jetbrains.kotlin.ir.declarations.* +import org.jetbrains.kotlin.ir.expressions.IrExpression +import org.jetbrains.kotlin.ir.expressions.IrStatementOrigin +import org.jetbrains.kotlin.ir.expressions.impl.* +import org.jetbrains.kotlin.ir.symbols.IrFunctionSymbol +import org.jetbrains.kotlin.ir.symbols.IrSimpleFunctionSymbol +import org.jetbrains.kotlin.ir.symbols.IrValueSymbol +import org.jetbrains.kotlin.ir.symbols.impl.IrAnonymousInitializerSymbolImpl +import org.jetbrains.kotlin.ir.symbols.impl.IrValueParameterSymbolImpl +import org.jetbrains.kotlin.ir.types.IrSimpleType +import org.jetbrains.kotlin.ir.types.IrType +import org.jetbrains.kotlin.ir.types.impl.IrSimpleTypeImpl +import org.jetbrains.kotlin.ir.types.makeNullable +import org.jetbrains.kotlin.ir.types.typeOrFail +import org.jetbrains.kotlin.ir.types.typeWith +import org.jetbrains.kotlin.ir.util.SYNTHETIC_OFFSET +import org.jetbrains.kotlin.ir.util.constructors +import org.jetbrains.kotlin.ir.util.defaultType +import org.jetbrains.kotlin.ir.util.isAnnotationWithEqualFqName +import org.jetbrains.kotlin.name.FqName +import org.jetbrains.kotlin.name.Name +import org.jetbrains.kotlin.name.SpecialNames + +val RPC_FQN = FqName.fromSegments(listOf("fleet", "rpc")) +val RPC_ANNOTATION_FQN = RPC_FQN.child("Rpc".name) + +val RPC_CORE_FQN: FqName = RPC_FQN.child("core".name) +val REMOTE_RESOURCE_FQN: FqName = RPC_CORE_FQN.child("RemoteResource".name) + +val String.name: Name + get() = Name.identifier(this) + +fun irGet(type: IrType, symbol: IrValueSymbol, origin: IrStatementOrigin?): IrExpression { + return IrGetValueImpl( + UNDEFINED_OFFSET, + UNDEFINED_OFFSET, + type, + symbol, + origin + ) +} + +fun irGet(variable: IrValueDeclaration, origin: IrStatementOrigin? = null): IrExpression { + return irGet(variable.type, variable.symbol, origin) +} + +fun irCall(symbol: IrFunctionSymbol, origin: IrStatementOrigin? = null): IrCallImpl { + return IrCallImpl.fromSymbolOwner( + startOffset = UNDEFINED_OFFSET, + endOffset = UNDEFINED_OFFSET, + type = symbol.owner.returnType, + symbol = symbol as IrSimpleFunctionSymbol, + origin = origin, + ) +} + +fun IrClass.hasRpcAnnotation() = + this.annotations.any { it.isAnnotationWithEqualFqName(RPC_ANNOTATION_FQN) } + + +fun IrPluginContext.buildClassBase(customizer: IrClassBuilder.() -> Unit) = + irFactory.buildClass { + startOffset = SYNTHETIC_OFFSET + endOffset = SYNTHETIC_OFFSET + origin = IrDeclarationOrigin.DEFINED + kind = ClassKind.CLASS + modality = Modality.FINAL + customizer() + }.also { + thisReceiver(it) + constructor(it) + initializer(it) + } + +private fun IrPluginContext.thisReceiver(irClass: IrClass): IrValueParameter = + irFactory.createValueParameter( + startOffset = SYNTHETIC_OFFSET, + endOffset = SYNTHETIC_OFFSET, + origin = IrDeclarationOrigin.INSTANCE_RECEIVER, + symbol = IrValueParameterSymbolImpl(), + name = SpecialNames.THIS, + type = IrSimpleTypeImpl(irClass.symbol, false, emptyList(), emptyList()), + varargElementType = null, + isCrossinline = false, + isNoinline = false, + isHidden = false, + isAssignable = false, + kind = IrParameterKind.DispatchReceiver + ).also { + it.parent = irClass + irClass.thisReceiver = it + } + +private fun IrPluginContext.constructor(irClass: IrClass): IrConstructor = + irClass.addConstructor { + isPrimary = true + returnType = irClass.defaultType + }.apply { + parent = irClass + + body = irFactory.createBlockBody(SYNTHETIC_OFFSET, SYNTHETIC_OFFSET).apply { + statements += IrDelegatingConstructorCallImpl.fromSymbolOwner( + SYNTHETIC_OFFSET, + SYNTHETIC_OFFSET, + irBuiltIns.anyType, + irBuiltIns.anyClass.constructors.first(), + typeArgumentsCount = 0, + ) + + statements += IrInstanceInitializerCallImpl( + SYNTHETIC_OFFSET, + SYNTHETIC_OFFSET, + irClass.symbol, + irBuiltIns.unitType + ) + } + } + +private fun IrPluginContext.initializer(irClass: IrClass): IrAnonymousInitializer = + irFactory.createAnonymousInitializer( + SYNTHETIC_OFFSET, SYNTHETIC_OFFSET, + origin = IrDeclarationOrigin.DEFINED, + symbol = IrAnonymousInitializerSymbolImpl(), + isStatic = false + ).apply { + parent = irClass + } + +fun IrPluginContext.getProxyType(): IrSimpleType { // typealias Proxy = suspend (String, Array) -> Any? + val stringType = irBuiltIns.stringType + val nullableAnyType = irBuiltIns.anyNType.makeNullable() + val anyArrayType = irBuiltIns.arrayClass.typeWith(nullableAnyType) + + return irBuiltIns.suspendFunctionN(2) + .typeWith(/*thisType, */stringType, anyArrayType, nullableAnyType) +} + +fun IrBuilderWithScope.createArrayOfExpression( + arrayElementType: IrType, + arrayElements: List, + pluginContext: IrPluginContext +): IrExpression { + val arrayType = pluginContext.irBuiltIns.arrayClass.typeWith(arrayElementType) + val varargOfElements = IrVarargImpl(startOffset, endOffset, arrayType, arrayElementType, arrayElements) + val typeArguments = listOf(arrayElementType) + + return irCall(pluginContext.irBuiltIns.arrayOf, arrayType, typeArguments = typeArguments).apply { + arguments[0] = varargOfElements + } +} + +fun IrPluginContext.addFunctionOverride(originalFunction: IrSimpleFunction, + parentClass: IrClass, + block: (IrSimpleFunction) -> Unit) { + irFactory.buildFun { + name = originalFunction.name + returnType = originalFunction.returnType + modality = Modality.FINAL + }.also { function -> + function.isSuspend = originalFunction.isSuspend + function.overriddenSymbols = listOf(originalFunction.symbol) + function.parent = parentClass + + for (typeParameter in originalFunction.typeParameters) { + function.addTypeParameter { + updateFrom(typeParameter) + superTypes += typeParameter.superTypes + } + } + + // Prepend dispatch receiver (should be first in parameter list) + function.parameters += function.buildReceiverParameter { + type = parentClass.defaultType + } + + for (parameter in originalFunction.parameters) { + // Dispatch receiver was specified right above + // TODO review: maybe this dispatch receiver is already present with proper type on original function? (then this can be simplified) + if (parameter.kind != IrParameterKind.DispatchReceiver) { + function.addValueParameter { + name = parameter.name + updateFrom(parameter) + } + } + } + + block(function) + + parentClass.declarations += function + } +} + +fun IrFunction.generateBody(context: CompilerPluginContext, block: IrBlockBodyBuilder.() -> Unit) { + body = DeclarationIrBuilder(context.pluginContext, symbol).irBlockBody { + block() + } +} + +/** + * Creates a lambda based on the provided function type. + * + * @param type type of the function + */ +internal fun IrPluginContext.irLambda( + symbol: IrSymbolOwner, + parent: IrDeclarationParent, + type: IrSimpleType, + // Default values assume the `type` parameter is `FunctionN` class type (arg types first, last one is return type) + parameterTypes: List = type.arguments.dropLast(1).map { it.typeOrFail }, + returnType: IrType = type.arguments.last().typeOrFail, + block: (IrFunction) -> Unit, +): IrFunctionExpressionImpl { + return IrFunctionExpressionImpl( + startOffset = symbol.startOffset, + endOffset = symbol.endOffset, + type = type, + function = irFactory.buildFun { + origin = IrDeclarationOrigin.LOCAL_FUNCTION_FOR_LAMBDA + name = SpecialNames.NO_NAME_PROVIDED + visibility = DescriptorVisibilities.LOCAL + this.returnType = returnType + modality = Modality.FINAL + isSuspend = true + }.also { + it.parent = parent + + parameterTypes.forEachIndexed { index, typeArgument -> + it.addValueParameter("p${index}", typeArgument.typeOrFail) + } + + block(it) + }, + origin = IrStatementOrigin.LAMBDA + ) +} diff --git a/fleet/compiler-plugins/rpc/src/main/resources/META-INF/services/org.jetbrains.kotlin.compiler.plugin.CommandLineProcessor b/fleet/compiler-plugins/rpc/src/main/resources/META-INF/services/org.jetbrains.kotlin.compiler.plugin.CommandLineProcessor new file mode 100644 index 000000000000..2f5379896442 --- /dev/null +++ b/fleet/compiler-plugins/rpc/src/main/resources/META-INF/services/org.jetbrains.kotlin.compiler.plugin.CommandLineProcessor @@ -0,0 +1 @@ +com.jetbrains.fleet.rpc.plugin.RpcCommandLineProcessor \ No newline at end of file diff --git a/fleet/compiler-plugins/rpc/src/main/resources/META-INF/services/org.jetbrains.kotlin.compiler.plugin.CompilerPluginRegistrar b/fleet/compiler-plugins/rpc/src/main/resources/META-INF/services/org.jetbrains.kotlin.compiler.plugin.CompilerPluginRegistrar new file mode 100644 index 000000000000..f20ce7a247b2 --- /dev/null +++ b/fleet/compiler-plugins/rpc/src/main/resources/META-INF/services/org.jetbrains.kotlin.compiler.plugin.CompilerPluginRegistrar @@ -0,0 +1 @@ +com.jetbrains.fleet.rpc.plugin.RpcComponentRegistrar \ No newline at end of file