diff --git a/platform/core-impl/src/com/intellij/openapi/progress/impl/ProgressRunner.java b/platform/core-impl/src/com/intellij/openapi/progress/impl/ProgressRunner.java
index 369d657af670..99182366438b 100644
--- a/platform/core-impl/src/com/intellij/openapi/progress/impl/ProgressRunner.java
+++ b/platform/core-impl/src/com/intellij/openapi/progress/impl/ProgressRunner.java
@@ -2,7 +2,10 @@
package com.intellij.openapi.progress.impl;
import com.intellij.codeWithMe.ClientId;
+import com.intellij.concurrency.ContextAwareRunnable;
+import com.intellij.concurrency.ThreadContext;
import com.intellij.openapi.Disposable;
+import com.intellij.openapi.application.AccessToken;
import com.intellij.openapi.application.ApplicationManager;
import com.intellij.openapi.application.ModalityState;
import com.intellij.openapi.application.ex.ApplicationManagerEx;
@@ -13,8 +16,14 @@ import com.intellij.openapi.progress.*;
import com.intellij.openapi.util.Disposer;
import com.intellij.openapi.wm.ex.ProgressIndicatorEx;
import com.intellij.util.concurrency.AppExecutorUtil;
+import com.intellij.util.concurrency.Propagation;
import com.intellij.util.concurrency.Semaphore;
import com.intellij.util.ui.EDT;
+import kotlin.Pair;
+import kotlin.Unit;
+import kotlin.coroutines.CoroutineContext;
+import kotlin.coroutines.EmptyCoroutineContext;
+import kotlinx.coroutines.CompletableJob;
import org.jetbrains.annotations.Contract;
import org.jetbrains.annotations.NonNls;
import org.jetbrains.annotations.NotNull;
@@ -23,6 +32,8 @@ import org.jetbrains.annotations.Nullable;
import java.util.concurrent.*;
import java.util.function.Function;
+import static com.intellij.openapi.application.ModalityKt.asContextElement;
+
/**
*
* A builder-like API for running tasks with {@link ProgressIndicator}.
@@ -438,6 +449,27 @@ public final class ProgressRunner {
@NotNull CompletableFuture extends @NotNull ProgressIndicator> progressIndicatorFuture
) {
CompletableFuture resultFuture = new CompletableFuture<>();
+ Pair childContextAndJob = Propagation.createChildContext();
+ CoroutineContext childContext = childContextAndJob.getFirst();
+ CompletableJob childJob = childContextAndJob.getSecond();
+ if (childJob != null) {
+ // cancellation of the Job cancels the future
+ childJob.invokeOnCompletion(true, true, (throwable) -> {
+ if (throwable != null) {
+ resultFuture.completeExceptionally(throwable);
+ }
+ return Unit.INSTANCE;
+ });
+ // completion of the future completes the Job
+ resultFuture.whenComplete((result, throwable) -> {
+ if (throwable == null) {
+ childJob.complete();
+ }
+ else {
+ childJob.completeExceptionally(throwable);
+ }
+ });
+ }
progressIndicatorFuture.whenComplete((progressIndicator, throwable) -> {
if (throwable != null) {
resultFuture.completeExceptionally(throwable);
@@ -451,13 +483,19 @@ public final class ProgressRunner {
resultFuture.completeExceptionally(e);
}
};
+ Runnable contextRunnable = childContext.equals(EmptyCoroutineContext.INSTANCE) ? runnable : (ContextAwareRunnable)() -> {
+ CoroutineContext effectiveContext = childContext.plus(asContextElement(progressIndicator.getModalityState()));
+ try (AccessToken ignored = ThreadContext.installThreadContext(effectiveContext, false)) {
+ runnable.run();
+ }
+ };
switch (myThreadToUse) {
case POOLED:
- AppExecutorUtil.getAppExecutorService().execute(runnable);
+ AppExecutorUtil.getAppExecutorService().execute(contextRunnable);
break;
case WRITE:
ModalityState processModality = progressIndicator.getModalityState();
- ApplicationManager.getApplication().invokeLaterOnWriteThread(runnable, processModality);
+ ApplicationManager.getApplication().invokeLaterOnWriteThread(contextRunnable, processModality);
break;
default:
throw new IllegalStateException("Unexpected value: " + myThreadToUse);
diff --git a/platform/platform-tests/testSrc/com/intellij/util/concurrency/ThreadContextPropagationTest.kt b/platform/platform-tests/testSrc/com/intellij/util/concurrency/ThreadContextPropagationTest.kt
index a04e325c8775..0dc5c8241abe 100644
--- a/platform/platform-tests/testSrc/com/intellij/util/concurrency/ThreadContextPropagationTest.kt
+++ b/platform/platform-tests/testSrc/com/intellij/util/concurrency/ThreadContextPropagationTest.kt
@@ -5,22 +5,22 @@ import com.intellij.concurrency.*
import com.intellij.openapi.application.ApplicationManager
import com.intellij.openapi.application.EDT
import com.intellij.openapi.application.ModalityState
-import com.intellij.openapi.progress.blockingContext
-import com.intellij.openapi.progress.timeoutWaitUp
+import com.intellij.openapi.application.currentThreadContextModality
+import com.intellij.openapi.progress.*
import com.intellij.openapi.util.Conditions
+import com.intellij.testFramework.junit5.SystemProperty
import com.intellij.testFramework.junit5.TestApplication
import com.intellij.util.timeoutRunBlocking
-import kotlinx.coroutines.Dispatchers
-import kotlinx.coroutines.launch
-import kotlinx.coroutines.suspendCancellableCoroutine
-import kotlinx.coroutines.withContext
+import kotlinx.coroutines.*
import org.junit.jupiter.api.Assertions
+import org.junit.jupiter.api.Assertions.assertFalse
import org.junit.jupiter.api.Assertions.assertSame
import org.junit.jupiter.api.Test
import org.junit.jupiter.api.extension.ExtendWith
import org.junit.jupiter.api.extension.ExtensionContext
import org.junit.jupiter.api.extension.InvocationInterceptor
import org.junit.jupiter.api.extension.ReflectiveInvocationContext
+import java.lang.Runnable
import java.lang.reflect.Method
import java.util.concurrent.ExecutorService
import java.util.concurrent.Future
@@ -182,4 +182,37 @@ class ThreadContextPropagationTest {
}
}
}
+
+ @SystemProperty("intellij.progress.task.ignoreHeadless", "true")
+ @Test
+ fun `Task Modal`(): Unit = timeoutRunBlocking {
+ doTest {
+ object : Task.Modal(null, "", true) {
+ override fun run(indicator: ProgressIndicator) {
+ it()
+ }
+ }.queue()
+ }
+ }
+
+ @SystemProperty("intellij.progress.task.ignoreHeadless", "true")
+ @Test
+ fun `Task Modal receives newly entered modality state in the context`(): Unit = timeoutRunBlocking {
+ val finished = CompletableDeferred()
+ withModalProgress(ModalTaskOwner.guess(), "", TaskCancellation.cancellable()) {
+ blockingContext {
+ assertSame(currentThreadContextModality(), ModalityState.defaultModalityState())
+ object : Task.Modal(null, "", true) {
+ override fun run(indicator: ProgressIndicator) {
+ finished.completeWith(runCatching {
+ assertFalse(currentThreadContextModality() == ModalityState.NON_MODAL)
+ assertSame(currentThreadContextModality(), indicator.modalityState)
+ assertSame(currentThreadContextModality(), ModalityState.defaultModalityState())
+ })
+ }
+ }.queue()
+ }
+ }
+ finished.await()
+ }
}
diff --git a/platform/util/src/com/intellij/util/concurrency/propagation.kt b/platform/util/src/com/intellij/util/concurrency/propagation.kt
index 3863feb21514..a564852f9bde 100644
--- a/platform/util/src/com/intellij/util/concurrency/propagation.kt
+++ b/platform/util/src/com/intellij/util/concurrency/propagation.kt
@@ -62,7 +62,8 @@ class BlockingJob(val blockingJob: Job) : AbstractCoroutineContextElement(Blocki
companion object : CoroutineContext.Key
}
-private fun createChildContext(): Pair {
+@Internal
+fun createChildContext(): Pair {
val currentThreadContext = currentThreadContext()
// Problem: a task may infinitely reschedule itself