diff --git a/platform/projectModel-api/src/org/jetbrains/concurrency/AsyncPromise.kt b/platform/projectModel-api/src/org/jetbrains/concurrency/AsyncPromise.kt index 83c17e923d92..eb4250dbdce3 100644 --- a/platform/projectModel-api/src/org/jetbrains/concurrency/AsyncPromise.kt +++ b/platform/projectModel-api/src/org/jetbrains/concurrency/AsyncPromise.kt @@ -14,16 +14,22 @@ import java.util.concurrent.atomic.AtomicReference private val LOG = Logger.getInstance(AsyncPromise::class.java) +private data class PromiseValue(val result: T? = null, val error: Throwable? = null) + open class AsyncPromise : Promise, Getter, CancellablePromise { private val doneRef = AtomicReference?>() private val rejectedRef = AtomicReference?>() - private val stateRef = AtomicReference(State.PENDING) + private val valueRef = AtomicReference?>(null) - // result object or error message - @Volatile private var result: Any? = null - - override fun getState() = stateRef.get()!! + override fun getState(): State { + val value = valueRef.get() + return when { + value == null -> State.PENDING + value.error == null -> State.FULFILLED + else -> State.REJECTED + } + } override fun done(done: Consumer): Promise { setHandler(doneRef, done, State.FULFILLED) @@ -35,16 +41,18 @@ open class AsyncPromise : Promise, Getter, CancellablePromise { return this } - @Suppress("UNCHECKED_CAST") - override fun get() = if (state == State.FULFILLED) result as T? else null + override fun get(): T? { + val value = valueRef.get() + return if (value != null && value.error == null) value.result else null + } override fun then(handler: Function): Promise { - @Suppress("UNCHECKED_CAST") - when (state) { - State.PENDING -> { + val value = valueRef.get() + when { + value == null -> { } - State.FULFILLED -> return DonePromise(handler.`fun`(result as T?)) - State.REJECTED -> return rejectedPromise(result as Throwable) + value.error == null -> return resolvedPromise(handler.`fun`(value.result)) + else -> return RejectedPromise(value.error) } val promise = AsyncPromise() @@ -62,12 +70,12 @@ open class AsyncPromise : Promise, Getter, CancellablePromise { } override fun thenAsync(handler: Function>): Promise { - @Suppress("UNCHECKED_CAST") - when (state) { - State.PENDING -> { + val value = valueRef.get() + when { + value == null -> { } - State.FULFILLED -> return handler.`fun`(result as T?) - State.REJECTED -> return rejectedPromise(result as Throwable) + value.error == null -> return handler.`fun`(value.result) + else -> return RejectedPromise(value.error) } val promise = AsyncPromise() @@ -87,17 +95,11 @@ open class AsyncPromise : Promise, Getter, CancellablePromise { return this } - when (state) { - State.PENDING -> { - addHandlers(Consumer({ child.catchError { child.setResult(it) } }), Consumer({ child.setError(it) })) - } - State.FULFILLED -> { - @Suppress("UNCHECKED_CAST") - child.setResult(result as T) - } - State.REJECTED -> { - child.setError((result as Throwable?)!!) - } + val value = valueRef.get() + when { + value == null -> addHandlers(Consumer({ child.catchError { child.setResult(it) } }), Consumer({ child.setError(it) })) + value.error == null -> child.setResult(value.result) + else -> child.setError(value.error) } return this } @@ -108,12 +110,10 @@ open class AsyncPromise : Promise, Getter, CancellablePromise { } fun setResult(result: T?) { - if (!stateRef.compareAndSet(State.PENDING, State.FULFILLED)) { + if (!valueRef.compareAndSet(null, PromiseValue(result = result))) { return } - this.result = result - val done = doneRef.getAndSet(null) rejectedRef.set(null) @@ -129,13 +129,11 @@ open class AsyncPromise : Promise, Getter, CancellablePromise { } open fun setError(error: Throwable): Boolean { - if (!stateRef.compareAndSet(State.PENDING, State.REJECTED)) { + if (!valueRef.compareAndSet(null, PromiseValue(error = error))) { LOG.errorIfNotMessage(error) return false } - result = error - val rejected = rejectedRef.getAndSet(null) doneRef.set(null) @@ -155,20 +153,23 @@ open class AsyncPromise : Promise, Getter, CancellablePromise { } override fun blockingGet(timeout: Int, timeUnit: TimeUnit): T? { - if (isPending) { + var value = valueRef.get() + if (value == null) { val latch = CountDownLatch(1) processed { latch.countDown() } if (!latch.await(timeout.toLong(), timeUnit)) { throw TimeoutException() } + + value = valueRef.get()!! } - @Suppress("UNCHECKED_CAST") - if (isRejected) { - throw (result as Throwable) + val error = value.error + if (error == null) { + return value.result } else { - return result as T? + throw error } } @@ -177,11 +178,8 @@ open class AsyncPromise : Promise, Getter, CancellablePromise { return } - if (state != State.PENDING) { - if (state == targetState) { - @Suppress("UNCHECKED_CAST") - newConsumer.consume(result as T?) - } + valueRef.get()?.let { + callConsumerIfTargeted(targetState, newConsumer, it) return } @@ -202,9 +200,8 @@ open class AsyncPromise : Promise, Getter, CancellablePromise { // clearHandlers was called - just execute newConsumer if (executed) { - if (state == targetState) { - @Suppress("UNCHECKED_CAST") - newConsumer.consume(result as T?) + valueRef.get()?.let { + callConsumerIfTargeted(targetState, newConsumer, it) } return } @@ -221,12 +218,19 @@ open class AsyncPromise : Promise, Getter, CancellablePromise { if (state == targetState) { ref.getAndSet(null)?.let { - @Suppress("UNCHECKED_CAST") - it.consume(result as T?) + callConsumerIfTargeted(targetState, it, valueRef.get()!!) } } } + private fun callConsumerIfTargeted(targetState: State, newConsumer: Consumer, value: PromiseValue) { + val currentState = if (value.error == null) State.FULFILLED else State.REJECTED + if (currentState == targetState) { + @Suppress("UNCHECKED_CAST") + newConsumer.consume(if (currentState == State.FULFILLED) value.result as C_T? else value.error as C_T) + } + } + override fun get(timeout: Long, unit: TimeUnit): T? { return blockingGet(timeout.toInt(), unit) } @@ -242,7 +246,7 @@ open class AsyncPromise : Promise, Getter, CancellablePromise { } override fun isCancelled(): Boolean { - return result == OBSOLETE_ERROR + return valueRef.get()?.error == OBSOLETE_ERROR } } diff --git a/platform/projectModel-api/src/org/jetbrains/concurrency/promise.kt b/platform/projectModel-api/src/org/jetbrains/concurrency/promise.kt index c351d3677f48..a41c7b8c9fc7 100644 --- a/platform/projectModel-api/src/org/jetbrains/concurrency/promise.kt +++ b/platform/projectModel-api/src/org/jetbrains/concurrency/promise.kt @@ -189,7 +189,12 @@ fun all(promises: Collection>, totalResult: T?, ignoreErrors: Boo val totalPromise = AsyncPromise() val done = CountDownConsumer(promises.size, totalPromise, totalResult) - val rejected = if (ignoreErrors) Consumer { done.consume(null) } else Consumer { totalPromise.setError(it) } + val rejected = if (ignoreErrors) { + Consumer { done.consume(null) } + } + else { + Consumer { totalPromise.setError(it) } + } for (promise in promises) { promise.done(done)