Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
37 changes: 19 additions & 18 deletions TypeWhisper/ViewModels/AudioRecorderViewModel.swift
Original file line number Diff line number Diff line change
Expand Up @@ -1115,15 +1115,7 @@ final class AudioRecorderViewModel: ObservableObject {
let result = if let liveSessionResult = request.liveSessionResult {
liveSessionResult
} else {
try await modelManager.transcribe(
audioSamples: buffer,
languageSelection: request.languageSelection,
task: effectiveTask,
engineOverrideId: request.providerId,
cloudModelOverride: request.modelOverrideId,
prompt: request.prompt,
dictionaryTermHints: request.dictionaryTermHints
)
try await transcribeFinalRecording(request, task: effectiveTask)
}
let text = result.text.trimmingCharacters(in: .whitespacesAndNewlines)
if !text.isEmpty {
Expand Down Expand Up @@ -1162,15 +1154,7 @@ final class AudioRecorderViewModel: ObservableObject {
let effectiveTask = resolvedTask(for: request)

do {
let result = try await modelManager.transcribe(
audioSamples: request.buffer,
languageSelection: request.languageSelection,
task: effectiveTask,
engineOverrideId: request.providerId,
cloudModelOverride: request.modelOverrideId,
prompt: request.prompt,
dictionaryTermHints: request.dictionaryTermHints
)
let result = try await transcribeFinalRecording(request, task: effectiveTask)
let text = result.text.trimmingCharacters(in: .whitespacesAndNewlines)
guard !text.isEmpty else {
let failure = makeTranscriptionFailure(
Expand Down Expand Up @@ -1200,6 +1184,23 @@ final class AudioRecorderViewModel: ObservableObject {
}
}

private func transcribeFinalRecording(
_ request: FinalTranscriptionRequest,
task: TranscriptionTask
) async throws -> TranscriptionResult {
try await modelManager.transcribe(
audioSamples: request.buffer,
languageSelection: request.languageSelection,
task: task,
engineOverrideId: request.providerId,
cloudModelOverride: request.modelOverrideId,
prompt: request.prompt,
dictionaryTermHints: request.dictionaryTermHints,
onProgress: { _ in true },
onSourceProgress: { _ in true }
)
}

private func resolvedTask(for request: FinalTranscriptionRequest) -> TranscriptionTask {
guard request.task == .translate,
let providerId = request.providerId,
Expand Down
47 changes: 45 additions & 2 deletions TypeWhisperTests/AudioRecorderViewModelTests.swift
Original file line number Diff line number Diff line change
Expand Up @@ -524,6 +524,7 @@ final class AudioRecorderViewModelTests: XCTestCase {
XCTAssertEqual(request.audioSampleCount, finalizedFileSamples.count)
XCTAssertEqual(request.firstAudioSample, finalizedFileSamples.first)
XCTAssertNotEqual(request.audioSampleCount, captureSamples.count)
XCTAssertTrue(request.usedFileTranscriptionPipeline)
let recording = try XCTUnwrap(viewModel.recordings.first)
XCTAssertEqual(recording.transcript, "complete meeting transcript")
XCTAssertEqual(recording.calendarEvent, metadata)
Expand Down Expand Up @@ -901,6 +902,7 @@ final class AudioRecorderViewModelTests: XCTestCase {
XCTAssertEqual(request.language, "de")
XCTAssertTrue(request.translate)
XCTAssertTrue(request.prompt?.contains("TypeWhisper") == true)
XCTAssertTrue(request.usedFileTranscriptionPipeline)
XCTAssertTrue(plugin.selectedModelOverrides.contains("universal-3-5-pro"))
}

Expand Down Expand Up @@ -1764,13 +1766,14 @@ private final class AudioRecorderRestorableTranscriptionPlugin: NSObject, Transc
}
}

private final class AudioRecorderMockTranscriptionPlugin: NSObject, TranscriptionEnginePlugin, @unchecked Sendable {
private final class AudioRecorderMockTranscriptionPlugin: NSObject, SourceProgressTranscriptionEnginePlugin, @unchecked Sendable {
struct Request: Sendable {
let language: String?
let translate: Bool
let prompt: String?
let audioSampleCount: Int
let firstAudioSample: Float?
let usedFileTranscriptionPipeline: Bool
}

enum TranscriptionBehavior {
Expand Down Expand Up @@ -1830,12 +1833,52 @@ private final class AudioRecorderMockTranscriptionPlugin: NSObject, Transcriptio
translate: Bool,
prompt: String?
) async throws -> PluginTranscriptionResult {
try performTranscription(
audio: audio,
language: language,
translate: translate,
prompt: prompt,
usedFileTranscriptionPipeline: false
)
}

func transcribe(
audio: AudioData,
language: String?,
translate: Bool,
prompt: String?,
onProgress: @Sendable @escaping (String) -> Bool,
onSourceProgress: @Sendable @escaping (PluginTranscriptionSourceProgress) -> Bool
) async throws -> PluginTranscriptionResult {
let result = try performTranscription(
audio: audio,
language: language,
translate: translate,
prompt: prompt,
usedFileTranscriptionPipeline: true
)
_ = onProgress(result.text)
_ = onSourceProgress(PluginTranscriptionSourceProgress(
processedDuration: audio.duration,
totalDuration: audio.duration
))
return result
}

private func performTranscription(
audio: AudioData,
language: String?,
translate: Bool,
prompt: String?,
usedFileTranscriptionPipeline: Bool
) throws -> PluginTranscriptionResult {
lastRequest = Request(
language: language,
translate: translate,
prompt: prompt,
audioSampleCount: audio.samples.count,
firstAudioSample: audio.samples.first
firstAudioSample: audio.samples.first,
usedFileTranscriptionPipeline: usedFileTranscriptionPipeline
)
return switch behavior {
case .success(let text):
Expand Down