feat(dictation): add model-native audio transcription provider (#10589)

This commit is contained in:
Joel Wirāmu Pauling
2026-08-19 21:37:10 +00:00
committed by GitHub
parent e31f2f4292
commit 6722a195e1
5 changed files with 402 additions and 5 deletions
+4 -1
View File
@@ -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"
+31 -1
View File
@@ -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<String> {
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<String> {
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::<String>("GOOSE_MODEL")
.ok(),
#[cfg(feature = "local-inference")]
DictationProvider::Local => Some(whisper::recommend_model().to_string()),
}
}
fn dictation_selected_model(config: &Config, provider: DictationProvider) -> Option<String> {
if provider == DictationProvider::ModelNative {
return crate::config::Config::global()
.get_param::<String>("GOOSE_MODEL")
.ok();
}
#[cfg(feature = "local-inference")]
if provider == DictationProvider::Local {
return config
@@ -456,6 +485,7 @@ fn dictation_available_models(provider: DictationProvider) -> Vec<DictationModel
label: "Scribe v1".to_string(),
description: "ElevenLabs' hosted speech-to-text model.".to_string(),
}],
DictationProvider::ModelNative => vec![],
#[cfg(feature = "local-inference")]
DictationProvider::Local => whisper::available_models()
.iter()
+365 -2
View File
@@ -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<HashMap<String, String>>,
}
#[cfg(feature = "local-inference")]
static LOCAL_TRANSCRIBER: once_cell::sync::Lazy<
Mutex<Option<(String, super::whisper::WhisperTranscriber)>>,
@@ -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::<String>(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<u8>, audio_format: &str) -> Result<String> {
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<ModelNativeResolved> {
// 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::<String>(&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::<String>(&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::<String>("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::<String>("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::<String>("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<String, String> = config
.get_secret::<String>("OPENAI_CUSTOM_HEADERS")
.ok()
.map(crate::providers::openai::parse_custom_headers)
.unwrap_or_default();
if let Ok(org) = config.get_param::<String>("OPENAI_ORGANIZATION") {
headers.insert("OpenAI-Organization".to_string(), org);
}
if let Ok(project) = config.get_param::<String>("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::<String>("OPENROUTER_API_KEY")
.unwrap_or_default();
let base_url = normalize_openrouter_base_url(
&config
.get_param::<String>("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::<String>("GROQ_API_KEY")
.unwrap_or_default();
let base_url = config
.get_param::<String>("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::<String>("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::<String>("GOOGLE_HOST").is_ok();
let api_key = if has_custom_host {
// Custom host may not require auth (e.g. local proxy)
config
.get_secret::<String>("GOOGLE_API_KEY")
.unwrap_or_default()
} else {
config.get_secret::<String>("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::<String>("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)
}
}
@@ -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);
};
+1 -1
View File
@@ -1 +1 @@
export type DictationProvider = 'openai' | 'elevenlabs' | 'groq' | 'local';
export type DictationProvider = 'openai' | 'elevenlabs' | 'groq' | 'local' | 'model';