Skip to content

Commit c797939

Browse files
committed
Fix transcription timeout handling
1 parent 5a3fbc2 commit c797939

1 file changed

Lines changed: 100 additions & 19 deletions

File tree

Sources/TranscriptionService.swift

Lines changed: 100 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -54,29 +54,41 @@ class TranscriptionService {
5454
}
5555

5656
let timeoutSeconds = transcriptionTimeoutSeconds
57-
return try await withThrowingTaskGroup(of: String.self) { group in
58-
group.addTask { [weak self] in
59-
guard let self else {
60-
throw TranscriptionError.transcriptionFailed("Transcription service deallocated")
57+
let raceState = TranscriptionTimeoutRaceState()
58+
59+
return try await withTaskCancellationHandler {
60+
try await withCheckedThrowingContinuation { continuation in
61+
raceState.setContinuation(continuation)
62+
63+
let transcriptionTask = Task { [weak self] in
64+
do {
65+
guard let self else {
66+
throw TranscriptionError.transcriptionFailed("Transcription service deallocated")
67+
}
68+
let result = try await self.transcribeAudio(fileURL: fileURL)
69+
raceState.finish(.success(result))
70+
} catch {
71+
raceState.finish(.failure(Self.transcriptionTimeoutErrorIfNeeded(
72+
error,
73+
timeoutSeconds: timeoutSeconds
74+
)))
75+
}
6176
}
62-
return try await self.transcribeAudio(fileURL: fileURL)
63-
}
64-
65-
group.addTask {
66-
try await Task.sleep(nanoseconds: UInt64(timeoutSeconds * 1_000_000_000))
67-
throw TranscriptionError.transcriptionTimedOut(timeoutSeconds)
68-
}
6977

70-
do {
71-
guard let result = try await group.next() else {
72-
throw TranscriptionError.transcriptionFailed("No transcription result")
78+
let timeoutTask = Task {
79+
do {
80+
try await Task.sleep(nanoseconds: UInt64(timeoutSeconds * 1_000_000_000))
81+
raceState.finish(.failure(TranscriptionError.transcriptionTimedOut(timeoutSeconds)))
82+
} catch is CancellationError {
83+
} catch {
84+
raceState.finish(.failure(error))
85+
}
7386
}
74-
group.cancelAll()
75-
return result
76-
} catch {
77-
group.cancelAll()
78-
throw error
87+
88+
raceState.setTasks([transcriptionTask, timeoutTask])
7989
}
90+
} onCancel: {
91+
raceState.cancel()
8092
}
8193
}
8294

@@ -228,6 +240,16 @@ class TranscriptionService {
228240
}
229241
}
230242

243+
private static func transcriptionTimeoutErrorIfNeeded(
244+
_ error: Error,
245+
timeoutSeconds: TimeInterval
246+
) -> Error {
247+
if let urlError = error as? URLError, urlError.code == .timedOut {
248+
return TranscriptionError.transcriptionTimedOut(timeoutSeconds)
249+
}
250+
return error
251+
}
252+
231253
private static func normalizedBaseURL(from baseURL: String) throws -> URL {
232254
let trimmed = baseURL.trimmingCharacters(in: .whitespacesAndNewlines)
233255
guard !trimmed.isEmpty else {
@@ -361,6 +383,65 @@ enum TranscriptionError: LocalizedError {
361383
}
362384
}
363385

386+
private final class TranscriptionTimeoutRaceState {
387+
private let lock = NSLock()
388+
private var didFinish = false
389+
private var continuation: CheckedContinuation<String, Error>?
390+
private var tasks: [Task<Void, Never>] = []
391+
392+
func setContinuation(_ continuation: CheckedContinuation<String, Error>) {
393+
lock.lock()
394+
if didFinish {
395+
lock.unlock()
396+
continuation.resume(throwing: CancellationError())
397+
return
398+
}
399+
400+
self.continuation = continuation
401+
lock.unlock()
402+
}
403+
404+
func setTasks(_ tasks: [Task<Void, Never>]) {
405+
lock.lock()
406+
if didFinish {
407+
lock.unlock()
408+
tasks.forEach { $0.cancel() }
409+
return
410+
}
411+
412+
self.tasks = tasks
413+
lock.unlock()
414+
}
415+
416+
func finish(_ result: Result<String, Error>) {
417+
lock.lock()
418+
guard !didFinish else {
419+
lock.unlock()
420+
return
421+
}
422+
423+
didFinish = true
424+
let continuation = self.continuation
425+
self.continuation = nil
426+
let tasks = self.tasks
427+
self.tasks = []
428+
lock.unlock()
429+
430+
tasks.forEach { $0.cancel() }
431+
432+
switch result {
433+
case .success(let value):
434+
continuation?.resume(returning: value)
435+
case .failure(let error):
436+
continuation?.resume(throwing: error)
437+
}
438+
}
439+
440+
func cancel() {
441+
finish(.failure(CancellationError()))
442+
}
443+
}
444+
364445
private struct PreparedUploadAudio {
365446
let fileURL: URL
366447
let deleteOnCleanup: Bool

0 commit comments

Comments
 (0)