Kotlin: async adapted methods are cancelable.

Test: ./gradlew connectedAndroidTest
Bug: b/453643094
Change-Id: Iecae738e5e5bebe29b43237952cddf9e095d05fa
Reviewed-on: https://dawn-review.googlesource.com/c/dawn/+/270494
Commit-Queue: Jim Blackler <jimblackler@google.com>
Reviewed-by: Mridul Goyal <mridulgoyal@google.com>
diff --git a/generator/templates/art/api_kotlin_async_helpers.kt b/generator/templates/art/api_kotlin_async_helpers.kt
index e064339..226b44b 100644
--- a/generator/templates/art/api_kotlin_async_helpers.kt
+++ b/generator/templates/art/api_kotlin_async_helpers.kt
@@ -82,7 +82,7 @@
             (arg.type.category == 'kotlin type' and arg.type.name.get() == 'java.util.concurrent.Executor')
         ) %}
             {{ kotlin_annotation(arg) }}  {{ as_varName(arg.name) }}: {{ kotlin_definition(arg) }},
-        {%- endfor %}): {{ return_name }} = suspendCoroutine {
+        {%- endfor %}): {{ return_name }} = suspendCancellableCoroutine {
         {{ method.name.camelCase() }}(
             {%- for arg in kotlin_record_members(method.arguments) %}
                 {%- if arg.type.category == 'kotlin type' and arg.type.name.get() == 'java.util.concurrent.Executor' -%}
@@ -90,12 +90,13 @@
                 {%- elif arg.name.get() == callback_arg.name.get() %}{
                     {%- for arg in kotlin_record_members(callback_arg.type.arguments) %}
                         {{- as_varName(arg.name) }},
-                    {%- endfor %} -> it.resume({{ return_name }}(
-                        //* We make an instance of the callback parameters -> return type wrapper.
-                        {%- for arg in result_args %}
-                            {{- as_varName(arg.name) }},
-                        {%- endfor %}
-                    ))},
+                    {%- endfor %} -> if (it.isActive) {
+                        it.resume({{ return_name }}(
+                            //* We make an instance of the callback parameters -> return type wrapper.
+                            {%- for arg in result_args %}
+                                {{- as_varName(arg.name) }},
+                            {%- endfor %}))
+                    }}
                 {%- else -%}
                     {{- as_varName(arg.name) }},
                 {%- endif %}
diff --git a/generator/templates/art/api_kotlin_object.kt b/generator/templates/art/api_kotlin_object.kt
index d432fa9..e461595 100644
--- a/generator/templates/art/api_kotlin_object.kt
+++ b/generator/templates/art/api_kotlin_object.kt
@@ -30,7 +30,7 @@
 import java.nio.ByteBuffer
 import java.util.concurrent.Executor
 import kotlin.coroutines.resume
-import kotlin.coroutines.suspendCoroutine
+import kotlinx.coroutines.suspendCancellableCoroutine
 
 {% from 'art/api_kotlin_types.kt' import kotlin_annotation, kotlin_declaration, kotlin_definition, check_if_doc_present, generate_kdoc, generate_simple_kdoc with context %}
 {% from 'art/api_kotlin_async_helpers.kt' import async_wrapper with context %}
diff --git a/tools/android/webgpu/src/androidTest/java/androidx/webgpu/AsyncHelperTest.kt b/tools/android/webgpu/src/androidTest/java/androidx/webgpu/AsyncHelperTest.kt
index d651e0e..588c849 100644
--- a/tools/android/webgpu/src/androidTest/java/androidx/webgpu/AsyncHelperTest.kt
+++ b/tools/android/webgpu/src/androidTest/java/androidx/webgpu/AsyncHelperTest.kt
@@ -2,20 +2,35 @@
 
 import androidx.test.ext.junit.runners.AndroidJUnit4
 import androidx.test.filters.SmallTest
+import androidx.webgpu.helper.WebGpu
 import androidx.webgpu.helper.createWebGpu
+import kotlinx.coroutines.launch
 import kotlinx.coroutines.runBlocking
 import org.junit.Assert.assertEquals
+import org.junit.Assume.assumeFalse
+import org.junit.Before
 import org.junit.Test
 import org.junit.runner.RunWith
+import java.util.concurrent.atomic.AtomicBoolean
 
 @RunWith(AndroidJUnit4::class)
 @SmallTest
 class AsyncHelperTest {
+
+    private lateinit var webGpu: WebGpu
+    private lateinit var device: GPUDevice
+
+    @Before
+    fun setup() {
+        runBlocking {
+            webGpu = createWebGpu()
+            device = webGpu.device
+        }
+    }
+
     @Test
     fun asyncMethodTest() {
         runBlocking {
-            val webGpu = createWebGpu()
-            val device = webGpu.device
             /* Set up a shader module to support the async call. */
             val shaderModule = device.createShaderModule(
                 ShaderModuleDescriptor(shaderSourceWGSL = ShaderSourceWGSL(""))
@@ -36,8 +51,6 @@
     @Test
     fun asyncMethodTestValidationPasses() {
         runBlocking {
-            val webGpu = createWebGpu()
-            val device = webGpu.device
             /* Set up a shader module to support the async call. */
             val shaderModule = device.createShaderModule(
                 ShaderModuleDescriptor(
@@ -71,4 +84,53 @@
 
         }
     }
+
+    private fun baseCancellationTest(doCancel: Boolean): Boolean {
+        val hasReturned = AtomicBoolean(false)
+
+        runBlocking {
+            val shaderModule = device.createShaderModule(
+                ShaderModuleDescriptor(shaderSourceWGSL = ShaderSourceWGSL(""))
+            )
+
+            /* Launch the function in a new coroutine, giving us a job handle we can cancel. */
+            val job = launch {
+                device.createRenderPipelineAsync(
+                    RenderPipelineDescriptor(vertex = VertexState(module = shaderModule))
+                )
+                hasReturned.set(true)
+            }
+            assumeFalse("The job completed before we could test it", hasReturned.get())
+
+            if (doCancel) {
+                job.cancel()
+            }
+            job.join()
+        }
+        return hasReturned.get()
+    }
+
+    /**
+     * Test that the async-based job will complete if it's not cancelled.
+     */
+    @Test
+    fun asyncMethodCancellationTestControl() {
+        assertEquals(
+            "The async job should have completed but it failed to do so.",
+            true,
+            baseCancellationTest(false)
+        )
+    }
+
+    /**
+     * Test that the async-based job will not complete if it is cancelled.
+     */
+    @Test
+    fun asyncMethodCancellationTest() {
+        assertEquals(
+            "The async job should have been cancelled but it completed.",
+            false,
+            baseCancellationTest(true)
+        )
+    }
 }