AsyncPromise — correctly execute handler if another handler added concurrently

This commit is contained in:
Vladimir Krivosheev
2016-09-15 17:26:31 +02:00
parent c5dd0fa9ca
commit cb1edab159
2 changed files with 158 additions and 29 deletions
@@ -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<String>()
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<Throwable>()
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)
}
@@ -28,8 +28,8 @@ private val LOG = Logger.getInstance(AsyncPromise::class.java)
private val OBSOLETE_ERROR = Promise.createError("Obsolete")
open class AsyncPromise<T> : Promise<T>(), Getter<T> {
@Volatile private var done: Consumer<in T>? = null
@Volatile private var rejected: Consumer<in Throwable>? = null
private val doneRef = AtomicReference<Consumer<in T>?>()
private val rejectedRef = AtomicReference<Consumer<in Throwable>?>()
private val state = AtomicReference(State.PENDING)
@@ -45,7 +45,7 @@ open class AsyncPromise<T> : Promise<T>(), Getter<T> {
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<T> : Promise<T>(), Getter<T> {
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<T> : Promise<T>(), Getter<T> {
}
private fun addHandlers(done: Consumer<T>, rejected: Consumer<Throwable>) {
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<T> : Promise<T>(), Getter<T> {
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<T> : Promise<T>(), Getter<T> {
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<T> : Promise<T>(), Getter<T> {
return true
}
private fun clearHandlers() {
done = null
rejected = null
private fun <T> getAndClearHandler(ref: AtomicReference<Consumer<in T>?>): Consumer<in T>? {
var handler: Consumer<in T>?
do {
handler = ref.get()
}
while (!ref.compareAndSet(handler, null))
return handler
}
override fun processed(processed: Consumer<in T>): Promise<T> {
@@ -212,29 +218,47 @@ open class AsyncPromise<T> : Promise<T>(), Getter<T> {
return this
}
private fun <T> setHandler(oldConsumer: Consumer<in T>?, newConsumer: Consumer<in T>, targetState: State): Consumer<in T>? = when (oldConsumer) {
null -> newConsumer
is CompoundConsumer<*> -> {
@Suppress("UNCHECKED_CAST")
val compoundConsumer = oldConsumer as CompoundConsumer<T>
synchronized(compoundConsumer) {
compoundConsumer.consumers.let {
if (it == null) {
// clearHandlers was called - just execute newConsumer
private fun <T> setHandler(ref: AtomicReference<Consumer<in T>?>, newConsumer: Consumer<in T>, targetState: State) {
while (true) {
val oldConsumer = ref.get()
val newEffectiveConsumer = when (oldConsumer) {
null -> newConsumer
is CompoundConsumer<*> -> {
@Suppress("UNCHECKED_CAST")
val compoundConsumer = oldConsumer as CompoundConsumer<T>
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 <T> AsyncPromise<*>.catchError(runnable: () -> T): T? {
private val cancelledPromise = RejectedPromise<Any?>(OBSOLETE_ERROR)
@Suppress("CAST_NEVER_SUCCEEDS")
@Suppress("UNCHECKED_CAST")
fun <T> cancelledPromise(): Promise<T> = cancelledPromise as Promise<T>
fun <T> rejectedPromise(error: Throwable): Promise<T> = Promise.reject(error)