[Kotlin]: Factor method argument conversion with KotlinRecord.

Adds support for function pointers (and userdata), void* and nullable
string to the macro generating the conversion helper.

Change handling of JNI methods to first put all of their arguments in
the KotlinRecord, convert that to a structure containing all the C
arguments, and then call the function with arguments taken from this
arguments structure.

Bug: 352711433
Change-Id: Iaadb5eda237c929b4dc2d2321d980a6303a9fe3e
Reviewed-on: https://dawn-review.googlesource.com/c/dawn/+/198255
Commit-Queue: Corentin Wallez <cwallez@chromium.org>
Reviewed-by: Sonakshi Saxena <nexa@google.com>
Reviewed-by: Jim Blackler <jimblackler@google.com>
Commit-Queue: Sonakshi Saxena <nexa@google.com>
diff --git a/generator/templates/art/kotlin_record_conversion.cpp b/generator/templates/art/kotlin_record_conversion.cpp
index e118ce6..c33b6b4 100644
--- a/generator/templates/art/kotlin_record_conversion.cpp
+++ b/generator/templates/art/kotlin_record_conversion.cpp
@@ -24,7 +24,7 @@
 //* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
 //* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
 //* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
-{% from 'art/api_jni_types.kt' import arg_to_jni_type with context %}
+{% from 'art/api_jni_types.kt' import arg_to_jni_type, convert_to_kotlin, jni_signature with context %}
 
 {% macro define_kotlin_record_structure(struct_name, members) %}
     struct {{struct_name}} {
@@ -81,6 +81,10 @@
                     out = reinterpret_cast<const uint32_t*>(c->GetIntArrayElements(in));
                     outLength = env->GetArrayLength(in);
 
+                {% elif member.type.name.get() == 'void' %}
+                    out = env->GetDirectBufferAddress(in);
+                    outLength = env->GetDirectBufferCapacity(in);
+
                 {% else %}
                     //* These container types are represented in Kotlin as arrays of objects.
                     outLength = env->GetArrayLength(in);
@@ -134,6 +138,44 @@
 
             {% elif member.type.category in ["native", "enum", "bitmask"] %}
                 out = static_cast<{{as_cType(member.type.name)}}>(in);
+
+            {% elif member.type.category == 'function pointer' %}
+                //* Function pointers themselves require each argument converting.
+                //* A custom native callback is generated to wrap the Kotlin callback.
+                out = [](
+                    {%- for callbackArg in member.type.arguments %}
+                        {{ as_annotated_cType(callbackArg) }}{{ ',' if not loop.last }}
+                    {%- endfor %}) {
+                    UserData* userData1 = static_cast<UserData *>(userdata);
+                    JNIEnv *env = userData1->env;
+                    if (env->ExceptionCheck()) {
+                        return;
+                    }
+
+                    {%- for callbackArg in kotlin_record_members(member.type.arguments) -%}
+                        {{ convert_to_kotlin(callbackArg.name.camelCase(),
+                                             '_' + callbackArg.name.camelCase(),
+                                             'input->' + callbackArg.length.name.camelCase() if callbackArg.length.name,
+                                             callbackArg) }}
+                    {% endfor %}
+
+                    //* Get the client (Kotlin) callback so we can call it.
+                    jmethodID callbackMethod = env->GetMethodID(
+                            env->FindClass("{{ jni_name(member.type) }}"), "callback", "(
+                        {%- for callbackArg in kotlin_record_members(member.type.arguments) -%}
+                            {{- jni_signature(callbackArg) -}}
+                        {%- endfor %})V");
+
+                    //* Call the callback with all converted parameters.
+                    env->CallVoidMethod(userData1->callback, callbackMethod
+                    {%- for callbackArg in kotlin_record_members(member.type.arguments) %}
+                         ,_{{ callbackArg.name.camelCase() }}
+                    {%- endfor %});
+                };
+                //* TODO(b/330293719): free associated resources.
+                outStruct->userdata = new UserData(
+                        {.env = env, .callback = env->NewGlobalRef(in)});
+
             {% else %}
                 {{ unreachable_code() }}
             {% endif %}
diff --git a/generator/templates/art/methods.cpp b/generator/templates/art/methods.cpp
index 4d55795..c48d96d 100644
--- a/generator/templates/art/methods.cpp
+++ b/generator/templates/art/methods.cpp
@@ -24,6 +24,7 @@
 //* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
 //* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
 //* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
+{% from 'art/kotlin_record_conversion.cpp' import define_kotlin_record_structure, define_kotlin_to_struct_conversion with context %}
 {% from 'art/api_jni_types.kt' import arg_to_jni_type, convert_to_kotlin, jni_signature, to_jni_type with context %}
 #include <jni.h>
 #include <stdlib.h>
@@ -66,13 +67,30 @@
 }
 
 {% macro render_method(method, object) %}
