use serde::{Deserialize, Serialize};
use crate::envelope::Adsr;
use crate::error::Result;
const MAX_OPERATORS: usize = 6;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub enum FmAlgorithm {
Serial2,
Parallel2,
Serial4,
Stack4,
Custom,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FmOperator {
frequency: f32,
phase: f32,
envelope: Adsr,
output_level: f32,
feedback: f32,
feedback_state: f32,
sample_rate: f32,
}
impl FmOperator {
pub fn new(frequency: f32, sample_rate: f32) -> Result<Self> {
let envelope = Adsr::with_sample_rate(0.01, 0.1, 0.8, 0.3, sample_rate)?;
Ok(Self {
frequency,
phase: 0.0,
envelope,
output_level: 1.0,
feedback: 0.0,
feedback_state: 0.0,
sample_rate,
})
}
pub fn set_frequency(&mut self, freq: f32) {
self.frequency = freq;
}
pub fn set_level(&mut self, level: f32) {
self.output_level = level.clamp(0.0, 1.0);
}
pub fn set_feedback(&mut self, amount: f32) {
self.feedback = amount.clamp(0.0, 1.0);
}
pub fn gate_on(&mut self) {
self.envelope.gate_on();
}
pub fn gate_off(&mut self) {
self.envelope.gate_off();
}
#[inline]
#[must_use]
pub fn next_sample(&mut self, phase_mod: f32) -> f32 {
let fb = self.feedback * self.feedback_state;
let mod_phase = self.phase + phase_mod + fb;
let out = (mod_phase * std::f32::consts::TAU).sin();
let env = self.envelope.next_value();
let result = out * env * self.output_level;
self.feedback_state = crate::flush_denormal(result);
let phase_inc = self.frequency / self.sample_rate;
self.phase += phase_inc;
self.phase -= self.phase.floor();
result
}
pub fn envelope_mut(&mut self) -> &mut Adsr {
&mut self.envelope
}
#[must_use]
pub fn envelope(&self) -> &Adsr {
&self.envelope
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FmSynthEngine {
operators: Vec<FmOperator>,
algorithm: FmAlgorithm,
sample_rate: f32,
}
impl FmSynthEngine {
pub fn new(num_operators: usize, sample_rate: f32) -> Result<Self> {
if num_operators == 0 || num_operators > MAX_OPERATORS {
return Err(crate::error::NaadError::InvalidParameter {
name: "num_operators".to_string(),
reason: format!("must be 1..={MAX_OPERATORS}"),
});
}
if sample_rate <= 0.0 || !sample_rate.is_finite() {
return Err(crate::error::NaadError::InvalidSampleRate { sample_rate });
}
let mut operators = Vec::with_capacity(num_operators);
for _ in 0..num_operators {
operators.push(FmOperator::new(440.0, sample_rate)?);
}
Ok(Self {
operators,
algorithm: FmAlgorithm::Serial2,
sample_rate,
})
}
pub fn set_algorithm(&mut self, algorithm: FmAlgorithm) {
self.algorithm = algorithm;
}
pub fn set_operator_freq(&mut self, index: usize, freq: f32) -> Option<()> {
let op = self.operators.get_mut(index)?;
op.set_frequency(freq);
Some(())
}
pub fn set_operator_level(&mut self, index: usize, level: f32) -> Option<()> {
let op = self.operators.get_mut(index)?;
op.set_level(level);
Some(())
}
pub fn note_on(&mut self) {
for op in &mut self.operators {
op.gate_on();
}
}
pub fn note_off(&mut self) {
for op in &mut self.operators {
op.gate_off();
}
}
pub fn operator_mut(&mut self, index: usize) -> Option<&mut FmOperator> {
self.operators.get_mut(index)
}
#[must_use]
pub fn operator(&self, index: usize) -> Option<&FmOperator> {
self.operators.get(index)
}
#[inline]
#[must_use]
pub fn next_sample(&mut self) -> f32 {
let n = self.operators.len();
match self.algorithm {
FmAlgorithm::Serial2 => self.process_serial2(n),
FmAlgorithm::Parallel2 => self.process_parallel2(n),
FmAlgorithm::Serial4 => self.process_serial4(n),
FmAlgorithm::Stack4 => self.process_stack4(n),
FmAlgorithm::Custom => {
let mut sum = 0.0f32;
for op in &mut self.operators {
sum += op.next_sample(0.0);
}
sum / n.max(1) as f32
}
}
}
#[inline]
pub fn fill_buffer(&mut self, buffer: &mut [f32]) {
for s in buffer.iter_mut() {
*s = self.next_sample();
}
}
#[inline]
fn process_serial2(&mut self, n: usize) -> f32 {
if n < 2 {
return self.operators[0].next_sample(0.0);
}
let (first, rest) = self.operators.split_at_mut(1);
let mod_out = first[0].next_sample(0.0);
rest[0].next_sample(mod_out)
}
#[inline]
fn process_parallel2(&mut self, n: usize) -> f32 {
if n < 2 {
return self.operators[0].next_sample(0.0);
}
let a = self.operators[0].next_sample(0.0);
let b = self.operators[1].next_sample(0.0);
(a + b) * 0.5
}
#[inline]
fn process_serial4(&mut self, n: usize) -> f32 {
let mut mod_signal = 0.0f32;
let count = n.min(4);
for i in 0..count {
mod_signal = self.operators[i].next_sample(mod_signal);
}
mod_signal
}
#[inline]
fn process_stack4(&mut self, n: usize) -> f32 {
if n < 4 {
return self.process_serial4(n);
}
let (first_pair, second_pair) = self.operators.split_at_mut(2);
let mod_a = first_pair[0].next_sample(0.0);
let mod_b = first_pair[1].next_sample(0.0);
let mod_sum = (mod_a + mod_b) * 0.5;
let car_a = second_pair[0].next_sample(mod_sum);
let car_b = second_pair[1].next_sample(mod_sum);
(car_a + car_b) * 0.5
}
#[must_use]
pub fn is_active(&self) -> bool {
self.operators.iter().any(|op| op.envelope.is_active())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_basic_output() {
let mut engine = FmSynthEngine::new(2, 44100.0).unwrap();
engine.set_operator_freq(0, 440.0);
engine.set_operator_freq(1, 440.0);
engine.set_algorithm(FmAlgorithm::Serial2);
engine.note_on();
let mut buf = [0.0f32; 1024];
engine.fill_buffer(&mut buf);
assert!(buf.iter().any(|&s| s.abs() > 0.01), "should produce output");
assert!(buf.iter().all(|s| s.is_finite()));
}
#[test]
fn test_serial_vs_parallel_differ() {
let mut serial = FmSynthEngine::new(2, 44100.0).unwrap();
serial.set_operator_freq(0, 200.0);
serial.set_operator_freq(1, 440.0);
serial.set_operator_level(0, 1.0);
serial.set_algorithm(FmAlgorithm::Serial2);
serial.note_on();
let mut parallel = FmSynthEngine::new(2, 44100.0).unwrap();
parallel.set_operator_freq(0, 200.0);
parallel.set_operator_freq(1, 440.0);
parallel.set_operator_level(0, 1.0);
parallel.set_algorithm(FmAlgorithm::Parallel2);
parallel.note_on();
let mut buf_s = [0.0f32; 512];
let mut buf_p = [0.0f32; 512];
serial.fill_buffer(&mut buf_s);
parallel.fill_buffer(&mut buf_p);
let diff: f32 = buf_s
.iter()
.zip(buf_p.iter())
.map(|(a, b)| (a - b).abs())
.sum();
assert!(
diff > 0.1,
"serial and parallel should produce different spectra"
);
}
#[test]
fn test_serde_roundtrip() {
let engine = FmSynthEngine::new(4, 44100.0).unwrap();
let json = serde_json::to_string(&engine).unwrap();
let back: FmSynthEngine = serde_json::from_str(&json).unwrap();
assert_eq!(engine.operators.len(), back.operators.len());
assert_eq!(engine.algorithm, back.algorithm);
}
#[test]
fn test_four_op_algorithms() {
let mut engine = FmSynthEngine::new(4, 44100.0).unwrap();
engine.set_algorithm(FmAlgorithm::Serial4);
engine.note_on();
let mut buf = [0.0f32; 256];
engine.fill_buffer(&mut buf);
assert!(buf.iter().any(|&s| s.abs() > 0.001));
let mut engine2 = FmSynthEngine::new(4, 44100.0).unwrap();
engine2.set_algorithm(FmAlgorithm::Stack4);
engine2.note_on();
let mut buf2 = [0.0f32; 256];
engine2.fill_buffer(&mut buf2);
assert!(buf2.iter().any(|&s| s.abs() > 0.001));
}
}