use std::panic::{catch_unwind, AssertUnwindSafe};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use flexaudio_core::backend::RawSink;
use flexaudio_core::clock::monotonic_now_ns;
use flexaudio_core::types::Error;
use windows::core::PCWSTR;
use windows::Win32::Foundation::{CloseHandle, HANDLE, WAIT_OBJECT_0};
use windows::Win32::Media::Audio::{
IAudioCaptureClient, IAudioClient, AUDCLNT_BUFFERFLAGS_SILENT, AUDCLNT_SHAREMODE_SHARED,
AUDCLNT_STREAMFLAGS_EVENTCALLBACK, AUDCLNT_STREAMFLAGS_LOOPBACK, WAVEFORMATEX,
WAVEFORMATEXTENSIBLE,
};
use windows::Win32::Media::KernelStreaming::WAVE_FORMAT_EXTENSIBLE;
use windows::Win32::Media::Multimedia::{KSDATAFORMAT_SUBTYPE_IEEE_FLOAT, WAVE_FORMAT_IEEE_FLOAT};
use windows::Win32::System::Com::{CoInitializeEx, CoUninitialize, COINIT_MULTITHREADED};
use windows::Win32::System::Threading::{CreateEventW, WaitForSingleObject};
pub(crate) fn classify_hr(code: i32) -> Option<Error> {
const E_ACCESSDENIED: i32 = 0x80070005u32 as i32;
const AUDCLNT_E_DEVICE_IN_USE: i32 = 0x8889000Au32 as i32;
const AUDCLNT_E_EXCLUSIVE_MODE_NOT_ALLOWED: i32 = 0x8889000Eu32 as i32;
const AUDCLNT_E_DEVICE_INVALIDATED: i32 = 0x88890004u32 as i32;
const E_NOTFOUND: i32 = 0x80070490u32 as i32;
match code {
E_ACCESSDENIED | AUDCLNT_E_DEVICE_IN_USE | AUDCLNT_E_EXCLUSIVE_MODE_NOT_ALLOWED => {
Some(Error::PermissionDenied)
}
AUDCLNT_E_DEVICE_INVALIDATED | E_NOTFOUND => Some(Error::DeviceNotFound),
_ => None,
}
}
pub(crate) fn map_hr(ctx: &str, e: windows::core::Error) -> Error {
if let Some(mapped) = classify_hr(e.code().0) {
return mapped;
}
Error::Backend(format!("{ctx}: {e}"))
}
pub(crate) fn now_ns() -> i64 {
monotonic_now_ns()
}
pub(crate) struct ComThread {
uninit_on_drop: bool,
}
impl ComThread {
pub(crate) fn new() -> Self {
let hr = unsafe { CoInitializeEx(None, COINIT_MULTITHREADED) };
let uninit_on_drop = hr.is_ok();
Self { uninit_on_drop }
}
}
impl Drop for ComThread {
fn drop(&mut self) {
if self.uninit_on_drop {
unsafe { CoUninitialize() };
}
}
}
pub(crate) unsafe fn parse_mix_format(pwfx: *const WAVEFORMATEX) -> Result<(u32, u16), Error> {
use core::ptr::addr_of;
if pwfx.is_null() {
return Err(Error::Backend("GetMixFormat returned null format".into()));
}
let format_tag = addr_of!((*pwfx).wFormatTag).read_unaligned();
let rate = addr_of!((*pwfx).nSamplesPerSec).read_unaligned();
let channels = addr_of!((*pwfx).nChannels).read_unaligned();
let bits = addr_of!((*pwfx).wBitsPerSample).read_unaligned();
let cb_size = addr_of!((*pwfx).cbSize).read_unaligned();
let tag = format_tag as u32;
let is_float = if tag == WAVE_FORMAT_IEEE_FLOAT {
true
} else if tag == WAVE_FORMAT_EXTENSIBLE {
if (cb_size as usize) >= 22 {
let sub = addr_of!((*(pwfx as *const WAVEFORMATEXTENSIBLE)).SubFormat).read_unaligned();
sub == KSDATAFORMAT_SUBTYPE_IEEE_FLOAT
} else {
false
}
} else {
false
};
if !is_float {
return Err(Error::Backend(format!(
"unsupported mix format (not IEEE float): tag={tag} bits={bits}"
)));
}
Ok((rate, channels))
}
pub(crate) unsafe fn capture_loop(
client: &IAudioClient,
capture: &IAudioCaptureClient,
event: HANDLE,
channels: u16,
mut sink: RawSink,
stop_flag: &Arc<AtomicBool>,
) {
let channels = channels.max(1) as usize;
let mut silence: Vec<f32> = Vec::new();
if let Ok(buffer_frames) = client.GetBufferSize() {
let max_silence = (buffer_frames as usize).saturating_mul(channels);
silence.resize(max_silence, 0.0);
}
if client.Start().is_err() {
let _ = CloseHandle(event);
return;
}
while !stop_flag.load(Ordering::SeqCst) {
let _ = WaitForSingleObject(event, 100);
if stop_flag.load(Ordering::SeqCst) {
break;
}
loop {
let packet = match capture.GetNextPacketSize() {
Ok(p) => p,
Err(_e) => {
stop_flag.store(true, Ordering::SeqCst);
break;
}
};
if packet == 0 {
break;
}
let mut pdata: *mut u8 = std::ptr::null_mut();
let mut frames: u32 = 0;
let mut flags: u32 = 0;
if capture
.GetBuffer(&mut pdata, &mut frames, &mut flags, None, None)
.is_err()
{
stop_flag.store(true, Ordering::SeqCst);
break;
}
let n = frames as usize * channels;
let _ = catch_unwind(AssertUnwindSafe(|| {
if (flags & AUDCLNT_BUFFERFLAGS_SILENT.0 as u32) != 0 {
if silence.len() < n {
silence.resize(n, 0.0);
}
if n > 0 {
sink.push(&silence[..n], now_ns());
}
} else if !pdata.is_null() && n > 0 {
let slice = std::slice::from_raw_parts(pdata as *const f32, n);
sink.push(slice, now_ns());
}
}));
let _ = capture.ReleaseBuffer(frames);
}
}
let _ = client.Stop();
let _ = CloseHandle(event);
}
pub(crate) unsafe fn init_loopback_capture(
client: &IAudioClient,
pwfx: *const WAVEFORMATEX,
) -> Result<(IAudioCaptureClient, HANDLE), Error> {
client
.Initialize(
AUDCLNT_SHAREMODE_SHARED,
AUDCLNT_STREAMFLAGS_LOOPBACK | AUDCLNT_STREAMFLAGS_EVENTCALLBACK,
0, 0, pwfx,
None,
)
.map_err(|e| map_hr("IAudioClient::Initialize", e))?;
let event =
CreateEventW(None, false, false, PCWSTR::null()).map_err(|e| map_hr("CreateEventW", e))?;
if let Err(e) = client.SetEventHandle(event) {
let _ = CloseHandle(event);
return Err(map_hr("IAudioClient::SetEventHandle", e));
}
let capture: IAudioCaptureClient = match client.GetService() {
Ok(c) => c,
Err(e) => {
let _ = CloseHandle(event);
return Err(map_hr("IAudioClient::GetService(IAudioCaptureClient)", e));
}
};
Ok((capture, event))
}
pub(crate) fn wait_event_signaled(handle: HANDLE, timeout_ms: u32) -> bool {
let r = unsafe { WaitForSingleObject(handle, timeout_ms) };
r == WAIT_OBJECT_0
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn classify_hr_maps_access_denied_to_permission_denied() {
assert!(matches!(
classify_hr(0x80070005u32 as i32),
Some(Error::PermissionDenied)
));
assert!(matches!(
classify_hr(0x8889000Au32 as i32),
Some(Error::PermissionDenied)
));
assert!(matches!(
classify_hr(0x8889000Eu32 as i32),
Some(Error::PermissionDenied)
));
}
#[test]
fn classify_hr_maps_device_codes_to_device_not_found() {
assert!(matches!(
classify_hr(0x88890004u32 as i32),
Some(Error::DeviceNotFound)
));
assert!(matches!(
classify_hr(0x80070490u32 as i32),
Some(Error::DeviceNotFound)
));
}
#[test]
fn classify_hr_unknown_is_none() {
assert!(classify_hr(0x80004005u32 as i32).is_none());
assert!(classify_hr(0).is_none());
assert!(classify_hr(0x88890008u32 as i32).is_none());
}
}