use std::path::PathBuf;
use serde_json::{json, Value};
pub mod parakeet;
pub trait SttHost: Send + Sync {
fn whisper_base_url(&self) -> String;
fn gateway_url(&self) -> String;
fn gateway_bearer(&self) -> Result<String, String>;
fn parakeet_model_dir(&self) -> PathBuf;
}
#[derive(Debug, Clone, Default, serde::Serialize)]
#[serde(rename_all = "camelCase")]
pub struct TranscriptSegment {
pub start_ms: u64,
pub end_ms: u64,
pub text: String,
}
#[derive(Debug, Clone, Default)]
pub struct Transcription {
pub text: String,
pub segments: Vec<TranscriptSegment>,
}
fn parse_verbose_segments(body: &Value) -> Vec<TranscriptSegment> {
body.get("segments")
.and_then(Value::as_array)
.map(|arr| {
arr.iter()
.filter_map(|s| {
let start = s.get("start").and_then(Value::as_f64)?;
let end = s.get("end").and_then(Value::as_f64)?;
let text = s
.get("text")
.and_then(Value::as_str)
.unwrap_or("")
.trim()
.to_string();
Some(TranscriptSegment {
start_ms: (start.max(0.0) * 1000.0) as u64,
end_ms: (end.max(0.0) * 1000.0) as u64,
text,
})
})
.collect()
})
.unwrap_or_default()
}
pub fn default_stt_engine() -> String {
if let Ok(env_engine) = std::env::var("RYU_STT_ENGINE") {
let trimmed = env_engine.trim();
if !trimmed.is_empty() {
return trimmed.to_string();
}
}
#[cfg(feature = "voice-parakeet")]
{
"parakeet".to_string()
}
#[cfg(not(feature = "voice-parakeet"))]
{
"whisper".to_string()
}
}
pub async fn transcribe_wav(
client: &reqwest::Client,
host: &dyn SttHost,
bytes: Vec<u8>,
filename: String,
engine: Option<&str>,
) -> Result<String, String> {
transcribe_wav_detailed(client, host, bytes, filename, engine)
.await
.map(|t| t.text)
}
pub async fn transcribe_wav_detailed(
client: &reqwest::Client,
host: &dyn SttHost,
bytes: Vec<u8>,
filename: String,
engine: Option<&str>,
) -> Result<Transcription, String> {
let engine = engine
.map(str::trim)
.filter(|s| !s.is_empty())
.map(str::to_string)
.unwrap_or_else(default_stt_engine);
if engine == "parakeet" {
return parakeet::transcribe(bytes, host.parakeet_model_dir())
.await
.map(|text| Transcription {
text,
segments: Vec::new(),
})
.map_err(|e| format!("parakeet transcription failed: {e:#}"));
}
if engine == "gateway" {
return transcribe_via_gateway(client, host, bytes).await;
}
let part = reqwest::multipart::Part::bytes(bytes).file_name(filename);
let form = reqwest::multipart::Form::new()
.part("file", part)
.text("response_format", "verbose_json");
let url = format!("{}/inference", host.whisper_base_url());
let resp = client
.post(&url)
.multipart(form)
.send()
.await
.map_err(|e| {
format!(
"whisper voice engine not reachable at {url}: {e}. \
Install + start `whispercpp` from the Store first."
)
})?;
if !resp.status().is_success() {
let status = resp.status();
let body = resp.text().await.unwrap_or_default();
return Err(format!("whisper returned {status}: {body}"));
}
let value: Value = resp
.json()
.await
.map_err(|e| format!("could not parse whisper response: {e}"))?;
let text = value
.get("text")
.and_then(Value::as_str)
.unwrap_or("")
.trim()
.to_string();
let segments = parse_verbose_segments(&value);
Ok(Transcription { text, segments })
}
async fn transcribe_via_gateway(
client: &reqwest::Client,
host: &dyn SttHost,
bytes: Vec<u8>,
) -> Result<Transcription, String> {
use base64::Engine as _;
let audio_b64 = base64::engine::general_purpose::STANDARD.encode(&bytes);
let provider = std::env::var("RYU_STT_GATEWAY_PROVIDER")
.ok()
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
.unwrap_or_else(|| "openai".to_string());
let model = std::env::var("RYU_STT_GATEWAY_MODEL")
.ok()
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
.unwrap_or_else(|| "whisper-large-v3".to_string());
let base = host.gateway_url();
let base = base.trim_end_matches('/');
let url = format!("{base}/v1/audio/transcriptions");
let bearer = host.gateway_bearer()?;
let payload = json!({
"model": model,
"file": audio_b64,
"response_format": "verbose_json",
});
let resp = client
.post(&url)
.bearer_auth(bearer)
.header("x-ryu-slot-stt-provider", &provider)
.header("x-ryu-slot-stt-model", &model)
.json(&payload)
.send()
.await
.map_err(|e| format!("gateway STT unreachable at {url}: {e}"))?;
if !resp.status().is_success() {
let status = resp.status();
let detail = resp.text().await.unwrap_or_default();
return Err(format!("gateway STT returned {status}: {detail}"));
}
let value: Value = resp
.json()
.await
.map_err(|e| format!("could not parse gateway STT response: {e}"))?;
let text = value
.get("text")
.and_then(Value::as_str)
.unwrap_or("")
.trim()
.to_string();
let segments = parse_verbose_segments(&value);
Ok(Transcription { text, segments })
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parses_verbose_segments_seconds_to_ms() {
let body = json!({
"text": "hello world",
"segments": [
{ "start": 0.0, "end": 1.5, "text": " hello" },
{ "start": 1.5, "end": 2.25, "text": " world " },
]
});
let segs = parse_verbose_segments(&body);
assert_eq!(segs.len(), 2);
assert_eq!(segs[0].start_ms, 0);
assert_eq!(segs[0].end_ms, 1500);
assert_eq!(segs[0].text, "hello");
assert_eq!(segs[1].start_ms, 1500);
assert_eq!(segs[1].end_ms, 2250);
assert_eq!(segs[1].text, "world");
}
#[test]
fn missing_or_malformed_segments_yield_empty() {
assert!(parse_verbose_segments(&json!({ "text": "x" })).is_empty());
assert!(parse_verbose_segments(&json!({ "segments": "not-an-array" })).is_empty());
let partial = json!({ "segments": [ { "text": "no timings" } ] });
assert!(parse_verbose_segments(&partial).is_empty());
}
#[test]
fn default_engine_env_override_wins() {
let prev = std::env::var("RYU_STT_ENGINE").ok();
std::env::set_var("RYU_STT_ENGINE", "gateway");
assert_eq!(default_stt_engine(), "gateway");
std::env::set_var("RYU_STT_ENGINE", " ");
let compiled = default_stt_engine();
assert!(compiled == "parakeet" || compiled == "whisper");
match prev {
Some(v) => std::env::set_var("RYU_STT_ENGINE", v),
None => std::env::remove_var("RYU_STT_ENGINE"),
}
}
#[test]
fn transcript_segment_serializes_camel_case() {
let seg = TranscriptSegment {
start_ms: 10,
end_ms: 20,
text: "hi".into(),
};
let v = serde_json::to_value(&seg).unwrap();
assert_eq!(v["startMs"], 10);
assert_eq!(v["endMs"], 20);
assert_eq!(v["text"], "hi");
}
}