Promise — atomic value object instead of two fields for better thread safety

possible issue fixed: result is set after state, so, client from another thread can get null
This commit is contained in:
Vladimir Krivosheev
2018-02-19 15:11:32 +01:00
parent 03ca5985e3
commit d85b26e235
2 changed files with 60 additions and 51 deletions
@@ -14,16 +14,22 @@ import java.util.concurrent.atomic.AtomicReference
private val LOG = Logger.getInstance(AsyncPromise::class.java)
private data class PromiseValue<out T>(val result: T? = null, val error: Throwable? = null)
open class AsyncPromise<T> : Promise<T>, Getter<T>, CancellablePromise<T> {
private val doneRef = AtomicReference<Consumer<in T>?>()
private val rejectedRef = AtomicReference<Consumer<in Throwable>?>()
private val stateRef = AtomicReference(State.PENDING)
private val valueRef = AtomicReference<PromiseValue<T>?>(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<in T>): Promise<T> {
setHandler(doneRef, done, State.FULFILLED)
@@ -35,16 +41,18 @@ open class AsyncPromise<T> : Promise<T>, Getter<T>, CancellablePromise<T> {
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 <SUB_RESULT> then(handler: Function<in T, out SUB_RESULT>): Promise<SUB_RESULT> {
@Suppress("UNCHECKED_CAST")
when (state) {
State.PENDING -> {
val value = valueRef.get()
when {
value == null -> {
}
State.FULFILLED -> return DonePromise<SUB_RESULT>(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<SUB_RESULT>()
@@ -62,12 +70,12 @@ open class AsyncPromise<T> : Promise<T>, Getter<T>, CancellablePromise<T> {
}
override fun <SUB_RESULT> thenAsync(handler: Function<in T, Promise<SUB_RESULT>>): Promise<SUB_RESULT> {
@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<SUB_RESULT>()
@@ -87,17 +95,11 @@ open class AsyncPromise<T> : Promise<T>, Getter<T>, CancellablePromise<T> {
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<T> : Promise<T>, Getter<T>, CancellablePromise<T> {
}
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<T> : Promise<T>, Getter<T>, CancellablePromise<T> {
}
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<T> : Promise<T>, Getter<T>, CancellablePromise<T> {
}
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<T> : Promise<T>, Getter<T>, CancellablePromise<T> {
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<T> : Promise<T>, Getter<T>, CancellablePromise<T> {
// 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<T> : Promise<T>, Getter<T>, CancellablePromise<T> {
if (state == targetState) {
ref.getAndSet(null)?.let {
@Suppress("UNCHECKED_CAST")
it.consume(result as T?)
callConsumerIfTargeted(targetState, it, valueRef.get()!!)
}
}
}
private fun <C_T> callConsumerIfTargeted(targetState: State, newConsumer: Consumer<in C_T>, value: PromiseValue<T>) {
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<T> : Promise<T>, Getter<T>, CancellablePromise<T> {
}
override fun isCancelled(): Boolean {
return result == OBSOLETE_ERROR
return valueRef.get()?.error == OBSOLETE_ERROR
}
}
@@ -189,7 +189,12 @@ fun <T> all(promises: Collection<Promise<*>>, totalResult: T?, ignoreErrors: Boo
val totalPromise = AsyncPromise<T>()
val done = CountDownConsumer(promises.size, totalPromise, totalResult)
val rejected = if (ignoreErrors) Consumer { done.consume(null) } else Consumer<Throwable> { totalPromise.setError(it) }
val rejected = if (ignoreErrors) {
Consumer { done.consume(null) }
}
else {
Consumer<Throwable> { totalPromise.setError(it) }
}
for (promise in promises) {
promise.done(done)