From 5cfb8ef648770f979e357616750f7313253763dd Mon Sep 17 00:00:00 2001 From: Thibault Wittemberg Date: Sat, 3 Oct 2026 11:46:18 +0200 Subject: [PATCH] Fix terminal cancellation in switchToLatest Record cancellation in the protected iterator state, resume outer waiters, and cancel producer tasks outside the lock. Reject late transitions and replacement child tasks while preserving normal latest-child draining. Add nine regression tests, including non-cooperative outer sequences and cancellation racing delivery, and document the fix in the changelog. --- CHANGELOG.md | 1 + .../AsyncSwitchToLatestSequence.swift | 61 +++- ...AsyncSwitchToLatestCancellationTests.swift | 307 ++++++++++++++++++ 3 files changed, 362 insertions(+), 7 deletions(-) create mode 100644 Tests/Operators/AsyncSwitchToLatestCancellationTests.swift diff --git a/CHANGELOG.md b/CHANGELOG.md index 5a02bc5..61f7ef1 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,6 @@ **Unreleased:** +- SwitchToLatest: finish cancelled collection while the latest channel or outer sequence remains open, and discard late producer results (https://github.com/sideeffect-io/AsyncExtensions/issues/53). - Subjects: fix a deadlock when sending values or termination concurrently with consumer cancellation (https://github.com/sideeffect-io/AsyncExtensions/issues/52). - Multicast: preserve upstream element order and prevent termination from overtaking values when consumers advance concurrently. diff --git a/Sources/Operators/AsyncSwitchToLatestSequence.swift b/Sources/Operators/AsyncSwitchToLatestSequence.swift index 30147b8..989a7fb 100644 --- a/Sources/Operators/AsyncSwitchToLatestSequence.swift +++ b/Sources/Operators/AsyncSwitchToLatestSequence.swift @@ -51,6 +51,14 @@ where Base.Element: AsyncSequence, Base: Sendable, Base.Element.Element: Sendabl case processingChildIterator(Result) case finished(Result?) case failed(Error) + case cancelled + + var isCancelled: Bool { + if case .cancelled = self { + return true + } + return false + } var isFinished: Bool { if case .finished = self { @@ -94,10 +102,11 @@ where Base.Element: AsyncSequence, Base: Sendable, Base.Element.Element: Sendabl enum BaseDecision { case resumeNext(UnsafeContinuation?, Never>, Task?) case cancelPreviousChildTask(Task?) + case stop } enum NextDecision { - case immediatelyResume(Task) + case immediatelyResume(Task?) case suspend } @@ -138,8 +147,10 @@ where Base.Element: AsyncSequence, Base: Sendable, Base.Element.Element: Sendabl self.baseTask = Task { [base, state] in do { for try await child in base { + guard !Task.isCancelled else { return } let childIterator = child.makeAsyncIterator() let decision = state.withCriticalRegion { state -> BaseDecision in + guard !state.base.isCancelled else { return .stop } switch state.base { case .waitingForChildIterator(let continuation): state.base = .processingChildIterator(.success(childIterator)) @@ -153,6 +164,8 @@ where Base.Element: AsyncSequence, Base: Sendable, Base.Element.Element: Sendabl } switch decision { + case .stop: + return case .cancelPreviousChildTask(let task): task?.cancel() case .resumeNext(let continuation, let childTask): @@ -161,6 +174,7 @@ where Base.Element: AsyncSequence, Base: Sendable, Base.Element.Element: Sendabl } let decision = state.withCriticalRegion { state -> BaseDecision in + guard !state.base.isCancelled else { return .stop } switch state.base { case .waitingForChildIterator(let continuation): state.base = .finished(nil) @@ -172,6 +186,8 @@ where Base.Element: AsyncSequence, Base: Sendable, Base.Element.Element: Sendabl } switch decision { + case .stop: + return case .cancelPreviousChildTask: break case .resumeNext(let continuation, let childTask): @@ -179,6 +195,7 @@ where Base.Element: AsyncSequence, Base: Sendable, Base.Element.Element: Sendabl } } catch { let decision = state.withCriticalRegion { state -> BaseDecision in + guard !state.base.isCancelled else { return .stop } switch state.base { case .waitingForChildIterator(let continuation): state.base = .failed(error) @@ -192,6 +209,8 @@ where Base.Element: AsyncSequence, Base: Sendable, Base.Element.Element: Sendabl } switch decision { + case .stop: + return case .cancelPreviousChildTask(let task): task?.cancel() case .resumeNext(let continuation, let childTask): @@ -217,8 +236,34 @@ where Base.Element: AsyncSequence, Base: Sendable, Base.Element.Element: Sendabl } } + static func cancel(baseTask: Task?, state: ManagedCriticalState) { + let cancellation = state.withCriticalRegion { state -> ( + continuation: UnsafeContinuation?, Never>?, + childTask: Task? + ) in + let continuation: UnsafeContinuation?, Never>? + if case .waitingForChildIterator(let waiting) = state.base { + continuation = waiting + } else { + continuation = nil + } + let childTask = state.childTask + // Claim the waiter and discard retained iterators before invoking any cancellation handlers. + state.base = .cancelled + state.childTask = nil + return (continuation, childTask) + } + + cancellation.continuation?.resume(returning: nil) + baseTask?.cancel() + cancellation.childTask?.cancel() + } + public mutating func next() async rethrows -> Element? { - guard !Task.isCancelled else { return nil } + guard !Task.isCancelled else { + Self.cancel(baseTask: self.baseTask, state: self.state) + return nil + } self.startBase() return try await withTaskCancellationHandler { @@ -226,6 +271,8 @@ where Base.Element: AsyncSequence, Base: Sendable, Base.Element.Element: Sendabl let childTask = await withUnsafeContinuation { [state] (continuation: UnsafeContinuation?, Never>) in let decision = state.withCriticalRegion { state -> NextDecision in switch state.base { + case .cancelled: + return .immediatelyResume(nil) case .newChildIteratorAvailable(let childIterator): state.base = .processingChildIterator(childIterator) let childTask = Self.makeChildTask(childIterator: childIterator) @@ -260,7 +307,10 @@ where Base.Element: AsyncSequence, Base: Sendable, Base.Element.Element: Sendabl let value = await childTask?.value let decision = state.withCriticalRegion { state -> PostElementDecision in - if state.base.isNewAvailableChildIterator { + state.childTask = nil + if state.base.isCancelled { + return .returnFinish + } else if state.base.isNewAvailableChildIterator { return .pass } else { switch value { @@ -299,10 +349,7 @@ where Base.Element: AsyncSequence, Base: Sendable, Base.Element.Element: Sendabl } } } onCancel: { [baseTask, state] in - baseTask?.cancel() - state.withCriticalRegion { - $0.childTask?.cancel() - } + Self.cancel(baseTask: baseTask, state: state) } } } diff --git a/Tests/Operators/AsyncSwitchToLatestCancellationTests.swift b/Tests/Operators/AsyncSwitchToLatestCancellationTests.swift new file mode 100644 index 0000000..4320d94 --- /dev/null +++ b/Tests/Operators/AsyncSwitchToLatestCancellationTests.swift @@ -0,0 +1,307 @@ +@testable import AsyncExtensions +import XCTest + +// The gate deliberately ignores cancellation, but can always be released by test cleanup. +private struct SwitchCancellationGate: Sendable { + struct State { + var continuation: UnsafeContinuation? + var isReleased = false + } + + let state = ManagedCriticalState(State()) + let suspended: XCTestExpectation + + func wait() async { + await withUnsafeContinuation { continuation in + let resume = state.withCriticalRegion { state -> Bool in + guard !state.isReleased else { return true } + state.continuation = continuation + return false + } + if resume { continuation.resume() } + suspended.fulfill() + } + } + + func release() { + let continuation = state.withCriticalRegion { state -> UnsafeContinuation? in + state.isReleased = true + defer { state.continuation = nil } + return state.continuation + } + continuation?.resume() + } +} + +private struct GatedSwitchSequence: AsyncSequence, Sendable { + let gate: SwitchCancellationGate + let result: Result + + func makeAsyncIterator() -> Iterator { Iterator(gate: gate, result: result) } + + struct Iterator: AsyncIteratorProtocol, Sendable { + let gate: SwitchCancellationGate + let result: Result + var hasReturned = false + + mutating func next() async throws -> Element? { + guard !hasReturned else { return nil } + hasReturned = true + await gate.wait() + return try result.get() + } + } +} + +private struct ObservedSwitchChild: AsyncSequence, Sendable { + let channel: AsyncBufferedChannel + let suspended: XCTestExpectation + + func makeAsyncIterator() -> Iterator { Iterator(channel: channel, suspended: suspended) } + + struct Iterator: AsyncIteratorProtocol, Sendable { + let channel: AsyncBufferedChannel + let suspended: XCTestExpectation + let didSuspend = ManagedCriticalState(false) + + func next() async -> Int? { + await channel.next(onSuspend: { + let firstSuspension = didSuspend.withCriticalRegion { didSuspend -> Bool in + defer { didSuspend = true } + return !didSuspend + } + if firstSuspension { suspended.fulfill() } + }) + } + } +} + +final class AsyncSwitchToLatestCancellationTests: XCTestCase { + func test_cancellation_while_latest_child_is_suspended_finishes_collection() async { + let (continuation, outer) = AsyncStream.pipe() + let received = (1...3).map { expectation(description: "Received \($0)") } + let suspended = (1...3).map { expectation(description: "Child \($0) suspended") } + let channels = (1...3).map { _ in AsyncBufferedChannel() } + let finished = expectation(description: "Collection finished") + let task = Task { + var values = [Int]() + for await value in outer.switchToLatest() { + values.append(value) + if (1...3).contains(value) { received[value - 1].fulfill() } + } + XCTAssertEqual(values, [1, 2, 3]) + finished.fulfill() + } + + for index in channels.indices { + channels[index].send(index + 1) + continuation.yield(ObservedSwitchChild(channel: channels[index], suspended: suspended[index])) + await fulfillment(of: [received[index], suspended[index]], timeout: 2) + } + task.cancel() + await fulfillment(of: [finished], timeout: 2) + + // Cleanup also lets the unfixed implementation finish after its timeout failure. + channels.forEach { $0.finish() } + continuation.finish() + await task.value + } + + func test_cancellation_while_waiting_for_non_cooperative_outer_ignores_late_child() async { + await assertCancellationWhileWaitingForOuter(result: .success(AsyncBufferedChannel())) + } + + func test_cancellation_while_waiting_for_non_cooperative_outer_ignores_late_error() async { + await assertCancellationWhileWaitingForOuter(result: .failure(MockError(code: 53))) + } + + private func assertCancellationWhileWaitingForOuter(result: Result?, Error>) async { + let gate = SwitchCancellationGate(suspended: expectation(description: "Outer suspended")) + var iterator = GatedSwitchSequence(gate: gate, result: result).switchToLatest().makeAsyncIterator() + let state = iterator.state + let finished = expectation(description: "Consumer finished before outer released") + let task = Task { + do { + let value = try await iterator.next() + XCTAssertNil(value) + finished.fulfill() + // Reuse from an uncancelled task must not resurrect the cancelled iterator. + return iterator + } catch { + XCTFail("Cancelled iteration delivered a late error: \(error)") + finished.fulfill() + return iterator + } + } + + await fulfillment(of: [gate.suspended], timeout: 2) + await waitUntil { + state.withCriticalRegion { + if case .waitingForChildIterator = $0.base { return true } + return false + } + } + task.cancel() + await fulfillment(of: [finished], timeout: 2) + gate.release() + if case .success(let channel) = result { channel?.finish() } + var cancelledIterator = await task.value + // Await the producer so the assertion also covers its late success/failure transition. + await cancelledIterator.baseTask?.value + do { + let value = try await cancelledIterator.next() + XCTAssertNil(value) + } catch { + XCTFail("Late outer result resurrected cancelled iteration: \(error)") + } + } + + func test_finite_outer_allows_latest_child_to_drain() async { + let channel = AsyncBufferedChannel() + let suspended = expectation(description: "Latest child suspended after outer finished") + let child = ObservedSwitchChild(channel: channel, suspended: suspended) + var iterator = [child].async.switchToLatest().makeAsyncIterator() + let state = iterator.state + let task = Task { () -> [Int] in + var values = [Int]() + while let value = await iterator.next() { values.append(value) } + return values + } + await fulfillment(of: [suspended], timeout: 2) + await waitUntil { state.withCriticalRegion { $0.base.isFinished } } + channel.send(10) + channel.send(20) + channel.finish() + let values = await task.value + XCTAssertEqual(values, [10, 20]) + } + + func test_cancellation_discards_late_inner_value() async { + await assertCancellationDiscardsInnerResult(.success(53)) + } + + func test_cancellation_discards_late_inner_error() async { + await assertCancellationDiscardsInnerResult(.failure(MockError(code: 53))) + } + + private func assertCancellationDiscardsInnerResult(_ result: Result) async { + let gate = SwitchCancellationGate(suspended: expectation(description: "Inner suspended")) + let child = GatedSwitchSequence(gate: gate, result: result) + var iterator = [child].async.switchToLatest().makeAsyncIterator() + let task = Task { + do { + let value = try await iterator.next() + XCTAssertNil(value) + } catch { + XCTFail("Cancelled iteration delivered a late inner error: \(error)") + } + return iterator + } + await fulfillment(of: [gate.suspended], timeout: 2) + task.cancel() + // A non-cooperative inner task must return before its task.value await can complete. + gate.release() + iterator = await task.value + XCTAssertTrue(iterator.state.criticalState.base.isCancelled) + XCTAssertNil(iterator.state.criticalState.childTask) + await iterator.baseTask?.value + } + + func test_cancellation_racing_outer_delivery_finishes_without_resurrecting_iteration() async { + for _ in 0..<100 { + let gate = SwitchCancellationGate(suspended: expectation(description: "Outer suspended")) + let child = AsyncBufferedChannel() + child.send(53) + child.finish() + var iterator = GatedSwitchSequence(gate: gate, result: .success(child)).switchToLatest().makeAsyncIterator() + let finished = expectation(description: "Racing consumer finished") + let task = Task { + do { + while let value = try await iterator.next() { XCTAssertEqual(value, 53) } + } catch { + XCTFail("Unexpected failure: \(error)") + } + finished.fulfill() + return iterator + } + await fulfillment(of: [gate.suspended], timeout: 2) + race({ task.cancel() }, { gate.release() }) + await fulfillment(of: [finished], timeout: 2) + iterator = await task.value + await iterator.baseTask?.value + do { + let value = try await iterator.next() + XCTAssertNil(value) + } catch { + XCTFail("Finished iterator delivered a late error: \(error)") + } + } + } + + func test_already_cancelled_task_does_not_start_outer_iteration() async { + let (continuation, outer) = AsyncStream>.pipe() + let gate = SwitchCancellationGate(suspended: expectation(description: "Consumer paused before next")) + let task = Task { + await gate.wait() + var iterator = outer.switchToLatest().makeAsyncIterator() + let value = await iterator.next() + XCTAssertNil(value) + XCTAssertNil(iterator.baseTask) + return iterator + } + await fulfillment(of: [gate.suspended], timeout: 2) + task.cancel() + gate.release() + var iterator = await task.value + let child = AsyncBufferedChannel() + child.send(53) + child.finish() + continuation.yield(child) + continuation.finish() + let value = await iterator.next() + XCTAssertNil(value) + XCTAssertNil(iterator.baseTask) + XCTAssertNil(iterator.state.criticalState.childTask) + } + + func test_cancellation_between_next_calls_cancels_existing_outer_task() async { + let (continuation, outer) = AsyncStream>.pipe() + let channel = AsyncBufferedChannel() + channel.send(1) + continuation.yield(channel) + var iterator = outer.switchToLatest().makeAsyncIterator() + let first = await iterator.next() + XCTAssertEqual(first, 1) + + let terminated = expectation(description: "Outer cancelled") + continuation.onTermination = { _ in terminated.fulfill() } + let gate = SwitchCancellationGate(suspended: expectation(description: "Consumer between next calls")) + let task = Task { + await gate.wait() + let value = await iterator.next() + XCTAssertNil(value) + return iterator + } + await fulfillment(of: [gate.suspended], timeout: 2) + task.cancel() + gate.release() + iterator = await task.value + await fulfillment(of: [terminated], timeout: 2) + channel.finish() + continuation.finish() + let value = await iterator.next() + XCTAssertNil(value) + XCTAssertNil(iterator.state.criticalState.childTask) + } + + // Observe the actual state rather than assume a fixed number of scheduler yields is enough. + private func waitUntil(_ condition: () -> Bool) async { + let deadline = DispatchTime.now().uptimeNanoseconds + 2_000_000_000 + while !condition() { + guard DispatchTime.now().uptimeNanoseconds < deadline else { + return XCTFail("Iterator did not reach the expected state") + } + await Task.yield() + } + } +}