diff --git a/Sources/AsyncSubjects/AsyncCurrentValueSubject.swift b/Sources/AsyncSubjects/AsyncCurrentValueSubject.swift index 5225105..ff651d7 100644 --- a/Sources/AsyncSubjects/AsyncCurrentValueSubject.swift +++ b/Sources/AsyncSubjects/AsyncCurrentValueSubject.swift @@ -91,23 +91,23 @@ public final class AsyncCurrentValueSubject: AsyncSubject where Element func handleNewConsumer() -> (iterator: AsyncBufferedChannel.Iterator, unregister: @Sendable () -> Void) { let asyncBufferedChannel = AsyncBufferedChannel() - let (terminalState, current) = self.state.withCriticalRegion { state -> (Termination?, Element) in - (state.terminalState, state.current) - } - - if let terminalState = terminalState, terminalState.isFinished { - asyncBufferedChannel.finish() - return (asyncBufferedChannel.makeAsyncIterator(), {}) - } + let consumerId = self.state.withCriticalRegion { state -> Int? in + if let terminalState = state.terminalState, terminalState.isFinished { + asyncBufferedChannel.finish() + return nil + } - asyncBufferedChannel.send(current) + asyncBufferedChannel.send(state.current) - let consumerId = self.state.withCriticalRegion { state -> Int in state.ids += 1 state.channels[state.ids] = asyncBufferedChannel return state.ids } + guard let consumerId = consumerId else { + return (asyncBufferedChannel.makeAsyncIterator(), {}) + } + let unregister = { @Sendable [state] in state.withCriticalRegion { state in state.channels[consumerId] = nil diff --git a/Sources/AsyncSubjects/AsyncPassthroughSubject.swift b/Sources/AsyncSubjects/AsyncPassthroughSubject.swift index 2badeb9..43813ce 100644 --- a/Sources/AsyncSubjects/AsyncPassthroughSubject.swift +++ b/Sources/AsyncSubjects/AsyncPassthroughSubject.swift @@ -75,21 +75,21 @@ public final class AsyncPassthroughSubject: AsyncSubject { func handleNewConsumer() -> (iterator: AsyncBufferedChannel.Iterator, unregister: @Sendable () -> Void) { let asyncBufferedChannel = AsyncBufferedChannel() - let terminalState = self.state.withCriticalRegion { state in - state.terminalState - } - - if let terminalState = terminalState, terminalState.isFinished { - asyncBufferedChannel.finish() - return (asyncBufferedChannel.makeAsyncIterator(), {}) - } + let consumerId = self.state.withCriticalRegion { state -> Int? in + if let terminalState = state.terminalState, terminalState.isFinished { + asyncBufferedChannel.finish() + return nil + } - let consumerId = self.state.withCriticalRegion { state -> Int in state.ids += 1 state.channels[state.ids] = asyncBufferedChannel return state.ids } + guard let consumerId = consumerId else { + return (asyncBufferedChannel.makeAsyncIterator(), {}) + } + let unregister = { @Sendable [state] in state.withCriticalRegion { state in state.channels[consumerId] = nil diff --git a/Sources/AsyncSubjects/AsyncReplaySubject.swift b/Sources/AsyncSubjects/AsyncReplaySubject.swift index f4e610e..eb806ca 100644 --- a/Sources/AsyncSubjects/AsyncReplaySubject.swift +++ b/Sources/AsyncSubjects/AsyncReplaySubject.swift @@ -75,25 +75,25 @@ public final class AsyncReplaySubject: AsyncSubject where Element: Send func handleNewConsumer() -> (iterator: AsyncBufferedChannel.Iterator, unregister: @Sendable () -> Void) { let asyncBufferedChannel = AsyncBufferedChannel() - let (terminalState, elements) = self.state.withCriticalRegion { state -> (Termination?, [Element]) in - (state.terminalState, state.buffer) - } - - if let terminalState = terminalState, terminalState.isFinished { - asyncBufferedChannel.finish() - return (asyncBufferedChannel.makeAsyncIterator(), {}) - } + let consumerId = self.state.withCriticalRegion { state -> Int? in + if let terminalState = state.terminalState, terminalState.isFinished { + asyncBufferedChannel.finish() + return nil + } - for element in elements { - asyncBufferedChannel.send(element) - } + for element in state.buffer { + asyncBufferedChannel.send(element) + } - let consumerId = self.state.withCriticalRegion { state -> Int in state.ids += 1 state.channels[state.ids] = asyncBufferedChannel return state.ids } + guard let consumerId = consumerId else { + return (asyncBufferedChannel.makeAsyncIterator(), {}) + } + let unregister = { @Sendable [state] in state.withCriticalRegion { state in state.channels[consumerId] = nil diff --git a/Sources/AsyncSubjects/AsyncThrowingCurrentValueSubject.swift b/Sources/AsyncSubjects/AsyncThrowingCurrentValueSubject.swift index 2294b09..7ee5e1b 100644 --- a/Sources/AsyncSubjects/AsyncThrowingCurrentValueSubject.swift +++ b/Sources/AsyncSubjects/AsyncThrowingCurrentValueSubject.swift @@ -97,28 +97,28 @@ public final class AsyncThrowingCurrentValueSubject: As ) -> (iterator: AsyncThrowingBufferedChannel.Iterator, unregister: @Sendable () -> Void) { let asyncBufferedChannel = AsyncThrowingBufferedChannel() - let (terminalState, current) = self.state.withCriticalRegion { state -> (Termination?, Element) in - (state.terminalState, state.current) - } - - if let terminalState = terminalState { - switch terminalState { - case .finished: - asyncBufferedChannel.finish() - case .failure(let error): - asyncBufferedChannel.fail(error) + let consumerId = self.state.withCriticalRegion { state -> Int? in + if let terminalState = state.terminalState { + switch terminalState { + case .finished: + asyncBufferedChannel.finish() + case .failure(let error): + asyncBufferedChannel.fail(error) + } + return nil } - return (asyncBufferedChannel.makeAsyncIterator(), {}) - } - asyncBufferedChannel.send(current) + asyncBufferedChannel.send(state.current) - let consumerId = self.state.withCriticalRegion { state -> Int in state.ids += 1 state.channels[state.ids] = asyncBufferedChannel return state.ids } + guard let consumerId = consumerId else { + return (asyncBufferedChannel.makeAsyncIterator(), {}) + } + let unregister = { @Sendable [state] in state.withCriticalRegion { state in state.channels[consumerId] = nil diff --git a/Sources/AsyncSubjects/AsyncThrowingPassthroughSubject.swift b/Sources/AsyncSubjects/AsyncThrowingPassthroughSubject.swift index c1da4a5..687e29c 100644 --- a/Sources/AsyncSubjects/AsyncThrowingPassthroughSubject.swift +++ b/Sources/AsyncSubjects/AsyncThrowingPassthroughSubject.swift @@ -83,26 +83,26 @@ public final class AsyncThrowingPassthroughSubject: Asy ) -> (iterator: AsyncThrowingBufferedChannel.Iterator, unregister: @Sendable () -> Void) { let asyncBufferedChannel = AsyncThrowingBufferedChannel() - let terminalState = self.state.withCriticalRegion { state in - state.terminalState - } - - if let terminalState = terminalState { - switch terminalState { - case .finished: - asyncBufferedChannel.finish() - case .failure(let error): - asyncBufferedChannel.fail(error) + let consumerId = self.state.withCriticalRegion { state -> Int? in + if let terminalState = state.terminalState { + switch terminalState { + case .finished: + asyncBufferedChannel.finish() + case .failure(let error): + asyncBufferedChannel.fail(error) + } + return nil } - return (asyncBufferedChannel.makeAsyncIterator(), {}) - } - let consumerId = self.state.withCriticalRegion { state -> Int in state.ids += 1 state.channels[state.ids] = asyncBufferedChannel return state.ids } + guard let consumerId = consumerId else { + return (asyncBufferedChannel.makeAsyncIterator(), {}) + } + let unregister = { @Sendable [state] in state.withCriticalRegion { state in state.channels[consumerId] = nil diff --git a/Sources/AsyncSubjects/AsyncThrowingReplaySubject.swift b/Sources/AsyncSubjects/AsyncThrowingReplaySubject.swift index c736d49..92666ad 100644 --- a/Sources/AsyncSubjects/AsyncThrowingReplaySubject.swift +++ b/Sources/AsyncSubjects/AsyncThrowingReplaySubject.swift @@ -80,30 +80,30 @@ public final class AsyncThrowingReplaySubject: AsyncSub ) -> (iterator: AsyncThrowingBufferedChannel.Iterator, unregister: @Sendable () -> Void) { let asyncBufferedChannel = AsyncThrowingBufferedChannel() - let (terminalState, elements) = self.state.withCriticalRegion { state -> (Termination?, [Element]) in - (state.terminalState, state.buffer) - } - - if let terminalState = terminalState { - switch terminalState { - case .finished: - asyncBufferedChannel.finish() - case .failure(let error): - asyncBufferedChannel.fail(error) + let consumerId = self.state.withCriticalRegion { state -> Int? in + if let terminalState = state.terminalState { + switch terminalState { + case .finished: + asyncBufferedChannel.finish() + case .failure(let error): + asyncBufferedChannel.fail(error) + } + return nil } - return (asyncBufferedChannel.makeAsyncIterator(), {}) - } - for element in elements { - asyncBufferedChannel.send(element) - } + for element in state.buffer { + asyncBufferedChannel.send(element) + } - let consumerId = self.state.withCriticalRegion { state -> Int in state.ids += 1 state.channels[state.ids] = asyncBufferedChannel return state.ids } + guard let consumerId = consumerId else { + return (asyncBufferedChannel.makeAsyncIterator(), {}) + } + let unregister = { @Sendable [state] in state.withCriticalRegion { state in state.channels[consumerId] = nil diff --git a/Tests/AsyncSubjets/AsyncCurrentValueSubjectTests.swift b/Tests/AsyncSubjets/AsyncCurrentValueSubjectTests.swift index 3d2d9e2..a7ccadd 100644 --- a/Tests/AsyncSubjets/AsyncCurrentValueSubjectTests.swift +++ b/Tests/AsyncSubjets/AsyncCurrentValueSubjectTests.swift @@ -204,4 +204,32 @@ final class AsyncCurrentValueSubjectTests: XCTestCase { XCTAssertEqual(receivedElementsA, expectedElements) XCTAssertEqual(receivedElementsB, expectedElements) } + + func test_subscription_racing_send_receives_sent_element() async { + for _ in 0..<10_000 { + let sut = AsyncCurrentValueSubject(0) + var iterator: AsyncCurrentValueSubject.Iterator? + + race({ iterator = sut.makeAsyncIterator() }, { sut.send(1) }) + + let drained = await drainBufferedElements(of: iterator!) + guard drained.elements.last == 1 else { + return XCTFail("Expected to receive the sent element, received \(drained.elements)") + } + } + } + + func test_subscription_racing_termination_is_terminated() async { + for _ in 0..<10_000 { + let sut = AsyncCurrentValueSubject(0) + var iterator: AsyncCurrentValueSubject.Iterator? + + race({ iterator = sut.makeAsyncIterator() }, { sut.send(.finished) }) + + let drained = await drainBufferedElements(of: iterator!) + guard drained.isTerminated else { + return XCTFail("Expected the subscription to be terminated") + } + } + } } diff --git a/Tests/AsyncSubjets/AsyncPassthroughSubjectTests.swift b/Tests/AsyncSubjets/AsyncPassthroughSubjectTests.swift index 728ed28..3f098a9 100644 --- a/Tests/AsyncSubjets/AsyncPassthroughSubjectTests.swift +++ b/Tests/AsyncSubjets/AsyncPassthroughSubjectTests.swift @@ -198,4 +198,18 @@ final class AsyncPassthroughSubjectTests: XCTestCase { XCTAssertEqual(receivedElementsA, expectedElements) XCTAssertEqual(receivedElementsB, expectedElements) } + + func test_subscription_racing_termination_is_terminated() async { + for _ in 0..<10_000 { + let sut = AsyncPassthroughSubject() + var iterator: AsyncPassthroughSubject.Iterator? + + race({ iterator = sut.makeAsyncIterator() }, { sut.send(.finished) }) + + let drained = await drainBufferedElements(of: iterator!) + guard drained.isTerminated else { + return XCTFail("Expected the subscription to be terminated") + } + } + } } diff --git a/Tests/AsyncSubjets/AsyncReplaySubjectTests.swift b/Tests/AsyncSubjets/AsyncReplaySubjectTests.swift index 0f824fc..aba3a58 100644 --- a/Tests/AsyncSubjets/AsyncReplaySubjectTests.swift +++ b/Tests/AsyncSubjets/AsyncReplaySubjectTests.swift @@ -227,4 +227,34 @@ final class AsyncReplaySubjectTests: XCTestCase { XCTAssertEqual(receivedElementsA, expectedElements) XCTAssertEqual(receivedElementsB, expectedElements) } + + func test_subscription_racing_send_receives_sent_element() async { + for _ in 0..<10_000 { + let sut = AsyncReplaySubject(bufferSize: 2) + sut.send(0) + var iterator: AsyncReplaySubject.Iterator? + + race({ iterator = sut.makeAsyncIterator() }, { sut.send(1) }) + + let drained = await drainBufferedElements(of: iterator!) + guard drained.elements.last == 1 else { + return XCTFail("Expected to receive the sent element, received \(drained.elements)") + } + } + } + + func test_subscription_racing_termination_is_terminated() async { + for _ in 0..<10_000 { + let sut = AsyncReplaySubject(bufferSize: 2) + sut.send(0) + var iterator: AsyncReplaySubject.Iterator? + + race({ iterator = sut.makeAsyncIterator() }, { sut.send(.finished) }) + + let drained = await drainBufferedElements(of: iterator!) + guard drained.isTerminated else { + return XCTFail("Expected the subscription to be terminated") + } + } + } } diff --git a/Tests/AsyncSubjets/AsyncThrowingCurrentValueSubjectTests.swift b/Tests/AsyncSubjets/AsyncThrowingCurrentValueSubjectTests.swift index c2cfba0..d933970 100644 --- a/Tests/AsyncSubjets/AsyncThrowingCurrentValueSubjectTests.swift +++ b/Tests/AsyncSubjets/AsyncThrowingCurrentValueSubjectTests.swift @@ -256,4 +256,32 @@ final class AsyncThrowingCurrentValueSubjectTests: XCTestCase { XCTAssertEqual(receivedElementsA, expectedElements) XCTAssertEqual(receivedElementsB, expectedElements) } + + func test_subscription_racing_send_receives_sent_element() async { + for _ in 0..<10_000 { + let sut = AsyncThrowingCurrentValueSubject(0) + var iterator: AsyncThrowingCurrentValueSubject.Iterator? + + race({ iterator = sut.makeAsyncIterator() }, { sut.send(1) }) + + let drained = await drainBufferedElements(of: iterator!) + guard drained.elements.last == 1 else { + return XCTFail("Expected to receive the sent element, received \(drained.elements)") + } + } + } + + func test_subscription_racing_termination_is_terminated() async { + for _ in 0..<10_000 { + let sut = AsyncThrowingCurrentValueSubject(0) + var iterator: AsyncThrowingCurrentValueSubject.Iterator? + + race({ iterator = sut.makeAsyncIterator() }, { sut.send(.failure(MockError(code: 1))) }) + + let drained = await drainBufferedElements(of: iterator!) + guard drained.isTerminated else { + return XCTFail("Expected the subscription to be terminated") + } + } + } } diff --git a/Tests/AsyncSubjets/AsyncThrowingPassthroughSubjectTests.swift b/Tests/AsyncSubjets/AsyncThrowingPassthroughSubjectTests.swift index e9838dc..b50c8ac 100644 --- a/Tests/AsyncSubjets/AsyncThrowingPassthroughSubjectTests.swift +++ b/Tests/AsyncSubjets/AsyncThrowingPassthroughSubjectTests.swift @@ -261,4 +261,18 @@ final class AsyncThrowingPassthroughSubjectTests: XCTestCase { XCTAssertEqual(receivedElementsA, expectedElements) XCTAssertEqual(receivedElementsB, expectedElements) } + + func test_subscription_racing_termination_is_terminated() async { + for _ in 0..<10_000 { + let sut = AsyncThrowingPassthroughSubject() + var iterator: AsyncThrowingPassthroughSubject.Iterator? + + race({ iterator = sut.makeAsyncIterator() }, { sut.send(.failure(MockError(code: 1))) }) + + let drained = await drainBufferedElements(of: iterator!) + guard drained.isTerminated else { + return XCTFail("Expected the subscription to be terminated") + } + } + } } diff --git a/Tests/AsyncSubjets/AsyncThrowingReplaySubjectTests.swift b/Tests/AsyncSubjets/AsyncThrowingReplaySubjectTests.swift index 222ef5a..b1d2d85 100644 --- a/Tests/AsyncSubjets/AsyncThrowingReplaySubjectTests.swift +++ b/Tests/AsyncSubjets/AsyncThrowingReplaySubjectTests.swift @@ -281,4 +281,34 @@ final class AsyncThrowingReplaySubjectTests: XCTestCase { XCTAssertEqual(receivedElementsA, expectedElements) XCTAssertEqual(receivedElementsB, expectedElements) } + + func test_subscription_racing_send_receives_sent_element() async { + for _ in 0..<10_000 { + let sut = AsyncThrowingReplaySubject(bufferSize: 2) + sut.send(0) + var iterator: AsyncThrowingReplaySubject.Iterator? + + race({ iterator = sut.makeAsyncIterator() }, { sut.send(1) }) + + let drained = await drainBufferedElements(of: iterator!) + guard drained.elements.last == 1 else { + return XCTFail("Expected to receive the sent element, received \(drained.elements)") + } + } + } + + func test_subscription_racing_termination_is_terminated() async { + for _ in 0..<10_000 { + let sut = AsyncThrowingReplaySubject(bufferSize: 2) + sut.send(0) + var iterator: AsyncThrowingReplaySubject.Iterator? + + race({ iterator = sut.makeAsyncIterator() }, { sut.send(.failure(MockError(code: 1))) }) + + let drained = await drainBufferedElements(of: iterator!) + guard drained.isTerminated else { + return XCTFail("Expected the subscription to be terminated") + } + } + } } diff --git a/Tests/Supporting/Helpers.swift b/Tests/Supporting/Helpers.swift index 355f799..2f7a921 100644 --- a/Tests/Supporting/Helpers.swift +++ b/Tests/Supporting/Helpers.swift @@ -5,6 +5,9 @@ // Created by Thibault Wittemberg on 11/09/2022. // +@testable import AsyncExtensions +import Dispatch + struct Indefinite: Sequence, IteratorProtocol, Sendable { let value: Element @@ -49,3 +52,29 @@ struct Tuple3: Equatable { self.value3 = values.2 } } + +func race(_ first: () -> Void, _ second: () -> Void) { + DispatchQueue.concurrentPerform(iterations: 2) { index in + index == 0 ? first() : second() + } +} + +func drainBufferedElements( + of iterator: Iterator +) async -> (elements: [Iterator.Element], isTerminated: Bool) { + var iterator = iterator + var elements = [Iterator.Element]() + + while iterator.hasBufferedElements { + do { + guard let element = try await iterator.next() else { + return (elements, true) + } + elements.append(element) + } catch { + return (elements, true) + } + } + + return (elements, false) +}