diff --git a/Sources/FluidAudio/Diarizer/Offline/Core/OfflineDiarizerManager.swift b/Sources/FluidAudio/Diarizer/Offline/Core/OfflineDiarizerManager.swift index 17f04669b..e618f980f 100644 --- a/Sources/FluidAudio/Diarizer/Offline/Core/OfflineDiarizerManager.swift +++ b/Sources/FluidAudio/Diarizer/Offline/Core/OfflineDiarizerManager.swift @@ -3,6 +3,12 @@ import Accelerate import Foundation import OSLog +@available(macOS 14.0, iOS 17.0, *) +private enum OfflinePreparationWorkerResult: Sendable { + case segmentation(SegmentationOutput, TimeInterval) + case embeddings([TimedEmbedding], TimeInterval) +} + @available(macOS 14.0, iOS 17.0, *) public final class OfflineDiarizerManager { private let logger = AppLogger(category: "OfflineDiarizer") @@ -187,64 +193,93 @@ public final class OfflineDiarizerManager { let capturedModels = models let capturedConfig = config - let segmentationTask = Task.detached(priority: .userInitiated) { - [capturedModels, capturedConfig] () throws -> (SegmentationOutput, TimeInterval) in - let processor = OfflineSegmentationProcessor() - let start = Date() - do { - let segmentation = try await processor.process( - audioSource: audioSource, - segmentationModel: capturedModels.segmentationModel, - config: capturedConfig, - chunkHandler: { chunk in - progressCallback?(chunk.chunkIndex + 1, totalChunks) - switch chunkContinuation.yield(chunk) { - case .enqueued, .dropped: - return .continue - case .terminated: - return .stop - @unknown default: - return .stop + let results = try await withThrowingTaskGroup( + of: OfflinePreparationWorkerResult.self, + returning: ( + segmentation: (SegmentationOutput, TimeInterval), + embeddings: ([TimedEmbedding], TimeInterval) + ).self + ) { group in + group.addTask(priority: .userInitiated) { [capturedModels, capturedConfig] in + let processor = OfflineSegmentationProcessor() + let start = Date() + do { + let segmentation = try await processor.process( + audioSource: audioSource, + segmentationModel: capturedModels.segmentationModel, + config: capturedConfig, + chunkHandler: { chunk in + progressCallback?(chunk.chunkIndex + 1, totalChunks) + switch chunkContinuation.yield(chunk) { + case .enqueued, .dropped: + return .continue + case .terminated: + return .stop + @unknown default: + return .stop + } } - } + ) + chunkContinuation.finish() + return .segmentation( + segmentation, + Date().timeIntervalSince(start) + ) + } catch { + chunkContinuation.finish(throwing: error) + throw error + } + } + + group.addTask(priority: .userInitiated) { [capturedModels, capturedConfig] in + let extractor = OfflineEmbeddingExtractor( + fbankModel: capturedModels.fbankModel, + embeddingModel: capturedModels.embeddingModel, + pldaTransform: PLDATransform( + pldaRhoModel: capturedModels.pldaRhoModel, + psi: capturedModels.pldaPsi + ), + config: capturedConfig + ) + let start = Date() + let embeddings = try await extractor.extractEmbeddings( + audioSource: audioSource, + segmentationStream: chunkStream + ) + return .embeddings( + embeddings, + Date().timeIntervalSince(start) ) - chunkContinuation.finish() - return (segmentation, Date().timeIntervalSince(start)) + } + + var segmentationResult: (SegmentationOutput, TimeInterval)? + var embeddingResult: ([TimedEmbedding], TimeInterval)? + + do { + while let result = try await group.next() { + switch result { + case .segmentation(let segmentation, let duration): + segmentationResult = (segmentation, duration) + case .embeddings(let embeddings, let duration): + embeddingResult = (embeddings, duration) + } + } } catch { + group.cancelAll() chunkContinuation.finish(throwing: error) throw error } - } - let embeddingTask = Task.detached(priority: .userInitiated) { - [capturedModels, capturedConfig] () throws -> ([TimedEmbedding], TimeInterval) in - let extractor = OfflineEmbeddingExtractor( - fbankModel: capturedModels.fbankModel, - embeddingModel: capturedModels.embeddingModel, - pldaTransform: PLDATransform(pldaRhoModel: capturedModels.pldaRhoModel, psi: capturedModels.pldaPsi), - config: capturedConfig - ) - let start = Date() - let embeddings = try await extractor.extractEmbeddings( - audioSource: audioSource, - segmentationStream: chunkStream - ) - return (embeddings, Date().timeIntervalSince(start)) + guard let segmentationResult, let embeddingResult else { + throw OfflineDiarizationError.processingFailed( + "Offline preparation workers ended without complete results" + ) + } + return (segmentationResult, embeddingResult) } - let segmentationResult: (SegmentationOutput, TimeInterval) - let embeddingResult: ([TimedEmbedding], TimeInterval) - do { - async let awaitedSegmentation = segmentationTask.value - async let awaitedEmbeddings = embeddingTask.value - segmentationResult = try await awaitedSegmentation - embeddingResult = try await awaitedEmbeddings - } catch { - segmentationTask.cancel() - embeddingTask.cancel() - chunkContinuation.finish(throwing: error) - throw error - } + let segmentationResult = results.segmentation + let embeddingResult = results.embeddings let (segmentation, segmentationTime) = segmentationResult logger.debug("Segmentation completed in \(segmentationTime)s (async)") diff --git a/Tests/FluidAudioTests/Diarizer/Offline/OfflineDiarizerCancellationTests.swift b/Tests/FluidAudioTests/Diarizer/Offline/OfflineDiarizerCancellationTests.swift new file mode 100644 index 000000000..8491ce7bf --- /dev/null +++ b/Tests/FluidAudioTests/Diarizer/Offline/OfflineDiarizerCancellationTests.swift @@ -0,0 +1,80 @@ +import os +import XCTest + +@testable import FluidAudio + +@available(macOS 14.0, iOS 17.0, *) +final class OfflineDiarizerCancellationTests: XCTestCase { + func testCancellingFileProcessingStopsAllInferenceWorkers() async throws { + try requireOfflineDiarizerModels() + + guard + let audioPath = ProcessInfo.processInfo.environment[ + "FLUIDAUDIO_CANCELLATION_TEST_AUDIO" + ], + FileManager.default.fileExists(atPath: audioPath) + else { + throw XCTSkip( + "Set FLUIDAUDIO_CANCELLATION_TEST_AUDIO to a real meeting recording" + ) + } + let audioURL = URL(fileURLWithPath: audioPath) + + let manager = OfflineDiarizerManager() + let progress = AsyncStream.makeStream(bufferingPolicy: .bufferingNewest(1)) + let progressCount = OSAllocatedUnfairLock(initialState: 0) + var progressIterator = progress.stream.makeAsyncIterator() + + let processingTask = Task { + try await manager.process(audioURL) { _, _ in + progressCount.withLock { $0 += 1 } + progress.continuation.yield(()) + } + } + + guard await progressIterator.next() != nil else { + processingTask.cancel() + XCTFail("Diarization ended before inference began") + return + } + + let cancellationStart = ContinuousClock.now + processingTask.cancel() + + do { + _ = try await processingTask.value + XCTFail("Cancelled diarization unexpectedly completed") + } catch is CancellationError { + // Expected. Returning from the structured task group also proves both + // inference workers have stopped rather than continuing in the background. + } catch { + XCTFail("Expected CancellationError, received \(error)") + } + + let cancellationDuration = cancellationStart.duration(to: .now) + XCTAssertLessThan( + cancellationDuration, + .seconds(2), + "Cancellation should stop segmentation and embedding between model calls" + ) + + let countWhenCancelledCallReturned = progressCount.withLock { $0 } + try await Task.sleep(for: .milliseconds(250)) + XCTAssertEqual( + progressCount.withLock { $0 }, + countWhenCancelledCallReturned, + "No inference progress may continue after process(URL) returns" + ) + } + + private func requireOfflineDiarizerModels() throws { + let repoDirectory = OfflineDiarizerModels.defaultModelsDirectory() + .appendingPathComponent(Repo.diarizer.folderName, isDirectory: true) + let allPresent = ModelNames.OfflineDiarizer.requiredModels.allSatisfy { + FileManager.default.fileExists(atPath: repoDirectory.appendingPathComponent($0).path) + } + guard allPresent else { + throw XCTSkip("Offline diarizer models not available") + } + } +}