diff --git a/CHANGELOG.md b/CHANGELOG.md index 43b7b23..ba628b4 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,3 +1,7 @@ +**Unreleased:** + +- Subjects: fix a deadlock when sending values or termination concurrently with consumer cancellation (https://github.com/sideeffect-io/AsyncExtensions/issues/52). + **v0.5.2 - Oxygen:** This version is a bug fix version. diff --git a/Sources/AsyncSubjects/AsyncCurrentValueSubject.swift b/Sources/AsyncSubjects/AsyncCurrentValueSubject.swift index ff651d7..482bc33 100644 --- a/Sources/AsyncSubjects/AsyncCurrentValueSubject.swift +++ b/Sources/AsyncSubjects/AsyncCurrentValueSubject.swift @@ -67,24 +67,27 @@ public final class AsyncCurrentValueSubject: AsyncSubject where Element /// Sends a value to all consumers /// - Parameter element: the value to send public func send(_ element: Element) { - self.state.withCriticalRegion { state in + let channels = self.state.withCriticalRegion { state in state.current = element - for channel in state.channels.values { - channel.send(element) - } + return Array(state.channels.values) + } + // Resuming a consumer must not hold the lock used by its cancellation handler. + for channel in channels { + channel.send(element) } } /// Finishes the async sequences with a normal ending. /// - Parameter termination: The termination to finish the subject. public func send(_ termination: Termination) { - self.state.withCriticalRegion { state in + let channels = self.state.withCriticalRegion { state in state.terminalState = termination let channels = Array(state.channels.values) state.channels.removeAll() - for channel in channels { - channel.finish() - } + return channels + } + for channel in channels { + channel.finish() } } diff --git a/Sources/AsyncSubjects/AsyncPassthroughSubject.swift b/Sources/AsyncSubjects/AsyncPassthroughSubject.swift index 43813ce..6f15f1c 100644 --- a/Sources/AsyncSubjects/AsyncPassthroughSubject.swift +++ b/Sources/AsyncSubjects/AsyncPassthroughSubject.swift @@ -52,23 +52,26 @@ public final class AsyncPassthroughSubject: AsyncSubject { /// Sends a value to all consumers /// - Parameter element: the value to send public func send(_ element: Element) { - self.state.withCriticalRegion { state in - for channel in state.channels.values { - channel.send(element) - } + let channels = self.state.withCriticalRegion { state in + return Array(state.channels.values) + } + // Resuming a consumer must not hold the lock used by its cancellation handler. + for channel in channels { + channel.send(element) } } /// Finishes the subject with a normal ending. /// - Parameter termination: The termination to finish the subject public func send(_ termination: Termination) { - self.state.withCriticalRegion { state in + let channels = self.state.withCriticalRegion { state in state.terminalState = termination let channels = Array(state.channels.values) state.channels.removeAll() - for channel in channels { - channel.finish() - } + return channels + } + for channel in channels { + channel.finish() } } diff --git a/Sources/AsyncSubjects/AsyncReplaySubject.swift b/Sources/AsyncSubjects/AsyncReplaySubject.swift index eb806ca..f3fc334 100644 --- a/Sources/AsyncSubjects/AsyncReplaySubject.swift +++ b/Sources/AsyncSubjects/AsyncReplaySubject.swift @@ -46,29 +46,32 @@ public final class AsyncReplaySubject: AsyncSubject where Element: Send /// Sends a value to all consumers /// - Parameter element: the value to send public func send(_ element: Element) { - self.state.withCriticalRegion { state in + let channels = self.state.withCriticalRegion { state in if state.buffer.count >= state.bufferSize && !state.buffer.isEmpty { state.buffer.removeFirst() } state.buffer.append(element) - for channel in state.channels.values { - channel.send(element) - } + return Array(state.channels.values) + } + // Resuming a consumer must not hold the lock used by its cancellation handler. + for channel in channels { + channel.send(element) } } /// Finishes the subject with a normal ending. /// - Parameter termination: The termination to finish the subject. public func send(_ termination: Termination) { - self.state.withCriticalRegion { state in + let channels = self.state.withCriticalRegion { state in state.terminalState = termination let channels = Array(state.channels.values) state.channels.removeAll() state.buffer.removeAll() state.bufferSize = 0 - for channel in channels { - channel.finish() - } + return channels + } + for channel in channels { + channel.finish() } } diff --git a/Sources/AsyncSubjects/AsyncThrowingCurrentValueSubject.swift b/Sources/AsyncSubjects/AsyncThrowingCurrentValueSubject.swift index 7ee5e1b..1b10921 100644 --- a/Sources/AsyncSubjects/AsyncThrowingCurrentValueSubject.swift +++ b/Sources/AsyncSubjects/AsyncThrowingCurrentValueSubject.swift @@ -67,28 +67,31 @@ public final class AsyncThrowingCurrentValueSubject: As /// Sends a value to all consumers /// - Parameter element: the value to send public func send(_ element: Element) { - self.state.withCriticalRegion { state in + let channels = self.state.withCriticalRegion { state in state.current = element - for channel in state.channels.values { - channel.send(element) - } + return Array(state.channels.values) + } + // Resuming a consumer must not hold the lock used by its cancellation handler. + for channel in channels { + channel.send(element) } } /// Finishes the subject with either a normal ending or an error. /// - Parameter termination: The termination to finish the subject. public func send(_ termination: Termination) { - self.state.withCriticalRegion { state in + let channels = self.state.withCriticalRegion { state in state.terminalState = termination let channels = Array(state.channels.values) state.channels.removeAll() - for channel in channels { - switch termination { - case .finished: - channel.finish() - case .failure(let error): - channel.fail(error) - } + return channels + } + for channel in channels { + switch termination { + case .finished: + channel.finish() + case .failure(let error): + channel.fail(error) } } } diff --git a/Sources/AsyncSubjects/AsyncThrowingPassthroughSubject.swift b/Sources/AsyncSubjects/AsyncThrowingPassthroughSubject.swift index 687e29c..f456db3 100644 --- a/Sources/AsyncSubjects/AsyncThrowingPassthroughSubject.swift +++ b/Sources/AsyncSubjects/AsyncThrowingPassthroughSubject.swift @@ -53,28 +53,30 @@ public final class AsyncThrowingPassthroughSubject: Asy /// Sends a value to all consumers /// - Parameter element: the value to send public func send(_ element: Element) { - self.state.withCriticalRegion { state in - for channel in state.channels.values { - channel.send(element) - } + let channels = self.state.withCriticalRegion { state in + return Array(state.channels.values) + } + // Resuming a consumer must not hold the lock used by its cancellation handler. + for channel in channels { + channel.send(element) } } /// Finishes the subject with either a normal ending or an error. /// - Parameter termination: The termination to finish the subject public func send(_ termination: Termination) { - self.state.withCriticalRegion { state in + let channels = self.state.withCriticalRegion { state in state.terminalState = termination let channels = Array(state.channels.values) state.channels.removeAll() - - for channel in channels { - switch termination { - case .finished: - channel.finish() - case .failure(let error): - channel.fail(error) - } + return channels + } + for channel in channels { + switch termination { + case .finished: + channel.finish() + case .failure(let error): + channel.fail(error) } } } diff --git a/Sources/AsyncSubjects/AsyncThrowingReplaySubject.swift b/Sources/AsyncSubjects/AsyncThrowingReplaySubject.swift index 92666ad..f3f21c8 100644 --- a/Sources/AsyncSubjects/AsyncThrowingReplaySubject.swift +++ b/Sources/AsyncSubjects/AsyncThrowingReplaySubject.swift @@ -45,33 +45,36 @@ public final class AsyncThrowingReplaySubject: AsyncSub /// Sends a value to all consumers /// - Parameter element: the value to send public func send(_ element: Element) { - self.state.withCriticalRegion { state in + let channels = self.state.withCriticalRegion { state in if state.buffer.count >= state.bufferSize && !state.buffer.isEmpty { state.buffer.removeFirst() } state.buffer.append(element) - for channel in state.channels.values { - channel.send(element) - } + return Array(state.channels.values) + } + // Resuming a consumer must not hold the lock used by its cancellation handler. + for channel in channels { + channel.send(element) } } /// Finishes the subject with either a normal ending or an error. /// - Parameter termination: The termination to finish the subject public func send(_ termination: Termination) { - self.state.withCriticalRegion { state in + let channels = self.state.withCriticalRegion { state in state.terminalState = termination let channels = Array(state.channels.values) state.channels.removeAll() state.buffer.removeAll() state.bufferSize = 0 - for channel in channels { - switch termination { - case .finished: - channel.finish() - case .failure(let error): - channel.fail(error) - } + return channels + } + for channel in channels { + switch termination { + case .finished: + channel.finish() + case .failure(let error): + channel.fail(error) } } } diff --git a/Tests/AsyncSubjets/AsyncSubjectCancellationTests.swift b/Tests/AsyncSubjets/AsyncSubjectCancellationTests.swift new file mode 100644 index 0000000..17a1308 --- /dev/null +++ b/Tests/AsyncSubjets/AsyncSubjectCancellationTests.swift @@ -0,0 +1,155 @@ +import Dispatch +@testable import AsyncExtensions +import XCTest + +final class AsyncSubjectCancellationTests: XCTestCase { + func test_AsyncPassthroughSubject_does_not_deadlock_during_concurrent_cancellation() async { + await assertCancellationRaces( + makeSubject: { AsyncPassthroughSubject() }, + isSuspended: { subject in + guard let channel = subject.state.criticalState.channels.values.first else { return false } + if case .awaiting = channel.state.criticalState { return true } + return false + }, + sends: [ + { $0.send(1) }, + { $0.send(.finished) } + ] + ) + } + + func test_AsyncCurrentValueSubject_does_not_deadlock_during_concurrent_cancellation() async { + await assertCancellationRaces( + makeSubject: { AsyncCurrentValueSubject(0) }, + isSuspended: { subject in + guard let channel = subject.state.criticalState.channels.values.first else { return false } + if case .awaiting = channel.state.criticalState { return true } + return false + }, + sends: [ + { $0.send(1) }, + { $0.send(.finished) } + ] + ) + } + + func test_AsyncReplaySubject_does_not_deadlock_during_concurrent_cancellation() async { + await assertCancellationRaces( + makeSubject: { AsyncReplaySubject(bufferSize: 1) }, + isSuspended: { subject in + guard let channel = subject.state.criticalState.channels.values.first else { return false } + if case .awaiting = channel.state.criticalState { return true } + return false + }, + sends: [ + { $0.send(1) }, + { $0.send(.finished) } + ] + ) + } + + func test_AsyncThrowingPassthroughSubject_does_not_deadlock_during_concurrent_cancellation() async { + await assertCancellationRaces( + makeSubject: { AsyncThrowingPassthroughSubject() }, + isSuspended: { subject in + guard let channel = subject.state.criticalState.channels.values.first else { return false } + if case .awaiting = channel.state.criticalState { return true } + return false + }, + sends: [ + { $0.send(1) }, + { $0.send(.finished) }, + { $0.send(.failure(MockError(code: 1))) } + ] + ) + } + + func test_AsyncThrowingCurrentValueSubject_does_not_deadlock_during_concurrent_cancellation() async { + await assertCancellationRaces( + makeSubject: { AsyncThrowingCurrentValueSubject(0) }, + isSuspended: { subject in + guard let channel = subject.state.criticalState.channels.values.first else { return false } + if case .awaiting = channel.state.criticalState { return true } + return false + }, + sends: [ + { $0.send(1) }, + { $0.send(.finished) }, + { $0.send(.failure(MockError(code: 1))) } + ] + ) + } + + func test_AsyncThrowingReplaySubject_does_not_deadlock_during_concurrent_cancellation() async { + await assertCancellationRaces( + makeSubject: { AsyncThrowingReplaySubject(bufferSize: 1) }, + isSuspended: { subject in + guard let channel = subject.state.criticalState.channels.values.first else { return false } + if case .awaiting = channel.state.criticalState { return true } + return false + }, + sends: [ + { $0.send(1) }, + { $0.send(.finished) }, + { $0.send(.failure(MockError(code: 1))) } + ] + ) + } + + private func assertCancellationRaces( + makeSubject: @escaping @Sendable () -> Subject, + isSuspended: @escaping @Sendable (Subject) -> Bool, + sends: [@Sendable (Subject) -> Void], + file: StaticString = #filePath, + line: UInt = #line + ) async { + for send in sends { + let completed = expectation(description: "Sending, cancellation and consumer exit all complete") + DispatchQueue.global().async { + for _ in 0..<1_000 { + let subject = makeSubject() + // Register before starting the race so passthrough values cannot be lost at setup. + let iterator = subject.makeAsyncIterator() + let consumerExited = DispatchSemaphore(value: 0) + let consumer = Task.detached { + defer { consumerExited.signal() } + var iterator = iterator + do { + while let _ = try await iterator.next() {} + } catch { + // A throwing subject may deliver its failure before cancellation wins. + } + } + + // Wait for an installed continuation, rather than racing an unstarted consumer. + let deadline = DispatchTime.now() + .seconds(5) + while !isSuspended(subject) { + if DispatchTime.now() >= deadline { + XCTFail("Consumer did not suspend", file: file, line: line) + consumer.cancel() + completed.fulfill() + return + } + Thread.sleep(forTimeInterval: 0.0001) + } + + DispatchQueue.concurrentPerform(iterations: 2) { index in + if index == 0 { + consumer.cancel() + } else { + send(subject) + } + } + guard consumerExited.wait(timeout: .now() + .seconds(5)) == .success else { + XCTFail("Consumer did not exit", file: file, line: line) + completed.fulfill() + return + } + } + completed.fulfill() + } + // A deadlocked concurrentPerform stays on a Dispatch worker, not the test executor. + await fulfillment(of: [completed], timeout: 30) + } + } +}