use crate::config::{Config, Provider};
use crate::error::TalkError;
use async_trait::async_trait;
use std::collections::BTreeMap;
use std::path::{Path, PathBuf};
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, clap::ValueEnum)]
pub enum UploadFormat {
#[default]
Wav,
Ogg,
}
pub(crate) enum TranscriptionBody {
File(PathBuf),
#[cfg_attr(not(feature = "capture"), allow(dead_code))]
Pipe {
chunks: tokio::sync::mpsc::Receiver<Vec<u8>>,
file_name: String,
},
}
pub(crate) fn normalize_file_for_upload(
path: &std::path::Path,
) -> Result<(Vec<u8>, String), TalkError> {
if !path.exists() {
return Err(TalkError::Transcription(format!(
"Audio file not found: {}",
path.display()
)));
}
let stem = path.file_stem().and_then(|s| s.to_str()).unwrap_or("audio");
let normalized_name = format!("{stem}.ogg");
match encode_16k_mono_ogg(path) {
Ok(bytes) => {
log::info!(
"upload normalization: {} -> 16kHz mono ogg ({} bytes)",
path.display(),
bytes.len()
);
Ok((bytes, normalized_name))
}
Err(err) => {
log::warn!(
"upload normalization failed for {} ({}); uploading original bytes",
path.display(),
err
);
let bytes = std::fs::read(path)
.map_err(|e| TalkError::Transcription(format!("Failed to read audio file: {e}")))?;
let file_name = path
.file_name()
.and_then(|name| name.to_str())
.unwrap_or("audio.wav")
.to_string();
Ok((bytes, file_name))
}
}
}
fn encode_16k_mono_ogg(path: &std::path::Path) -> Result<Vec<u8>, TalkError> {
use crate::audio::{AudioWriter, OggOpusWriter};
let pcm = crate::record::audio::read_audio_as_i16(path)?;
let mut writer = OggOpusWriter::new(crate::config::AudioConfig::new())?;
let mut out = writer.header()?;
out.extend_from_slice(&writer.write_pcm(&pcm)?);
out.extend_from_slice(&writer.finalize()?);
Ok(out)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum RequestTimeoutPolicy {
#[default]
Proportional,
UserAttended,
}
#[derive(Debug, Clone, Default)]
pub struct TranscribeOptions {
pub allow_api: bool,
pub policy: RequestTimeoutPolicy,
pub cancel_token: Option<tokio_util::sync::CancellationToken>,
pub skip_legacy_lock: bool,
}
pub mod jobs;
pub mod mistral;
pub mod model_suggestions;
pub mod openai;
pub mod openai_realtime;
#[cfg(feature = "parakeet")]
pub mod parakeet;
pub mod realtime;
pub mod transport;
pub use mistral::MistralOneShotTranscriber;
pub use openai::OpenAIOneShotTranscriber;
pub use openai_realtime::OpenAIRealtimeTranscriber;
#[cfg(feature = "parakeet")]
pub use parakeet::ParakeetOneShotTranscriber;
pub use realtime::{MistralRealtimeTranscriber, OrderedItemTranscript, TranscriptionEvent};
#[derive(Debug, Clone, Default)]
pub struct TranscriptionResult {
pub text: String,
pub metadata: TranscriptionMetadata,
pub diarization: Option<Vec<DiarizationSegment>>,
pub segments: Option<Vec<TranscriptSegment>>,
}
#[derive(Debug, Clone, PartialEq)]
pub struct DiarizationSegment {
pub speaker: String,
pub start: f64,
pub end: f64,
pub text: String,
}
#[derive(Debug, Clone, PartialEq)]
pub struct TranscriptSegment {
pub start: f64,
pub end: f64,
pub text: String,
}
pub(crate) fn parse_transcript_segments(
segments: &[serde_json::Value],
) -> Option<Vec<TranscriptSegment>> {
let mut result = Vec::new();
for seg in segments {
let start = seg.get("start").and_then(|v| v.as_f64());
let end = seg.get("end").and_then(|v| v.as_f64());
let text = seg.get("text").and_then(|v| v.as_str()).unwrap_or("");
if let (Some(start), Some(end)) = (start, end) {
if !text.is_empty() {
result.push(TranscriptSegment {
start,
end,
text: text.to_string(),
});
}
}
}
if result.is_empty() {
None
} else {
Some(result)
}
}
pub fn format_transcription_output(result: &TranscriptionResult, timestamp: bool) -> String {
let Some(ref segments) = result.diarization else {
return result.text.clone();
};
if segments.is_empty() {
return result.text.clone();
}
let mut lines = Vec::new();
let mut current_speaker: Option<&str> = None;
let mut current_texts: Vec<&str> = Vec::new();
let mut current_start: f64 = 0.0;
for seg in segments {
if current_speaker == Some(seg.speaker.as_str()) {
current_texts.push(seg.text.trim());
} else {
if let Some(speaker) = current_speaker {
if timestamp {
lines.push(format!(
"[{}] {} {}",
format_timestamp(current_start),
speaker,
current_texts.join(" ")
));
} else {
lines.push(format!("[{}] {}", speaker, current_texts.join(" ")));
}
}
current_speaker = Some(&seg.speaker);
current_start = seg.start;
current_texts.clear();
current_texts.push(seg.text.trim());
}
}
if let Some(speaker) = current_speaker {
if timestamp {
lines.push(format!(
"[{}] {} {}",
format_timestamp(current_start),
speaker,
current_texts.join(" ")
));
} else {
lines.push(format!("[{}] {}", speaker, current_texts.join(" ")));
}
}
lines.join("\n")
}
fn format_timestamp(seconds: f64) -> String {
let total_secs = seconds as u64;
let hours = total_secs / 3600;
let minutes = (total_secs % 3600) / 60;
let secs = total_secs % 60;
format!("{:02}:{:02}:{:02}", hours, minutes, secs)
}
#[derive(Debug, Clone, Default)]
pub struct TranscriptionMetadata {
pub request_latency_ms: Option<u64>,
pub session_elapsed_ms: Option<u64>,
pub request_id: Option<String>,
pub provider_processing_ms: Option<u64>,
pub detected_language: Option<String>,
pub audio_seconds: Option<f64>,
pub segment_count: Option<usize>,
pub word_count: Option<usize>,
pub token_usage: Option<TokenUsage>,
pub provider_specific: Option<ProviderSpecificMetadata>,
}
#[derive(Debug, Clone, Default)]
pub struct TokenUsage {
pub input_tokens: Option<u64>,
pub output_tokens: Option<u64>,
pub total_tokens: Option<u64>,
}
#[derive(Debug, Clone)]
pub enum ProviderSpecificMetadata {
OpenAI(OpenAIProviderMetadata),
Mistral(MistralProviderMetadata),
}
#[derive(Debug, Clone, Default)]
pub struct OpenAIProviderMetadata {
pub model: Option<String>,
pub usage_raw: Option<serde_json::Value>,
pub rate_limit_headers: BTreeMap<String, String>,
pub unknown_event_types: Vec<String>,
pub realtime: Option<OpenAIRealtimeMetadata>,
}
#[derive(Debug, Clone, Default)]
pub struct OpenAIRealtimeMetadata {
pub session_id: Option<String>,
pub conversation_id: Option<String>,
pub event_counts: BTreeMap<String, u64>,
pub last_rate_limits: Option<serde_json::Value>,
pub ws_upgrade_headers: BTreeMap<String, String>,
}
#[derive(Debug, Clone, Default)]
pub struct MistralProviderMetadata {
pub model: Option<String>,
pub usage_raw: Option<serde_json::Value>,
pub unknown_event_types: Vec<String>,
}
#[async_trait]
pub(crate) trait OneShotTranscriber: Send + Sync {
async fn validate(&self) -> Result<(), TalkError>;
async fn fetch_transcription(
&self,
body: TranscriptionBody,
) -> Result<TranscriptionResult, TalkError>;
fn set_sink(&mut self, _sink: std::sync::Arc<dyn crate::telemetry::TelemetrySink>) {}
fn set_cancel_token(&mut self, _token: tokio_util::sync::CancellationToken) {}
}
#[cfg_attr(not(all(feature = "capture", feature = "ui")), allow(dead_code))]
#[async_trait]
pub(crate) trait RealtimeTranscriber: Send + Sync {
async fn validate(&self) -> Result<(), TalkError>;
async fn transcribe_realtime(
&self,
audio_rx: tokio::sync::mpsc::Receiver<Vec<i16>>,
) -> Result<tokio::sync::mpsc::Receiver<TranscriptionEvent>, TalkError>;
fn set_sink(&mut self, _sink: std::sync::Arc<dyn crate::telemetry::TelemetrySink>) {}
fn set_cancel_token(&mut self, _token: tokio_util::sync::CancellationToken) {}
}
pub fn is_model_error(provider: Provider, error: &TalkError) -> bool {
use crate::error::PipelineFailureKind;
if let TalkError::Pipeline(pf) = error {
if matches!(pf.kind, PipelineFailureKind::ModelRejected { .. }) {
return true;
}
}
match provider {
Provider::Mistral => mistral::is_model_error(error),
Provider::OpenAI => openai::is_model_error(error),
Provider::Parakeet => false,
}
}
pub async fn enrich_model_error(
config: &Config,
provider: Provider,
model: Option<&str>,
error: TalkError,
) -> TalkError {
match provider {
Provider::Mistral => {
let Some(ref cfg) = config.providers.mistral else {
return error;
};
let model_name = model.unwrap_or(&cfg.model);
let api_base = cfg.url.as_deref().unwrap_or(mistral::API_BASE);
mistral::enrich_model_error(error, &cfg.api_key, model_name, api_base).await
}
Provider::OpenAI => {
let Some(ref cfg) = config.providers.openai else {
return error;
};
let model_name = model.unwrap_or(&cfg.model);
let api_base = cfg.url.as_deref().unwrap_or(openai::API_BASE);
openai::enrich_model_error(error, &cfg.api_key, model_name, api_base).await
}
Provider::Parakeet => error,
}
}
pub(crate) fn create_oneshot_transcriber(
config: &Config,
provider: Provider,
model: Option<&str>,
diarize: bool,
policy: RequestTimeoutPolicy,
) -> Result<Box<dyn OneShotTranscriber>, TalkError> {
match provider {
Provider::Mistral => {
let mut cfg = config.providers.mistral.clone().ok_or_else(|| {
TalkError::Config(
"Mistral provider selected but providers.mistral is not configured".to_string(),
)
})?;
if cfg.api_key.is_empty() {
return Err(TalkError::Config(
"providers.mistral.api_key is required".to_string(),
));
}
if let Some(m) = model {
cfg.model = m.to_string();
}
Ok(Box::new(MistralOneShotTranscriber::with_policy(
cfg, diarize, policy,
)?))
}
Provider::OpenAI => {
let cfg = config.providers.openai.as_ref().ok_or_else(|| {
TalkError::Config(
"OpenAI provider selected but providers.openai is not configured".to_string(),
)
})?;
if cfg.api_key.is_empty() {
return Err(TalkError::Config(
"providers.openai.api_key is required".to_string(),
));
}
let cfg = override_openai_batch_model(cfg, model);
Ok(Box::new(OpenAIOneShotTranscriber::with_policy(
cfg, policy,
)?))
}
#[cfg(feature = "parakeet")]
Provider::Parakeet => {
let mut cfg = config.providers.parakeet.clone().unwrap_or_default();
if let Some(m) = model {
cfg.model = Some(m.to_string());
}
Ok(Box::new(ParakeetOneShotTranscriber::with_policy(
cfg, policy,
)?))
}
#[cfg(not(feature = "parakeet"))]
Provider::Parakeet => Err(TalkError::Config(
"talk-rs was built without the 'parakeet' feature; rebuild without \
--no-default-features to enable the local Parakeet backend"
.to_string(),
)),
}
}
fn override_openai_batch_model(
config: &crate::config::OpenAIConfig,
model: Option<&str>,
) -> crate::config::OpenAIConfig {
let mut config = config.clone();
if let Some(model) = model {
config.model = model.to_string();
}
config
}
pub fn read_cached_transcript(audio_path: &std::path::Path, config: &Config) -> Option<String> {
use crate::recording_cache::{get_transcript, TranscriptStatus, TranscriptionCache};
match get_transcript(audio_path) {
TranscriptStatus::Available(text) => return Some(text),
TranscriptStatus::InProgress | TranscriptStatus::NotAvailable => {}
}
let provider = config
.transcription
.as_ref()
.map(|t| t.default_provider)
.unwrap_or(Provider::Mistral);
let effective_model = resolve_effective_model(config, provider, None);
TranscriptionCache::get(audio_path, provider, &effective_model).map(|r| r.text)
}
pub async fn produce_transcript(
audio_path: &std::path::Path,
config: &Config,
provider: Provider,
model: Option<&str>,
sink: &std::sync::Arc<dyn crate::telemetry::TelemetrySink>,
) -> Result<String, TalkError> {
use crate::recording_cache::{self, TranscriptStatus};
match recording_cache::get_transcript(audio_path) {
TranscriptStatus::Available(text) => return Ok(text),
TranscriptStatus::InProgress => return Err(TalkError::TranscriptInProgress),
TranscriptStatus::NotAvailable => {}
}
recording_cache::acquire_pick_lock(audio_path)?;
let effective_model = resolve_effective_model(config, provider, model);
let result = transcribe_audio(
audio_path,
config,
provider,
model,
false,
TranscribeOptions {
allow_api: true,
policy: RequestTimeoutPolicy::Proportional,
cancel_token: None,
skip_legacy_lock: false,
},
sink,
)
.await;
let final_result = match result {
Ok(r) => {
let text = r.text.trim().to_string();
if let Err(e) = recording_cache::write_pick(
audio_path,
&provider.to_string(),
&effective_model,
false,
&text,
) {
log::warn!("failed to write pick file: {}", e);
}
Ok(text)
}
Err(e) => Err(e),
};
if let Err(e) = recording_cache::release_pick_lock(audio_path) {
log::warn!("failed to release pick lock: {}", e);
}
final_result
}
pub async fn transcribe_audio(
audio_path: &Path,
config: &Config,
provider: Provider,
model: Option<&str>,
diarize: bool,
options: TranscribeOptions,
sink: &std::sync::Arc<dyn crate::telemetry::TelemetrySink>,
) -> Result<TranscriptionResult, TalkError> {
let TranscribeOptions {
allow_api,
policy,
cancel_token,
skip_legacy_lock,
} = options;
use crate::recording_cache::{self, TranscriptionCache};
let effective_model = resolve_effective_model(config, provider, model);
if !diarize {
if let Some(cached) = TranscriptionCache::get(audio_path, provider, &effective_model) {
log::info!(
"transcription cache hit for {}:{} on {}",
provider,
effective_model,
audio_path.display()
);
return Ok(cached);
}
}
if !allow_api {
log::debug!(
"transcription cache miss for {}:{} on {} — API call forbidden",
provider,
effective_model,
audio_path.display()
);
return Err(TalkError::CacheOnly);
}
if !skip_legacy_lock {
recording_cache::acquire_model_lock(audio_path, provider, &effective_model, false)?;
}
log::info!(
"transcription cache miss for {}:{} on {} — calling API",
provider,
effective_model,
audio_path.display()
);
let api_result = async {
let mut transcriber = create_oneshot_transcriber(config, provider, model, diarize, policy)?;
transcriber.set_sink(sink.clone());
if let Some(token) = cancel_token {
transcriber.set_cancel_token(token);
}
transcriber.validate().await?;
transcriber
.fetch_transcription(TranscriptionBody::File(audio_path.to_path_buf()))
.await
}
.await;
let release_lock = |context: &str| {
if skip_legacy_lock {
return;
}
if let Err(e) =
recording_cache::release_model_lock(audio_path, provider, &effective_model, false)
{
log::warn!("failed to release model lock {}: {}", context, e);
}
};
match api_result {
Ok(result) => {
if let Err(e) =
TranscriptionCache::store(audio_path, provider, &effective_model, false, &result)
{
log::warn!("failed to cache transcription result: {}", e);
}
release_lock("after success");
Ok(result)
}
Err(e) => {
release_lock("after error");
Err(e)
}
}
}
fn resolve_effective_model(config: &Config, provider: Provider, model: Option<&str>) -> String {
if let Some(m) = model {
return m.to_string();
}
match provider {
Provider::Mistral => config
.providers
.mistral
.as_ref()
.map(|c| c.model.clone())
.unwrap_or_else(|| "voxtral-mini-latest".to_string()),
Provider::OpenAI => config
.providers
.openai
.as_ref()
.map(|c| c.model.clone())
.unwrap_or_else(|| "gpt-transcribe".to_string()),
Provider::Parakeet => config
.providers
.parakeet
.as_ref()
.map(|c| c.resolved_model_name())
.unwrap_or_else(|| "parakeet-tdt-0.6b-v3-int8".to_string()),
}
}
#[cfg_attr(not(all(feature = "capture", feature = "ui")), allow(dead_code))]
pub(crate) fn create_realtime_transcriber(
config: &Config,
provider: Provider,
model: Option<&str>,
) -> Result<Box<dyn RealtimeTranscriber>, TalkError> {
match provider {
Provider::Mistral => {
let cfg = config.providers.mistral.clone().ok_or_else(|| {
TalkError::Config(
"Mistral provider selected but providers.mistral is not configured".to_string(),
)
})?;
if cfg.api_key.is_empty() {
return Err(TalkError::Config(
"providers.mistral.api_key is required".to_string(),
));
}
Ok(Box::new(MistralRealtimeTranscriber::new(cfg)))
}
Provider::OpenAI => {
let mut cfg = config.providers.openai.clone().ok_or_else(|| {
TalkError::Config(
"OpenAI provider selected but providers.openai is not configured".to_string(),
)
})?;
if cfg.api_key.is_empty() {
return Err(TalkError::Config(
"providers.openai.api_key is required".to_string(),
));
}
if let Some(m) = model {
cfg.realtime_model = m.to_string();
}
Ok(Box::new(OpenAIRealtimeTranscriber::new(cfg)))
}
Provider::Parakeet => Err(TalkError::Config(
"parakeet provider has no realtime mode; use one-shot transcription (omit --realtime)"
.to_string(),
)),
}
}
pub struct MockOneShotTranscriber {
pub response_text: String,
}
impl MockOneShotTranscriber {
pub fn new(response_text: impl Into<String>) -> Self {
Self {
response_text: response_text.into(),
}
}
}
#[async_trait]
impl OneShotTranscriber for MockOneShotTranscriber {
async fn validate(&self) -> Result<(), TalkError> {
Ok(())
}
async fn fetch_transcription(
&self,
body: TranscriptionBody,
) -> Result<TranscriptionResult, TalkError> {
if let TranscriptionBody::Pipe { mut chunks, .. } = body {
while chunks.recv().await.is_some() {}
}
Ok(TranscriptionResult {
text: self.response_text.clone(),
metadata: TranscriptionMetadata::default(),
diarization: None,
segments: None,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
fn fixture(name: &str) -> PathBuf {
let mut p = PathBuf::from(env!("CARGO_MANIFEST_DIR"));
p.push("tests");
p.push("fixtures");
p.push(name);
p
}
fn opus_head_channels(ogg: &[u8]) -> u8 {
let pos = ogg
.windows(8)
.position(|w| w == b"OpusHead")
.expect("OpusHead present");
ogg[pos + 9]
}
#[test]
fn test_normalize_stereo_44k_m4a_to_mono_ogg() {
let path = fixture("sine_440_0.5s_stereo.m4a");
assert!(path.exists(), "fixture missing: {}", path.display());
let (bytes, name) = normalize_file_for_upload(&path).expect("normalization should succeed");
assert!(name.ends_with(".ogg"), "expected .ogg name, got {name}");
assert_eq!(&bytes[0..4], b"OggS", "should be an OGG stream");
assert_eq!(opus_head_channels(&bytes), 1, "must be downmixed to mono");
let mut decoder = opus::Decoder::new(16_000, opus::Channels::Mono)
.expect("decoder creation should succeed");
let mut packet_reader =
ogg::reading::PacketReader::new(std::io::Cursor::new(bytes.clone()));
let _ = packet_reader.read_packet_expected().expect("OpusHead");
let _ = packet_reader.read_packet_expected().expect("OpusTags");
let first_audio = packet_reader
.read_packet_expected()
.expect("at least one audio packet");
let mut out = vec![0i16; 16_000]; let decoded = decoder
.decode(&first_audio.data, &mut out, false)
.expect("decode should succeed");
assert!(decoded > 0, "decoded frame should be non-empty");
}
#[test]
fn test_normalize_mono_m4a_stays_mono_ogg() {
let path = fixture("sine_440_0.5s_mono.m4a");
assert!(path.exists(), "fixture missing: {}", path.display());
let (bytes, name) = normalize_file_for_upload(&path).expect("normalization should succeed");
assert!(name.ends_with(".ogg"));
assert_eq!(&bytes[0..4], b"OggS");
assert_eq!(opus_head_channels(&bytes), 1);
}
#[test]
fn test_normalize_missing_file_errors() {
let err = normalize_file_for_upload(&PathBuf::from("/nonexistent/audio.m4a"))
.expect_err("missing file should error");
assert!(
err.to_string().contains("not found"),
"expected not-found error, got: {err}"
);
}
#[test]
fn test_normalize_undecodable_falls_back_to_raw_bytes() {
let dir = tempfile::TempDir::new().expect("tmp dir");
let path = dir.path().join("weird.bin");
std::fs::write(&path, b"not really audio").expect("write");
let (bytes, name) = normalize_file_for_upload(&path).expect("fallback should succeed");
assert_eq!(bytes, b"not really audio");
assert_eq!(name, "weird.bin", "fallback keeps original file name");
}
#[tokio::test]
async fn test_mock_transcriber_returns_text() {
let mock = MockOneShotTranscriber::new("Hello, world!");
let result = mock
.fetch_transcription(TranscriptionBody::File(PathBuf::from("/tmp/test.wav")))
.await;
assert!(result.is_ok());
assert_eq!(result.unwrap().text, "Hello, world!");
}
#[tokio::test]
async fn test_mock_transcriber_stream() {
let mock = MockOneShotTranscriber::new("Streamed transcription");
let (tx, rx) = tokio::sync::mpsc::channel(4);
tx.send(vec![0u8; 100]).await.unwrap();
tx.send(vec![1u8; 200]).await.unwrap();
drop(tx);
let result = mock
.fetch_transcription(TranscriptionBody::Pipe {
chunks: rx,
file_name: "test.wav".to_string(),
})
.await;
assert!(result.is_ok());
assert_eq!(result.unwrap().text, "Streamed transcription");
}
#[tokio::test]
async fn test_mock_transcriber_ignores_path() {
let mock = MockOneShotTranscriber::new("Fixed response");
let result1 = mock
.fetch_transcription(TranscriptionBody::File(PathBuf::from("/path/one.wav")))
.await;
let result2 = mock
.fetch_transcription(TranscriptionBody::File(PathBuf::from("/path/two.wav")))
.await;
assert_eq!(result1.unwrap().text, "Fixed response");
assert_eq!(result2.unwrap().text, "Fixed response");
}
#[test]
fn test_provider_from_str() {
assert_eq!("mistral".parse::<Provider>().unwrap(), Provider::Mistral);
assert_eq!("openai".parse::<Provider>().unwrap(), Provider::OpenAI);
assert_eq!("OpenAI".parse::<Provider>().unwrap(), Provider::OpenAI);
assert!("unknown".parse::<Provider>().is_err());
}
#[test]
fn test_provider_display() {
assert_eq!(Provider::Mistral.to_string(), "mistral");
assert_eq!(Provider::OpenAI.to_string(), "openai");
}
#[test]
fn test_format_plain_text_without_diarization() {
let result = TranscriptionResult {
text: "Hello world.".to_string(),
metadata: Default::default(),
diarization: None,
segments: None,
};
assert_eq!(format_transcription_output(&result, false), "Hello world.");
}
#[test]
fn test_format_diarized_output() {
let result = TranscriptionResult {
text: "Hello. I am fine.".to_string(),
metadata: Default::default(),
diarization: Some(vec![
DiarizationSegment {
speaker: "SPEAKER_00".to_string(),
start: 0.0,
end: 1.5,
text: "Hello.".to_string(),
},
DiarizationSegment {
speaker: "SPEAKER_01".to_string(),
start: 1.5,
end: 3.0,
text: "I am fine.".to_string(),
},
]),
segments: None,
};
assert_eq!(
format_transcription_output(&result, false),
"[SPEAKER_00] Hello.\n[SPEAKER_01] I am fine."
);
}
#[test]
fn test_format_diarized_merges_same_speaker() {
let result = TranscriptionResult {
text: "Hello. How are you? I am fine.".to_string(),
metadata: Default::default(),
diarization: Some(vec![
DiarizationSegment {
speaker: "SPEAKER_00".to_string(),
start: 0.0,
end: 1.0,
text: "Hello.".to_string(),
},
DiarizationSegment {
speaker: "SPEAKER_00".to_string(),
start: 1.0,
end: 2.0,
text: "How are you?".to_string(),
},
DiarizationSegment {
speaker: "SPEAKER_01".to_string(),
start: 2.0,
end: 3.5,
text: "I am fine.".to_string(),
},
]),
segments: None,
};
assert_eq!(
format_transcription_output(&result, false),
"[SPEAKER_00] Hello. How are you?\n[SPEAKER_01] I am fine."
);
}
#[test]
fn test_format_diarized_empty_segments() {
let result = TranscriptionResult {
text: "Hello world.".to_string(),
metadata: Default::default(),
diarization: Some(vec![]),
segments: None,
};
assert_eq!(format_transcription_output(&result, false), "Hello world.");
}
#[test]
fn test_format_diarized_with_timestamps() {
let result = TranscriptionResult {
text: "Hello. I am fine.".to_string(),
metadata: Default::default(),
diarization: Some(vec![
DiarizationSegment {
speaker: "speaker_1".to_string(),
start: 0.0,
end: 1.5,
text: "Hello.".to_string(),
},
DiarizationSegment {
speaker: "speaker_2".to_string(),
start: 1.5,
end: 3.0,
text: "I am fine.".to_string(),
},
]),
segments: None,
};
assert_eq!(
format_transcription_output(&result, true),
"[00:00:00] speaker_1 Hello.\n[00:00:01] speaker_2 I am fine."
);
}
#[test]
fn test_format_diarized_with_timestamps_merges_same_speaker() {
let result = TranscriptionResult {
text: "Hello. How are you? I am fine.".to_string(),
metadata: Default::default(),
diarization: Some(vec![
DiarizationSegment {
speaker: "speaker_1".to_string(),
start: 0.0,
end: 1.0,
text: "Hello.".to_string(),
},
DiarizationSegment {
speaker: "speaker_1".to_string(),
start: 1.0,
end: 2.0,
text: "How are you?".to_string(),
},
DiarizationSegment {
speaker: "speaker_2".to_string(),
start: 2.0,
end: 3.5,
text: "I am fine.".to_string(),
},
]),
segments: None,
};
assert_eq!(
format_transcription_output(&result, true),
"[00:00:00] speaker_1 Hello. How are you?\n[00:00:02] speaker_2 I am fine."
);
}
#[test]
fn test_format_timestamp_helper() {
assert_eq!(format_timestamp(0.0), "00:00:00");
assert_eq!(format_timestamp(1.5), "00:00:01");
assert_eq!(format_timestamp(61.0), "00:01:01");
assert_eq!(format_timestamp(3661.0), "01:01:01");
}
#[test]
fn test_parse_transcript_segments_voxtral_shape() {
let raw = serde_json::json!([
{
"text": "Du coup je viens de corriger.",
"start": 1.2,
"end": 10.7,
"type": "transcription_segment",
"speaker_id": null
},
{
"text": " Le premier c'était un bug.",
"start": 11.9,
"end": 24.5,
"type": "transcription_segment",
"speaker_id": null
}
]);
let slice = raw.as_array().expect("array literal is array");
let parsed = parse_transcript_segments(slice).expect("some segments");
assert_eq!(parsed.len(), 2);
assert_eq!(parsed[0].start, 1.2);
assert_eq!(parsed[0].end, 10.7);
assert_eq!(parsed[0].text, "Du coup je viens de corriger.");
assert_eq!(parsed[1].start, 11.9);
assert_eq!(parsed[1].end, 24.5);
}
#[test]
fn test_parse_transcript_segments_whisper_verbose_json_shape() {
let raw = serde_json::json!([
{
"id": 0,
"seek": 0,
"start": 0.0,
"end": 3.2,
"text": " Hello world.",
"tokens": [50364, 2425, 1002, 13, 50524],
"temperature": 0.0,
"avg_logprob": -0.3,
"compression_ratio": 1.1,
"no_speech_prob": 0.01
}
]);
let slice = raw.as_array().expect("array literal is array");
let parsed = parse_transcript_segments(slice).expect("some segments");
assert_eq!(parsed.len(), 1);
assert_eq!(parsed[0].start, 0.0);
assert_eq!(parsed[0].end, 3.2);
assert_eq!(parsed[0].text, " Hello world.");
}
#[test]
fn test_parse_transcript_segments_skips_malformed() {
let raw = serde_json::json!([
{ "text": "no timing here" },
{ "start": 0.0, "text": "no end here" },
{ "start": 0.0, "end": 1.0, "text": "" },
{ "start": 5.0, "end": 6.0 }
]);
let slice = raw.as_array().expect("array literal is array");
assert!(parse_transcript_segments(slice).is_none());
}
#[test]
fn test_parse_transcript_segments_mixed_good_and_bad() {
let raw = serde_json::json!([
{ "text": "no timing" },
{ "start": 1.0, "end": 2.5, "text": "valid" },
{ "start": 0.0, "end": 1.0, "text": "" }
]);
let slice = raw.as_array().expect("array literal is array");
let parsed = parse_transcript_segments(slice).expect("some segments");
assert_eq!(parsed.len(), 1);
assert_eq!(parsed[0].text, "valid");
}
#[test]
fn test_parse_transcript_segments_empty_input() {
let parsed = parse_transcript_segments(&[]);
assert!(parsed.is_none());
}
#[test]
fn test_transcription_result_default_has_no_segments() {
let result = TranscriptionResult::default();
assert!(result.segments.is_none());
assert!(result.diarization.is_none());
assert_eq!(result.text, "");
}
#[test]
fn openai_batch_override_clones_config_and_replaces_only_model() {
let original = crate::config::OpenAIConfig {
api_key: "key".to_string(),
url: Some("https://example.test".to_string()),
model: "gpt-transcribe".to_string(),
realtime_model: "gpt-live-transcribe".to_string(),
prompt: Some("prompt".to_string()),
keywords: Some(vec!["keyword".to_string()]),
languages: Some(vec!["fr".to_string()]),
realtime_delay: Some(crate::config::OpenAIRealtimeDelay::Low),
};
let overridden = override_openai_batch_model(&original, Some("whisper-1"));
assert_eq!(overridden.model, "whisper-1");
assert_eq!(overridden.api_key, original.api_key);
assert_eq!(overridden.url, original.url);
assert_eq!(overridden.realtime_model, original.realtime_model);
assert_eq!(overridden.prompt, original.prompt);
assert_eq!(overridden.keywords, original.keywords);
assert_eq!(overridden.languages, original.languages);
assert_eq!(overridden.realtime_delay, original.realtime_delay);
}
fn minimal_config() -> Config {
use tempfile::NamedTempFile;
let yaml = r#"
output_dir: /tmp/test-output
providers:
mistral:
api_key: test
"#;
let mut file = NamedTempFile::new().expect("tmp file");
std::io::Write::write_all(&mut file, yaml.as_bytes()).expect("write");
Config::load(Some(file.path())).expect("load")
}
#[tokio::test]
async fn test_produce_transcript_returns_existing_pick() {
let dir = tempfile::TempDir::new().expect("tmp dir");
let audio_path = dir.path().join("has-pick.ogg");
std::fs::write(&audio_path, b"fake").expect("write audio");
crate::recording_cache::write_pick(
&audio_path,
"openai",
"whisper-1",
false,
"cached text",
)
.expect("write pick");
let config = minimal_config();
let result = produce_transcript(
&audio_path,
&config,
Provider::OpenAI,
None,
&(std::sync::Arc::new(crate::telemetry::NoOpSink)
as std::sync::Arc<dyn crate::telemetry::TelemetrySink>),
)
.await;
assert_eq!(result.unwrap(), "cached text");
}
#[tokio::test]
async fn test_produce_transcript_returns_in_progress_when_locked() {
let dir = tempfile::TempDir::new().expect("tmp dir");
let audio_path = dir.path().join("locked.ogg");
std::fs::write(&audio_path, b"fake").expect("write audio");
crate::recording_cache::acquire_pick_lock(&audio_path).expect("lock");
let config = minimal_config();
let result = produce_transcript(
&audio_path,
&config,
Provider::OpenAI,
None,
&(std::sync::Arc::new(crate::telemetry::NoOpSink)
as std::sync::Arc<dyn crate::telemetry::TelemetrySink>),
)
.await;
assert!(matches!(result, Err(TalkError::TranscriptInProgress)));
}
#[tokio::test]
async fn test_produce_transcript_prefers_pick_over_lock_only_when_lock_absent() {
let dir = tempfile::TempDir::new().expect("tmp dir");
let audio_path = dir.path().join("pick-only.ogg");
std::fs::write(&audio_path, b"fake").expect("write audio");
crate::recording_cache::write_pick(&audio_path, "openai", "whisper-1", false, "x")
.expect("write pick");
let config = minimal_config();
let result = produce_transcript(
&audio_path,
&config,
Provider::OpenAI,
None,
&(std::sync::Arc::new(crate::telemetry::NoOpSink)
as std::sync::Arc<dyn crate::telemetry::TelemetrySink>),
)
.await;
assert_eq!(result.unwrap(), "x");
}
#[cfg(feature = "parakeet")]
#[test]
fn create_oneshot_transcriber_parakeet_constructs_without_io() {
use crate::config::{ParakeetConfig, ParakeetVariant, ProvidersConfig};
let tmp = tempfile::TempDir::new().expect("tmp dir");
let parakeet_cfg = ParakeetConfig {
variant: ParakeetVariant::Int8,
model_dir: Some(tmp.path().to_path_buf()),
num_threads: 1,
model: None,
};
let config = Config {
output_dir: tmp.path().to_path_buf(),
providers: ProvidersConfig {
mistral: None,
openai: None,
parakeet: Some(parakeet_cfg),
kokoro: None,
},
indicators: None,
transcription: None,
speak: None,
paste: None,
audio: None,
recording: None,
};
let t = create_oneshot_transcriber(
&config,
Provider::Parakeet,
None,
false,
RequestTimeoutPolicy::Proportional,
)
.expect("parakeet construction must succeed");
let _: Box<dyn OneShotTranscriber> = t;
}
#[test]
fn encode_16k_mono_ogg_normalizes_a_long_recording_quickly() {
use crate::audio::{AudioWriter, OggOpusWriter};
const SECONDS: usize = 20 * 60;
const RATE: usize = 16_000;
let dir = tempfile::TempDir::new().expect("tmp dir");
let path = dir.path().join("talkrs-perf-check.ogg");
let pcm: Vec<i16> = (0..SECONDS * RATE)
.map(|i| {
let t = i as f32 / RATE as f32;
((t * 440.0 * std::f32::consts::TAU).sin() * 12000.0) as i16
})
.collect();
let mut writer =
OggOpusWriter::new(crate::config::AudioConfig::new()).expect("writer construction");
let mut fixture = writer.header().expect("header");
fixture.extend_from_slice(&writer.write_pcm(&pcm).expect("encode fixture"));
fixture.extend_from_slice(&writer.finalize().expect("finalize fixture"));
std::fs::write(&path, &fixture).expect("write fixture");
let started = std::time::Instant::now();
let out = encode_16k_mono_ogg(&path).expect("normalization must succeed");
let elapsed = started.elapsed();
assert!(!out.is_empty(), "normalization produced no bytes");
assert_eq!(
&out[0..4],
b"OggS",
"normalized output must be an OGG stream"
);
eprintln!(
"encode_16k_mono_ogg: {SECONDS}s recording ({} bytes in, {} bytes out) in {:?}",
fixture.len(),
out.len(),
elapsed
);
assert!(
elapsed < std::time::Duration::from_secs(60),
"normalizing a {SECONDS}s recording took {elapsed:?}; the pre-upload step has \
regressed to the quadratic-buffer regime"
);
}
}