From 6722a195e11e00e5ccebbb162d94ded4ca17f5e0 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Joel=20Wir=C4=81mu=20Pauling?= Date: Wed, 19 Aug 2026 21:37:10 +0000 Subject: [PATCH] feat(dictation): add model-native audio transcription provider (#10589) --- crates/goose/src/acp/server/config.rs | 5 +- crates/goose/src/acp/server/dictation.rs | 32 +- crates/goose/src/dictation/providers.rs | 367 +++++++++++++++++- .../settings/dictation/DictationSettings.tsx | 1 + ui/desktop/src/types/dictation.ts | 2 +- 5 files changed, 402 insertions(+), 5 deletions(-) diff --git a/crates/goose/src/acp/server/config.rs b/crates/goose/src/acp/server/config.rs index 22a0b58b4..d45e10ca4 100644 --- a/crates/goose/src/acp/server/config.rs +++ b/crates/goose/src/acp/server/config.rs @@ -395,7 +395,10 @@ fn prepare_voice_dictation_preferred_mic( } fn is_supported_voice_dictation_provider(value: &str) -> bool { - matches!(value, "openai" | "groq" | "elevenlabs" | "__disabled__") || { + matches!( + value, + "openai" | "groq" | "elevenlabs" | "model" | "__disabled__" + ) || { #[cfg(feature = "local-inference")] { value == "local" diff --git a/crates/goose/src/acp/server/dictation.rs b/crates/goose/src/acp/server/dictation.rs index 9444423d5..22d73fd50 100644 --- a/crates/goose/src/acp/server/dictation.rs +++ b/crates/goose/src/acp/server/dictation.rs @@ -2,7 +2,8 @@ use super::*; #[cfg(feature = "local-inference")] use crate::dictation::providers::transcribe_local; use crate::dictation::providers::{ - all_providers, get_provider_def, is_configured, transcribe_with_provider, DictationProvider, + all_providers, get_provider_def, is_configured, transcribe_with_model, + transcribe_with_provider, DictationProvider, }; #[cfg(feature = "local-inference")] use crate::dictation::whisper; @@ -61,6 +62,17 @@ impl GooseAcpAgent { let text = match provider { #[cfg(feature = "local-inference")] DictationProvider::Local => transcribe_local(audio_bytes).await, + DictationProvider::ModelNative => { + let audio_format = match extension { + "wav" => "wav", + "mp3" => "mp3", + "webm" => "webm", + "mp4" => "mp4", + "m4a" => "m4a", + _ => "wav", + }; + transcribe_with_model(audio_bytes, audio_format).await + } remote => { let (model_param, default_model) = dictation_transcribe_params(remote); let model = dictation_selected_model(config, remote) @@ -336,6 +348,7 @@ impl GooseAcpAgent { DictationProvider::OpenAI => OPENAI_TRANSCRIPTION_MODEL_CONFIG_KEY, DictationProvider::Groq => GROQ_TRANSCRIPTION_MODEL_CONFIG_KEY, DictationProvider::ElevenLabs => ELEVENLABS_TRANSCRIPTION_MODEL_CONFIG_KEY, + DictationProvider::ModelNative => return Ok(EmptyResponse {}), #[cfg(feature = "local-inference")] DictationProvider::Local => { let model = whisper::get_model(&req.model_id).ok_or_else(|| { @@ -368,6 +381,11 @@ fn parse_dictation_provider( fn dictation_secret_config_key( provider: DictationProvider, ) -> Result<&'static str, agent_client_protocol::Error> { + if provider == DictationProvider::ModelNative { + return Err(agent_client_protocol::Error::invalid_params() + .data("Model-native provider uses the active chat provider's credentials.")); + } + let def = get_provider_def(provider); if def.uses_provider_config { return Err(agent_client_protocol::Error::invalid_params().data( @@ -391,6 +409,7 @@ fn dictation_model_config_key(provider: DictationProvider) -> Option { DictationProvider::ElevenLabs => { Some(ELEVENLABS_TRANSCRIPTION_MODEL_CONFIG_KEY.to_string()) } + DictationProvider::ModelNative => None, #[cfg(feature = "local-inference")] DictationProvider::Local => Some(whisper::LOCAL_WHISPER_MODEL_CONFIG_KEY.to_string()), } @@ -403,6 +422,7 @@ fn dictation_transcribe_params(provider: DictationProvider) -> (&'static str, &' DictationProvider::OpenAI => ("model", OPENAI_TRANSCRIPTION_MODEL), DictationProvider::Groq => ("model", GROQ_TRANSCRIPTION_MODEL), DictationProvider::ElevenLabs => ("model_id", ELEVENLABS_TRANSCRIPTION_MODEL), + DictationProvider::ModelNative => ("", ""), #[cfg(feature = "local-inference")] DictationProvider::Local => ("", ""), } @@ -413,12 +433,21 @@ fn dictation_default_model(provider: DictationProvider) -> Option { DictationProvider::OpenAI => Some(OPENAI_TRANSCRIPTION_MODEL.to_string()), DictationProvider::Groq => Some(GROQ_TRANSCRIPTION_MODEL.to_string()), DictationProvider::ElevenLabs => Some(ELEVENLABS_TRANSCRIPTION_MODEL.to_string()), + DictationProvider::ModelNative => crate::config::Config::global() + .get_param::("GOOSE_MODEL") + .ok(), #[cfg(feature = "local-inference")] DictationProvider::Local => Some(whisper::recommend_model().to_string()), } } fn dictation_selected_model(config: &Config, provider: DictationProvider) -> Option { + if provider == DictationProvider::ModelNative { + return crate::config::Config::global() + .get_param::("GOOSE_MODEL") + .ok(); + } + #[cfg(feature = "local-inference")] if provider == DictationProvider::Local { return config @@ -456,6 +485,7 @@ fn dictation_available_models(provider: DictationProvider) -> Vec vec![], #[cfg(feature = "local-inference")] DictationProvider::Local => whisper::available_models() .iter() diff --git a/crates/goose/src/dictation/providers.rs b/crates/goose/src/dictation/providers.rs index 120423009..015c73707 100644 --- a/crates/goose/src/dictation/providers.rs +++ b/crates/goose/src/dictation/providers.rs @@ -5,7 +5,9 @@ use crate::dictation::whisper::LOCAL_WHISPER_MODEL_CONFIG_KEY; use crate::providers::api_client::{ApiClient, AuthMethod}; use crate::providers::openai::parse_openai_base_url; use anyhow::Result; +use base64::{engine::general_purpose::STANDARD as BASE64_STD, Engine as _}; use serde::{Deserialize, Serialize}; +use std::collections::HashMap; #[cfg(feature = "local-inference")] use std::sync::Mutex; use std::time::Duration; @@ -14,6 +16,12 @@ const REQUEST_TIMEOUT: Duration = Duration::from_secs(30); const OPENAI_VERSIONLESS_TRANSCRIPTIONS_PATH: &str = "audio/transcriptions"; type OpenAiDictationTarget = (String, Vec<(String, String)>, String); +struct ModelNativeResolved { + api_key: String, + base_url: String, + headers: Option>, +} + #[cfg(feature = "local-inference")] static LOCAL_TRANSCRIBER: once_cell::sync::Lazy< Mutex>, @@ -28,6 +36,8 @@ pub enum DictationProvider { OpenAI, ElevenLabs, Groq, + #[serde(rename = "model")] + ModelNative, #[cfg(feature = "local-inference")] Local, } @@ -88,16 +98,30 @@ pub const LOCAL_PROVIDER_DEF: DictationProviderDef = DictationProviderDef { settings_path: None, }; +pub const MODEL_NATIVE_PROVIDER_DEF: DictationProviderDef = DictationProviderDef { + provider: DictationProvider::ModelNative, + config_key: "", + default_base_url: "", + endpoint_path: "", + host_key: None, + description: "Uses your active chat model for transcription. Supports models with native audio input (e.g. Gemini, GPT-4o-audio, Gemma4). No separate API key needed.", + uses_provider_config: true, + settings_path: Some("Settings > Models"), +}; + /// Returns all provider definitions, including Local when the `local-inference` feature is enabled. pub fn all_providers() -> Vec<&'static DictationProviderDef> { #[cfg(not(feature = "local-inference"))] { - PROVIDERS.iter().collect() + let mut all: Vec<&DictationProviderDef> = PROVIDERS.iter().collect(); + all.push(&MODEL_NATIVE_PROVIDER_DEF); + all } #[cfg(feature = "local-inference")] { let mut all: Vec<&DictationProviderDef> = PROVIDERS.iter().collect(); all.push(&LOCAL_PROVIDER_DEF); + all.push(&MODEL_NATIVE_PROVIDER_DEF); all } } @@ -107,6 +131,9 @@ pub fn get_provider_def(provider: DictationProvider) -> &'static DictationProvid if provider == DictationProvider::Local { return &LOCAL_PROVIDER_DEF; } + if provider == DictationProvider::ModelNative { + return &MODEL_NATIVE_PROVIDER_DEF; + } PROVIDERS .iter() .find(|def| def.provider == provider) @@ -124,6 +151,15 @@ pub fn is_configured(provider: DictationProvider) -> bool { .and_then(|v| v.as_str().map(|s| s.to_string())) .and_then(|id| super::whisper::get_model(&id)) .is_some_and(|m| m.is_downloaded()), + DictationProvider::ModelNative => { + // Only configured if the active provider can be resolved for + // model-native dictation (OpenAI-compatible endpoint required). + if let Some(name) = crate::config::providers::get_active_provider(config) { + resolve_model_native_config(config, &name).is_ok() + } else { + false + } + } _ => { let def = get_provider_def(provider); config.get_secret::(def.config_key).is_ok() @@ -247,6 +283,9 @@ fn build_api_client(provider: DictationProvider) -> Result<(ApiClient, String)> header_name: "xi-api-key".to_string(), key: api_key, }, + DictationProvider::ModelNative => { + anyhow::bail!("ModelNative does not use the dictation API client") + } #[cfg(feature = "local-inference")] DictationProvider::Local => anyhow::bail!("Local provider should not use API client"), }; @@ -322,12 +361,300 @@ pub async fn transcribe_with_provider( Ok(text) } +const MODEL_TRANSCRIPTION_TIMEOUT: Duration = Duration::from_secs(60); +const TRANSCRIPTION_SYSTEM_PROMPT: &str = + "Transcribe the following audio exactly as spoken. Output only the transcription text, with no commentary, labels, formatting, or explanation."; + +pub async fn transcribe_with_model(audio_bytes: Vec, audio_format: &str) -> Result { + let config = Config::global(); + + let provider_name = crate::config::providers::get_active_provider(config) + .ok_or_else(|| anyhow::anyhow!("No active provider configured"))?; + + let model_name = crate::config::providers::get_active_model(config) + .ok_or_else(|| anyhow::anyhow!("No active model configured"))?; + + let resolved = resolve_model_native_config(config, &provider_name)?; + let api_key = &resolved.api_key; + let base_url = &resolved.base_url; + let audio_base64 = BASE64_STD.encode(&audio_bytes); + + let request_body = serde_json::json!({ + "model": model_name, + "messages": [{ + "role": "user", + "content": [ + { "type": "text", "text": TRANSCRIPTION_SYSTEM_PROMPT }, + { + "type": "input_audio", + "input_audio": { + "data": audio_base64, + "format": audio_format + } + } + ] + }] + }); + + let tls = provider_tls_config_from_config(config)?; + let auth = if api_key.is_empty() { + AuthMethod::NoAuth + } else { + AuthMethod::BearerToken(resolved.api_key.clone()) + }; + + let (host, query_params, has_v1) = parse_openai_base_url(base_url)?; + let endpoint = if has_v1 { + "v1/chat/completions" + } else { + "chat/completions" + }; + + let mut client = ApiClient::with_timeout_and_tls(host, auth, MODEL_TRANSCRIPTION_TIMEOUT, tls)?; + if !query_params.is_empty() { + client = client.with_query(query_params); + } + if let Some(ref custom_headers) = resolved.headers { + for (k, v) in custom_headers { + client = client.with_header(k, v)?; + } + } + + let response = client + .response_post(endpoint, &request_body) + .await + .map_err(|e| { + tracing::error!("Model-native transcription request failed: {}", e); + anyhow::anyhow!("Transcription request failed: {}", e) + })?; + + if !response.status().is_success() { + let status = response.status(); + let body = response.text().await.unwrap_or_default(); + if status.as_u16() == 401 { + anyhow::bail!("Invalid API key"); + } + anyhow::bail!("Chat completions error ({}): {}", status, body); + } + + let data: serde_json::Value = response.json().await.map_err(|e| { + tracing::error!("Failed to parse chat completions response: {}", e); + anyhow::anyhow!(e) + })?; + + let text = data["choices"][0]["message"]["content"] + .as_str() + .ok_or_else(|| anyhow::anyhow!("No content in chat completions response"))? + .to_string(); + + Ok(text) +} + +/// Normalize an OpenRouter base URL to ensure the `/api/v1` path is present. +/// +/// OpenRouter's chat completions endpoint lives at `api/v1/chat/completions`. +/// Users may configure just the host, the host with `/api`, or the full +/// `/api/v1` path. This function ensures the URL always ends with `/api/v1` +/// so that `parse_openai_base_url` detects the `/v1` segment correctly. +fn normalize_openrouter_base_url(base_url: &str) -> String { + if !base_url.contains("/api") { + format!("{}/api/v1", base_url.trim_end_matches('/')) + } else if base_url.ends_with("/api") { + format!("{}/v1", base_url) + } else { + base_url.to_string() + } +} + +fn resolve_model_native_config( + config: &Config, + provider_name: &str, +) -> Result { + // Try loading the declarative/custom provider config first — this handles + // custom_* providers whose base_url lives in a JSON file, not in env vars. + if let Ok(loaded) = crate::config::declarative_providers::load_provider(provider_name) { + let mut cfg = loaded.config; + // Only OpenAI-compatible engines support the input_audio content type. + // Anthropic uses a different API shape and does not accept input_audio. + use goose_providers::declarative::ProviderEngine; + match cfg.engine { + ProviderEngine::OpenAI | ProviderEngine::Ollama => {} + ProviderEngine::Anthropic => { + anyhow::bail!( + "Provider '{}' uses the Anthropic engine which does not support \ + the input_audio content type for model-native dictation", + provider_name + ) + } + } + // Resolve env var placeholders (e.g. ${LMSTUDIO_HOST}) in base_url + if let Some(ref env_vars) = cfg.env_vars { + cfg.base_url = + crate::config::declarative_providers::expand_env_vars(&cfg.base_url, env_vars)?; + } + let api_key = if cfg.api_key_env.is_empty() { + String::new() + } else if cfg.requires_auth { + config.get_secret::(&cfg.api_key_env).map_err(|_| { + anyhow::anyhow!( + "API key '{}' required for model-native dictation but not configured", + cfg.api_key_env + ) + })? + } else { + // Auth is optional (e.g. LM Studio, llama-swap) — use the key + // if configured, but tolerate it being absent. + config + .get_secret::(&cfg.api_key_env) + .unwrap_or_default() + }; + let headers = cfg.headers.clone(); + return Ok(ModelNativeResolved { + api_key, + base_url: cfg.base_url, + headers, + }); + } + + // Fallback: well-known providers resolved from env vars. + // + // OpenAI resolution mirrors openai_def.rs::resolve_base_url(): + // 1. OPENAI_HOST env var (session override, deprecated but honoured) + // 2. OPENAI_BASE_URL (env or config) - ecosystem-standard + // 3. OPENAI_HOST from config - persisted by goose configure + // 4. Default https://api.openai.com + match provider_name { + "openai" => { + let api_key = config + .get_secret::("OPENAI_API_KEY") + .unwrap_or_default(); + let base_url = if let Ok(h) = std::env::var("OPENAI_HOST") { + h + } else if let Ok(u) = config.get_param::("OPENAI_BASE_URL") { + let trimmed = u.trim().to_string(); + if trimmed.is_empty() { + "https://api.openai.com".to_string() + } else { + trimmed + } + } else { + config + .get_param::("OPENAI_HOST") + .unwrap_or_else(|_| "https://api.openai.com".to_string()) + }; + // Forward OPENAI_CUSTOM_HEADERS, OPENAI_ORGANIZATION, and + // OPENAI_PROJECT so that org/project-scoped and proxy setups + // work identically to normal chat (see openai_def.rs). + let mut headers: std::collections::HashMap = config + .get_secret::("OPENAI_CUSTOM_HEADERS") + .ok() + .map(crate::providers::openai::parse_custom_headers) + .unwrap_or_default(); + if let Ok(org) = config.get_param::("OPENAI_ORGANIZATION") { + headers.insert("OpenAI-Organization".to_string(), org); + } + if let Ok(project) = config.get_param::("OPENAI_PROJECT") { + headers.insert("OpenAI-Project".to_string(), project); + } + let headers = if headers.is_empty() { + None + } else { + Some(headers) + }; + Ok(ModelNativeResolved { + api_key, + base_url, + headers, + }) + } + "openrouter" => { + let api_key = config + .get_secret::("OPENROUTER_API_KEY") + .unwrap_or_default(); + let base_url = normalize_openrouter_base_url( + &config + .get_param::("OPENROUTER_HOST") + .unwrap_or_else(|_| "https://openrouter.ai/api/v1".to_string()), + ); + Ok(ModelNativeResolved { + api_key, + base_url, + headers: None, + }) + } + "groq" => { + let api_key = config + .get_secret::("GROQ_API_KEY") + .unwrap_or_default(); + let base_url = config + .get_param::("GROQ_HOST") + .unwrap_or_else(|_| "https://api.groq.com/openai".to_string()); + Ok(ModelNativeResolved { + api_key, + base_url, + headers: None, + }) + } + "ollama" => { + let base_url = config + .get_param::("OLLAMA_HOST") + .unwrap_or_else(|_| "http://localhost:11434".to_string()); + Ok(ModelNativeResolved { + api_key: String::new(), + base_url, + headers: None, + }) + } + "google" => { + // Google Gemini OpenAI-compatible endpoint lives at /v1beta/openai. + // The default includes this path so /chat/completions is appended + // correctly by the caller. + let has_custom_host = config.get_param::("GOOGLE_HOST").is_ok(); + let api_key = if has_custom_host { + // Custom host may not require auth (e.g. local proxy) + config + .get_secret::("GOOGLE_API_KEY") + .unwrap_or_default() + } else { + config.get_secret::("GOOGLE_API_KEY").map_err(|_| { + anyhow::anyhow!( + "GOOGLE_API_KEY required for model-native dictation \ + with the hosted Gemini endpoint" + ) + })? + }; + let base_url = config + .get_param::("GOOGLE_HOST") + .unwrap_or_else(|_| { + "https://generativelanguage.googleapis.com/v1beta/openai".to_string() + }); + Ok(ModelNativeResolved { + api_key, + base_url, + headers: None, + }) + } + other => { + // Providers that reach this branch have no declarative config + // (load_provider failed above) and are not in the known + // OpenAI-compatible set. Reject rather than sending an + // input_audio payload to an incompatible endpoint. + anyhow::bail!( + "Provider '{}' is not supported for model-native dictation. \ + Use a provider with an OpenAI-compatible chat completions endpoint.", + other + ) + } + } +} #[cfg(test)] mod tests { use super::{ - openai_dictation_target, resolve_openai_base_url_target, + all_providers, build_api_client, get_provider_def, normalize_openrouter_base_url, + openai_dictation_target, resolve_openai_base_url_target, DictationProvider, OPENAI_VERSIONLESS_TRANSCRIPTIONS_PATH, }; + use test_case::test_case; #[test] fn openai_dictation_target_preserves_prefix_and_query_params() { @@ -367,4 +694,40 @@ mod tests { .unwrap() .is_none()); } + + #[test] + fn model_native_serde_roundtrip() { + let json = r#""model""#; + let p: DictationProvider = serde_json::from_str(json).unwrap(); + assert_eq!(p, DictationProvider::ModelNative); + assert_eq!(serde_json::to_string(&p).unwrap(), r#""model""#); + } + + #[test] + fn model_native_provider_def_uses_provider_config() { + let def = get_provider_def(DictationProvider::ModelNative); + assert!(def.uses_provider_config); + assert!(def.config_key.is_empty()); + assert_eq!(def.provider, DictationProvider::ModelNative); + } + + #[test] + fn all_providers_includes_model_native() { + assert!(all_providers() + .iter() + .any(|d| d.provider == DictationProvider::ModelNative)); + } + + #[test] + fn build_api_client_rejects_model_native() { + assert!(build_api_client(DictationProvider::ModelNative).is_err()); + } + + #[test_case("https://openrouter.ai" => "https://openrouter.ai/api/v1" ; "bare host gets api v1")] + #[test_case("https://openrouter.ai/api" => "https://openrouter.ai/api/v1" ; "api without v1")] + #[test_case("https://openrouter.ai/api/v1" => "https://openrouter.ai/api/v1" ; "already correct")] + #[test_case("https://custom.proxy/api/v1" => "https://custom.proxy/api/v1" ; "custom proxy already correct")] + fn test_normalize_openrouter_base_url(input: &str) -> String { + normalize_openrouter_base_url(input) + } } diff --git a/ui/desktop/src/components/settings/dictation/DictationSettings.tsx b/ui/desktop/src/components/settings/dictation/DictationSettings.tsx index 9b317ef46..0507ca89a 100644 --- a/ui/desktop/src/components/settings/dictation/DictationSettings.tsx +++ b/ui/desktop/src/components/settings/dictation/DictationSettings.tsx @@ -177,6 +177,7 @@ export const DictationSettings = () => { const getProviderLabel = (p: DictationProvider | null): string => { if (!p) return intl.formatMessage(i18n.disabled); + if (p === "model") return "Model (Native Audio)"; return p.charAt(0).toUpperCase() + p.slice(1); }; diff --git a/ui/desktop/src/types/dictation.ts b/ui/desktop/src/types/dictation.ts index 35efe3fce..afbc28aa3 100644 --- a/ui/desktop/src/types/dictation.ts +++ b/ui/desktop/src/types/dictation.ts @@ -1 +1 @@ -export type DictationProvider = 'openai' | 'elevenlabs' | 'groq' | 'local'; +export type DictationProvider = 'openai' | 'elevenlabs' | 'groq' | 'local' | 'model';