+    {% set ObjectName = object.name.CamelCase() if object else "FunctionsKt" %}
+    {% set FunctionSuffix = ObjectName + "_" +  method.name.camelCase() %}
+    {% set KotlinRecord = FunctionSuffix + "KotlinRecord" %}
+    {% set ArgsStruct = FunctionSuffix + "ArgsStruct" %}
+
+    //* Define the helper structs to perform most of the conversion.
+    struct {{KotlinRecord}} {
+        {% for arg in kotlin_record_members(method.arguments) %}
+            {{ arg_to_jni_type(arg) }} {{ as_varName(arg.name) }};
+        {% endfor %}
+    };
+    struct {{ArgsStruct}} {
+        {% for arg in method.arguments %}
+            {{ as_annotated_cType(arg) }};
+        {% endfor %}
+    };
+    {{ define_kotlin_to_struct_conversion("ConvertInternal", KotlinRecord, ArgsStruct, method.arguments)}}
+
     {% set _kotlin_return = kotlin_return(method) %}
     //*  A JNI-external method is built with the JNI signature expected to match the host Kotlin.
     DEFAULT extern "C"
-    {{ arg_to_jni_type(_kotlin_return) }} Java_{{ kotlin_package.replace('.', '_') }}_
-            {{- object.name.CamelCase() if object else 'FunctionsKt' -}}
-            _{{ method.name.camelCase() }}(JNIEnv *env
-                    {{ ', jobject obj' if object else ', jclass clazz' -}}
+    {{ arg_to_jni_type(_kotlin_return) }}
+    Java_{{ kotlin_package.replace('.', '_') }}_{{ FunctionSuffix }}
+            (JNIEnv *env{{ ', jobject obj' if object else ', jclass clazz' -}}
 
     //* Make the signature for each argument in turn.
     {% for arg in kotlin_record_members(method.arguments) %},
@@ -82,131 +100,13 @@
     // * Helper context for the duration of this method call.
     JNIContext c(env);
 
-    //*  A variable is declared for each parameter of the native method.
-    {% for arg in method.arguments %}
-        {{ as_annotated_cType(arg) }};
-    {% endfor %}
-
-    //* Each parameter is converted from the JNI parameter to the expected form of the native
-    //* parameter.
+    //* Perform the conversion of arguments.
+    {{KotlinRecord}} kotlinRecord;
     {% for arg in kotlin_record_members(method.arguments) %}
-        {% if arg.length == 'strlen' %}
-            if (_{{ as_varName(arg.name) }}) {  //* Don't convert null strings.
-                {{ as_varName(arg.name) }} = c.GetStringUTFChars(_{{ as_varName(arg.name) }});
-            } else {
-                {{ as_varName(arg.name) }} = nullptr;
-            }
-        {% elif arg.constant_length == 1 %}
-            //*  Optional structure.
-            {% if arg.type.category == 'structure' %}
-                if (_{{ as_varName(arg.name) }}) {
-                    auto convertedMember = c.Alloc<{{ as_cType(arg.type.name) }}>();
-                    ToNative(&c, env, _{{ as_varName(arg.name) }}, convertedMember);
-                    {{ as_varName(arg.name) }} = convertedMember;
-                } else {
-                    {{ as_varName(arg.name) }} = nullptr;
-                }
-            {% else %}
-                {{ unreachable_code() }}
-            {% endif %}
-        {% elif arg.length %}
-            //*  Container types.
-            {% if arg.type.name.get() == 'uint32_t' %}
-                {{ as_varName(arg.name) }} =
-                        reinterpret_cast<const {{ as_cType(arg.type.name) }}*>(
-                               c.GetIntArrayElements(_{{ as_varName(arg.name) }}));
-                {{ arg.length.name.camelCase() }} =
-                       env->GetArrayLength(_{{ as_varName(arg.name) }});
-            {% elif arg.type.name.get() == 'void' %}
-                {{ as_varName(arg.name) }} =
-                        env->GetDirectBufferAddress(_{{ as_varName(arg.name) }});
-                {{ arg.length.name.camelCase() }} =
-                        env->GetDirectBufferCapacity(_{{ as_varName(arg.name) }});
-            {% else %} {
-                size_t length = env->GetArrayLength(_{{ as_varName(arg.name) }});
-                auto out = c.AllocArray<{{ as_cType(arg.type.name) }}>(length);
-                {% if arg.type.category in ['bitmask', 'enum'] %} {
-                    jclass memberClass = env->FindClass("{{ jni_name(arg.type) }}");
-                    jmethodID getValue = env->GetMethodID(memberClass, "getValue", "()I");
-                    for (int idx = 0; idx != length; idx++) {
-                        jobject element =
-                                env->GetObjectArrayElement(_{{ as_varName(arg.name) }}, idx);
-                        out[idx] = static_cast<{{ as_cType(arg.type.name) }}>(
-                                env->CallIntMethod(element, getValue));
-                    }
-                } {% elif arg.type.category == 'object' %} {
-                    jclass memberClass = env->FindClass("{{ jni_name(arg.type) }}");
-                    jmethodID getHandle = env->GetMethodID(memberClass, "getHandle", "()J");
-                    for (int idx = 0; idx != length; idx++) {
-                        jobject element =
-                                env->GetObjectArrayElement(_{{ as_varName(arg.name) }}, idx);
-                        out[idx] = reinterpret_cast<{{ as_cType(arg.type.name) }}>(
-                                env->CallLongMethod(element, getHandle));
-                    }
-                } {% else %}
-                    {{ unreachable_code() }}
-                {% endif %}
-                {{ as_varName(arg.name) }} = out;
-                {{ arg.length.name.camelCase() }} = length;
-            } {% endif %}
-
-        //*  Single value types.
-        {% elif arg.type.category == 'object' %}
-            if (_{{ as_varName(arg.name) }}) {
-                jclass memberClass = env->FindClass("{{ jni_name(arg.type) }}");
-                jmethodID getHandle = env->GetMethodID(memberClass, "getHandle", "()J");
-                {{ as_varName(arg.name) }} =
-                        reinterpret_cast<{{ as_cType(arg.type.name) }}>(
-                                env->CallLongMethod(_{{ as_varName(arg.name) }}, getHandle));
-            } else {
-                {{ as_varName(arg.name) }} = nullptr;
-            }
-        {% elif arg.type.name.get() in ['int32_t', 'size_t', 'uint32_t', 'uint64_t'] or arg.type.category in ['bitmask', 'enum'] %}
-            {{ as_varName(arg.name) }} =
-                    static_cast<{{ as_cType(arg.type.name) }}>(_{{ as_varName(arg.name) }});
-        {% elif arg.type.name.get() in ['float', 'int'] %}
-            {{ as_varName(arg.name) }} = _{{ as_varName(arg.name) }};
-
-        {% elif arg.type.category == 'function pointer' %} {
-            //* Function pointers themselves require each argument converting.
-            //* A custom native callback is generated to wrap the Kotlin callback.
-            {{ as_varName(arg.name) }} = [](
-                {%- for callbackArg in arg.type.arguments %}
-                    {{ as_annotated_cType(callbackArg) }}{{ ',' if not loop.last }}
-                {%- endfor %}) {
-                UserData* userData1 = static_cast<UserData *>(userdata);
-                JNIEnv *env = userData1->env;
-                if (env->ExceptionCheck()) {
-                    return;
-                }
-
-                {%- for callbackArg in kotlin_record_members(arg.type.arguments) -%}
-                    {{ convert_to_kotlin(callbackArg.name.camelCase(),
-                                         '_' + callbackArg.name.camelCase(),
-                                         'input->' + callbackArg.length.name.camelCase() if callbackArg.length.name,
-                                         callbackArg) }}
-                {% endfor %}
-
-                //* Get the client (Kotlin) callback so we can call it.
-                jmethodID callbackMethod = env->GetMethodID(
-                        env->FindClass("{{ jni_name(arg.type) }}"), "callback", "(
-                    {%- for callbackArg in kotlin_record_members(arg.type.arguments) -%}
-                        {{- jni_signature(callbackArg) -}}
-                    {%- endfor %})V");
-
-                //* Call the callback with all converted parameters.
-                env->CallVoidMethod(userData1->callback, callbackMethod
-                {%- for callbackArg in kotlin_record_members(arg.type.arguments) %}
-                    ,_{{ callbackArg.name.camelCase() }}
-                {%- endfor %});
-            };
-            //* TODO(b/330293719): free associated resources.
-            userdata = new UserData(
-                    {.env = env, .callback = env->NewGlobalRef(_{{ as_varName(arg.name) }})});
-        } {% else %}
-            {{ unreachable_code() }}
-        {% endif %}
+        kotlinRecord.{{ as_varName(arg.name) }} = _{{ as_varName(arg.name) }};
     {% endfor %}
+    {{ArgsStruct}} args;
+    ConvertInternal(&c, kotlinRecord, &args);
 
     {% if object %}
         jclass memberClass = env->FindClass("{{ jni_name(object) }}");
@@ -222,7 +122,7 @@
         size_t size = wgpu{{ object.name.CamelCase() }}{{ method.name.CamelCase() }}(handle
             {% for arg in method.arguments -%},
                 //* The replaced output parameter is set to nullptr on the first call.
-                {{ 'nullptr' if arg.annotation == '*' else as_varName(arg.name) -}}
+                {{ 'nullptr' if arg.annotation == '*' else "args." + as_varName(arg.name) -}}
             {% endfor %}
         );
         //* Allocate the native container
@@ -234,7 +134,7 @@
         wgpu{{ object.name.CamelCase() }}{{ method.name.CamelCase() }}(handle
             {% for arg in method.arguments -%}
                 {{- ', ' if object or not loop.first -}}
-                {{- 'returnAllocation.get()' if arg == _kotlin_return else as_varName(arg.name) -}}
+                {{- 'returnAllocation.get()' if arg == _kotlin_return else "args." + as_varName(arg.name) -}}
             {% endfor %}
         );
         if (env->ExceptionCheck()) {  //* Early out if client (Kotlin) callback threw an exception.
@@ -244,7 +144,7 @@
         {% if _kotlin_return.annotation == '*' %}
             //* Make a native container to accept the data output via parameter.
             {{ as_cType(_kotlin_return.type.name) }} out;
-            {{ _kotlin_return.name.get() }} = &out;
+            args.{{ _kotlin_return.name.get() }} = &out;
         {% endif %}
         {{ 'auto result =' if _kotlin_return.type.name.get() != 'void' }}
         {% if object %}
@@ -253,7 +153,7 @@
             wgpu{{ method.name.CamelCase() }}(
         {% endif %}
             {% for arg in method.arguments -%}
-                {{- ',' if object or not loop.first }}{{ as_varName(arg.name) -}}
+                {{- ',' if object or not loop.first }}args.{{ as_varName(arg.name) -}}
             {% endfor %}
         );
         if (env->ExceptionCheck()) {  //* Early out if client (Kotlin) callback threw an exception.
@@ -261,7 +161,10 @@
         }
     {% endif %}
     {% if _kotlin_return.type.name.get() != 'void' %}
-        {{ convert_to_kotlin(_kotlin_return.name.get() if _kotlin_return.annotation == '*' else 'result',
+        {% if _kotlin_return.type.name.get() in ['void const *', 'void *'] %}
+            size_t size = args.size;
+        {% endif %}
+        {{ convert_to_kotlin("args." + _kotlin_return.name.get() if _kotlin_return.annotation == '*' else 'result',
                              'result_kt',
                              'size' if _kotlin_return.type.name.get() in ['void const *', 'void *'] or _kotlin_return.length == 'size_t',
                              _kotlin_return) }}