use crate::audio::ced::{
WINDOW_SAMPLES,
error::{Error, Result, WinditError},
};
pub use windit::plan::Span;
#[cfg(test)]
mod tests;
pub const DEFAULT_HOP_SAMPLES: u32 = WINDOW_SAMPLES as u32;
pub const DEFAULT_MAX_WINDOWS: u32 = 100_000;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct DropBelowMin {
min_samples: u32,
}
impl DropBelowMin {
#[inline(always)]
pub const fn new(min_samples: u32) -> Self {
Self { min_samples }
}
#[inline(always)]
pub const fn min_samples(&self) -> u32 {
self.min_samples
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum TailPolicy {
#[default]
Pad,
DropBelowMin(DropBelowMin),
}
#[cfg(feature = "serde")]
mod tail_policy_serde {
use super::{DropBelowMin, TailPolicy};
#[derive(serde::Serialize, serde::Deserialize)]
#[serde(rename_all = "snake_case", tag = "kind", content = "value")]
pub(super) enum Document {
Pad,
DropBelowMin(DropBelowMin),
}
#[derive(serde::Serialize, serde::Deserialize)]
pub(super) enum Binary {
Pad,
DropBelowMin(DropBelowMin),
}
impl From<TailPolicy> for Document {
fn from(policy: TailPolicy) -> Self {
match policy {
TailPolicy::Pad => Self::Pad,
TailPolicy::DropBelowMin(d) => Self::DropBelowMin(d),
}
}
}
impl From<Document> for TailPolicy {
fn from(doc: Document) -> Self {
match doc {
Document::Pad => Self::Pad,
Document::DropBelowMin(d) => Self::DropBelowMin(d),
}
}
}
impl From<TailPolicy> for Binary {
fn from(policy: TailPolicy) -> Self {
match policy {
TailPolicy::Pad => Self::Pad,
TailPolicy::DropBelowMin(d) => Self::DropBelowMin(d),
}
}
}
impl From<Binary> for TailPolicy {
fn from(binary: Binary) -> Self {
match binary {
Binary::Pad => Self::Pad,
Binary::DropBelowMin(d) => Self::DropBelowMin(d),
}
}
}
}
#[cfg(feature = "serde")]
impl serde::Serialize for TailPolicy {
fn serialize<S: serde::Serializer>(
&self,
serializer: S,
) -> core::result::Result<S::Ok, S::Error> {
if serializer.is_human_readable() {
serde::Serialize::serialize(&tail_policy_serde::Document::from(*self), serializer)
} else {
serde::Serialize::serialize(&tail_policy_serde::Binary::from(*self), serializer)
}
}
}
#[cfg(feature = "serde")]
impl<'de> serde::Deserialize<'de> for TailPolicy {
fn deserialize<D: serde::Deserializer<'de>>(
deserializer: D,
) -> core::result::Result<Self, D::Error> {
if deserializer.is_human_readable() {
<tail_policy_serde::Document as serde::Deserialize>::deserialize(deserializer).map(Self::from)
} else {
<tail_policy_serde::Binary as serde::Deserialize>::deserialize(deserializer).map(Self::from)
}
}
}
const fn check_hop_samples(v: u32) -> bool {
v > 0 && v as usize <= WINDOW_SAMPLES
}
const fn check_max_windows(v: u32) -> bool {
v > 0
}
const fn check_tail(tail: TailPolicy) -> bool {
match tail {
TailPolicy::Pad => true,
TailPolicy::DropBelowMin(d) => {
d.min_samples() > 0 && d.min_samples() as usize <= WINDOW_SAMPLES
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[cfg_attr(feature = "serde", serde(try_from = "WindowPlanRepr"))]
pub struct WindowPlan {
hop_samples: u32,
tail: TailPolicy,
max_windows: u32,
}
#[cfg(feature = "serde")]
#[derive(serde::Deserialize)]
#[serde(deny_unknown_fields)]
struct WindowPlanRepr {
#[serde(default = "default_hop_samples")]
hop_samples: u32,
#[serde(default)]
tail: TailPolicy,
#[serde(default = "default_max_windows")]
max_windows: u32,
}
#[cfg(feature = "serde")]
fn default_hop_samples() -> u32 {
DEFAULT_HOP_SAMPLES
}
#[cfg(feature = "serde")]
fn default_max_windows() -> u32 {
DEFAULT_MAX_WINDOWS
}
#[cfg(feature = "serde")]
impl TryFrom<WindowPlanRepr> for WindowPlan {
type Error = String;
fn try_from(r: WindowPlanRepr) -> core::result::Result<Self, Self::Error> {
if !check_hop_samples(r.hop_samples) {
return Err(format!(
"hop_samples ({}) must be > 0 and <= WINDOW_SAMPLES ({WINDOW_SAMPLES})",
r.hop_samples
));
}
if !check_tail(r.tail) {
let tail: &dyn core::fmt::Debug = match &r.tail {
TailPolicy::DropBelowMin(min) => min,
TailPolicy::Pad => &r.tail,
};
return Err(format!(
"tail DropBelowMin.min_samples must be > 0 and <= WINDOW_SAMPLES ({WINDOW_SAMPLES}), got {tail:?}"
));
}
if !check_max_windows(r.max_windows) {
return Err(format!("max_windows ({}) must be > 0", r.max_windows));
}
Ok(Self {
hop_samples: r.hop_samples,
tail: r.tail,
max_windows: r.max_windows,
})
}
}
impl Default for WindowPlan {
fn default() -> Self {
Self::new()
}
}
impl WindowPlan {
pub const fn new() -> Self {
Self {
hop_samples: DEFAULT_HOP_SAMPLES,
tail: TailPolicy::Pad,
max_windows: DEFAULT_MAX_WINDOWS,
}
}
#[inline]
pub const fn hop_samples(&self) -> u32 {
self.hop_samples
}
#[inline]
pub const fn tail_policy(&self) -> TailPolicy {
self.tail
}
#[must_use]
pub const fn with_hop_samples(mut self, hop_samples: u32) -> Self {
self.set_hop_samples(hop_samples);
self
}
pub const fn set_hop_samples(&mut self, hop_samples: u32) -> &mut Self {
assert!(
check_hop_samples(hop_samples),
"hop_samples must be > 0 and <= WINDOW_SAMPLES (160_000)"
);
self.hop_samples = hop_samples;
self
}
#[must_use]
pub const fn with_tail_policy(mut self, tail: TailPolicy) -> Self {
self.set_tail_policy(tail);
self
}
pub const fn set_tail_policy(&mut self, tail: TailPolicy) -> &mut Self {
assert!(
check_tail(tail),
"TailPolicy::DropBelowMin.min_samples must be > 0 and <= WINDOW_SAMPLES (160_000)"
);
self.tail = tail;
self
}
#[inline]
pub const fn max_windows(&self) -> u32 {
self.max_windows
}
#[must_use]
pub const fn with_max_windows(mut self, max_windows: u32) -> Self {
self.set_max_windows(max_windows);
self
}
pub const fn set_max_windows(&mut self, max_windows: u32) -> &mut Self {
assert!(check_max_windows(max_windows), "max_windows must be > 0");
self.max_windows = max_windows;
self
}
fn windit_options(&self) -> windit::plan::WindowOptions {
windit::plan::WindowOptions::new(WINDOW_SAMPLES)
.with_hop(self.hop_samples as usize)
.with_tail(match self.tail {
TailPolicy::Pad => windit::plan::TailPolicy::PadFull,
TailPolicy::DropBelowMin(d) => {
windit::plan::TailPolicy::DropBelowMin(d.min_samples() as usize)
}
})
.with_max_windows(self.max_windows as usize)
}
fn planned_windows(&self, total_samples: usize) -> usize {
if total_samples == 0 {
return 0;
}
if total_samples <= WINDOW_SAMPLES {
return 1;
}
let hop = self.hop_samples as usize;
match self.tail {
TailPolicy::Pad => total_samples.div_ceil(hop),
TailPolicy::DropBelowMin(d) => (total_samples - d.min_samples() as usize) / hop + 1,
}
}
pub fn spans(&self, total_samples: usize) -> Result<Vec<Span>> {
let planned = self.planned_windows(total_samples);
let max = self.max_windows as usize;
if planned > max {
return Err(Error::Windowing(WinditError::TooManyWindows {
got: planned,
max,
}));
}
if total_samples == 0 {
return Ok(Vec::new());
}
if total_samples <= WINDOW_SAMPLES {
return Ok(vec![Span::new(0, total_samples, WINDOW_SAMPLES)]);
}
let mut spans = windit::plan::WindowPlan::spans(&self.windit_options(), total_samples)?;
let hop = self.hop_samples as usize;
let min_keep = match self.tail {
TailPolicy::Pad => 1,
TailPolicy::DropBelowMin(d) => d.min_samples() as usize,
};
let extra = planned - spans.len();
spans
.try_reserve_exact(extra)
.map_err(|_| Error::Windowing(WinditError::AllocFailed { elements: extra }))?;
let first_tail_start = (total_samples - WINDOW_SAMPLES).div_ceil(hop) * hop;
let mut start = first_tail_start + hop;
while start < total_samples {
let len = total_samples - start; if len >= min_keep {
spans.push(Span::new(start, len, WINDOW_SAMPLES));
}
start += hop;
}
debug_assert_eq!(
spans.len(),
planned,
"planned_windows drifted from construction"
);
Ok(spans)
}
}