use super::hrtf_data::{HrirMeasurement, HrtfDatabase, HrtfManager, MAX_HRIR_LENGTH};
use crate::{AudioError, AudioResult};
use rustfft::FftPlanner;
use std::f32::consts::PI;
use std::sync::Arc;
const SPEED_OF_SOUND: f32 = 343.0;
const MIN_DISTANCE: f32 = 0.1;
#[derive(Debug, Clone, Copy)]
pub struct SourcePosition {
pub x: f32,
pub y: f32,
pub z: f32,
}
impl SourcePosition {
pub fn new(x: f32, y: f32, z: f32) -> Self {
Self { x, y, z }
}
pub fn from_spherical(azimuth: f32, elevation: f32, distance: f32) -> Self {
let x = distance * elevation.cos() * azimuth.sin();
let y = distance * elevation.cos() * azimuth.cos();
let z = distance * elevation.sin();
Self { x, y, z }
}
pub fn to_spherical(&self) -> (f32, f32, f32) {
let distance = (self.x * self.x + self.y * self.y + self.z * self.z).sqrt();
let azimuth = self.x.atan2(self.y);
let elevation = if distance > 0.0 {
(self.z / distance).asin()
} else {
0.0
};
(azimuth, elevation, distance)
}
pub fn distance(&self) -> f32 {
(self.x * self.x + self.y * self.y + self.z * self.z).sqrt()
}
}
#[derive(Debug, Clone, Copy)]
pub struct ListenerOrientation {
pub yaw: f32,
pub pitch: f32,
pub roll: f32,
}
impl ListenerOrientation {
pub fn new(yaw: f32, pitch: f32, roll: f32) -> Self {
Self { yaw, pitch, roll }
}
pub fn default() -> Self {
Self {
yaw: 0.0,
pitch: 0.0,
roll: 0.0,
}
}
pub fn transform(&self, pos: &SourcePosition) -> SourcePosition {
let cos_yaw = self.yaw.cos();
let sin_yaw = self.yaw.sin();
let x1 = pos.x * cos_yaw - pos.y * sin_yaw;
let y1 = pos.x * sin_yaw + pos.y * cos_yaw;
let z1 = pos.z;
let cos_pitch = self.pitch.cos();
let sin_pitch = self.pitch.sin();
let y2 = y1 * cos_pitch - z1 * sin_pitch;
let z2 = y1 * sin_pitch + z1 * cos_pitch;
let x2 = x1;
let cos_roll = self.roll.cos();
let sin_roll = self.roll.sin();
let x3 = x2 * cos_roll - z2 * sin_roll;
let z3 = x2 * sin_roll + z2 * cos_roll;
SourcePosition::new(x3, y2, z3)
}
}
pub struct HrtfConvolver {
fft_size: usize,
hop_size: usize,
input_buffer: Vec<f32>,
output_buffer_left: Vec<f32>,
output_buffer_right: Vec<f32>,
hrir_left: Vec<f32>,
hrir_right: Vec<f32>,
fft_planner: FftPlanner<f32>,
sample_counter: usize,
}
impl HrtfConvolver {
pub fn new(fft_size: usize) -> Self {
let hop_size = fft_size / 2;
Self {
fft_size,
hop_size,
input_buffer: vec![0.0; fft_size],
output_buffer_left: vec![0.0; fft_size * 2],
output_buffer_right: vec![0.0; fft_size * 2],
hrir_left: vec![0.0; MAX_HRIR_LENGTH],
hrir_right: vec![0.0; MAX_HRIR_LENGTH],
fft_planner: FftPlanner::new(),
sample_counter: 0,
}
}
pub fn set_hrir(&mut self, hrir: &HrirMeasurement) {
let len = hrir.len().min(MAX_HRIR_LENGTH);
self.hrir_left[..len].copy_from_slice(&hrir.left[..len]);
self.hrir_right[..len].copy_from_slice(&hrir.right[..len]);
}
pub fn process_sample(&mut self, input: f32) -> (f32, f32) {
self.input_buffer.rotate_left(1);
let buffer_len = self.input_buffer.len();
self.input_buffer[buffer_len - 1] = input;
let mut left = 0.0;
let mut right = 0.0;
let hrir_len = MAX_HRIR_LENGTH.min(buffer_len);
for i in 0..hrir_len {
let buffer_idx = buffer_len - 1 - i;
left += self.input_buffer[buffer_idx] * self.hrir_left[i];
right += self.input_buffer[buffer_idx] * self.hrir_right[i];
}
(left, right)
}
pub fn process_buffer(
&mut self,
input: &[f32],
output_left: &mut [f32],
output_right: &mut [f32],
) -> AudioResult<()> {
if input.len() != output_left.len() || input.len() != output_right.len() {
return Err(AudioError::InvalidParameter(
"Buffer size mismatch".to_string(),
));
}
for i in 0..input.len() {
let (left, right) = self.process_sample(input[i]);
output_left[i] = left;
output_right[i] = right;
}
Ok(())
}
}
#[derive(Debug, Clone, Copy)]
pub enum DistanceModel {
None,
Linear {
min: f32,
max: f32,
},
Inverse {
min: f32,
rolloff: f32,
},
Exponential {
min: f32,
rolloff: f32,
},
}
impl DistanceModel {
pub fn calculate_gain(&self, distance: f32) -> f32 {
match self {
DistanceModel::None => 1.0,
DistanceModel::Linear { min, max } => {
if distance <= *min {
1.0
} else if distance >= *max {
0.0
} else {
1.0 - (distance - min) / (max - min)
}
}
DistanceModel::Inverse { min, rolloff } => {
let distance = distance.max(*min);
min / (min + rolloff * (distance - min))
}
DistanceModel::Exponential { min, rolloff } => {
let distance = distance.max(*min);
(distance / min).powf(-rolloff)
}
}
}
}
pub struct DopplerProcessor {
sample_rate: f32,
previous_distance: f32,
speed_of_sound: f32,
}
impl DopplerProcessor {
pub fn new(sample_rate: u32) -> Self {
Self {
sample_rate: sample_rate as f32,
previous_distance: 1.0,
speed_of_sound: SPEED_OF_SOUND,
}
}
pub fn calculate_shift(&mut self, current_distance: f32, delta_time: f32) -> f32 {
let velocity = (current_distance - self.previous_distance) / delta_time;
self.previous_distance = current_distance;
let shift = self.speed_of_sound / (self.speed_of_sound + velocity);
shift.clamp(0.5, 2.0) }
}
pub struct BinauralRenderer {
sample_rate: u32,
hrtf_manager: Arc<HrtfManager>,
hrtf_database: Arc<HrtfDatabase>,
convolver: HrtfConvolver,
distance_model: DistanceModel,
doppler: DopplerProcessor,
listener_orientation: ListenerOrientation,
enable_doppler: bool,
}
impl BinauralRenderer {
pub fn new(sample_rate: u32) -> AudioResult<Self> {
let hrtf_manager = Arc::new(HrtfManager::default());
let hrtf_database = hrtf_manager.get_default()?;
Ok(Self {
sample_rate,
hrtf_manager: hrtf_manager.clone(),
hrtf_database,
convolver: HrtfConvolver::new(512),
distance_model: DistanceModel::Inverse {
min: 1.0,
rolloff: 1.0,
},
doppler: DopplerProcessor::new(sample_rate),
listener_orientation: ListenerOrientation::default(),
enable_doppler: false,
})
}
pub fn set_hrtf_database(&mut self, database_name: &str) -> AudioResult<()> {
self.hrtf_database = self.hrtf_manager.get_database(database_name)?;
Ok(())
}
pub fn set_distance_model(&mut self, model: DistanceModel) {
self.distance_model = model;
}
pub fn set_listener_orientation(&mut self, orientation: ListenerOrientation) {
self.listener_orientation = orientation;
}
pub fn set_doppler_enabled(&mut self, enabled: bool) {
self.enable_doppler = enabled;
}
fn update_hrtf(&mut self, position: &SourcePosition) -> AudioResult<()> {
let relative_pos = self.listener_orientation.transform(position);
let (azimuth, elevation, _distance) = relative_pos.to_spherical();
let hrir = self
.hrtf_database
.interpolate(azimuth, elevation)
.ok_or_else(|| AudioError::Internal("Failed to get HRTF".to_string()))?;
self.convolver.set_hrir(&hrir);
Ok(())
}
pub fn render(
&mut self,
input: &[f32],
position: &SourcePosition,
output_left: &mut [f32],
output_right: &mut [f32],
) -> AudioResult<()> {
if input.len() != output_left.len() || input.len() != output_right.len() {
return Err(AudioError::InvalidParameter(
"Buffer size mismatch".to_string(),
));
}
self.update_hrtf(position)?;
let distance = position.distance().max(MIN_DISTANCE);
let distance_gain = self.distance_model.calculate_gain(distance);
let mut attenuated_input = vec![0.0; input.len()];
for (i, &sample) in input.iter().enumerate() {
attenuated_input[i] = sample * distance_gain;
}
self.convolver
.process_buffer(&attenuated_input, output_left, output_right)?;
Ok(())
}
pub fn render_with_motion(
&mut self,
input: &[f32],
position: &SourcePosition,
_velocity: &SourcePosition,
output_left: &mut [f32],
output_right: &mut [f32],
) -> AudioResult<()> {
if !self.enable_doppler {
return self.render(input, position, output_left, output_right);
}
let distance = position.distance();
let delta_time = input.len() as f32 / self.sample_rate as f32;
let _doppler_shift = self.doppler.calculate_shift(distance, delta_time);
self.render(input, position, output_left, output_right)
}
pub fn render_multi_source(
&mut self,
sources: &[(Vec<f32>, SourcePosition)],
output_left: &mut [f32],
output_right: &mut [f32],
) -> AudioResult<()> {
if sources.is_empty() {
return Ok(());
}
let buffer_size = output_left.len();
output_left.fill(0.0);
output_right.fill(0.0);
let mut temp_left = vec![0.0; buffer_size];
let mut temp_right = vec![0.0; buffer_size];
for (input, position) in sources {
if input.len() != buffer_size {
return Err(AudioError::InvalidParameter(
"Buffer size mismatch".to_string(),
));
}
temp_left.fill(0.0);
temp_right.fill(0.0);
self.render(input, position, &mut temp_left, &mut temp_right)?;
for i in 0..buffer_size {
output_left[i] += temp_left[i];
output_right[i] += temp_right[i];
}
}
Ok(())
}
}
#[derive(Debug, Clone, Copy)]
pub enum BinauralPreset {
Front,
Back,
Left,
Right,
Above,
Below,
}
impl BinauralPreset {
pub fn position(&self) -> SourcePosition {
match self {
BinauralPreset::Front => SourcePosition::from_spherical(0.0, 0.0, 1.0),
BinauralPreset::Back => SourcePosition::from_spherical(PI, 0.0, 1.0),
BinauralPreset::Left => SourcePosition::from_spherical(-PI / 2.0, 0.0, 1.0),
BinauralPreset::Right => SourcePosition::from_spherical(PI / 2.0, 0.0, 1.0),
BinauralPreset::Above => SourcePosition::from_spherical(0.0, PI / 2.0, 1.0),
BinauralPreset::Below => SourcePosition::from_spherical(0.0, -PI / 2.0, 1.0),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_source_position() {
let pos = SourcePosition::new(1.0, 0.0, 0.0);
let (azimuth, elevation, distance) = pos.to_spherical();
assert!((azimuth - PI / 2.0).abs() < 0.01);
assert!(elevation.abs() < 0.01);
assert!((distance - 1.0).abs() < 0.01);
}
#[test]
fn test_listener_orientation() {
let orientation = ListenerOrientation::new(PI / 2.0, 0.0, 0.0);
let pos = SourcePosition::new(1.0, 0.0, 0.0);
let transformed = orientation.transform(&pos);
assert!((transformed.y - 1.0).abs() < 0.01);
}
#[test]
fn test_distance_model() {
let model = DistanceModel::Inverse {
min: 1.0,
rolloff: 1.0,
};
let gain1 = model.calculate_gain(1.0);
let gain2 = model.calculate_gain(2.0);
assert!(gain1 > gain2);
assert!(gain1 <= 1.0);
}
#[test]
fn test_hrtf_convolver() {
let mut convolver = HrtfConvolver::new(512);
let hrir = HrirMeasurement::new(vec![1.0, 0.5, 0.25], vec![0.8, 0.4, 0.2], 0.0, 0.0);
convolver.set_hrir(&hrir);
for _ in 0..10 {
let _ = convolver.process_sample(1.0);
}
let (left, right) = convolver.process_sample(1.0);
assert!(left != 0.0 || right != 0.0);
}
#[test]
fn test_binaural_renderer() {
let mut renderer = BinauralRenderer::new(44100).unwrap();
let input = vec![1.0; 512];
let mut output_left = vec![0.0; 512];
let mut output_right = vec![0.0; 512];
let position = SourcePosition::from_spherical(0.0, 0.0, 2.0);
let result = renderer.render(&input, &position, &mut output_left, &mut output_right);
assert!(result.is_ok());
assert!(output_left.iter().any(|&x| x != 0.0));
assert!(output_right.iter().any(|&x| x != 0.0));
}
#[test]
fn test_binaural_presets() {
let pos = BinauralPreset::Front.position();
let (azimuth, _elevation, _distance) = pos.to_spherical();
assert!(azimuth.abs() < 0.01);
let pos = BinauralPreset::Left.position();
let (azimuth, _elevation, _distance) = pos.to_spherical();
assert!((azimuth + PI / 2.0).abs() < 0.01);
}
}