From cb1edab159b6da19addccf9c09b2de97e874f653 Mon Sep 17 00:00:00 2001 From: Vladimir Krivosheev Date: Thu, 15 Sep 2016 17:25:33 +0200 Subject: [PATCH] =?UTF-8?q?AsyncPromise=20=E2=80=94=20correctly=20execute?= =?UTF-8?q?=20handler=20if=20another=20handler=20added=20concurrently?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../jetbrains/concurrency/AsyncPromiseTest.kt | 105 ++++++++++++++++++ .../org/jetbrains/concurrency/AsyncPromise.kt | 82 +++++++++----- 2 files changed, 158 insertions(+), 29 deletions(-) create mode 100644 platform/platform-tests/testSrc/org/jetbrains/concurrency/AsyncPromiseTest.kt diff --git a/platform/platform-tests/testSrc/org/jetbrains/concurrency/AsyncPromiseTest.kt b/platform/platform-tests/testSrc/org/jetbrains/concurrency/AsyncPromiseTest.kt new file mode 100644 index 000000000000..7cc462e1caec --- /dev/null +++ b/platform/platform-tests/testSrc/org/jetbrains/concurrency/AsyncPromiseTest.kt @@ -0,0 +1,105 @@ +/* + * Copyright 2000-2016 JetBrains s.r.o. + * + * 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 + * + * http://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. + */ +package org.jetbrains.concurrency + +import com.intellij.util.containers.ContainerUtil +import com.intellij.util.lang.CompoundRuntimeException +import org.assertj.core.api.Assertions.assertThat +import org.junit.Test +import java.util.concurrent.CountDownLatch +import java.util.concurrent.Executors +import java.util.concurrent.TimeUnit +import java.util.concurrent.atomic.AtomicInteger + +class AsyncPromiseTest { + @Test + fun done() { + doHandlerTest(false) + } + + @Test + fun rejected() { + doHandlerTest(true) + } + + fun doHandlerTest(reject: Boolean) { + val promise = AsyncPromise() + val count = AtomicInteger() + + val r = { + if (reject) { + promise.rejected { count.incrementAndGet() } + } + else { + promise.done { count.incrementAndGet() } + } + } + + val numThreads = 30 + assertConcurrent(*Array(numThreads, { r })) + + if (reject) { + promise.setError("test") + } + else { + promise.setResult("test") + } + + assertThat(count.get()).isEqualTo(numThreads) + if (!reject) { + assertThat(promise.get()).isEqualTo("test") + } + + r() + assertThat(count.get()).isEqualTo(numThreads + 1) + } +} + +fun assertConcurrent(vararg runnables: () -> Any, maxTimeoutSeconds: Int = 5) { + val numThreads = runnables.size + val exceptions = ContainerUtil.createLockFreeCopyOnWriteList() + val threadPool = Executors.newFixedThreadPool(numThreads) + try { + val allExecutorThreadsReady = CountDownLatch(numThreads) + val afterInitBlocker = CountDownLatch(1) + val allDone = CountDownLatch(numThreads) + for (submittedTestRunnable in runnables) { + threadPool.submit { + allExecutorThreadsReady.countDown() + try { + afterInitBlocker.await() + submittedTestRunnable() + } + catch (e: Throwable) { + exceptions.add(e) + } + finally { + allDone.countDown() + } + } + } + + // wait until all threads are ready + assertThat(allExecutorThreadsReady.await((runnables.size * 10).toLong(), TimeUnit.MILLISECONDS)).isTrue() + // start all test runners + afterInitBlocker.countDown() + assertThat(allDone.await(maxTimeoutSeconds.toLong(), TimeUnit.SECONDS)).isTrue() + } + finally { + threadPool.shutdownNow() + } + CompoundRuntimeException.throwIfNotEmpty(exceptions) +} \ No newline at end of file diff --git a/platform/projectModel-api/src/org/jetbrains/concurrency/AsyncPromise.kt b/platform/projectModel-api/src/org/jetbrains/concurrency/AsyncPromise.kt index 9ef2eaa1e536..2ba5688ae259 100644 --- a/platform/projectModel-api/src/org/jetbrains/concurrency/AsyncPromise.kt +++ b/platform/projectModel-api/src/org/jetbrains/concurrency/AsyncPromise.kt @@ -28,8 +28,8 @@ private val LOG = Logger.getInstance(AsyncPromise::class.java) private val OBSOLETE_ERROR = Promise.createError("Obsolete") open class AsyncPromise : Promise(), Getter { - @Volatile private var done: Consumer? = null - @Volatile private var rejected: Consumer? = null + private val doneRef = AtomicReference?>() + private val rejectedRef = AtomicReference?>() private val state = AtomicReference(State.PENDING) @@ -45,7 +45,7 @@ open class AsyncPromise : Promise(), Getter { when (state.get()!!) { State.PENDING -> { - this.done = setHandler(this.done, done, State.FULFILLED) + setHandler(doneRef, done, State.FULFILLED) } State.FULFILLED -> { @Suppress("UNCHECKED_CAST") @@ -65,7 +65,7 @@ open class AsyncPromise : Promise(), Getter { when (state.get()!!) { State.PENDING -> { - this.rejected = setHandler(this.rejected, rejected, State.REJECTED) + setHandler(rejectedRef, rejected, State.REJECTED) } State.FULFILLED -> { } @@ -159,8 +159,8 @@ open class AsyncPromise : Promise(), Getter { } private fun addHandlers(done: Consumer, rejected: Consumer) { - this.done = setHandler(this.done, done, State.FULFILLED) - this.rejected = setHandler(this.rejected, rejected, State.REJECTED) + setHandler(doneRef, done, State.FULFILLED) + setHandler(rejectedRef, rejected, State.REJECTED) } fun setResult(result: T?) { @@ -170,8 +170,9 @@ open class AsyncPromise : Promise(), Getter { this.result = result - val done = this.done - clearHandlers() + val done = getAndClearHandler(doneRef) + rejectedRef.set(null) + if (done != null && !isObsolete(done)) { done.consume(result) } @@ -190,8 +191,9 @@ open class AsyncPromise : Promise(), Getter { result = error - val rejected = this.rejected - clearHandlers() + val rejected = getAndClearHandler(rejectedRef) + doneRef.set(null) + if (rejected == null) { Promise.logError(LOG, error) } @@ -201,9 +203,13 @@ open class AsyncPromise : Promise(), Getter { return true } - private fun clearHandlers() { - done = null - rejected = null + private fun getAndClearHandler(ref: AtomicReference?>): Consumer? { + var handler: Consumer? + do { + handler = ref.get() + } + while (!ref.compareAndSet(handler, null)) + return handler } override fun processed(processed: Consumer): Promise { @@ -212,29 +218,47 @@ open class AsyncPromise : Promise(), Getter { return this } - private fun setHandler(oldConsumer: Consumer?, newConsumer: Consumer, targetState: State): Consumer? = when (oldConsumer) { - null -> newConsumer - is CompoundConsumer<*> -> { - @Suppress("UNCHECKED_CAST") - val compoundConsumer = oldConsumer as CompoundConsumer - synchronized(compoundConsumer) { - compoundConsumer.consumers.let { - if (it == null) { - // clearHandlers was called - just execute newConsumer + private fun setHandler(ref: AtomicReference?>, newConsumer: Consumer, targetState: State) { + while (true) { + val oldConsumer = ref.get() + val newEffectiveConsumer = when (oldConsumer) { + null -> newConsumer + is CompoundConsumer<*> -> { + @Suppress("UNCHECKED_CAST") + val compoundConsumer = oldConsumer as CompoundConsumer + var executed = true + synchronized(compoundConsumer) { + compoundConsumer.consumers?.let { + it.add(newConsumer) + executed = false + } + } + + // clearHandlers was called - just execute newConsumer + if (executed) { if (state.get() == targetState) { @Suppress("UNCHECKED_CAST") newConsumer.consume(result as T?) } - return null - } - else { - it.add(newConsumer) - return compoundConsumer + return } + + compoundConsumer } + else -> CompoundConsumer(oldConsumer, newConsumer) + } + + if (ref.compareAndSet(oldConsumer, newEffectiveConsumer)) { + break + } + } + + if (state.get() == targetState) { + getAndClearHandler(ref)?.let { + @Suppress("UNCHECKED_CAST") + it.consume(result as T?) } } - else -> CompoundConsumer(oldConsumer, newConsumer) } } @@ -288,7 +312,7 @@ inline fun AsyncPromise<*>.catchError(runnable: () -> T): T? { private val cancelledPromise = RejectedPromise(OBSOLETE_ERROR) -@Suppress("CAST_NEVER_SUCCEEDS") +@Suppress("UNCHECKED_CAST") fun cancelledPromise(): Promise = cancelledPromise as Promise fun rejectedPromise(error: Throwable): Promise = Promise.reject(error) \ No newline at end of file