use std::sync::mpsc::{Receiver, RecvTimeoutError, Sender, SyncSender, sync_channel};
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use cpal::traits::{DeviceTrait, StreamTrait};
use super::mixer::{self, Mixer};
use super::sink::{Registration, Sink};
use crate::Error;
const RETRY_MIN: Duration = Duration::from_millis(500);
const RETRY_MAX: Duration = Duration::from_secs(4);
const UNDERRUN_LIMIT: u32 = 20;
const ERROR_LIMIT: u32 = 3;
const ERROR_WINDOW: Duration = Duration::from_secs(5);
const COMMAND_QUEUE: usize = 2 * mixer::MAX_SINKS;
const SYNC_ATTEMPTS: u32 = 8;
const SYNC_DELAY: Duration = Duration::from_millis(4);
const SCRATCH_FRAMES: usize = 2048;
#[derive(Default)]
pub(super) struct Shared {
state: Mutex<State>,
}
#[derive(Default)]
struct State {
rate: u32,
mixer: Option<SyncSender<mixer::Command>>,
sinks: Vec<Registration>,
detaching: Vec<u64>,
next_id: u64,
}
impl Shared {
pub(super) fn add<F>(&self, build: F) -> Result<Sink, Error>
where
F: FnOnce(u64, u32) -> Result<(Sink, Registration), Error>,
{
let mut state = self.state.lock().unwrap();
if state.sinks.len() >= mixer::MAX_SINKS {
return Err(Error::Unsupported(format!(
"at most {} playback sinks per device",
mixer::MAX_SINKS
)));
}
let rate = if state.rate == 0 { 48_000 } else { state.rate };
let (sink, mut registration) = build(state.next_id, rate)?;
state.next_id += 1;
if let Some(mixer) = &state.mixer {
registration.attach(mixer);
}
state.sinks.push(registration);
Ok(sink)
}
pub(super) fn remove(&self, id: u64) {
let mut state = self.state.lock().unwrap();
state.sinks.retain(|s| s.id != id);
let Some(mixer) = &state.mixer else { return };
if mixer.try_send(mixer::Command::Remove { id }).is_err() {
state.detaching.push(id);
}
}
pub(super) fn sync(&self) -> bool {
let mut state = self.state.lock().unwrap();
let Some(mixer) = state.mixer.clone() else {
state.detaching.clear();
return true;
};
state
.detaching
.retain(|id| mixer.try_send(mixer::Command::Remove { id: *id }).is_err());
for sink in &mut state.sinks {
sink.attach(&mixer);
}
state.detaching.is_empty() && state.sinks.iter().all(|s| s.attached())
}
fn rebind(&self, rate: u32, mixer: SyncSender<mixer::Command>) {
let mut state = self.state.lock().unwrap();
for sink in &mut state.sinks {
sink.rebuild(rate);
sink.attach(&mixer);
}
state.rate = rate;
state.mixer = Some(mixer);
state.detaching.clear();
}
fn unbind(&self) {
self.state.lock().unwrap().mixer = None;
}
}
pub(super) enum Command {
Switch {
device: Option<String>,
reply: tokio::sync::oneshot::Sender<Result<(), Error>>,
},
Failed { generation: u64, error: cpal::Error },
Sync,
Shutdown,
}
pub(super) fn run(
commands: Receiver<Command>,
failures: Sender<Command>,
shared: Arc<Shared>,
device: Option<String>,
opened: tokio::sync::oneshot::Sender<Result<(), Error>>,
) {
let mut driver = Driver {
shared,
failures,
device,
stream: None,
retired: None,
generation: 0,
retry: RETRY_MIN,
retry_at: None,
underruns: 0,
unclassified: 0,
window: Instant::now(),
};
let first = driver.start();
let started = first.is_ok();
if opened.send(first).is_err() || !started {
return;
}
loop {
let command = match driver.retry_at {
Some(at) => driver.commands_until(&commands, at),
None => commands.recv().map_err(|_| Timeout::Disconnected),
};
match command {
Ok(Command::Switch { device, reply }) => {
driver.device = device;
let _ = reply.send(driver.restart());
}
Ok(Command::Failed { generation, error }) => {
if driver.should_restart(generation, &error) {
let _ = driver.restart();
}
}
Ok(Command::Sync) => driver.sync(),
Ok(Command::Shutdown) => break,
Err(Timeout::Elapsed) => {
if driver.restart().is_ok() {
tracing::info!("audio output recovered");
}
}
Err(Timeout::Disconnected) => break,
}
}
}
enum Timeout {
Elapsed,
Disconnected,
}
struct Driver {
shared: Arc<Shared>,
failures: Sender<Command>,
device: Option<String>,
stream: Option<cpal::Stream>,
retired: Option<Receiver<mixer::Entry>>,
generation: u64,
retry: Duration,
retry_at: Option<Instant>,
underruns: u32,
unclassified: u32,
window: Instant,
}
impl Driver {
fn commands_until(&self, commands: &Receiver<Command>, at: Instant) -> Result<Command, Timeout> {
match commands.recv_timeout(at.saturating_duration_since(Instant::now())) {
Ok(command) => Ok(command),
Err(RecvTimeoutError::Timeout) => Err(Timeout::Elapsed),
Err(RecvTimeoutError::Disconnected) => Err(Timeout::Disconnected),
}
}
fn start(&mut self) -> Result<(), Error> {
let device = super::device::open(self.device.as_deref())?;
let supported = super::device::negotiate(&device)?;
let format = supported.sample_format();
let config: cpal::StreamConfig = supported.into();
let rate = config.sample_rate;
let channels = config.channels as usize;
if rate == 0 || channels == 0 {
return Err(Error::Playback(format!(
"output device negotiated an empty format ({rate} Hz, {channels} channels)"
)));
}
let (tx, rx) = sync_channel(COMMAND_QUEUE);
let (retired_tx, retired_rx) = sync_channel(mixer::MAX_SINKS);
let mixer = Mixer::new(rx, retired_tx, rate, channels);
self.generation += 1;
let stream = self.build(&device, config, format, mixer)?;
stream
.play()
.map_err(|err| Error::Playback(format!("cannot start output stream: {err}")))?;
self.shared.rebind(rate, tx);
self.stream = Some(stream);
self.retired = Some(retired_rx);
self.retry = RETRY_MIN;
tracing::info!(rate, channels, ?format, "opened audio output");
Ok(())
}
fn build(
&self,
device: &cpal::Device,
config: cpal::StreamConfig,
format: cpal::SampleFormat,
mixer: Mixer,
) -> Result<cpal::Stream, Error> {
match format {
cpal::SampleFormat::F32 => self.build_as::<f32>(device, config, mixer),
cpal::SampleFormat::I16 => self.build_as::<i16>(device, config, mixer),
cpal::SampleFormat::U16 => self.build_as::<u16>(device, config, mixer),
cpal::SampleFormat::I32 => self.build_as::<i32>(device, config, mixer),
other => Err(Error::Unsupported(format!("output sample format {other:?}"))),
}
}
fn build_as<T>(
&self,
device: &cpal::Device,
config: cpal::StreamConfig,
mut mixer: Mixer,
) -> Result<cpal::Stream, Error>
where
T: cpal::SizedSample + cpal::FromSample<f32>,
{
let failures = self.failures.clone();
let generation = self.generation;
let mut scratch = vec![0.0f32; SCRATCH_FRAMES * config.channels as usize];
device
.build_output_stream::<T, _, _>(
config,
move |data, _| {
for chunk in data.chunks_mut(scratch.len()) {
let scratch = &mut scratch[..chunk.len()];
mixer.fill(scratch);
for (out, sample) in chunk.iter_mut().zip(scratch.iter()) {
*out = T::from_sample(*sample);
}
}
},
move |error| {
let _ = failures.send(Command::Failed { generation, error });
},
None,
)
.map_err(|err| Error::Playback(format!("cannot open output stream: {err}")))
}
fn restart(&mut self) -> Result<(), Error> {
self.stop();
let result = self.start();
self.retry_at = match &result {
Ok(()) => None,
Err(err) => {
tracing::debug!(%err, "audio output unavailable");
Some(self.schedule())
}
};
result
}
fn stop(&mut self) {
self.shared.unbind();
self.stream = None;
self.retired = None;
}
fn sync(&mut self) {
if let Some(retired) = &self.retired {
while retired.try_recv().is_ok() {}
}
for _ in 0..SYNC_ATTEMPTS {
if self.shared.sync() {
return;
}
std::thread::sleep(SYNC_DELAY);
}
tracing::warn!("audio output is not keeping up with sink changes");
}
fn schedule(&mut self) -> Instant {
let at = Instant::now() + self.retry;
self.retry = (self.retry * 2).min(RETRY_MAX);
at
}
fn should_restart(&mut self, generation: u64, error: &cpal::Error) -> bool {
if generation != self.generation {
tracing::debug!(%error, generation, "ignoring an error from a replaced audio output");
return false;
}
self.fatal(error)
}
fn fatal(&mut self, err: &cpal::Error) -> bool {
self.roll_window();
match err.kind() {
cpal::ErrorKind::DeviceNotAvailable | cpal::ErrorKind::StreamInvalidated => {
tracing::warn!(%err, "audio output lost");
true
}
cpal::ErrorKind::DeviceChanged | cpal::ErrorKind::RealtimeDenied => {
tracing::debug!(%err, "audio output changed underneath us");
false
}
cpal::ErrorKind::Xrun => {
self.underruns += 1;
let restart = self.underruns > UNDERRUN_LIMIT;
if restart {
self.underruns = 0;
tracing::warn!("restarting audio output after repeated underruns");
}
restart
}
_ => {
tracing::warn!(%err, "audio output error");
self.unclassified += 1;
let restart = self.unclassified >= ERROR_LIMIT;
if restart {
self.unclassified = 0;
tracing::warn!("restarting audio output after repeated unclassified errors");
}
restart
}
}
}
fn roll_window(&mut self) {
if self.window.elapsed() > ERROR_WINDOW {
self.underruns = 0;
self.unclassified = 0;
self.window = Instant::now();
}
}
}
#[cfg(test)]
mod tests {
use std::sync::mpsc::channel;
use super::*;
use crate::playback::sink::{self, Input};
struct Wired {
shared: Arc<Shared>,
handle: Arc<super::super::Handle>,
mixer: Receiver<mixer::Command>,
driver: Receiver<Command>,
}
fn wired(depth: usize) -> Wired {
let shared = Arc::new(Shared::default());
let (commands, driver) = channel();
let handle = Arc::new(super::super::Handle { commands });
let (tx, mixer) = sync_channel(depth);
shared.rebind(48_000, tx);
Wired {
shared,
handle,
mixer,
driver,
}
}
fn add(shared: &Arc<Shared>, handle: &Arc<super::super::Handle>) -> Result<Sink, Error> {
shared.add(|id, rate| sink::new(id, rate, Input::default(), shared.clone(), handle.clone()))
}
fn syncs(driver: &Receiver<Command>) -> usize {
std::iter::from_fn(|| driver.try_recv().ok())
.filter(|c| matches!(c, Command::Sync))
.count()
}
#[test]
fn registrations_survive_a_full_mixer_queue() {
let depth = 4;
let w = wired(depth);
let sinks: Vec<_> = (0..depth + 1).map(|_| add(&w.shared, &w.handle).unwrap()).collect();
assert_eq!(sinks.len(), depth + 1);
assert!(!w.shared.sync(), "expected a sink to be waiting on the queue");
while w.mixer.try_recv().is_ok() {}
assert!(w.shared.sync(), "the waiting sink was never re-sent");
let state = w.shared.state.lock().unwrap();
assert!(state.sinks.iter().all(|s| s.attached()), "a sink is still unattached");
}
#[test]
fn adding_and_dropping_a_sink_wakes_the_driver() {
let w = wired(8);
let engine = super::super::Engine {
shared: w.shared.clone(),
handle: w.handle.clone(),
};
let sink = engine.sink(Input::default()).unwrap();
assert_eq!(syncs(&w.driver), 1, "adding a sink did not wake the driver");
drop(sink);
assert_eq!(syncs(&w.driver), 1, "dropping a sink did not wake the driver");
}
#[test]
fn removals_survive_a_full_mixer_queue() {
let w = wired(1);
let sink = add(&w.shared, &w.handle).unwrap();
let id = w.shared.state.lock().unwrap().sinks[0].id;
drop(sink);
assert_eq!(w.shared.state.lock().unwrap().detaching, vec![id]);
while w.mixer.try_recv().is_ok() {}
assert!(w.shared.sync());
assert!(
w.shared.state.lock().unwrap().detaching.is_empty(),
"the removal was lost"
);
}
#[test]
fn refuses_more_sinks_than_the_mixer_can_hold() {
let w = wired(4 * mixer::MAX_SINKS);
let sinks: Vec<_> = (0..mixer::MAX_SINKS)
.map(|_| add(&w.shared, &w.handle).unwrap())
.collect();
assert!(matches!(add(&w.shared, &w.handle), Err(Error::Unsupported(_))));
drop(sinks.into_iter().next_back());
add(&w.shared, &w.handle).expect("a slot freed by the dropped sink");
}
#[test]
fn ignores_errors_from_a_replaced_stream() {
let w = wired(8);
let (failures, _requests) = channel();
let mut driver = Driver {
shared: w.shared,
failures,
device: None,
stream: None,
retired: None,
generation: 7,
retry: RETRY_MIN,
retry_at: None,
underruns: 0,
unclassified: 0,
window: Instant::now(),
};
let lost = cpal::Error::new(cpal::ErrorKind::DeviceNotAvailable);
assert!(!driver.should_restart(6, &lost), "acted on a retired stream's error");
assert!(driver.should_restart(7, &lost), "ignored the live stream's error");
}
}