From ca0664ab3e82d67b1dccde7640fcd3345071c763 Mon Sep 17 00:00:00 2001 From: Vladimir Krivosheev Date: Fri, 16 Sep 2016 14:56:14 +0200 Subject: [PATCH] =?UTF-8?q?simplify=20=E2=80=94=20use=20getAndSet?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../jetbrains/concurrency/AsyncPromiseTest.kt | 27 ++++++++++++++++++- .../org/jetbrains/concurrency/AsyncPromise.kt | 15 +++-------- 2 files changed, 29 insertions(+), 13 deletions(-) diff --git a/platform/platform-tests/testSrc/org/jetbrains/concurrency/AsyncPromiseTest.kt b/platform/platform-tests/testSrc/org/jetbrains/concurrency/AsyncPromiseTest.kt index 7cc462e1caec..feceefc60997 100644 --- a/platform/platform-tests/testSrc/org/jetbrains/concurrency/AsyncPromiseTest.kt +++ b/platform/platform-tests/testSrc/org/jetbrains/concurrency/AsyncPromiseTest.kt @@ -35,6 +35,31 @@ class AsyncPromiseTest { doHandlerTest(true) } + @Test + fun state() { + val promise = AsyncPromise() + val count = AtomicInteger() + + val r = { + promise.done { count.incrementAndGet() } + } + + val s = { + promise.setResult("test") + } + + val numThreads = 30 + assertConcurrent(*Array(numThreads, { + if (it and 1 === 0) r else s + })) + + assertThat(count.get()).isEqualTo(numThreads / 2) + assertThat(promise.get()).isEqualTo("test") + + r() + assertThat(count.get()).isEqualTo((numThreads / 2) + 1) + } + fun doHandlerTest(reject: Boolean) { val promise = AsyncPromise() val count = AtomicInteger() @@ -68,7 +93,7 @@ class AsyncPromiseTest { } } -fun assertConcurrent(vararg runnables: () -> Any, maxTimeoutSeconds: Int = 5) { +fun assertConcurrent(vararg runnables: () -> Any?, maxTimeoutSeconds: Int = 5) { val numThreads = runnables.size val exceptions = ContainerUtil.createLockFreeCopyOnWriteList() val threadPool = Executors.newFixedThreadPool(numThreads) diff --git a/platform/projectModel-api/src/org/jetbrains/concurrency/AsyncPromise.kt b/platform/projectModel-api/src/org/jetbrains/concurrency/AsyncPromise.kt index 6c1d7f3baf12..3048371aab52 100644 --- a/platform/projectModel-api/src/org/jetbrains/concurrency/AsyncPromise.kt +++ b/platform/projectModel-api/src/org/jetbrains/concurrency/AsyncPromise.kt @@ -171,7 +171,7 @@ open class AsyncPromise : Promise, Getter { this.result = result - val done = getAndClearHandler(doneRef) + val done = doneRef.getAndSet(null) rejectedRef.set(null) if (done != null && !isObsolete(done)) { @@ -192,7 +192,7 @@ open class AsyncPromise : Promise, Getter { result = error - val rejected = getAndClearHandler(rejectedRef) + val rejected = rejectedRef.getAndSet(null) doneRef.set(null) if (rejected == null) { @@ -204,15 +204,6 @@ open class AsyncPromise : Promise, Getter { return true } - 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 { done(processed) rejected { processed.consume(null) } @@ -271,7 +262,7 @@ open class AsyncPromise : Promise, Getter { } if (state == targetState) { - getAndClearHandler(ref)?.let { + ref.getAndSet(null)?.let { @Suppress("UNCHECKED_CAST") it.consume(result as T?) }