diff --git a/Tests/Operators/AsyncMulticastSequenceTests.swift b/Tests/Operators/AsyncMulticastSequenceTests.swift index 8693e11..c94bd39 100644 --- a/Tests/Operators/AsyncMulticastSequenceTests.swift +++ b/Tests/Operators/AsyncMulticastSequenceTests.swift @@ -5,7 +5,7 @@ // Created by Thibault Wittemberg on 21/02/2022. // -import AsyncExtensions +@testable import AsyncExtensions import XCTest private class SpyAsyncSequenceForNumberOfIterators: AsyncSequence { @@ -39,7 +39,142 @@ private class SpyAsyncSequenceForNumberOfIterators: AsyncSequence { } } +/// An upstream the test feeds by hand; `demanded` fires each time a consumer enters `next()`. +private struct GatedUpstream: AsyncSequence, Sendable { + typealias Element = Int + + let stream: AsyncThrowingStream + let demanded: @Sendable () -> Void + + func makeAsyncIterator() -> Iterator { + Iterator(base: self.stream.makeAsyncIterator(), demanded: self.demanded) + } + + struct Iterator: AsyncIteratorProtocol, @unchecked Sendable { + var base: AsyncThrowingStream.Iterator + let demanded: @Sendable () -> Void + + mutating func next() async throws -> Int? { + self.demanded() + return try await self.base.next() + } + } +} + +private enum SubscriberExit { + case cancelTask + case releaseIterator +} + final class AsyncMulticastSequenceTests: XCTestCase { + func test_cancelling_one_subscriber_task_keeps_upstream_for_the_other_when_connected() async { + await self.assertSurvivingSubscriberCompletes(autoconnect: false, exit: .cancelTask) + } + + func test_cancelling_one_subscriber_task_keeps_upstream_for_the_other_when_autoconnected() async { + await self.assertSurvivingSubscriberCompletes(autoconnect: true, exit: .cancelTask) + } + + func test_releasing_one_subscriber_iterator_keeps_upstream_for_the_other_when_connected() async { + await self.assertSurvivingSubscriberCompletes(autoconnect: false, exit: .releaseIterator) + } + + func test_releasing_one_subscriber_iterator_keeps_upstream_for_the_other_when_autoconnected() async { + await self.assertSurvivingSubscriberCompletes(autoconnect: true, exit: .releaseIterator) + } + + /// Subscriber A is the one pulling on the shared upstream when it leaves; subscriber B, already + /// registered and waiting on the subject, must still receive every element and finish normally. + private func assertSurvivingSubscriberCompletes(autoconnect: Bool, exit: SubscriberExit) async { + let (stream, continuation) = AsyncThrowingStream.makeStream() + let demanded = expectation(description: "A consumer is awaiting the upstream") + demanded.assertForOverFulfill = false + let upstream = GatedUpstream(stream: stream) { demanded.fulfill() } + + let subject = AsyncThrowingPassthroughSubject() + let multicasted = upstream.multicast(subject) + let sut = autoconnect ? multicasted.autoconnect() : multicasted + let subscriberCount = { subject.state.withCriticalRegion { $0.channels.count } } + + // B registers with the subject before any element flows. + let iteratorB = sut.makeAsyncIterator() + + let aHasLeft = expectation(description: "Subscriber A has left") + let bHasFinished = expectation(description: "Subscriber B has finished") + + // A's iterator exists only inside this task, so no copy of it outlives A. + let taskA = Task { () throws -> [Int] in + defer { aHasLeft.fulfill() } + var iterator = sut.makeAsyncIterator() + var received = [Int]() + switch exit { + case .cancelTask: + while let element = try await iterator.next() { + received.append(element) + } + case .releaseIterator: + if let element = try await iterator.next() { + received.append(element) + } + } + return received + } + + if !autoconnect { + sut.connect() + } + + // A is now inside the upstream's next(); only then does B ask for elements, so A is the puller. + await fulfillment(of: [demanded], timeout: 1) + XCTAssertEqual(subscriberCount(), 2) + + let taskB = Task { () throws -> [Int] in + defer { bHasFinished.fulfill() } + var iterator = iteratorB + var received = [Int]() + while let element = try await iterator.next() { + received.append(element) + } + return received + } + + let expectedForA: [Int] + switch exit { + case .cancelTask: + // A must finish while upstream is silent; any element sent now could wake it. + taskA.cancel() + expectedForA = [] + case .releaseIterator: + continuation.yield(1) + expectedForA = [1] + } + await fulfillment(of: [aHasLeft], timeout: 1) + XCTAssertEqual(subscriberCount(), 1, "Subscriber A is still registered with the subject") + + if exit == .cancelTask { + continuation.yield(1) + } + continuation.yield(2) + continuation.yield(3) + continuation.finish() + + await fulfillment(of: [bHasFinished], timeout: 1) + taskB.cancel() + // A's result was settled when it left, before any of these elements were sent. + do { + let received = try await taskA.value + XCTAssertEqual(received, expectedForA) + } catch { + XCTFail("Subscriber A should leave without an error, got \(error)") + } + do { + let received = try await taskB.value + XCTAssertEqual(received, [1, 2, 3]) + } catch { + XCTFail("Subscriber B should finish normally, got \(error)") + } + } + func test_concurrent_consumers_receive_all_elements_in_order_before_finish() async { await assertConcurrentDelivery() }