#![allow(unsafe_code)]
use std::collections::VecDeque;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
use std::thread::{self, JoinHandle};
use crate::{CaptureError, Select};
use mediaway_common::{AudioFrame, Bytes, CodecKind, SampleFormat, StreamInfo};
use crate::windows_audio::wasapi_config::{
WasapiCaptureConfig as AudioCaptureConfig, WasapiProcessTreeScope as ProcessTreeScope,
WasapiSource as AudioCaptureSource,
};
use windows::Win32::Devices::FunctionDiscovery::{
PKEY_Device_FriendlyName, PKEY_DeviceInterface_FriendlyName,
};
use windows::Win32::Foundation::PROPERTYKEY;
use windows::Win32::Media::Audio::{
AUDCLNT_BUFFERFLAGS_SILENT, AUDCLNT_SHAREMODE_SHARED, AUDCLNT_STREAMFLAGS_LOOPBACK,
DEVICE_STATE_ACTIVE, EDataFlow, IAudioCaptureClient, IAudioClient, IMMDevice,
IMMDeviceEnumerator, MMDeviceEnumerator, WAVEFORMATEX, WAVEFORMATEXTENSIBLE, eCapture,
eConsole, eRender,
};
use windows::Win32::System::Com::StructuredStorage::{PropVariantClear, PropVariantToStringAlloc};
use windows::Win32::System::Com::{
CLSCTX_ALL, CLSCTX_INPROC_SERVER, COINIT_MULTITHREADED, CoCreateInstance, CoInitializeEx,
CoTaskMemFree, CoUninitialize, STGM_READ,
};
use windows::Win32::UI::Shell::PropertiesSystem::IPropertyStore;
use windows::core::GUID;
use crate::windows_audio::wasapi_process;
const WAVE_FORMAT_IEEE_FLOAT: u16 = 0x0003;
const WAVE_FORMAT_EXTENSIBLE: u16 = 0xFFFE;
const KSDATAFORMAT_SUBTYPE_IEEE_FLOAT: GUID =
GUID::from_u128(0x0000_0003_0000_0010_8000_00aa_0038_9b71);
const PCM_QUEUE_CAP: usize = 64;
struct SharedQueue {
frames: Mutex<VecDeque<AudioFrame>>,
stop: AtomicBool,
device_lost: AtomicBool,
}
pub struct WindowsWasapiCapture {
inner: Option<WasapiSession>,
}
struct WasapiSession {
stream_info: StreamInfo,
queue: Arc<SharedQueue>,
worker: Option<JoinHandle<()>>,
}
impl WindowsWasapiCapture {
pub fn open(config: &AudioCaptureConfig) -> Result<Self, CaptureError> {
if config.sample_format != SampleFormat::F32 {
return Err(CaptureError::Unsupported);
}
if config.time_base.den == 0 {
return Err(CaptureError::InvalidInput);
}
let queue = Arc::new(SharedQueue {
frames: Mutex::new(VecDeque::new()),
stop: AtomicBool::new(false),
device_lost: AtomicBool::new(false),
});
let queue_worker = Arc::clone(&queue);
let source = config.source.clone();
let time_base = config.time_base;
let (tx_info, rx_info) = std::sync::mpsc::sync_channel(1);
let worker = thread::Builder::new()
.name("mediaway-wasapi".into())
.spawn(move || {
let result = run_wasapi_worker(source, time_base, &queue_worker, &tx_info);
if let Err(e) = result {
let _ = tx_info.send(Err(e));
}
})
.map_err(|_| CaptureError::Backend)?;
let stream_info = rx_info.recv().map_err(|_| CaptureError::Backend)??;
Ok(Self {
inner: Some(WasapiSession {
stream_info,
queue,
worker: Some(worker),
}),
})
}
}
impl WindowsWasapiCapture {
#[must_use]
pub 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()
}
}
pub fn poll_frame(&mut self) -> Result<Option<AudioFrame>, CaptureError> {
let Some(session) = self.inner.as_ref() else {
return Err(CaptureError::Closed);
};
{
let mut q = session
.queue
.frames
.lock()
.map_err(|_| CaptureError::Backend)?;
if let Some(frame) = q.pop_front() {
return Ok(Some(frame));
}
}
if session.queue.device_lost.load(Ordering::Relaxed) {
return Err(CaptureError::DeviceLost);
}
Ok(None)
}
pub fn close(&mut self) -> Result<(), CaptureError> {
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 WindowsWasapiCapture {
fn drop(&mut self) {
let _ = self.close();
}
}
impl WindowsWasapiCapture {
pub fn open_microphone(
config: &crate::audio::AudioCaptureConfig,
) -> Result<Self, CaptureError> {
Self::open(&AudioCaptureConfig {
source: AudioCaptureSource::Microphone {
select: config.select.clone(),
},
time_base: config.time_base,
sample_format: config.sample_format,
})
}
}
impl crate::audio::AudioCapture for WindowsWasapiCapture {
fn stream_info(&self) -> &StreamInfo {
Self::stream_info(self)
}
fn poll_frame(&mut self) -> Result<Option<AudioFrame>, CaptureError> {
Self::poll_frame(self)
}
fn close(&mut self) -> Result<(), CaptureError> {
Self::close(self)
}
}
#[allow(clippy::redundant_pub_crate)]
pub(crate) fn closed_audio_info() -> &'static StreamInfo {
use std::sync::OnceLock;
static INFO: OnceLock<StreamInfo> = OnceLock::new();
INFO.get_or_init(|| StreamInfo::Audio {
id: 0,
codec: CodecKind::RawAudio,
time_base: mediaway_common::Rational::new(1, 48_000),
sample_rate: 0,
channels: 0,
extra_data: Bytes::new(),
})
}
fn run_wasapi_worker(
source: AudioCaptureSource,
time_base: mediaway_common::Rational,
queue: &SharedQueue,
tx_info: &std::sync::mpsc::SyncSender<Result<StreamInfo, CaptureError>>,
) -> Result<(), CaptureError> {
let hr = unsafe { CoInitializeEx(None, COINIT_MULTITHREADED) };
if hr.is_err() {
return Err(notify_err(tx_info, CaptureError::Backend));
}
let _com = ComGuard;
let (audio_client, capture, sample_rate, channels) = open_wasapi_client(source, tx_info)?;
let info = StreamInfo::Audio {
id: 0,
codec: CodecKind::RawAudio,
time_base,
extra_data: Bytes::new(),
sample_rate,
channels,
};
let _ = tx_info.send(Ok(info));
pump_capture_loop(&audio_client, &capture, sample_rate, channels, queue);
let _ = unsafe { audio_client.Stop() };
Ok(())
}
fn open_wasapi_client(
source: AudioCaptureSource,
tx_info: &std::sync::mpsc::SyncSender<Result<StreamInfo, CaptureError>>,
) -> Result<(IAudioClient, IAudioCaptureClient, u32, u16), CaptureError> {
let (data_flow, loopback, select) = match source {
AudioCaptureSource::Microphone { select } => (eCapture, false, select),
AudioCaptureSource::Loopback { select } => (eRender, true, select),
AudioCaptureSource::ProcessLoopback {
process_id,
tree_scope,
} => {
let include_tree = matches!(tree_scope, ProcessTreeScope::IncludeChildren);
return wasapi_process::open_process_loopback_client(process_id, include_tree)
.map_err(|e| notify_err(tx_info, e));
}
};
let enumerator: IMMDeviceEnumerator =
unsafe { CoCreateInstance(&MMDeviceEnumerator, None, CLSCTX_INPROC_SERVER) }
.map_err(|_| CaptureError::Backend)?;
let device =
resolve_endpoint(&enumerator, data_flow, &select).map_err(|e| notify_err(tx_info, e))?;
let audio_client: IAudioClient = unsafe { device.Activate::<IAudioClient>(CLSCTX_ALL, None) }
.map_err(|_| CaptureError::Backend)?;
let format_ptr = unsafe { audio_client.GetMixFormat() }.map_err(|_| CaptureError::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, CaptureError::Unsupported));
}
let stream_flags = if loopback {
AUDCLNT_STREAMFLAGS_LOOPBACK
} else {
0
};
let init = unsafe {
audio_client.Initialize(
AUDCLNT_SHAREMODE_SHARED,
stream_flags,
10_000_000,
0,
format_ptr,
None,
)
};
unsafe { CoTaskMemFree(Some(format_ptr.cast())) };
init.map_err(|_| CaptureError::Backend)?;
let capture: IAudioCaptureClient =
unsafe { audio_client.GetService() }.map_err(|_| CaptureError::Backend)?;
unsafe { audio_client.Start() }.map_err(|_| CaptureError::Backend)?;
Ok((audio_client, capture, sample_rate, channels))
}
fn pump_capture_loop(
_audio_client: &IAudioClient,
capture: &IAudioCaptureClient,
sample_rate: u32,
channels: u16,
queue: &SharedQueue,
) {
let mut pts: i64 = 0;
while !queue.stop.load(Ordering::Relaxed) {
let Ok(packet_length) = (unsafe { capture.GetNextPacketSize() }) else {
queue.device_lost.store(true, Ordering::SeqCst);
break;
};
if packet_length == 0 {
thread::sleep(std::time::Duration::from_millis(5));
continue;
}
let mut data_ptr: *mut u8 = std::ptr::null_mut();
let mut num_frames = 0u32;
let mut flags = 0u32;
if unsafe {
capture.GetBuffer(
&raw mut data_ptr,
&raw mut num_frames,
&raw mut flags,
None,
None,
)
}
.is_err()
{
queue.device_lost.store(true, Ordering::SeqCst);
break;
}
let silent = (flags & AUDCLNT_BUFFERFLAGS_SILENT.0 as u32) != 0;
if num_frames > 0 && !data_ptr.is_null() && !silent {
let samples = num_frames as usize * channels as usize;
let bytes = samples * 4;
let pcm = unsafe { copy_pcm_buffer(data_ptr, bytes) };
let frame = AudioFrame {
pts,
duration: u64::from(num_frames),
sample_rate,
channels,
format: SampleFormat::F32,
data: Bytes::from(pcm),
};
pts = pts.saturating_add(i64::from(num_frames));
if let Ok(mut q) = queue.frames.lock() {
if q.len() >= PCM_QUEUE_CAP {
let _ = q.pop_front();
}
q.push_back(frame);
}
}
let _ = unsafe { capture.ReleaseBuffer(num_frames) };
}
}
unsafe fn copy_pcm_buffer(src: *const u8, len: usize) -> Vec<u8> {
let mut pcm: Vec<u8> = Vec::with_capacity(len);
unsafe {
std::ptr::copy_nonoverlapping(src, pcm.as_mut_ptr(), len);
pcm.set_len(len);
}
pcm
}
fn notify_err(
tx: &std::sync::mpsc::SyncSender<Result<StreamInfo, CaptureError>>,
err: CaptureError,
) -> CaptureError {
let _ = tx.send(Err(err.clone()));
err
}
#[allow(clippy::redundant_pub_crate)]
pub(crate) unsafe fn read_float_mix(format_ptr: *mut WAVEFORMATEX) -> (u32, u16, bool) {
if format_ptr.is_null() {
return (0, 0, false);
}
let sample_rate = unsafe { std::ptr::addr_of!((*format_ptr).nSamplesPerSec).read_unaligned() };
let channels = unsafe { std::ptr::addr_of!((*format_ptr).nChannels).read_unaligned() };
let tag = unsafe { std::ptr::addr_of!((*format_ptr).wFormatTag).read_unaligned() };
let valid = if tag == WAVE_FORMAT_IEEE_FLOAT {
true
} else if tag == WAVE_FORMAT_EXTENSIBLE {
let ext = format_ptr.cast::<WAVEFORMATEXTENSIBLE>();
let sub = unsafe { std::ptr::addr_of!((*ext).SubFormat).read_unaligned() };
sub == KSDATAFORMAT_SUBTYPE_IEEE_FLOAT
} else {
false
};
(sample_rate, channels, valid)
}
pub fn resolve_endpoint(
enumerator: &IMMDeviceEnumerator,
data_flow: EDataFlow,
select: &Select,
) -> Result<IMMDevice, CaptureError> {
match select {
Select::Default => {
unsafe { enumerator.GetDefaultAudioEndpoint(data_flow, eConsole) }
.map_err(|_| CaptureError::AccessDenied)
}
Select::Id(id) => {
let endpoint_id_str = id
.as_wasapi_endpoint_id()
.ok_or(CaptureError::Unsupported)?;
find_endpoint(enumerator, data_flow, |candidate| {
endpoint_id(candidate).as_deref() == Some(endpoint_id_str)
})
}
Select::NameContains(needle) => {
let needle = needle.to_lowercase();
find_endpoint(enumerator, data_flow, |candidate| {
endpoint_friendly_name(candidate)
.is_some_and(|name| name.to_lowercase().contains(&needle))
})
}
}
}
fn find_endpoint(
enumerator: &IMMDeviceEnumerator,
data_flow: EDataFlow,
mut matches: impl FnMut(&IMMDevice) -> bool,
) -> Result<IMMDevice, CaptureError> {
let collection = unsafe { enumerator.EnumAudioEndpoints(data_flow, DEVICE_STATE_ACTIVE) }
.map_err(|_| CaptureError::Backend)?;
let count = unsafe { collection.GetCount() }.map_err(|_| CaptureError::Backend)?;
for index in 0..count {
let Ok(device) = (unsafe { collection.Item(index) }) else {
continue;
};
if matches(&device) {
return Ok(device);
}
}
Err(CaptureError::InvalidInput)
}
pub fn endpoint_id(device: &IMMDevice) -> Option<String> {
let raw = unsafe { device.GetId() }.ok()?;
if raw.is_null() {
return None;
}
let id = unsafe { raw.to_string() }.ok();
unsafe { CoTaskMemFree(Some(raw.0.cast())) };
id
}
pub fn endpoint_friendly_name(device: &IMMDevice) -> Option<String> {
let store = unsafe { device.OpenPropertyStore(STGM_READ) }.ok()?;
let endpoint_name = property_string(&store, PKEY_Device_FriendlyName);
let interface_name = property_string(&store, PKEY_DeviceInterface_FriendlyName);
combine_endpoint_and_interface_names(endpoint_name, interface_name)
}
fn combine_endpoint_and_interface_names(
endpoint_name: Option<String>,
interface_name: Option<String>,
) -> Option<String> {
match (endpoint_name, interface_name) {
(Some(endpoint_name), Some(interface_name))
if !endpoint_name
.to_lowercase()
.contains(&interface_name.to_lowercase()) =>
{
Some(format!("{endpoint_name} ({interface_name})"))
}
(Some(endpoint_name), _) => Some(endpoint_name),
(None, Some(interface_name)) => Some(interface_name),
(None, None) => None,
}
}
fn property_string(store: &IPropertyStore, key: PROPERTYKEY) -> Option<String> {
let mut value = unsafe { store.GetValue(&raw const key) }.ok()?;
let raw = unsafe { PropVariantToStringAlloc(&raw const value) }.ok();
unsafe {
let _ = PropVariantClear(&raw mut value);
}
let raw = raw?;
if raw.is_null() {
return None;
}
let name = unsafe { raw.to_string() }.ok();
unsafe { CoTaskMemFree(Some(raw.0.cast())) };
name
}
pub struct ComGuard;
impl Drop for ComGuard {
fn drop(&mut self) {
unsafe {
CoUninitialize();
}
}
}
#[cfg(test)]
#[path = "wasapi_tests.rs"]
mod tests;