From 1a205720ffc17de33f324d1245276b457ff76063 Mon Sep 17 00:00:00 2001 From: John Tennant Date: Wed, 26 Aug 2026 19:47:59 +0000 Subject: [PATCH 01/17] feat(voice): add OpenAI speech-to-text --- src-tauri/Cargo.toml | 2 + src-tauri/src/commands/native_voice.rs | 489 +++++++++++++++++- src-tauri/src/commands/openai_audio.rs | 30 +- src/app/AppShell.tsx | 4 +- src/features/chat/ui/ChatInputToolbar.tsx | 3 +- src/features/chat/ui/ChatView.tsx | 3 +- .../api/voiceConversation.ts | 2 +- .../lib/voiceInputPreference.ts | 7 +- .../lib/voiceSetupReadiness.test.ts | 4 +- .../lib/voiceSetupReadiness.ts | 16 +- .../ui/VoiceSettings.test.tsx | 10 +- .../voice-conversation/ui/VoiceSettings.tsx | 33 +- 12 files changed, 575 insertions(+), 28 deletions(-) diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 0e0abf2aa..ac898a9b9 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -95,6 +95,8 @@ tauri-plugin-updater = "2" tauri-plugin-window-state = "2" time = { version = "0.3", features = ["formatting"] } tokio = { version = "1.50.0", features = ["full"] } +tokio-tungstenite = { version = "0.27", features = ["rustls-tls-webpki-roots"] } +rustls = { version = "0.23", default-features = false, features = ["aws_lc_rs", "std", "tls12"] } url = "2" uuid = { version = "1", features = ["v4", "serde"] } zip = { version = "2", default-features = false, features = ["deflate"] } diff --git a/src-tauri/src/commands/native_voice.rs b/src-tauri/src/commands/native_voice.rs index d633a372e..3fbd28368 100644 --- a/src-tauri/src/commands/native_voice.rs +++ b/src-tauri/src/commands/native_voice.rs @@ -50,6 +50,7 @@ const STT_WORKER_SHUTDOWN_TIMEOUT_SECONDS: u64 = pub enum VoiceInputBackend { Parakeet, Macos, + Openai, } fn active_vad_threshold_for_speech( @@ -828,6 +829,61 @@ impl SttPipeline { )) } + fn new_openai( + input_muted: Arc, + input_mute_epoch: Arc, + assistant_speaking: Arc, + assistant_vad_threshold: Arc, + speech_vad_threshold: f32, + ) -> Result<(Self, tokio_mpsc::Receiver), String> { + if !super::openai_audio::openai_api_key_available()? { + return Err( + "OpenAI speech-to-text is not configured. Configure the OpenAI provider in Berd." + .to_string(), + ); + } + let (audio_tx, audio_rx) = mpsc::sync_channel(AUDIO_QUEUE_DEPTH); + let (event_tx, event_rx) = tokio_mpsc::channel(64); + let shutdown = Arc::new(AtomicBool::new(false)); + let discard_on_shutdown = Arc::new(AtomicBool::new(false)); + let shutdown_mute_epoch = Arc::new(AtomicU64::new(0)); + let worker_shutdown = Arc::clone(&shutdown); + let worker_discard_on_shutdown = Arc::clone(&discard_on_shutdown); + let worker_input_muted = Arc::clone(&input_muted); + let worker_input_mute_epoch = Arc::clone(&input_mute_epoch); + let worker_shutdown_mute_epoch = Arc::clone(&shutdown_mute_epoch); + let thread = thread::Builder::new() + .name("berd-openai-stt".into()) + .spawn(move || { + openai_stt_worker( + audio_rx, + event_tx, + worker_shutdown, + worker_discard_on_shutdown, + worker_input_muted, + worker_input_mute_epoch, + worker_shutdown_mute_epoch, + assistant_speaking, + assistant_vad_threshold, + speech_vad_threshold, + ) + }) + .map_err(|error| format!("start OpenAI speech recognition: {error}"))?; + Ok(( + Self { + audio_tx, + audio_seen: AtomicBool::new(false), + shutdown, + discard_on_shutdown, + input_muted, + input_mute_epoch, + shutdown_mute_epoch, + thread: Some(thread), + }, + event_rx, + )) + } + fn new_macos( input_muted: Arc, input_mute_epoch: Arc, @@ -1074,7 +1130,8 @@ where } async fn status(app: &AppHandle, state: &NativeVoiceState) -> NativeVoiceStatus { - let parakeet_available = parakeet_model_dir(app).is_ok(); + let parakeet_available = parakeet_model_dir(app).is_ok() + || super::openai_audio::openai_api_key_available().unwrap_or(false); #[cfg(target_os = "macos")] let macos_status = || async { mac_speech::status_async() @@ -1272,6 +1329,13 @@ pub async fn start_native_voice_conversation( Arc::clone(&state.assistant_vad_threshold), speech_vad_threshold, ), + VoiceInputBackend::Openai => SttPipeline::new_openai( + Arc::clone(&state.input_muted), + Arc::clone(&state.input_mute_epoch), + Arc::clone(&state.assistant_speaking), + Arc::clone(&state.assistant_vad_threshold), + speech_vad_threshold, + ), }; let (pipeline, mut events) = match pipeline { Ok(result) => result, @@ -2549,6 +2613,429 @@ fn macos_stt_worker( let _ = forward_macos_events(&mut recognition_events, &event_tx, Some(delivery_deadline)); } +#[derive(Debug, PartialEq, Eq)] +struct OpenAiCommittedTurn { + item_id: String, + mute_epoch: u64, +} + +fn record_openai_transcription_event( + value: &serde_json::Value, + current_mute_epoch: u64, + pending_commit_epochs: &mut VecDeque, + committed: &mut VecDeque, + completed: &mut HashMap, +) -> Result, String> { + let mut newly_committed = None; + match value.get("type").and_then(|value| value.as_str()) { + Some("input_audio_buffer.committed") => { + if let Some(item_id) = value.get("item_id").and_then(|value| value.as_str()) { + let turn = OpenAiCommittedTurn { + item_id: item_id.to_string(), + mute_epoch: pending_commit_epochs + .pop_front() + .unwrap_or(current_mute_epoch), + }; + if turn.mute_epoch == current_mute_epoch { + committed.push_back(OpenAiCommittedTurn { + item_id: turn.item_id.clone(), + mute_epoch: turn.mute_epoch, + }); + } + newly_committed = Some(turn); + } + } + Some("conversation.item.input_audio_transcription.completed") => { + if let (Some(item_id), Some(transcript)) = ( + value.get("item_id").and_then(|value| value.as_str()), + value.get("transcript").and_then(|value| value.as_str()), + ) { + if committed + .iter() + .any(|turn| turn.item_id == item_id && turn.mute_epoch == current_mute_epoch) + { + completed.insert(item_id.to_string(), transcript.trim().to_string()); + } + } + } + Some("conversation.item.input_audio_transcription.failed") | Some("error") => { + return Err(value + .pointer("/error/message") + .and_then(|value| value.as_str()) + .unwrap_or("OpenAI realtime transcription failed.") + .to_string()); + } + _ => {} + } + Ok(newly_committed) +} + +fn deliver_completed_openai_turns( + committed: &mut VecDeque, + completed: &mut HashMap, + event_tx: &tokio_mpsc::Sender, + final_item_id: Option<&str>, + final_delivery: &mut Option>, +) { + while committed + .front() + .is_some_and(|turn| completed.contains_key(&turn.item_id)) + { + let turn = committed.pop_front().expect("checked front"); + let text = completed.remove(&turn.item_id).unwrap_or_default(); + let delivered = (Some(turn.item_id.as_str()) == final_item_id) + .then(|| final_delivery.take()) + .flatten(); + deliver_recognition_result(text, event_tx, delivered); + } +} + +#[allow(clippy::too_many_arguments)] // Worker boundary keeps channel and mute lifecycle inputs explicit. +fn openai_stt_worker( + audio_rx: Receiver, + event_tx: tokio_mpsc::Sender, + shutdown: Arc, + discard_on_shutdown: Arc, + input_muted: Arc, + input_mute_epoch: Arc, + shutdown_mute_epoch: Arc, + assistant_speaking: Arc, + assistant_vad_threshold: Arc, + speech_vad_threshold: f32, +) { + use base64::{engine::general_purpose::STANDARD as BASE64, Engine}; + use futures_util::{SinkExt, StreamExt}; + use rubato::{Fft, FixedSync, Resampler}; + use tokio_tungstenite::tungstenite::{client::IntoClientRequest, Message}; + + let runtime = match tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + { + Ok(runtime) => runtime, + Err(error) => { + let _ = event_tx.blocking_send(SttMessage::Failed(format!( + "Could not initialize OpenAI realtime transcription: {error}" + ))); + return; + } + }; + let key = match super::openai_audio::api_key() { + Ok(key) => key, + Err(error) => { + let _ = event_tx.blocking_send(SttMessage::Failed(error)); + return; + } + }; + let model = super::openai_audio::transcription_model(); + let endpoint = match super::openai_audio::realtime_endpoint() { + Ok(endpoint) => endpoint, + Err(error) => { + let _ = event_tx.blocking_send(SttMessage::Failed(error)); + return; + } + }; + if let Err(existing) = rustls::crypto::aws_lc_rs::default_provider().install_default() { + // Another dependency may have installed the same process-wide provider first. + drop(existing); + } + let connection = runtime.block_on(async { + let mut request = endpoint + .into_client_request() + .map_err(|error| format!("prepare OpenAI realtime connection: {error}"))?; + request.headers_mut().insert( + "Authorization", + format!("Bearer {key}") + .parse() + .map_err(|_| "OpenAI API key is not a valid header value".to_string())?, + ); + tokio_tungstenite::connect_async(request) + .await + .map(|(socket, _)| socket) + .map_err(|error| format!("connect to OpenAI realtime transcription: {error}")) + }); + let mut socket = match connection { + Ok(socket) => socket, + Err(error) => { + let _ = event_tx.blocking_send(SttMessage::Failed(error)); + return; + } + }; + let session_update = serde_json::json!({ + "type": "session.update", + "session": { + "type": "transcription", + "audio": { "input": { + "format": { "type": "audio/pcm", "rate": 24000 }, + "transcription": { "model": model, "delay": "low" }, + "turn_detection": null + }} + } + }); + if let Err(error) = + runtime.block_on(socket.send(Message::Text(session_update.to_string().into()))) + { + let _ = event_tx.blocking_send(SttMessage::Failed(format!( + "configure OpenAI realtime transcription: {error}" + ))); + return; + } + + let mut resampler = match Fft::::new(48_000, 24_000, 960, 2, 1, FixedSync::Input) { + Ok(resampler) => resampler, + Err(error) => { + let _ = event_tx.blocking_send(SttMessage::Failed(format!( + "Could not initialize OpenAI audio resampling: {error}" + ))); + return; + } + }; + let chunk_in = resampler.input_frames_next(); + let mut vad = earshot::Detector::new(earshot::DefaultPredictor::new()); + let mut input_48k = Vec::new(); + let mut vad_16k = Vec::new(); + let mut silence_frames = 0_usize; + let mut in_speech = false; + let mut turn_has_audio = false; + let mut observed_mute_epoch = input_mute_epoch.load(Ordering::Acquire); + let mut pending_commit_epochs = VecDeque::::new(); + let mut committed_turns = VecDeque::::new(); + let mut completed_turns = HashMap::::new(); + + loop { + while let Some(event) = runtime.block_on(async { + tokio::time::timeout(Duration::from_millis(1), socket.next()) + .await + .ok() + .flatten() + }) { + let message = match event { + Ok(Message::Text(text)) => text, + Ok(Message::Close(_)) => { + let _ = event_tx.blocking_send(SttMessage::Failed( + "OpenAI realtime transcription disconnected.".to_string(), + )); + return; + } + Ok(_) => continue, + Err(error) => { + let _ = event_tx.blocking_send(SttMessage::Failed(format!( + "OpenAI realtime transcription failed: {error}" + ))); + return; + } + }; + let Ok(value) = serde_json::from_str::(&message) else { + continue; + }; + if let Err(error) = record_openai_transcription_event( + &value, + observed_mute_epoch, + &mut pending_commit_epochs, + &mut committed_turns, + &mut completed_turns, + ) { + let _ = event_tx.blocking_send(SttMessage::Failed(error)); + return; + } + deliver_completed_openai_turns( + &mut committed_turns, + &mut completed_turns, + &event_tx, + None, + &mut None, + ); + } + + let batch = match audio_rx.recv_timeout(Duration::from_millis(20)) { + Ok(batch) => Some(batch), + Err(mpsc::RecvTimeoutError::Timeout) => None, + Err(mpsc::RecvTimeoutError::Disconnected) => break, + }; + let (shutting_down, current_mute_epoch) = + sample_effective_mute_epoch(&input_mute_epoch, &shutdown, &shutdown_mute_epoch); + if current_mute_epoch != observed_mute_epoch { + observed_mute_epoch = current_mute_epoch; + input_48k.clear(); + vad_16k.clear(); + silence_frames = 0; + turn_has_audio = false; + committed_turns.clear(); + completed_turns.clear(); + if std::mem::take(&mut in_speech) { + let _ = event_tx.blocking_send(SttMessage::Speaking(false)); + } + let _ = runtime.block_on( + socket.send(Message::Text( + serde_json::json!({"type":"input_audio_buffer.clear"}) + .to_string() + .into(), + )), + ); + } + if shutting_down && (discard_on_shutdown.load(Ordering::Acquire) || batch.is_none()) { + break; + } + if !shutting_down && input_muted.load(Ordering::Acquire) { + continue; + } + let Some(batch) = batch else { continue }; + if batch.mute_epoch != observed_mute_epoch { + continue; + } + input_48k.extend( + batch + .bytes + .chunks_exact(4) + .map(|sample| f32::from_le_bytes([sample[0], sample[1], sample[2], sample[3]])), + ); + while input_48k.len() >= chunk_in { + let chunk: Vec = input_48k.drain(..chunk_in).collect(); + let pcm_24k = resample(&mut resampler, &chunk); + let pcm_bytes: Vec = pcm_24k + .iter() + .flat_map(|sample| { + ((sample.clamp(-1.0, 1.0) * i16::MAX as f32).round() as i16).to_le_bytes() + }) + .collect(); + let append = serde_json::json!({ + "type": "input_audio_buffer.append", + "audio": BASE64.encode(pcm_bytes) + }); + if let Err(error) = + runtime.block_on(socket.send(Message::Text(append.to_string().into()))) + { + let _ = event_tx.blocking_send(SttMessage::Failed(format!( + "stream audio to OpenAI realtime transcription: {error}" + ))); + return; + } + turn_has_audio = true; + // Earshot requires 16 kHz; use every third 48 kHz source sample for activity only. + vad_16k.extend(chunk.iter().step_by(3).copied()); + while vad_16k.len() >= VAD_FRAME_SAMPLES { + let frame: Vec = vad_16k.drain(..VAD_FRAME_SAMPLES).collect(); + let threshold = active_vad_threshold_for_speech( + &assistant_speaking, + &assistant_vad_threshold, + speech_vad_threshold, + ); + if vad.predict_f32(&frame) > threshold { + silence_frames = 0; + if !in_speech { + in_speech = true; + let _ = event_tx.blocking_send(SttMessage::Speaking(true)); + } + } else if in_speech { + silence_frames += 1; + if silence_frames >= SILENCE_FLUSH_FRAMES { + let commit = serde_json::json!({"type":"input_audio_buffer.commit"}); + pending_commit_epochs.push_back(observed_mute_epoch); + if let Err(error) = + runtime.block_on(socket.send(Message::Text(commit.to_string().into()))) + { + let _ = event_tx.blocking_send(SttMessage::Failed(format!( + "commit OpenAI transcription turn: {error}" + ))); + return; + } + turn_has_audio = false; + silence_frames = 0; + in_speech = false; + let _ = event_tx.blocking_send(SttMessage::Speaking(false)); + } + } + } + } + } + if discard_on_shutdown.load(Ordering::Acquire) { + let _ = runtime.block_on(socket.close(None)); + return; + } + let mut final_item_id = None::; + let mut final_delivery = None::>; + let mut final_delivered = None::>; + if turn_has_audio { + let (delivered_tx, delivered_rx) = mpsc::sync_channel(0); + final_delivery = Some(delivered_tx); + final_delivered = Some(delivered_rx); + pending_commit_epochs.push_back(observed_mute_epoch); + let commit = serde_json::json!({"type":"input_audio_buffer.commit"}); + if runtime + .block_on(socket.send(Message::Text(commit.to_string().into()))) + .is_err() + { + let _ = runtime.block_on(socket.close(None)); + return; + } + } + + // Keep receiving until the shutdown commit itself completes. Earlier turns + // are delivered from the queue first, and already committed transcripts are + // drained even when there was no partial turn to commit at shutdown. + let deadline = std::time::Instant::now() + FINAL_TRANSCRIPT_DELIVERY_TIMEOUT; + while std::time::Instant::now() < deadline { + if final_delivered + .as_ref() + .is_some_and(|receiver| receiver.try_recv().is_ok()) + || (final_delivered.is_none() + && committed_turns.is_empty() + && pending_commit_epochs.is_empty()) + { + break; + } + let Some(message) = runtime.block_on(async { + tokio::time::timeout(Duration::from_millis(50), socket.next()) + .await + .ok() + .flatten() + }) else { + continue; + }; + let Ok(Message::Text(text)) = message else { + continue; + }; + let Ok(value) = serde_json::from_str::(&text) else { + continue; + }; + let event_type = value.get("type").and_then(|value| value.as_str()); + match record_openai_transcription_event( + &value, + observed_mute_epoch, + &mut pending_commit_epochs, + &mut committed_turns, + &mut completed_turns, + ) { + Ok(Some(turn)) + if turn.mute_epoch == observed_mute_epoch + && pending_commit_epochs.is_empty() + && final_delivered.is_some() => + { + final_item_id = Some(turn.item_id); + } + Ok(_) => {} + Err(error) => { + let _ = event_tx.blocking_send(SttMessage::Failed(error)); + break; + } + } + deliver_completed_openai_turns( + &mut committed_turns, + &mut completed_turns, + &event_tx, + final_item_id.as_deref(), + &mut final_delivery, + ); + if event_type == Some("conversation.item.input_audio_transcription.completed") + && final_item_id.as_deref() == value.get("item_id").and_then(|value| value.as_str()) + && final_delivery.is_none() + { + break; + } + } + let _ = runtime.block_on(socket.close(None)); +} + #[allow(clippy::too_many_arguments)] // Worker boundary keeps channel and mute lifecycle inputs explicit. fn stt_worker( model_dir: PathBuf, diff --git a/src-tauri/src/commands/openai_audio.rs b/src-tauri/src/commands/openai_audio.rs index c69956006..3ab4713bb 100644 --- a/src-tauri/src/commands/openai_audio.rs +++ b/src-tauri/src/commands/openai_audio.rs @@ -1,4 +1,4 @@ -//! OpenAI streaming speech playback for voice conversations. +//! OpenAI realtime transcription configuration and streaming speech playback. use std::sync::{ atomic::{AtomicBool, Ordering}, @@ -36,6 +36,8 @@ use std::time::Instant; #[cfg(target_os = "macos")] const DEFAULT_BASE_URL: &str = "https://api.openai.com/v1"; +const DEFAULT_REALTIME_SESSION_MODEL: &str = "gpt-realtime-2.1"; +const DEFAULT_TRANSCRIPTION_MODEL: &str = "gpt-live-transcribe"; const DEFAULT_TTS_MODEL: &str = "gpt-4o-mini-tts"; const DEFAULT_TTS_VOICE: &str = "marin"; #[cfg(target_os = "macos")] @@ -95,6 +97,7 @@ enum OpenAiStreamCommand { #[serde(rename_all = "camelCase")] pub struct OpenAiVoiceStatus { configured: bool, + transcription_model: String, speech_model: String, speech_voice: String, playback_speed: f32, @@ -270,6 +273,30 @@ fn base_url() -> Result { Ok(DEFAULT_BASE_URL.to_string()) } +pub(crate) fn realtime_endpoint() -> Result { + let mut url = reqwest::Url::parse(&endpoint("realtime")?) + .map_err(|error| format!("OpenAI realtime endpoint is invalid: {error}"))?; + url.query_pairs_mut() + .append_pair("model", DEFAULT_REALTIME_SESSION_MODEL); + match url.scheme() { + "http" => url.set_scheme("ws").expect("compatible scheme"), + "https" => url.set_scheme("wss").expect("compatible scheme"), + "ws" | "wss" => {} + scheme => { + return Err(format!( + "OpenAI realtime endpoint has unsupported scheme: {scheme}" + )) + } + } + Ok(url.to_string()) +} + +pub(crate) fn transcription_model() -> String { + env_trimmed("OPENAI_TRANSCRIPTION_MODEL") + .or_else(|| env_trimmed("OPENAI_STT_MODEL")) + .unwrap_or_else(|| DEFAULT_TRANSCRIPTION_MODEL.to_string()) +} + fn speech_model() -> String { env_trimmed("OPENAI_TTS_MODEL").unwrap_or_else(|| DEFAULT_TTS_MODEL.to_string()) } @@ -365,6 +392,7 @@ pub fn get_openai_voice_status( let tts_available = cfg!(target_os = "macos"); Ok(OpenAiVoiceStatus { configured, + transcription_model: transcription_model(), speech_model: speech_model(), speech_voice: speech_voice(), playback_speed, diff --git a/src/app/AppShell.tsx b/src/app/AppShell.tsx index 15d6d2ba4..ae559725b 100644 --- a/src/app/AppShell.tsx +++ b/src/app/AppShell.tsx @@ -744,7 +744,9 @@ export function AppShell({ ); const globalVoiceOutput = useVoiceOutputPreference(); const globalOpenAiVoiceSetup = useOpenAiVoiceSetup( - capabilities.voiceConversation && globalVoiceOutput.backend === "openai", + capabilities.voiceConversation && + (globalVoiceInput.backend === "openai" || + globalVoiceOutput.backend === "openai"), ); const globalSiriVoiceSetup = useSiriVoiceSetup( capabilities.voiceConversation && globalVoiceOutput.backend === "siri", diff --git a/src/features/chat/ui/ChatInputToolbar.tsx b/src/features/chat/ui/ChatInputToolbar.tsx index 3050e94f7..3d2f8aeaa 100644 --- a/src/features/chat/ui/ChatInputToolbar.tsx +++ b/src/features/chat/ui/ChatInputToolbar.tsx @@ -561,7 +561,8 @@ export function ChatInputToolbar({ } tooltip={voiceConversationTooltip} > - {ownsActiveVoiceConversation ? ( + {ownsActiveVoiceConversation || + voiceConversationState === "error" ? (