use std::time::Duration;
use anyhow::{anyhow, Result};
#[derive(Debug, Clone)]
pub struct SttEndpoint {
pub name: String,
pub url: String,
pub headers: Vec<(String, String)>,
pub content_type: String,
pub timeout_ms: u64,
pub extract: TranscriptExtract,
}
#[derive(Debug, Clone)]
pub enum TranscriptExtract {
RawBody,
JsonPointer(String),
}
impl SttEndpoint {
pub fn local_whisper(url: impl Into<String>) -> Self {
Self {
name: "local_whisper".to_string(),
url: url.into(),
headers: Vec::new(),
content_type: "audio/mpeg".to_string(),
timeout_ms: 60_000,
extract: TranscriptExtract::RawBody,
}
}
pub fn openai_whisper(api_key: impl Into<String>) -> Self {
Self {
name: "openai_whisper".to_string(),
url: "https://api.openai.com/v1/audio/transcriptions".to_string(),
headers: vec![
("Authorization".to_string(), format!("Bearer {}", api_key.into())),
],
content_type: "audio/mpeg".to_string(),
timeout_ms: 30_000,
extract: TranscriptExtract::JsonPointer("/text".to_string()),
}
}
pub fn deepgram(api_key: impl Into<String>) -> Self {
Self {
name: "deepgram".to_string(),
url: "https://api.deepgram.com/v1/listen".to_string(),
headers: vec![(
"Authorization".to_string(),
format!("Token {}", api_key.into()),
)],
content_type: "audio/mpeg".to_string(),
timeout_ms: 30_000,
extract: TranscriptExtract::JsonPointer(
"/results/channels/0/alternatives/0/transcript".to_string(),
),
}
}
pub fn speechmatics(api_key: impl Into<String>) -> Self {
Self {
name: "speechmatics".to_string(),
url: "https://asr.api.speechmatics.com/v2/jobs".to_string(),
headers: vec![(
"Authorization".to_string(),
format!("Bearer {}", api_key.into()),
)],
content_type: "audio/mpeg".to_string(),
timeout_ms: 30_000,
extract: TranscriptExtract::JsonPointer(
"/results/0/alternatives/0/content".to_string(),
),
}
}
pub fn assemblyai(api_key: impl Into<String>) -> Self {
Self {
name: "assemblyai".to_string(),
url: "https://api.assemblyai.com/v2/transcript".to_string(),
headers: vec![("Authorization".to_string(), api_key.into())],
content_type: "audio/mpeg".to_string(),
timeout_ms: 30_000,
extract: TranscriptExtract::JsonPointer("/text".to_string()),
}
}
pub fn groq_whisper(api_key: impl Into<String>) -> Self {
Self {
name: "groq_whisper".to_string(),
url: "https://api.groq.com/openai/v1/audio/transcriptions".to_string(),
headers: vec![(
"Authorization".to_string(),
format!("Bearer {}", api_key.into()),
)],
content_type: "audio/mpeg".to_string(),
timeout_ms: 15_000,
extract: TranscriptExtract::JsonPointer("/text".to_string()),
}
}
}
#[derive(Debug, Clone, Default)]
pub struct SttPipeline {
pub endpoints: Vec<SttEndpoint>,
}
impl SttPipeline {
pub fn new() -> Self {
Self::default()
}
pub fn with(mut self, endpoint: SttEndpoint) -> Self {
self.endpoints.push(endpoint);
self
}
pub async fn transcribe(&self, audio_bytes: Vec<u8>) -> Result<String> {
if self.endpoints.is_empty() {
return Err(anyhow!("SttPipeline has no endpoints configured"));
}
let mut last_err: Option<String> = None;
for ep in &self.endpoints {
let client = match reqwest::Client::builder()
.timeout(Duration::from_millis(ep.timeout_ms))
.build()
{
Ok(c) => c,
Err(e) => {
last_err = Some(format!("{}: client build: {e}", ep.name));
continue;
}
};
let mut req = client
.post(&ep.url)
.header("Content-Type", ep.content_type.clone())
.body(audio_bytes.clone());
for (k, v) in &ep.headers {
req = req.header(k.as_str(), v.as_str());
}
let resp = match req.send().await {
Ok(r) => r,
Err(e) => {
last_err = Some(format!("{}: send: {e}", ep.name));
continue;
}
};
if !resp.status().is_success() {
last_err = Some(format!("{}: HTTP {}", ep.name, resp.status().as_u16()));
continue;
}
let body = match resp.text().await {
Ok(b) => b,
Err(e) => {
last_err = Some(format!("{}: read body: {e}", ep.name));
continue;
}
};
match extract_transcript(&body, &ep.extract) {
Ok(t) if !t.trim().is_empty() => return Ok(t),
Ok(_) => {
last_err = Some(format!("{}: empty transcript", ep.name));
}
Err(e) => {
last_err = Some(format!("{}: extract: {e}", ep.name));
}
}
}
Err(anyhow!(
"all {} STT endpoints failed; last error: {}",
self.endpoints.len(),
last_err.unwrap_or_else(|| "no error captured".to_string()),
))
}
}
pub fn extract_transcript(body: &str, extract: &TranscriptExtract) -> Result<String> {
match extract {
TranscriptExtract::RawBody => Ok(body.trim().to_string()),
TranscriptExtract::JsonPointer(ptr) => {
let parsed: serde_json::Value = serde_json::from_str(body)
.map_err(|e| anyhow!("response is not valid JSON: {e}"))?;
let v = parsed.pointer(ptr).ok_or_else(|| {
anyhow!("JSON pointer {ptr} did not resolve in response")
})?;
v.as_str()
.map(|s| s.trim().to_string())
.ok_or_else(|| anyhow!("pointer {ptr} did not resolve to a string"))
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn vendor_endpoints_have_distinct_urls() {
let key = "test-key";
let endpoints = [
SttEndpoint::openai_whisper(key),
SttEndpoint::deepgram(key),
SttEndpoint::speechmatics(key),
SttEndpoint::assemblyai(key),
SttEndpoint::groq_whisper(key),
];
let urls: std::collections::HashSet<&str> =
endpoints.iter().map(|e| e.url.as_str()).collect();
assert_eq!(
urls.len(),
endpoints.len(),
"vendor STT endpoints must have distinct URLs"
);
}
#[test]
fn vendor_endpoints_have_distinct_names() {
let key = "test-key";
let endpoints = [
SttEndpoint::openai_whisper(key),
SttEndpoint::deepgram(key),
SttEndpoint::speechmatics(key),
SttEndpoint::assemblyai(key),
SttEndpoint::groq_whisper(key),
];
let names: std::collections::HashSet<&str> =
endpoints.iter().map(|e| e.name.as_str()).collect();
assert_eq!(names.len(), endpoints.len());
}
#[test]
fn assemblyai_uses_bare_authorization_header() {
let ep = SttEndpoint::assemblyai("ak_xyz");
let auth = ep
.headers
.iter()
.find(|(k, _)| k == "Authorization")
.expect("assemblyai must set Authorization");
assert_eq!(auth.1, "ak_xyz", "assemblyai must NOT prefix with Bearer");
}
#[test]
fn deepgram_uses_token_prefix_authorization() {
let ep = SttEndpoint::deepgram("dg_key");
let auth = ep
.headers
.iter()
.find(|(k, _)| k == "Authorization")
.expect("deepgram must set Authorization");
assert_eq!(auth.1, "Token dg_key");
}
#[test]
fn raw_body_returns_trimmed_input() {
assert_eq!(
extract_transcript(" hello world\n", &TranscriptExtract::RawBody).unwrap(),
"hello world",
);
}
#[test]
fn json_pointer_extracts_string_at_pointer() {
let body = r#"{"text": "two three four"}"#;
let extract = TranscriptExtract::JsonPointer("/text".to_string());
assert_eq!(
extract_transcript(body, &extract).unwrap(),
"two three four",
);
}
#[test]
fn json_pointer_handles_nested_paths() {
let body = r#"{"results":[{"alternatives":[{"transcript":"one two three"}]}]}"#;
let extract = TranscriptExtract::JsonPointer(
"/results/0/alternatives/0/transcript".to_string(),
);
assert_eq!(
extract_transcript(body, &extract).unwrap(),
"one two three",
);
}
#[test]
fn json_pointer_errors_on_missing_path() {
let body = r#"{"text": "x"}"#;
let extract = TranscriptExtract::JsonPointer("/wrong".to_string());
assert!(extract_transcript(body, &extract).is_err());
}
#[test]
fn json_pointer_errors_on_non_string_value() {
let body = r#"{"text": 42}"#;
let extract = TranscriptExtract::JsonPointer("/text".to_string());
assert!(extract_transcript(body, &extract).is_err());
}
#[test]
fn json_pointer_errors_on_invalid_json() {
let body = "not json";
let extract = TranscriptExtract::JsonPointer("/text".to_string());
assert!(extract_transcript(body, &extract).is_err());
}
#[test]
fn local_whisper_constructor_defaults() {
let e = SttEndpoint::local_whisper("http://localhost:9000/asr");
assert_eq!(e.name, "local_whisper");
assert_eq!(e.url, "http://localhost:9000/asr");
assert_eq!(e.content_type, "audio/mpeg");
assert!(e.headers.is_empty());
assert!(matches!(e.extract, TranscriptExtract::RawBody));
}
#[test]
fn openai_whisper_constructor_sets_auth() {
let e = SttEndpoint::openai_whisper("sk-test-1234");
assert!(e
.headers
.iter()
.any(|(k, v)| k == "Authorization" && v == "Bearer sk-test-1234"));
assert_eq!(e.url, "https://api.openai.com/v1/audio/transcriptions");
if let TranscriptExtract::JsonPointer(p) = &e.extract {
assert_eq!(p, "/text");
} else {
panic!("expected JsonPointer extract");
}
}
#[test]
fn pipeline_builder_chains_endpoints_in_order() {
let p = SttPipeline::new()
.with(SttEndpoint::local_whisper("http://a"))
.with(SttEndpoint::openai_whisper("k"));
assert_eq!(p.endpoints.len(), 2);
assert_eq!(p.endpoints[0].name, "local_whisper");
assert_eq!(p.endpoints[1].name, "openai_whisper");
}
#[tokio::test]
async fn empty_pipeline_returns_error() {
let p = SttPipeline::new();
let err = p.transcribe(vec![1, 2, 3]).await.unwrap_err().to_string();
assert!(err.contains("no endpoints"));
}
#[tokio::test]
async fn pipeline_returns_aggregated_error_when_all_fail() {
let p = SttPipeline::new()
.with(SttEndpoint::local_whisper(
"http://127.0.0.1:1/unreachable-1",
))
.with(SttEndpoint::local_whisper(
"http://127.0.0.1:1/unreachable-2",
));
let err = p.transcribe(vec![1, 2, 3]).await.unwrap_err().to_string();
assert!(err.contains("all 2 STT endpoints failed"), "got: {err}");
}
}