use core::{
num::NonZeroU64,
sync::atomic::{AtomicBool, AtomicU64, Ordering},
};
use std::collections::HashMap;
use asry::{
Lang, TimeRange,
emissions::{
OovDecision, OovDetection, OovEvent, OovKind, OovResolution, OutputClock, UnalignedCause,
UnitAlignment, default_oov_policy,
},
};
use crate::audio::align::{
aligner::Aligner,
error::{AlignError, DecisionLanguage, ForeignResolution, MisroutedResolution},
};
#[derive(Clone, PartialEq, Eq, Hash, Debug, derive_more::IsVariant)]
#[non_exhaustive]
pub enum AlignerKey {
Lang(Lang),
Any,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
#[error("unknown alignment fallback policy name")]
pub struct ParseAlignmentFallbackError(());
macro_rules! define_alignment_fallback {
(
$(#[$enum_meta:meta])*
$vis:vis enum $Name:ident {
$(
$(#[$variant_meta:meta])*
$Variant:ident => $spelling:literal
),+ $(,)?
}
) => {
$(#[$enum_meta])*
$vis enum $Name {
$(
$(#[$variant_meta])*
#[cfg_attr(feature = "serde", serde(rename = $spelling))]
$Variant,
)+
}
impl $Name {
#[inline(always)]
pub const fn as_str(&self) -> &'static str {
match self {
$( Self::$Variant => $spelling, )+
}
}
#[cfg(test)]
pub(crate) const ALL: &'static [Self] = &[$( Self::$Variant, )+];
}
impl core::str::FromStr for $Name {
type Err = ParseAlignmentFallbackError;
fn from_str(s: &str) -> Result<Self, Self::Err> {
Ok(match s {
$( $spelling => Self::$Variant, )+
_ => return Err(ParseAlignmentFallbackError(())),
})
}
}
};
}
define_alignment_fallback! {
#[derive(
Copy, Clone, PartialEq, Eq, Debug, Default, derive_more::Display, derive_more::IsVariant,
)]
#[display("{}", self.as_str())]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[non_exhaustive]
pub enum AlignmentFallback {
#[default]
SkipChunk => "skip_chunk",
Error => "error",
}
}
enum AlignmentLookup<'a> {
Hit(&'a Aligner),
AnyFallback(&'a Aligner),
Miss(AlignmentFallback),
}
impl AlignmentLookup<'_> {
fn binding(&self) -> AlignmentBinding {
match self {
Self::Hit(_) => AlignmentBinding::Exact,
Self::AnyFallback(aligner) => AlignmentBinding::AnyFallback(aligner.language_ref().clone()),
Self::Miss(fallback) => AlignmentBinding::Miss(*fallback),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct SetId(NonZeroU64);
impl SetId {
fn next() -> Self {
static COUNTER: AtomicU64 = AtomicU64::new(1);
let raw = COUNTER.fetch_add(1, Ordering::Relaxed);
Self(NonZeroU64::new(raw).expect("SetId counter overflowed u64"))
}
}
impl core::fmt::Display for SetId {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
write!(f, "alignment set #{}", self.0)
}
}
#[derive(Clone, PartialEq, Eq, Debug, derive_more::IsVariant)]
pub enum AlignmentBinding {
Exact,
AnyFallback(Lang),
Miss(AlignmentFallback),
}
pub struct AlignmentHandle<'a> {
set: &'a AlignmentSet,
language: Lang,
}
pub struct AlignmentSet {
aligners: HashMap<AlignerKey, Aligner>,
fallback: AlignmentFallback,
id: SetId,
}
impl AlignmentSet {
#[must_use]
pub const fn id(&self) -> SetId {
self.id
}
#[must_use]
pub const fn fallback(&self) -> AlignmentFallback {
self.fallback
}
#[must_use]
pub fn len(&self) -> usize {
self.aligners.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.aligners.is_empty()
}
#[must_use]
pub fn resolve<'a>(&'a self, language: &Lang) -> AlignmentHandle<'a> {
AlignmentHandle {
set: self,
language: language.clone(),
}
}
fn binding(&self, language: &Lang) -> AlignmentBinding {
self.lookup(language).binding()
}
#[must_use]
fn lookup<'a>(&'a self, language: &Lang) -> AlignmentLookup<'a> {
let lang_key = AlignerKey::Lang(language.clone());
if let Some(aligner) = self.aligners.get(&lang_key) {
return AlignmentLookup::Hit(aligner);
}
if let Some(aligner) = self.aligners.get(&AlignerKey::Any) {
return AlignmentLookup::AnyFallback(aligner);
}
AlignmentLookup::Miss(self.fallback)
}
pub fn detect_oov(&self, text: &str, language: &Lang) -> Result<SetDetection, AlignError> {
let detection = match self.lookup(language) {
AlignmentLookup::Hit(aligner) | AlignmentLookup::AnyFallback(aligner) => {
Some(aligner.detect_oov(text)?)
}
AlignmentLookup::Miss(_) => None,
};
Ok(SetDetection {
made_by: self.id,
language: language.clone(),
detection,
})
}
#[allow(clippy::too_many_arguments)]
#[cfg_attr(
feature = "tracing",
tracing::instrument(
name = "alignkit.registry.align_chunk",
level = "debug",
skip_all,
fields(
requested_language = ?language,
route = ?self.binding(language),
),
)
)]
pub fn align_chunk(
&self,
language: &Lang,
samples: &[f32],
sub_segments: &[TimeRange],
text: &str,
clock: OutputClock,
abort_flag: &AtomicBool,
resolution: SetResolution,
) -> Result<UnitAlignment, AlignError> {
if resolution.made_by != self.id {
return Err(AlignError::ForeignResolution(ForeignResolution::new(
resolution.made_by,
self.id,
)));
}
if resolution.language != *language {
return Err(AlignError::DecisionLanguage(DecisionLanguage::new(
language.clone(),
resolution.language,
)));
}
match (self.lookup(language), resolution.resolution) {
(AlignmentLookup::Hit(aligner) | AlignmentLookup::AnyFallback(aligner), Some(resolution)) => {
aligner
.align_chunk(samples, sub_segments, text, clock, abort_flag, resolution)
.map_err(|error| for_request(error, language))
}
(AlignmentLookup::Miss(fallback), None) => match fallback {
AlignmentFallback::SkipChunk => Ok(UnitAlignment::Unaligned(UnalignedCause::Skipped)),
AlignmentFallback::Error => Err(AlignError::LanguageUnsupported(language.clone())),
},
(lookup, decided) => Err(AlignError::MisroutedResolution(MisroutedResolution::new(
lookup.binding(),
decided.is_some(),
))),
}
}
}
#[must_use = "a detection does nothing until it is decided"]
pub struct SetDetection {
made_by: SetId,
language: Lang,
detection: Option<OovDetection>,
}
impl SetDetection {
#[must_use]
pub const fn made_by(&self) -> SetId {
self.made_by
}
#[must_use]
pub const fn language(&self) -> &Lang {
&self.language
}
#[must_use]
pub fn events(&self) -> Option<Vec<SetOovEvent<'_>>> {
self.detection.as_ref().map(|detection| {
detection
.events()
.iter()
.map(|event| SetOovEvent::new(event, &self.language))
.collect()
})
}
pub fn decide(self, mut policy: impl FnMut(&SetOovEvent<'_>) -> OovDecision) -> SetResolution {
let language = self.language;
let resolution = self
.detection
.map(|detection| detection.decide(|event| policy(&SetOovEvent::new(event, &language))));
SetResolution {
made_by: self.made_by,
language,
resolution,
}
}
}
impl core::fmt::Debug for SetDetection {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("SetDetection")
.field("made_by", &self.made_by)
.field("language", &self.language)
.field("events", &self.events())
.finish()
}
}
#[derive(Clone, Copy)]
pub struct SetOovEvent<'a> {
event: &'a OovEvent,
language: &'a Lang,
}
impl<'a> SetOovEvent<'a> {
const fn new(event: &'a OovEvent, language: &'a Lang) -> Self {
Self { event, language }
}
#[must_use]
pub const fn kind(&self) -> &'a OovKind {
self.event.kind()
}
#[must_use]
pub const fn char_index(&self) -> usize {
self.event.char_index()
}
#[must_use]
pub const fn word_index(&self) -> usize {
self.event.word_index()
}
#[must_use]
pub fn char(&self) -> Option<char> {
self.event.char()
}
#[must_use]
pub const fn language(&self) -> &'a Lang {
self.language
}
#[must_use]
pub fn default_decision(&self) -> OovDecision {
default_oov_policy(self.event)
}
}
impl PartialEq for SetOovEvent<'_> {
fn eq(&self, other: &Self) -> bool {
self.event.matches_position(other.event) && self.language == other.language
}
}
impl Eq for SetOovEvent<'_> {}
impl core::fmt::Debug for SetOovEvent<'_> {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("SetOovEvent")
.field("kind", self.kind())
.field("char_index", &self.char_index())
.field("word_index", &self.word_index())
.field("language", self.language)
.finish()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct SetResolvedOov<'a> {
event: SetOovEvent<'a>,
decision: OovDecision,
}
impl<'a> SetResolvedOov<'a> {
#[must_use]
pub const fn event(&self) -> SetOovEvent<'a> {
self.event
}
#[must_use]
pub const fn decision(&self) -> OovDecision {
self.decision
}
}
#[must_use = "a resolution does nothing until alignment applies it"]
pub struct SetResolution {
made_by: SetId,
language: Lang,
resolution: Option<OovResolution>,
}
impl SetResolution {
#[must_use]
pub const fn made_by(&self) -> SetId {
self.made_by
}
#[must_use]
pub const fn language(&self) -> &Lang {
&self.language
}
#[must_use]
pub fn resolved(&self) -> Option<Vec<SetResolvedOov<'_>>> {
self.resolution.as_ref().map(|resolution| {
resolution
.resolved()
.iter()
.map(|resolved| SetResolvedOov {
event: SetOovEvent::new(resolved.event(), &self.language),
decision: resolved.decision(),
})
.collect()
})
}
}
impl core::fmt::Debug for SetResolution {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("SetResolution")
.field("made_by", &self.made_by)
.field("language", &self.language)
.field("resolved", &self.resolved())
.finish()
}
}
fn for_request(error: AlignError, language: &Lang) -> AlignError {
match error {
AlignError::Refused(refusal) => AlignError::Refused(refusal.under(language)),
other => other,
}
}
impl AlignmentHandle<'_> {
#[must_use]
pub const fn language(&self) -> &Lang {
&self.language
}
#[must_use]
pub fn binding(&self) -> AlignmentBinding {
self.set.binding(&self.language)
}
pub fn detect_oov(&self, text: &str) -> Result<SetDetection, AlignError> {
self.set.detect_oov(text, &self.language)
}
#[allow(clippy::too_many_arguments)]
pub fn align_chunk(
&self,
samples: &[f32],
sub_segments: &[TimeRange],
text: &str,
clock: OutputClock,
abort_flag: &AtomicBool,
resolution: SetResolution,
) -> Result<UnitAlignment, AlignError> {
self.set.align_chunk(
&self.language,
samples,
sub_segments,
text,
clock,
abort_flag,
resolution,
)
}
}
pub struct AlignmentSetBuilder {
aligners: HashMap<AlignerKey, Aligner>,
fallback: AlignmentFallback,
}
impl AlignmentSetBuilder {
#[must_use]
pub fn new() -> Self {
Self {
aligners: HashMap::new(),
fallback: AlignmentFallback::SkipChunk,
}
}
#[must_use]
pub const fn with_fallback(mut self, fallback: AlignmentFallback) -> Self {
self.fallback = fallback;
self
}
pub const fn set_fallback(&mut self, fallback: AlignmentFallback) {
self.fallback = fallback;
}
#[must_use]
pub fn register(mut self, key: AlignerKey, aligner: Aligner) -> Self {
if let AlignerKey::Lang(ref key_lang) = key {
assert_eq!(
aligner.language_ref(),
key_lang,
"AlignerKey::Lang({key_lang:?}) cannot accept an aligner built for {actual:?}; \
register it under AlignerKey::Lang({actual:?}) or AlignerKey::Any, or rebuild the \
aligner for the desired language",
actual = aligner.language_ref(),
);
}
self.aligners.insert(key, aligner);
self
}
#[must_use]
pub fn len(&self) -> usize {
self.aligners.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.aligners.is_empty()
}
#[must_use]
pub fn build(self) -> AlignmentSet {
AlignmentSet {
aligners: self.aligners,
fallback: self.fallback,
id: SetId::next(),
}
}
}
impl Default for AlignmentSetBuilder {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests;