use super::text::flush_sentences;
use crate::audio::bt_profile;
use crate::audio::recording_feedback::{RecordingBadgeTeardown, RecordingFeedback};
use crate::audio::{AudioCapture, AudioWriter, OggOpusWriter};
use crate::config::{AudioConfig, Config, Provider};
use crate::error::TalkError;
use crate::transcription::{
self, MistralProviderMetadata, OpenAIProviderMetadata, OpenAIRealtimeMetadata,
OrderedItemTranscript, ProviderSpecificMetadata, TranscriptSegment, TranscriptionEvent,
TranscriptionMetadata, TranscriptionResult,
};
use crate::x11::visualizer::VisualizerHandle;
use std::collections::VecDeque;
use std::path::PathBuf;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use tokio::io::AsyncWriteExt;
use tokio_util::sync::CancellationToken;
pub(super) struct AudioBuffer {
chunks: tokio::sync::Mutex<Vec<Vec<i16>>>,
notify: tokio::sync::Notify,
closed: AtomicBool,
}
impl AudioBuffer {
pub(super) fn new() -> Self {
Self {
chunks: tokio::sync::Mutex::new(Vec::new()),
notify: tokio::sync::Notify::new(),
closed: AtomicBool::new(false),
}
}
pub(super) async fn push(&self, chunk: Vec<i16>) {
self.chunks.lock().await.push(chunk);
self.notify.notify_waiters();
}
pub(super) fn close(&self) {
self.closed.store(true, Ordering::Release);
self.notify.notify_waiters();
}
pub(super) async fn is_empty(&self) -> bool {
self.chunks.lock().await.is_empty()
}
pub(super) async fn read_from(&self, cursor: usize) -> (Vec<Vec<i16>>, usize) {
loop {
{
let buf = self.chunks.lock().await;
if buf.len() > cursor {
let new_chunks = buf[cursor..].to_vec();
return (new_chunks, buf.len());
}
if self.closed.load(Ordering::Acquire) {
return (Vec::new(), cursor);
}
}
self.notify.notified().await;
}
}
}
pub(super) async fn ogg_recording_task(
mut source: tokio::sync::mpsc::Receiver<Vec<i16>>,
ogg_path: PathBuf,
audio_config: AudioConfig,
buffer: Arc<AudioBuffer>,
) -> Result<(), TalkError> {
let mut writer = OggOpusWriter::new(audio_config)?;
let header = writer.header()?;
let mut file = tokio::fs::File::create(&ogg_path)
.await
.map_err(TalkError::Io)?;
file.write_all(&header).await.map_err(TalkError::Io)?;
let mut total_samples: u64 = 0;
while let Some(pcm_chunk) = source.recv().await {
total_samples += pcm_chunk.len() as u64;
let encoded_bytes = writer.write_pcm(&pcm_chunk)?;
if !encoded_bytes.is_empty() {
file.write_all(&encoded_bytes)
.await
.map_err(TalkError::Io)?;
}
buffer.push(pcm_chunk).await;
}
buffer.close();
let trailing_bytes = writer.finalize()?;
if !trailing_bytes.is_empty() {
file.write_all(&trailing_bytes)
.await
.map_err(TalkError::Io)?;
}
file.sync_all().await.map_err(TalkError::Io)?;
log::info!(
"cache OGG: {} samples ({:.1}s) saved to {}",
total_samples,
total_samples as f64 / 16000.0,
ogg_path.display()
);
Ok(())
}
pub(super) async fn buffer_feeder(
buffer: Arc<AudioBuffer>,
fwd_tx: tokio::sync::mpsc::Sender<Vec<i16>>,
start_cursor: usize,
) {
let mut cursor = start_cursor;
loop {
let (chunks, new_cursor) = buffer.read_from(cursor).await;
if chunks.is_empty() {
break;
}
for chunk in chunks {
if fwd_tx.send(chunk).await.is_err() {
log::warn!(
"transcriber channel closed at chunk {} — feeder stopping",
cursor
);
return;
}
cursor += 1;
}
cursor = new_cursor;
}
}
#[derive(Debug, Default, PartialEq, Eq)]
struct NormalTranscriptUpdate {
live_text: String,
segments_to_send: Vec<String>,
}
#[derive(Debug, Default, PartialEq, Eq)]
struct FinishedNormalTranscript {
text: String,
segments_to_send: Vec<String>,
}
#[derive(Debug, Default)]
struct NormalTranscriptAccumulator {
generic_segments: Vec<String>,
current_line: String,
item_text: OrderedItemTranscript,
item_segments: Vec<String>,
replay_prefix: VecDeque<String>,
}
impl NormalTranscriptAccumulator {
fn apply(&mut self, event: TranscriptionEvent) -> NormalTranscriptUpdate {
let segments_to_send = match event {
TranscriptionEvent::TextDelta { text } => {
self.current_line.push_str(&text);
let previous_len = self.generic_segments.len();
flush_sentences(&mut self.current_line, &mut self.generic_segments);
self.generic_segments[previous_len..].to_vec()
}
TranscriptionEvent::SegmentDelta { text, .. } => {
let segment = text.trim().to_string();
self.current_line.clear();
if segment.is_empty() {
Vec::new()
} else {
self.generic_segments.push(segment.clone());
vec![segment]
}
}
TranscriptionEvent::ItemCreated {
item_id,
previous_item_id,
} => {
self.item_text
.item_created(&item_id, previous_item_id.as_deref());
let drained = self.item_text.drain_completed_prefix();
self.accept_item_drain(drained)
}
TranscriptionEvent::ItemTextDelta {
item_id,
content_index,
text,
} => {
self.item_text.append_delta(&item_id, content_index, &text);
Vec::new()
}
TranscriptionEvent::ItemTextCompleted {
item_id,
content_index,
transcript,
} => {
self.item_text
.complete(&item_id, content_index, &transcript);
let drained = self.item_text.drain_completed_prefix();
self.accept_item_drain(drained)
}
_ => Vec::new(),
};
NormalTranscriptUpdate {
live_text: self.live_text(),
segments_to_send,
}
}
fn finish(&mut self) -> FinishedNormalTranscript {
if !self.item_text.is_empty() || !self.item_segments.is_empty() {
let drained = self.item_text.drain_terminal();
let segments_to_send = self.accept_item_drain(drained);
return FinishedNormalTranscript {
text: self.item_segments.join(" "),
segments_to_send,
};
}
let trailing = self.current_line.trim().to_string();
let segments_to_send = if trailing.is_empty() {
Vec::new()
} else {
self.generic_segments.push(trailing.clone());
vec![trailing]
};
self.current_line.clear();
FinishedNormalTranscript {
text: self.generic_segments.join(" "),
segments_to_send,
}
}
fn live_text(&self) -> String {
if !self.item_text.is_empty() {
return self.item_text.snapshot();
}
if !self.item_segments.is_empty() {
return self.item_segments.join(" ");
}
let mut live = self.generic_segments.join(" ");
if !live.is_empty() && !self.current_line.is_empty() {
live.push(' ');
}
live.push_str(&self.current_line);
live
}
fn text(&self) -> String {
if self.item_segments.is_empty() {
self.generic_segments.join(" ")
} else {
self.item_segments.join(" ")
}
}
fn segment_count(&self) -> usize {
if self.item_segments.is_empty() {
self.generic_segments.len()
} else {
self.item_segments.len()
}
}
fn reset_item_generation_for_replay(&mut self) {
self.item_text.reset_generation();
self.replay_prefix = self.item_segments.clone().into();
}
fn accept_item_drain(&mut self, drained: Vec<String>) -> Vec<String> {
let mut segments_to_send = Vec::new();
for segment in drained {
if self.replay_prefix.front() == Some(&segment) {
self.replay_prefix.pop_front();
continue;
}
if !self.replay_prefix.is_empty() {
self.replay_prefix.clear();
}
self.item_segments.push(segment.clone());
segments_to_send.push(segment);
}
segments_to_send
}
}
#[allow(clippy::too_many_arguments)]
pub(crate) async fn dictate_realtime(
config: Config,
provider: Provider,
model: Option<&str>,
cache_ogg_path: &std::path::Path,
audio_rx: tokio::sync::mpsc::Receiver<Vec<i16>>,
capture: &mut dyn AudioCapture,
from_file: bool,
feedback: &mut RecordingFeedback,
segment_tx: Option<tokio::sync::mpsc::Sender<String>>,
visualizer: Option<&VisualizerHandle>,
shutdown: &CancellationToken,
mut bt_guard: bt_profile::HeadsetGuard,
) -> Result<TranscriptionResult, TalkError> {
log::info!("caching audio to: {}", cache_ogg_path.display());
let buffer = Arc::new(AudioBuffer::new());
let ogg_task = tokio::spawn(ogg_recording_task(
audio_rx,
cache_ogg_path.to_path_buf(),
AudioConfig::new(),
Arc::clone(&buffer),
));
let transcriber = transcription::create_realtime_transcriber(&config, provider, model)?;
transcriber.validate().await?;
let (fwd_tx, fwd_rx) = tokio::sync::mpsc::channel::<Vec<i16>>(100);
let mut feeder_handle = tokio::spawn(buffer_feeder(Arc::clone(&buffer), fwd_tx, 0));
let mut event_rx = transcriber.transcribe_realtime(fwd_rx).await?;
let started = std::time::Instant::now();
if from_file {
log::info!("transcribing audio file (realtime)...");
} else {
log::info!("recording (realtime)... press Ctrl+C to stop");
}
let capture_stop = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false));
let capture_stop_clone = capture_stop.clone();
let shutdown_clone = shutdown.clone();
let ctrlc_task = tokio::spawn(async move {
log::warn!("[DBG] dictate_realtime: waiting on shutdown token");
shutdown_clone.cancelled().await;
log::warn!("[DBG] dictate_realtime: shutdown token fired, setting capture_stop");
capture_stop_clone.store(true, std::sync::atomic::Ordering::Release);
});
let mut transcript = NormalTranscriptAccumulator::default();
let mut timed_segments: Vec<TranscriptSegment> = Vec::new();
let mut detected_language: Option<String> = None;
let mut unknown_event_types: Vec<String> = Vec::new();
let mut event_counts: std::collections::BTreeMap<String, u64> =
std::collections::BTreeMap::new();
let mut api_segment_count: usize = 0;
let mut session_id: Option<String> = None;
let mut conversation_id: Option<String> = None;
let mut last_rate_limits: Option<serde_json::Value> = None;
let mut ws_upgrade_headers: std::collections::BTreeMap<String, String> =
std::collections::BTreeMap::new();
let bump = |key: &str, counts: &mut std::collections::BTreeMap<String, u64>| {
let entry = counts.entry(key.to_string()).or_insert(0);
*entry += 1;
};
loop {
if capture_stop.load(std::sync::atomic::Ordering::Acquire) {
log::info!("stopping recording");
feedback.teardown_recording(RecordingBadgeTeardown::KeepVisible);
feedback.play_stop_now();
capture.stop()?;
bt_guard.restore_now_async();
capture_stop.store(false, std::sync::atomic::Ordering::Release);
}
tokio::select! {
event = event_rx.recv() => {
match event {
Some(TranscriptionEvent::TextDelta { text }) => {
bump("text_delta", &mut event_counts);
let update = transcript.apply(TranscriptionEvent::TextDelta { text });
eprint!("\r{}", update.live_text);
if let Some(viz) = visualizer {
viz.set_text(&update.live_text);
}
if let Some(ref tx) = segment_tx {
for segment in update.segments_to_send {
let _ = tx.send(segment).await;
}
}
}
Some(TranscriptionEvent::SegmentDelta { text, start, end }) => {
bump("segment_delta", &mut event_counts);
api_segment_count += 1;
let segment_text = text.trim().to_string();
if !segment_text.is_empty() {
if let (Some(start), Some(end)) = (start, end) {
timed_segments.push(TranscriptSegment {
start,
end,
text: segment_text.clone(),
});
}
}
let update = transcript.apply(TranscriptionEvent::SegmentDelta {
text,
start,
end,
});
for segment in update.segments_to_send {
println!("{}", segment);
if let Some(ref tx) = segment_tx {
let _ = tx.send(segment).await;
}
}
if let Some(viz) = visualizer {
viz.set_text(&update.live_text);
}
}
Some(event @ TranscriptionEvent::ItemCreated { .. }) => {
bump("item_created", &mut event_counts);
let update = transcript.apply(event);
for segment in update.segments_to_send {
println!("{}", segment);
if let Some(ref tx) = segment_tx {
let _ = tx.send(segment).await;
}
}
}
Some(event @ TranscriptionEvent::ItemTextDelta { .. }) => {
bump("item_text_delta", &mut event_counts);
let update = transcript.apply(event);
eprint!("\r{}", update.live_text);
if let Some(viz) = visualizer {
viz.set_text(&update.live_text);
}
}
Some(event @ TranscriptionEvent::ItemTextCompleted { .. }) => {
bump("item_text_completed", &mut event_counts);
api_segment_count += 1;
let update = transcript.apply(event);
eprint!("\r{}", update.live_text);
if let Some(viz) = visualizer {
viz.set_text(&update.live_text);
}
for segment in update.segments_to_send {
println!("{}", segment);
if let Some(ref tx) = segment_tx {
let _ = tx.send(segment).await;
}
}
}
Some(TranscriptionEvent::Done) => {
bump("done", &mut event_counts);
break;
}
Some(TranscriptionEvent::Error { message }) => {
bump("error", &mut event_counts);
let msg = format!(
"Transcription error: {} — reconnecting",
message
);
log::warn!("{}", msg);
if let Some(viz) = visualizer {
viz.push_message(&msg);
}
feeder_handle.abort();
match transcription::create_realtime_transcriber(&config, provider, model)
{
Ok(new_transcriber) => {
let (new_fwd_tx, new_fwd_rx) =
tokio::sync::mpsc::channel::<Vec<i16>>(100);
feeder_handle = tokio::spawn(buffer_feeder(
Arc::clone(&buffer),
new_fwd_tx,
0,
));
match new_transcriber.transcribe_realtime(new_fwd_rx).await {
Ok(new_rx) => {
log::info!("realtime transcription reconnected");
transcript.reset_item_generation_for_replay();
event_rx = new_rx;
continue;
}
Err(e) => {
let msg = format!(
"Reconnect failed: {}",
e
);
log::warn!("{}", msg);
if let Some(viz) = visualizer {
viz.push_message(&msg);
}
break;
}
}
}
Err(e) => {
let msg = format!(
"Reconnect failed: {}",
e
);
log::warn!("{}", msg);
if let Some(viz) = visualizer {
viz.push_message(&msg);
}
break;
}
}
}
Some(TranscriptionEvent::SessionCreated) => {
bump("session_created", &mut event_counts);
log::debug!("session created event received");
}
Some(TranscriptionEvent::SessionInfo { session_id: sid, conversation_id: cid }) => {
bump("session_info", &mut event_counts);
if sid.is_some() {
session_id = sid;
}
if cid.is_some() {
conversation_id = cid;
}
}
Some(TranscriptionEvent::RateLimitsUpdated { raw }) => {
bump("rate_limits_updated", &mut event_counts);
last_rate_limits = Some(raw);
}
Some(TranscriptionEvent::TransportMetadata { headers }) => {
bump("transport_metadata", &mut event_counts);
ws_upgrade_headers.extend(headers);
}
Some(TranscriptionEvent::Language { language }) => {
bump("language", &mut event_counts);
log::info!("detected language: {}", language);
detected_language = Some(language);
}
Some(TranscriptionEvent::Unknown { event_type, .. }) => {
bump("unknown", &mut event_counts);
if let Some(kind) = event_type {
bump(&format!("event:{kind}"), &mut event_counts);
if !unknown_event_types.contains(&kind) {
unknown_event_types.push(kind);
}
}
}
None => {
bump("channel_closed", &mut event_counts);
log::warn!("realtime event channel closed — attempting reconnect");
if let Some(viz) = visualizer {
viz.push_message("Connection lost — reconnecting");
}
feeder_handle.abort();
match transcription::create_realtime_transcriber(&config, provider, model)
{
Ok(new_transcriber) => {
let (new_fwd_tx, new_fwd_rx) =
tokio::sync::mpsc::channel::<Vec<i16>>(100);
feeder_handle = tokio::spawn(buffer_feeder(
Arc::clone(&buffer),
new_fwd_tx,
0,
));
match new_transcriber.transcribe_realtime(new_fwd_rx).await {
Ok(new_rx) => {
log::info!("realtime transcription reconnected");
transcript.reset_item_generation_for_replay();
event_rx = new_rx;
continue;
}
Err(e) => {
let msg = format!("Reconnect failed: {}", e);
log::warn!("{}", msg);
if let Some(viz) = visualizer {
viz.push_message(&msg);
}
}
}
}
Err(e) => {
let msg = format!("Reconnect failed: {}", e);
log::warn!("{}", msg);
if let Some(viz) = visualizer {
viz.push_message(&msg);
}
}
}
break;
}
}
}
_ = tokio::time::sleep(std::time::Duration::from_millis(50)) => {
}
}
}
let finished = transcript.finish();
for segment in finished.segments_to_send {
println!("{}", segment);
if let Some(ref tx) = segment_tx {
let _ = tx.send(segment).await;
}
}
eprintln!();
ctrlc_task.abort();
feeder_handle.abort();
match ogg_task.await {
Ok(Ok(())) => log::debug!("cache OGG saved"),
Ok(Err(e)) => log::warn!("cache OGG write error: {}", e),
Err(e) => log::warn!("cache OGG task panicked: {}", e),
}
let provider_specific = match provider {
Provider::OpenAI => Some(ProviderSpecificMetadata::OpenAI(OpenAIProviderMetadata {
model: model.map(str::to_string),
usage_raw: None,
rate_limit_headers: std::collections::BTreeMap::new(),
unknown_event_types,
realtime: Some(OpenAIRealtimeMetadata {
session_id,
conversation_id,
event_counts,
last_rate_limits,
ws_upgrade_headers: ws_upgrade_headers.clone(),
}),
})),
Provider::Mistral => Some(ProviderSpecificMetadata::Mistral(MistralProviderMetadata {
model: model.map(str::to_string),
usage_raw: None,
unknown_event_types,
})),
Provider::Parakeet => None,
};
Ok(TranscriptionResult {
text: transcript.text(),
metadata: TranscriptionMetadata {
request_latency_ms: None,
session_elapsed_ms: Some(started.elapsed().as_millis() as u64),
request_id: ws_upgrade_headers.get("x-request-id").cloned(),
provider_processing_ms: ws_upgrade_headers
.get("openai-processing-ms")
.and_then(|s| s.parse::<u64>().ok()),
detected_language,
audio_seconds: None,
segment_count: Some(if api_segment_count > 0 {
api_segment_count
} else {
transcript.segment_count()
}),
word_count: None,
token_usage: None,
provider_specific,
},
diarization: None,
segments: if timed_segments.is_empty() {
None
} else {
Some(timed_segments)
},
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::AudioConfig;
#[test]
fn normal_openai_completion_emits_incrementally_and_finish_does_not_resend() {
let mut transcript = NormalTranscriptAccumulator::default();
transcript.apply(TranscriptionEvent::ItemCreated {
item_id: "item-1".to_string(),
previous_item_id: None,
});
let delta = transcript.apply(TranscriptionEvent::ItemTextDelta {
item_id: "item-1".to_string(),
content_index: 0,
text: "Hello world.".to_string(),
});
assert_eq!(delta.live_text, "Hello world.");
assert!(delta.segments_to_send.is_empty());
let completed = transcript.apply(TranscriptionEvent::ItemTextCompleted {
item_id: "item-1".to_string(),
content_index: 0,
transcript: "Hello, corrected world.".to_string(),
});
assert_eq!(completed.live_text, "Hello, corrected world.");
assert_eq!(
completed.segments_to_send,
vec!["Hello, corrected world.".to_string()]
);
let finished = transcript.finish();
assert_eq!(finished.text, "Hello, corrected world.");
assert!(finished.segments_to_send.is_empty());
}
#[test]
fn normal_openai_reverse_completion_emits_in_conversation_order() {
let mut transcript = NormalTranscriptAccumulator::default();
transcript.apply(TranscriptionEvent::ItemCreated {
item_id: "item-1".to_string(),
previous_item_id: None,
});
transcript.apply(TranscriptionEvent::ItemCreated {
item_id: "item-2".to_string(),
previous_item_id: Some("item-1".to_string()),
});
let second = transcript.apply(TranscriptionEvent::ItemTextCompleted {
item_id: "item-2".to_string(),
content_index: 0,
transcript: "second".to_string(),
});
assert!(second.segments_to_send.is_empty());
let first = transcript.apply(TranscriptionEvent::ItemTextCompleted {
item_id: "item-1".to_string(),
content_index: 0,
transcript: "first".to_string(),
});
assert_eq!(
first.segments_to_send,
vec!["first".to_string(), "second".to_string()]
);
assert_eq!(transcript.finish().text, "first second");
}
#[test]
fn normal_openai_late_item_created_event_unblocks_incremental_emission() {
let mut transcript = NormalTranscriptAccumulator::default();
let second = transcript.apply(TranscriptionEvent::ItemTextCompleted {
item_id: "item-2".to_string(),
content_index: 0,
transcript: "second".to_string(),
});
assert!(second.segments_to_send.is_empty());
transcript.apply(TranscriptionEvent::ItemCreated {
item_id: "item-1".to_string(),
previous_item_id: None,
});
let first = transcript.apply(TranscriptionEvent::ItemTextCompleted {
item_id: "item-1".to_string(),
content_index: 0,
transcript: "first".to_string(),
});
assert!(first.segments_to_send.is_empty());
let ordered = transcript.apply(TranscriptionEvent::ItemCreated {
item_id: "item-2".to_string(),
previous_item_id: Some("item-1".to_string()),
});
assert_eq!(
ordered.segments_to_send,
vec!["first".to_string(), "second".to_string()]
);
}
#[test]
fn normal_openai_finish_preserves_provisional_terminal_text() {
let mut transcript = NormalTranscriptAccumulator::default();
transcript.apply(TranscriptionEvent::ItemCreated {
item_id: "item-1".to_string(),
previous_item_id: None,
});
transcript.apply(TranscriptionEvent::ItemTextDelta {
item_id: "item-1".to_string(),
content_index: 0,
text: "provisional terminal text".to_string(),
});
let finished = transcript.finish();
assert_eq!(finished.text, "provisional terminal text");
assert_eq!(
finished.segments_to_send,
vec!["provisional terminal text".to_string()]
);
}
#[test]
fn normal_openai_replay_reset_deduplicates_emitted_prefix() {
let mut transcript = NormalTranscriptAccumulator::default();
transcript.apply(TranscriptionEvent::ItemCreated {
item_id: "old-item".to_string(),
previous_item_id: None,
});
let old = transcript.apply(TranscriptionEvent::ItemTextCompleted {
item_id: "old-item".to_string(),
content_index: 0,
transcript: "old text".to_string(),
});
assert_eq!(old.segments_to_send, vec!["old text".to_string()]);
transcript.reset_item_generation_for_replay();
transcript.apply(TranscriptionEvent::ItemCreated {
item_id: "fresh-1".to_string(),
previous_item_id: None,
});
let replayed = transcript.apply(TranscriptionEvent::ItemTextCompleted {
item_id: "fresh-1".to_string(),
content_index: 0,
transcript: "old text".to_string(),
});
assert!(replayed.segments_to_send.is_empty());
transcript.apply(TranscriptionEvent::ItemCreated {
item_id: "fresh-2".to_string(),
previous_item_id: Some("fresh-1".to_string()),
});
let suffix = transcript.apply(TranscriptionEvent::ItemTextCompleted {
item_id: "fresh-2".to_string(),
content_index: 0,
transcript: "new text".to_string(),
});
assert_eq!(suffix.segments_to_send, vec!["new text".to_string()]);
let finished = transcript.finish();
assert_eq!(finished.text, "old text new text");
assert!(finished.segments_to_send.is_empty());
}
#[test]
fn normal_generic_segments_remain_additive() {
let mut transcript = NormalTranscriptAccumulator::default();
let first = transcript.apply(TranscriptionEvent::SegmentDelta {
text: "first".to_string(),
start: None,
end: None,
});
let second = transcript.apply(TranscriptionEvent::SegmentDelta {
text: "second".to_string(),
start: None,
end: None,
});
assert_eq!(first.segments_to_send, vec!["first".to_string()]);
assert_eq!(second.segments_to_send, vec!["second".to_string()]);
assert_eq!(transcript.finish().text, "first second");
}
#[tokio::test]
async fn audio_buffer_push_then_read_returns_chunks() {
let buf = AudioBuffer::new();
buf.push(vec![1, 2, 3]).await;
buf.push(vec![4, 5, 6]).await;
let (chunks, cursor) = buf.read_from(0).await;
assert_eq!(chunks.len(), 2);
assert_eq!(chunks[0], vec![1, 2, 3]);
assert_eq!(chunks[1], vec![4, 5, 6]);
assert_eq!(cursor, 2);
}
#[tokio::test]
async fn audio_buffer_read_from_cursor_skips_earlier() {
let buf = AudioBuffer::new();
buf.push(vec![10]).await;
buf.push(vec![20]).await;
buf.push(vec![30]).await;
let (chunks, cursor) = buf.read_from(2).await;
assert_eq!(chunks.len(), 1);
assert_eq!(chunks[0], vec![30]);
assert_eq!(cursor, 3);
}
#[tokio::test]
async fn audio_buffer_close_unblocks_empty_read() {
let buf = Arc::new(AudioBuffer::new());
buf.push(vec![1]).await;
let (_chunks, cursor) = buf.read_from(0).await;
assert_eq!(cursor, 1);
let buf2 = Arc::clone(&buf);
tokio::spawn(async move {
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
buf2.close();
});
let (chunks, cursor) = buf.read_from(1).await;
assert!(chunks.is_empty());
assert_eq!(cursor, 1);
}
#[tokio::test]
async fn audio_buffer_push_after_close_still_accessible() {
let buf = AudioBuffer::new();
buf.push(vec![42]).await;
buf.close();
let (chunks, _) = buf.read_from(0).await;
assert_eq!(chunks, vec![vec![42]]);
}
#[tokio::test]
async fn buffer_feeder_replays_from_cursor_zero() {
let buf = Arc::new(AudioBuffer::new());
buf.push(vec![1, 2]).await;
buf.push(vec![3, 4]).await;
buf.close();
let (tx, mut rx) = tokio::sync::mpsc::channel(10);
buffer_feeder(buf, tx, 0).await;
let c1 = rx.recv().await;
let c2 = rx.recv().await;
let c3 = rx.recv().await;
assert_eq!(c1, Some(vec![1, 2]));
assert_eq!(c2, Some(vec![3, 4]));
assert!(c3.is_none()); }
#[tokio::test]
async fn buffer_feeder_stops_when_receiver_dropped() {
let buf = Arc::new(AudioBuffer::new());
buf.push(vec![10]).await;
buf.push(vec![20]).await;
let (tx, rx) = tokio::sync::mpsc::channel(1);
drop(rx);
let handle = tokio::spawn(buffer_feeder(buf, tx, 0));
tokio::time::timeout(std::time::Duration::from_secs(2), handle)
.await
.expect("feeder should finish promptly")
.expect("feeder should not panic");
}
#[tokio::test]
async fn buffer_feeder_starts_from_nonzero_cursor() {
let buf = Arc::new(AudioBuffer::new());
buf.push(vec![100]).await;
buf.push(vec![200]).await;
buf.push(vec![300]).await;
buf.close();
let (tx, mut rx) = tokio::sync::mpsc::channel(10);
buffer_feeder(buf, tx, 2).await;
let c1 = rx.recv().await;
let c2 = rx.recv().await;
assert_eq!(c1, Some(vec![300]));
assert!(c2.is_none());
}
fn read_ogg_packets(path: &std::path::Path) -> Vec<Vec<u8>> {
let file = std::fs::File::open(path).expect("open ogg");
let mut reader = ogg::reading::PacketReader::new(std::io::BufReader::new(file));
let mut packets = Vec::new();
while let Some(packet) = reader.read_packet().expect("read packet") {
packets.push(packet.data);
}
packets
}
#[tokio::test]
async fn ogg_recording_task_writes_complete_ogg() {
let dir = tempfile::tempdir().expect("create temp dir");
let ogg_path = dir.path().join("test.ogg");
let audio_config = AudioConfig::new();
let buffer = Arc::new(AudioBuffer::new());
let (tx, rx) = tokio::sync::mpsc::channel(10);
let buf_clone = Arc::clone(&buffer);
let path_clone = ogg_path.clone();
let handle = tokio::spawn(ogg_recording_task(rx, path_clone, audio_config, buf_clone));
for i in 0..5u16 {
let chunk: Vec<i16> = (0..320)
.map(|s| (s as i16).wrapping_mul(i as i16))
.collect();
tx.send(chunk).await.expect("send chunk");
}
drop(tx);
handle.await.expect("task join").expect("ogg write");
let data = std::fs::read(&ogg_path).expect("read ogg");
assert_eq!(&data[0..4], b"OggS");
let packets = read_ogg_packets(&ogg_path);
assert_eq!(&packets[0][..8], b"OpusHead");
assert_eq!(&packets[1][..8], b"OpusTags");
assert_eq!(packets.len(), 7);
}
#[tokio::test]
async fn ogg_recording_task_populates_buffer() {
let dir = tempfile::tempdir().expect("create temp dir");
let ogg_path = dir.path().join("test.ogg");
let audio_config = AudioConfig::new();
let buffer = Arc::new(AudioBuffer::new());
let (tx, rx) = tokio::sync::mpsc::channel(10);
let buf_clone = Arc::clone(&buffer);
let handle = tokio::spawn(ogg_recording_task(rx, ogg_path, audio_config, buf_clone));
tx.send(vec![1, 2, 3]).await.expect("send");
tx.send(vec![4, 5, 6]).await.expect("send");
drop(tx);
handle.await.expect("join").expect("ogg");
let (chunks, _) = buffer.read_from(0).await;
assert_eq!(chunks.len(), 2);
assert_eq!(chunks[0], vec![1, 2, 3]);
assert_eq!(chunks[1], vec![4, 5, 6]);
let (empty, _) = buffer.read_from(2).await;
assert!(empty.is_empty());
}
#[tokio::test]
async fn ogg_recording_independent_of_feeder_failure() {
let dir = tempfile::tempdir().expect("create temp dir");
let ogg_path = dir.path().join("test.ogg");
let audio_config = AudioConfig::new();
let buffer = Arc::new(AudioBuffer::new());
let (tx, rx) = tokio::sync::mpsc::channel(10);
let buf_clone = Arc::clone(&buffer);
let path_clone = ogg_path.clone();
let ogg_handle = tokio::spawn(ogg_recording_task(rx, path_clone, audio_config, buf_clone));
let (fwd_tx, fwd_rx) = tokio::sync::mpsc::channel(10);
let feeder = tokio::spawn(buffer_feeder(Arc::clone(&buffer), fwd_tx, 0));
tx.send(vec![10; 320]).await.expect("send");
tx.send(vec![20; 320]).await.expect("send");
drop(fwd_rx);
let _ = tokio::time::timeout(std::time::Duration::from_secs(1), feeder).await;
tx.send(vec![30; 320])
.await
.expect("send after feeder death");
drop(tx);
ogg_handle.await.expect("join").expect("ogg");
let packets = read_ogg_packets(&ogg_path);
assert_eq!(packets.len(), 5);
let (chunks, _) = buffer.read_from(0).await;
assert_eq!(chunks.len(), 3);
}
}