use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Duration;
use cpal::traits::{DeviceTrait, HostTrait, StreamTrait};
use cpal::{BufferSize, SampleFormat, SizedSample, StreamConfig, SupportedBufferSize};
use ringbuf::traits::{Consumer, Observer, Producer, Split};
use ringbuf::{HeapCons, HeapProd, HeapRb};
use tokio::sync::mpsc;
use crate::error::{AudioError, Result};
use crate::tts::TtsClip;
use super::resample::{resample_offline, LinearResampler};
use super::{config_ranges, device_name, negotiate_config, pick_device};
pub const CLIP_SAMPLE_RATE: u32 = crate::tts::TTS_SAMPLE_RATE;
const PLAYBACK_RING_CAPACITY: usize = 48_000;
const CLIPS_CHANNEL_CAPACITY: usize = 8;
const RETRY_INTERVAL: Duration = Duration::from_millis(5);
const OUTPUT_BUFFER_FRAMES: u32 = 512;
const FALLBACK_PERIOD_FRAMES: usize = 8192;
#[derive(Debug, Clone, Default)]
pub struct AudioOutputConfig {
pub device_name: Option<String>,
}
pub struct Playback {
worker: std::thread::JoinHandle<()>,
stop: Arc<AtomicBool>,
}
impl Playback {
pub fn start(cfg: &AudioOutputConfig) -> Result<(Playback, PlaybackHandle)> {
let host = cpal::default_host();
let device = pick_device(
cfg.device_name.as_deref(),
host.output_devices(),
host.default_output_device(),
)?;
tracing::info!(
device = device_name(&device).as_deref().unwrap_or("<unknown>"),
"playing to output device"
);
let preferred_rate = device
.default_output_config()
.map(|cfg| cfg.sample_rate())
.unwrap_or(48_000);
let ranges = config_ranges("output", device.supported_output_configs())?;
let negotiated = negotiate_config(&ranges, preferred_rate).ok_or_else(|| {
AudioError::StreamConfig(
"device offers no f32/i16/u16 output configuration".to_string(),
)
})?;
let sample_format = negotiated.sample_format();
let mut config = negotiated.config();
if let SupportedBufferSize::Range { min, max } = *negotiated.buffer_size() {
config.buffer_size = BufferSize::Fixed(OUTPUT_BUFFER_FRAMES.clamp(min, max));
}
let device_rate = config.sample_rate;
let max_period_frames = match config.buffer_size {
BufferSize::Fixed(n) => n as usize,
BufferSize::Default => FALLBACK_PERIOD_FRAMES,
};
tracing::debug!(
?sample_format,
channels = config.channels,
device_rate,
buffer_size = ?config.buffer_size,
"negotiated output stream config"
);
let (prod, cons) = HeapRb::new(PLAYBACK_RING_CAPACITY).split();
let (clips_tx, clips_rx) = mpsc::channel::<TtsClip>(CLIPS_CHANNEL_CAPACITY);
let flush_epoch = Arc::new(AtomicU64::new(0));
let is_playing = Arc::new(AtomicBool::new(false));
let stop = Arc::new(AtomicBool::new(false));
let clips_queued = Arc::new(AtomicU64::new(0));
let clips_pushed = Arc::new(AtomicU64::new(0));
let samples_pushed = Arc::new(AtomicU64::new(0));
let samples_consumed = Arc::new(AtomicU64::new(0));
let pump = OutputPump::new(
cons,
device_rate,
max_period_frames,
Arc::clone(&flush_epoch),
Arc::clone(&is_playing),
Arc::clone(&samples_consumed),
);
let stream = match sample_format {
SampleFormat::F32 => {
build_output::<f32>(&device, config, max_period_frames, pump, |s| s)
}
SampleFormat::I16 => {
build_output::<i16>(&device, config, max_period_frames, pump, |s| {
(s.clamp(-1.0, 1.0) * 32767.0) as i16
})
}
SampleFormat::U16 => {
build_output::<u16>(&device, config, max_period_frames, pump, |s| {
((s.clamp(-1.0, 1.0) + 1.0) * 32767.5) as u16
})
}
other => {
return Err(AudioError::StreamConfig(format!(
"unsupported output sample format {other:?}"
))
.into());
}
}?;
stream
.play()
.map_err(|err| AudioError::StreamBuild(err.to_string()))?;
let worker = std::thread::Builder::new()
.name("skadoosh-playback".to_string())
.spawn({
let flush_epoch = Arc::clone(&flush_epoch);
let stop = Arc::clone(&stop);
let clips_pushed = Arc::clone(&clips_pushed);
let samples_pushed = Arc::clone(&samples_pushed);
move || {
let _stream = stream;
clip_pump_loop(
clips_rx,
prod,
&flush_epoch,
&stop,
&clips_pushed,
&samples_pushed,
);
}
})
.map_err(|err| {
AudioError::StreamBuild(format!("failed to spawn playback thread: {err}"))
})?;
Ok((
Playback {
worker,
stop: Arc::clone(&stop),
},
PlaybackHandle {
clips_tx,
flush_epoch,
is_playing,
sample_rate: CLIP_SAMPLE_RATE,
stop,
clips_queued,
clips_pushed,
samples_pushed,
samples_consumed,
},
))
}
pub fn stop(self) {
self.stop.store(true, Ordering::Relaxed);
let _ = self.worker.join();
}
}
#[derive(Clone)]
pub struct PlaybackHandle {
clips_tx: mpsc::Sender<TtsClip>,
flush_epoch: Arc<AtomicU64>,
is_playing: Arc<AtomicBool>,
sample_rate: u32,
stop: Arc<AtomicBool>,
clips_queued: Arc<AtomicU64>,
clips_pushed: Arc<AtomicU64>,
samples_pushed: Arc<AtomicU64>,
samples_consumed: Arc<AtomicU64>,
}
impl PlaybackHandle {
pub async fn queue_clip(&self, clip: TtsClip) -> Result<()> {
self.clips_tx
.send(clip)
.await
.map_err(|_| AudioError::StreamBuild("playback thread exited".to_string()))?;
self.clips_queued.fetch_add(1, Ordering::SeqCst);
Ok(())
}
pub fn flush(&self) {
self.flush_epoch.fetch_add(1, Ordering::Release);
}
pub fn is_playing(&self) -> bool {
self.is_playing.load(Ordering::Acquire)
}
pub fn sample_rate(&self) -> u32 {
self.sample_rate
}
pub async fn wait_buffered(&self) -> bool {
let epoch = self.flush_epoch.load(Ordering::Acquire);
let target = self.clips_queued.load(Ordering::SeqCst);
loop {
if self.drain_aborted(epoch) {
return false;
}
if self.clips_pushed.load(Ordering::SeqCst) >= target {
return true;
}
tokio::time::sleep(RETRY_INTERVAL).await;
}
}
pub async fn wait_drained(&self) -> bool {
let epoch = self.flush_epoch.load(Ordering::Acquire);
let target_clips = self.clips_queued.load(Ordering::SeqCst);
let target_samples = loop {
if self.drain_aborted(epoch) {
return false;
}
let pushed = self.clips_pushed.load(Ordering::SeqCst);
let samples = self.samples_pushed.load(Ordering::SeqCst);
if pushed >= target_clips {
break samples;
}
tokio::time::sleep(RETRY_INTERVAL).await;
};
loop {
if self.drain_aborted(epoch) {
return false;
}
if self.samples_consumed.load(Ordering::SeqCst) >= target_samples {
return true;
}
tokio::time::sleep(RETRY_INTERVAL).await;
}
}
fn drain_aborted(&self, epoch: u64) -> bool {
self.stop.load(Ordering::Relaxed) || self.flush_epoch.load(Ordering::Acquire) != epoch
}
}
pub struct OutputPump {
cons: HeapCons<f32>,
resampler: LinearResampler,
flush_epoch: Arc<AtomicU64>,
seen_epoch: u64,
is_playing: Arc<AtomicBool>,
samples_consumed: Arc<AtomicU64>,
max_period_frames: usize,
pending: Vec<f32>,
ring_scratch: Vec<f32>,
block_scratch: Vec<f32>,
}
impl OutputPump {
pub fn new(
cons: HeapCons<f32>,
device_rate: u32,
max_period_frames: usize,
flush_epoch: Arc<AtomicU64>,
is_playing: Arc<AtomicBool>,
samples_consumed: Arc<AtomicU64>,
) -> Self {
let seen_epoch = flush_epoch.load(Ordering::Acquire);
let resampler = LinearResampler::new(CLIP_SAMPLE_RATE, device_rate);
let frames = max_period_frames as u64;
let clip = u64::from(CLIP_SAMPLE_RATE);
let device = u64::from(device_rate);
let ring_cap = (frames * clip / device + 2) as usize + 8;
let block_cap = ((frames * clip / device + 2) * device / clip + 2) as usize + 8;
Self {
cons,
resampler,
flush_epoch,
seen_epoch,
is_playing,
samples_consumed,
max_period_frames,
pending: Vec::with_capacity(block_cap),
ring_scratch: Vec::with_capacity(ring_cap),
block_scratch: Vec::with_capacity(block_cap),
}
}
pub fn render(&mut self, out: &mut [f32]) {
let epoch = self.flush_epoch.load(Ordering::Acquire);
if epoch != self.seen_epoch {
self.seen_epoch = epoch;
self.cons.clear();
self.pending.clear();
}
let frames = out.len();
debug_assert!(
frames <= self.max_period_frames,
"output period of {frames} frames exceeds the preallocated {}-frame scratch; \
this render allocates on the RT thread",
self.max_period_frames
);
let mut written = 0usize;
let mut popped = 0u64;
let take = self.pending.len().min(frames);
out[..take].copy_from_slice(&self.pending[..take]);
self.pending.drain(..take);
written += take;
while written < frames {
let remaining = frames - written;
let want = (remaining * CLIP_SAMPLE_RATE as usize / self.resampler.dst_rate() as usize
+ 2)
.min(self.cons.occupied_len());
if want == 0 {
break; }
self.ring_scratch.clear();
self.ring_scratch.resize(want, 0.0);
let got = self.cons.pop_slice(&mut self.ring_scratch);
self.ring_scratch.truncate(got);
popped += got as u64;
self.resampler
.process(&self.ring_scratch, &mut self.block_scratch);
if self.block_scratch.is_empty() {
break; }
let take = self.block_scratch.len().min(frames - written);
out[written..written + take].copy_from_slice(&self.block_scratch[..take]);
written += take;
if take < self.block_scratch.len() {
self.pending.extend_from_slice(&self.block_scratch[take..]);
}
}
if popped > 0 {
self.samples_consumed.fetch_add(popped, Ordering::SeqCst);
}
out[written..].fill(0.0);
self.is_playing.store(written > 0, Ordering::Release);
}
}
pub fn push_clip_blocking(
prod: &mut HeapProd<f32>,
samples: &[f32],
flush_epoch: &AtomicU64,
seen_epoch: &mut u64,
stop: &AtomicBool,
) -> bool {
let mut rest = samples;
while !rest.is_empty() {
let epoch = flush_epoch.load(Ordering::Acquire);
if epoch != *seen_epoch {
*seen_epoch = epoch;
return false;
}
if stop.load(Ordering::Relaxed) {
return false;
}
let pushed = prod.push_slice(rest);
rest = &rest[pushed..];
if !rest.is_empty() {
std::thread::sleep(RETRY_INTERVAL);
}
}
true
}
fn clip_pump_loop(
mut clips_rx: mpsc::Receiver<TtsClip>,
mut prod: HeapProd<f32>,
flush_epoch: &AtomicU64,
stop: &AtomicBool,
clips_pushed: &AtomicU64,
samples_pushed: &AtomicU64,
) {
let mut seen_epoch = flush_epoch.load(Ordering::Acquire);
while !stop.load(Ordering::Relaxed) {
match clips_rx.try_recv() {
Ok(clip) => {
let samples = if clip.sample_rate == CLIP_SAMPLE_RATE {
clip.samples
} else {
tracing::warn!(
clip_rate = clip.sample_rate,
"clip sample rate violates the {CLIP_SAMPLE_RATE} Hz contract; resampling"
);
resample_offline(&clip.samples, clip.sample_rate, CLIP_SAMPLE_RATE)
};
if push_clip_blocking(&mut prod, &samples, flush_epoch, &mut seen_epoch, stop) {
samples_pushed.fetch_add(samples.len() as u64, Ordering::SeqCst);
clips_pushed.fetch_add(1, Ordering::SeqCst);
} else {
while clips_rx.try_recv().is_ok() {}
}
}
Err(mpsc::error::TryRecvError::Empty) => std::thread::sleep(RETRY_INTERVAL),
Err(mpsc::error::TryRecvError::Disconnected) => break,
}
}
}
fn build_output<T>(
device: &cpal::Device,
config: StreamConfig,
max_period_frames: usize,
mut pump: OutputPump,
convert: fn(f32) -> T,
) -> Result<cpal::Stream>
where
T: SizedSample + 'static,
{
let channels = (config.channels as usize).max(1);
let mut mono: Vec<f32> = Vec::with_capacity(max_period_frames);
device
.build_output_stream(
config,
move |out: &mut [T], _: &cpal::OutputCallbackInfo| {
let frames = out.len() / channels;
debug_assert!(
frames <= mono.capacity(),
"output period of {frames} frames exceeds the preallocated {}-frame scratch; \
this callback allocates on the RT thread",
mono.capacity()
);
mono.clear();
mono.resize(frames, 0.0);
pump.render(&mut mono);
for (i, slot) in out.iter_mut().enumerate() {
*slot = convert(mono[i / channels]);
}
},
move |err| {
tracing::warn!(%err, "playback stream error");
},
None,
)
.map_err(|err| AudioError::StreamBuild(err.to_string()).into())
}
#[cfg(test)]
mod tests {
use super::*;
fn test_pump(device_rate: u32, max_period_frames: usize) -> (HeapProd<f32>, OutputPump) {
let (prod, cons) = HeapRb::<f32>::new(PLAYBACK_RING_CAPACITY).split();
let pump = OutputPump::new(
cons,
device_rate,
max_period_frames,
Arc::new(AtomicU64::new(0)),
Arc::new(AtomicBool::new(false)),
Arc::new(AtomicU64::new(0)),
);
(prod, pump)
}
#[test]
fn render_within_period_hint_does_not_grow_scratch() {
let (mut prod, mut pump) = test_pump(48_000, 512);
let ring_cap = pump.ring_scratch.capacity();
let block_cap = pump.block_scratch.capacity();
let pending_cap = pump.pending.capacity();
prod.push_slice(&vec![0.5f32; PLAYBACK_RING_CAPACITY]);
let mut period = [0.0f32; 512];
for _ in 0..16 {
pump.render(&mut period);
assert_eq!(pump.ring_scratch.capacity(), ring_cap, "ring scratch grew");
assert_eq!(
pump.block_scratch.capacity(),
block_cap,
"block scratch grew"
);
assert_eq!(pump.pending.capacity(), pending_cap, "pending grew");
assert!(period.iter().all(|&s| s == 0.5));
}
}
#[test]
fn heavy_downsample_big_period_does_not_grow_scratch() {
let (mut prod, mut pump) = test_pump(8_000, 4096);
let ring_cap = pump.ring_scratch.capacity();
assert!(ring_cap > 8192, "hint-based sizing must exceed 8192 here");
let block_cap = pump.block_scratch.capacity();
prod.push_slice(&vec![0.5f32; PLAYBACK_RING_CAPACITY]);
let mut period = [1.0f32; 4096];
pump.render(&mut period);
pump.render(&mut period);
assert_eq!(pump.ring_scratch.capacity(), ring_cap);
assert_eq!(pump.block_scratch.capacity(), block_cap);
assert!(period.iter().all(|&s| s == 0.5));
}
fn drain_fixture(
period_ms: u64,
) -> (
PlaybackHandle,
DrainCounters,
Arc<AtomicBool>,
Option<std::thread::JoinHandle<()>>,
std::thread::JoinHandle<()>,
) {
let (prod, cons) = HeapRb::<f32>::new(PLAYBACK_RING_CAPACITY).split();
let (clips_tx, clips_rx) = mpsc::channel::<TtsClip>(CLIPS_CHANNEL_CAPACITY);
let flush_epoch = Arc::new(AtomicU64::new(0));
let is_playing = Arc::new(AtomicBool::new(false));
let stop = Arc::new(AtomicBool::new(false));
let counters = DrainCounters::default();
let pump_thread = std::thread::spawn({
let flush_epoch = Arc::clone(&flush_epoch);
let stop = Arc::clone(&stop);
let counters = counters.clone();
move || {
clip_pump_loop(
clips_rx,
prod,
&flush_epoch,
&stop,
&counters.clips_pushed,
&counters.samples_pushed,
);
}
});
let device_thread = (period_ms != u64::MAX).then(|| {
std::thread::spawn({
let flush_epoch = Arc::clone(&flush_epoch);
let is_playing = Arc::clone(&is_playing);
let stop = Arc::clone(&stop);
let consumed = Arc::clone(&counters.samples_consumed);
move || {
let mut pump = OutputPump::new(
cons,
CLIP_SAMPLE_RATE,
480,
flush_epoch,
is_playing,
consumed,
);
let mut period = [0.0f32; 480];
while !stop.load(Ordering::Relaxed) {
pump.render(&mut period);
if period_ms > 0 {
std::thread::sleep(Duration::from_millis(period_ms));
}
}
}
})
});
let handle = PlaybackHandle {
clips_tx,
flush_epoch,
is_playing,
sample_rate: CLIP_SAMPLE_RATE,
stop: Arc::clone(&stop),
clips_queued: Arc::clone(&counters.clips_queued),
clips_pushed: Arc::clone(&counters.clips_pushed),
samples_pushed: Arc::clone(&counters.samples_pushed),
samples_consumed: Arc::clone(&counters.samples_consumed),
};
(handle, counters, stop, device_thread, pump_thread)
}
#[derive(Clone, Default)]
struct DrainCounters {
clips_queued: Arc<AtomicU64>,
clips_pushed: Arc<AtomicU64>,
samples_pushed: Arc<AtomicU64>,
samples_consumed: Arc<AtomicU64>,
}
fn clip(samples: usize) -> TtsClip {
TtsClip {
samples: vec![0.5f32; samples],
sample_rate: CLIP_SAMPLE_RATE,
}
}
#[tokio::test]
async fn wait_drained_returns_after_all_samples_consumed() {
let (handle, counters, stop, device, pump) = drain_fixture(20);
let total = 3 * 6_000u64; let started = std::time::Instant::now();
for _ in 0..3 {
handle.queue_clip(clip(6_000)).await.expect("queue");
}
let drained = tokio::time::timeout(Duration::from_secs(10), handle.wait_drained())
.await
.expect("graceful drain must not hang");
assert!(drained, "drained (not aborted)");
assert!(
counters.samples_consumed.load(Ordering::SeqCst) >= total,
"all queued samples consumed before finish returned: {}",
counters.samples_consumed.load(Ordering::SeqCst)
);
assert_eq!(counters.clips_pushed.load(Ordering::SeqCst), 3);
assert!(
started.elapsed() >= Duration::from_millis(500),
"wait blocked for roughly the playback duration: {:?}",
started.elapsed()
);
stop.store(true, Ordering::Relaxed);
device.expect("device thread").join().expect("device join");
pump.join().expect("pump join");
}
#[tokio::test]
async fn wait_drained_aborts_promptly_on_flush() {
let (handle, counters, stop, device, pump) = drain_fixture(u64::MAX);
handle.queue_clip(clip(6_000)).await.expect("queue");
for _ in 0..100 {
if counters.clips_pushed.load(Ordering::SeqCst) == 1 {
break;
}
tokio::time::sleep(Duration::from_millis(5)).await;
}
assert_eq!(counters.clips_pushed.load(Ordering::SeqCst), 1);
let h2 = handle.clone();
let waiter = tokio::spawn(async move { h2.wait_drained().await });
tokio::time::sleep(Duration::from_millis(50)).await; assert!(
!waiter.is_finished(),
"sanity: wedged device parks the wait"
);
handle.flush();
let drained = tokio::time::timeout(Duration::from_secs(2), waiter)
.await
.expect("flush must abort the wait promptly")
.expect("wait task panicked");
assert!(!drained, "aborted by the flush");
stop.store(true, Ordering::Relaxed);
pump.join().expect("pump join");
assert!(device.is_none());
}
#[tokio::test]
async fn wait_buffered_covers_queue_only() {
let (handle, counters, stop, device, pump) = drain_fixture(u64::MAX);
handle.queue_clip(clip(6_000)).await.expect("queue");
let buffered = tokio::time::timeout(Duration::from_secs(2), handle.wait_buffered())
.await
.expect("buffered wait must not hang");
assert!(buffered);
assert_eq!(counters.clips_pushed.load(Ordering::SeqCst), 1);
assert_eq!(counters.samples_consumed.load(Ordering::SeqCst), 0);
stop.store(true, Ordering::Relaxed);
pump.join().expect("pump join");
assert!(device.is_none());
}
}