#![allow(unsafe_code)]
#![allow(
clippy::redundant_pub_crate,
reason = "see wasapi.rs's identical allow"
)]
use std::collections::VecDeque;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use std::thread::{self, JoinHandle};
use std::time::Duration;
use crate::Select;
use crate::audio::{AudioPlayback, AudioPlaybackConfig, PlaybackError};
use mediaway_common::{AudioFrame, Bytes, CodecKind, Rational, SampleFormat, StreamInfo};
use windows::Win32::Media::Audio::{
AUDCLNT_BUFFERFLAGS_SILENT, AUDCLNT_SHAREMODE_SHARED, IAudioClient, IAudioRenderClient,
IMMDeviceEnumerator, MMDeviceEnumerator, eRender,
};
use windows::Win32::System::Com::{
CLSCTX_ALL, CLSCTX_INPROC_SERVER, COINIT_MULTITHREADED, CoCreateInstance, CoInitializeEx,
CoTaskMemFree,
};
use crate::windows_audio::wasapi::{ComGuard, closed_audio_info, read_float_mix, resolve_endpoint};
const PLAYBACK_QUEUE_CAP: usize = 64;
const BYTES_PER_SAMPLE_F32: usize = 4;
struct PlaybackQueueState {
frames: VecDeque<AudioFrame>,
cursor: usize,
}
struct PlaybackSharedQueue {
state: Mutex<PlaybackQueueState>,
stop: AtomicBool,
underrun_count: AtomicU64,
device_lost: AtomicBool,
}
pub struct WindowsWasapiPlayback {
inner: Option<PlaybackSession>,
}
struct PlaybackSession {
stream_info: StreamInfo,
queue: Arc<PlaybackSharedQueue>,
worker: Option<JoinHandle<()>>,
}
impl WindowsWasapiPlayback {
pub fn open(config: &AudioPlaybackConfig) -> Result<Self, PlaybackError> {
if config.sample_format != SampleFormat::F32 {
return Err(PlaybackError::Unsupported);
}
let queue = Arc::new(PlaybackSharedQueue {
state: Mutex::new(PlaybackQueueState {
frames: VecDeque::new(),
cursor: 0,
}),
stop: AtomicBool::new(false),
underrun_count: AtomicU64::new(0),
device_lost: AtomicBool::new(false),
});
let queue_worker = Arc::clone(&queue);
let select = config.select.clone();
let (tx_info, rx_info) = std::sync::mpsc::sync_channel(1);
let worker = thread::Builder::new()
.name("mediaway-wasapi-playback".into())
.spawn(move || {
let result = run_wasapi_playback_worker(&select, &queue_worker, &tx_info);
if let Err(e) = result {
let _ = tx_info.send(Err(e));
}
})
.map_err(|_| PlaybackError::Backend)?;
let stream_info = rx_info.recv().map_err(|_| PlaybackError::Backend)??;
Ok(Self {
inner: Some(PlaybackSession {
stream_info,
queue,
worker: Some(worker),
}),
})
}
}
impl AudioPlayback for WindowsWasapiPlayback {
fn stream_info(&self) -> &StreamInfo {
#[allow(
clippy::option_if_let_else,
reason = "map_or_else forces 'static vs 'self lifetime clash"
)]
if let Some(s) = self.inner.as_ref() {
&s.stream_info
} else {
closed_audio_info()
}
}
fn write_frame(&mut self, frame: AudioFrame) -> Result<(), PlaybackError> {
let Some(session) = self.inner.as_ref() else {
return Err(PlaybackError::Closed);
};
if session.queue.device_lost.load(Ordering::Relaxed) {
return Err(PlaybackError::DeviceLost);
}
let expected_rate = session.stream_info.sample_rate().unwrap_or(0);
let expected_channels = session.stream_info.channels().unwrap_or(0);
if frame.format != SampleFormat::F32
|| frame.sample_rate != expected_rate
|| frame.channels != expected_channels
{
return Err(PlaybackError::InvalidInput);
}
let Ok(mut state) = session.queue.state.lock() else {
return Err(PlaybackError::Backend);
};
if state.frames.len() >= PLAYBACK_QUEUE_CAP {
drop(state);
return Err(PlaybackError::QueueFull(frame));
}
state.frames.push_back(frame);
Ok(())
}
fn underrun_count(&self) -> u64 {
self.inner
.as_ref()
.map_or(0, |s| s.queue.underrun_count.load(Ordering::Relaxed))
}
fn close(&mut self) -> Result<(), PlaybackError> {
let Some(mut session) = self.inner.take() else {
return Ok(());
};
session.queue.stop.store(true, Ordering::SeqCst);
if let Some(h) = session.worker.take() {
let _ = h.join();
}
Ok(())
}
}
impl Drop for WindowsWasapiPlayback {
fn drop(&mut self) {
let _ = self.close();
}
}
fn run_wasapi_playback_worker(
select: &Select,
queue: &PlaybackSharedQueue,
tx_info: &std::sync::mpsc::SyncSender<Result<StreamInfo, PlaybackError>>,
) -> Result<(), PlaybackError> {
let hr = unsafe { CoInitializeEx(None, COINIT_MULTITHREADED) };
if hr.is_err() {
return Err(notify_err(tx_info, PlaybackError::Backend));
}
let _com = ComGuard;
let (audio_client, render, sample_rate, channels, buffer_frame_count) =
open_wasapi_render_client(select, tx_info)?;
let info = StreamInfo::Audio {
id: 0,
codec: CodecKind::RawAudio,
time_base: Rational::new(1, sample_rate.max(1)),
extra_data: Bytes::new(),
sample_rate,
channels,
};
let _ = tx_info.send(Ok(info));
pump_playback_loop(&audio_client, &render, channels, buffer_frame_count, queue);
let _ = unsafe { audio_client.Stop() };
Ok(())
}
fn open_wasapi_render_client(
select: &Select,
tx_info: &std::sync::mpsc::SyncSender<Result<StreamInfo, PlaybackError>>,
) -> Result<(IAudioClient, IAudioRenderClient, u32, u16, u32), PlaybackError> {
let enumerator: IMMDeviceEnumerator =
unsafe { CoCreateInstance(&MMDeviceEnumerator, None, CLSCTX_INPROC_SERVER) }
.map_err(|_| PlaybackError::Backend)?;
let device = resolve_endpoint(&enumerator, eRender, select).map_err(|e| {
notify_err(
tx_info,
match e {
crate::CaptureError::AccessDenied => PlaybackError::AccessDenied,
crate::CaptureError::Unsupported | crate::CaptureError::InvalidInput => {
PlaybackError::Unsupported
}
_ => PlaybackError::Backend,
},
)
})?;
let audio_client: IAudioClient = unsafe { device.Activate::<IAudioClient>(CLSCTX_ALL, None) }
.map_err(|_| PlaybackError::Backend)?;
let format_ptr = unsafe { audio_client.GetMixFormat() }.map_err(|_| PlaybackError::Backend)?;
let (sample_rate, channels, valid) = unsafe { read_float_mix(format_ptr) };
if !valid {
unsafe { CoTaskMemFree(Some(format_ptr.cast())) };
return Err(notify_err(tx_info, PlaybackError::Unsupported));
}
let init = unsafe {
audio_client.Initialize(AUDCLNT_SHAREMODE_SHARED, 0, 10_000_000, 0, format_ptr, None)
};
unsafe { CoTaskMemFree(Some(format_ptr.cast())) };
init.map_err(|_| PlaybackError::Backend)?;
let buffer_frame_count =
unsafe { audio_client.GetBufferSize() }.map_err(|_| PlaybackError::Backend)?;
let render: IAudioRenderClient =
unsafe { audio_client.GetService() }.map_err(|_| PlaybackError::Backend)?;
unsafe { audio_client.Start() }.map_err(|_| PlaybackError::Backend)?;
Ok((
audio_client,
render,
sample_rate,
channels,
buffer_frame_count,
))
}
fn pump_playback_loop(
audio_client: &IAudioClient,
render: &IAudioRenderClient,
channels: u16,
buffer_frame_count: u32,
queue: &PlaybackSharedQueue,
) {
let frame_size = usize::from(channels) * BYTES_PER_SAMPLE_F32;
while !queue.stop.load(Ordering::Relaxed) {
let Ok(padding) = (unsafe { audio_client.GetCurrentPadding() }) else {
queue.device_lost.store(true, Ordering::SeqCst);
break;
};
let available_frames = buffer_frame_count.saturating_sub(padding);
if available_frames == 0 || frame_size == 0 {
thread::sleep(Duration::from_millis(5));
continue;
}
let Ok(data_ptr) = (unsafe { render.GetBuffer(available_frames) }) else {
queue.device_lost.store(true, Ordering::SeqCst);
break;
};
let need_bytes = available_frames as usize * frame_size;
let dst = unsafe { std::slice::from_raw_parts_mut(data_ptr, need_bytes) };
let written = queue.state.lock().map_or(0, |mut state| {
let PlaybackQueueState { frames, cursor } = &mut *state;
fill_from_queue(frames, cursor, dst)
});
let flags = if written == 0 {
queue.underrun_count.fetch_add(1, Ordering::Relaxed);
AUDCLNT_BUFFERFLAGS_SILENT.0 as u32
} else {
if written < need_bytes {
dst[written..].fill(0);
queue.underrun_count.fetch_add(1, Ordering::Relaxed);
}
0
};
let _ = unsafe { render.ReleaseBuffer(available_frames, flags) };
}
}
fn fill_from_queue(queue: &mut VecDeque<AudioFrame>, cursor: &mut usize, dst: &mut [u8]) -> usize {
let mut written = 0usize;
while written < dst.len() {
let Some(front) = queue.front() else {
break;
};
if *cursor >= front.data.len() {
queue.pop_front();
*cursor = 0;
continue;
}
let remaining_in_frame = front.data.len() - *cursor;
let need = dst.len() - written;
let take = remaining_in_frame.min(need);
dst[written..written + take].copy_from_slice(&front.data[*cursor..*cursor + take]);
written += take;
*cursor += take;
if *cursor >= front.data.len() {
queue.pop_front();
*cursor = 0;
}
}
written
}
fn notify_err(
tx: &std::sync::mpsc::SyncSender<Result<StreamInfo, PlaybackError>>,
err: PlaybackError,
) -> PlaybackError {
let _ = tx.send(Err(err.clone()));
err
}
#[cfg(test)]
#[path = "wasapi_playback_tests.rs"]
mod tests;