use eredu_core::{RealtimeFrameForcing, RealtimeInputFrame, RealtimeSpeechConfig};
use crate::TokenDomain;
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct RealtimeIngressContract {
schedule: RealtimeSpeechConfig,
text: TokenDomain,
audio: TokenDomain,
}
impl RealtimeIngressContract {
pub fn new(
schedule: RealtimeSpeechConfig,
text: TokenDomain,
audio: TokenDomain,
) -> Result<Self, RealtimeIngressError> {
validate_token(schedule.text_padding_token(), text, RealtimeTokenKind::Text)?;
validate_token(
schedule.audio_padding_token(),
audio,
RealtimeTokenKind::Audio,
)?;
Ok(Self {
schedule,
text,
audio,
})
}
pub const fn schedule(&self) -> &RealtimeSpeechConfig {
&self.schedule
}
pub const fn text_domain(&self) -> TokenDomain {
self.text
}
pub const fn audio_domain(&self) -> TokenDomain {
self.audio
}
pub fn validate<'a>(
&'a self,
frame: &'a RealtimeInputFrame,
) -> Result<ValidatedRealtimeInput<'a>, RealtimeIngressError> {
let batch = frame.batch();
if batch == 0 {
return Err(RealtimeIngressError::EmptyBatch);
}
let input_columns = self.schedule.input_audio_codebooks();
validate_shape(
RealtimePayloadKind::InputAudio,
frame.input_audio_tokens().len(),
batch,
input_columns,
)?;
for &token in frame.input_audio_tokens() {
validate_token(token, self.audio, RealtimeTokenKind::Audio)?;
}
let generated = self.schedule.generated_audio_codebooks();
let forced_audio = frame.forced_generated_audio_tokens();
let forcing = match (forced_audio, frame.forced_generated_audio_codebooks()) {
(None, None) => vec![false; generated],
(None, Some(_)) => return Err(RealtimeIngressError::ForcingMaskWithoutPayload),
(Some(tokens), mask) => {
validate_shape(
RealtimePayloadKind::ForcedAudio,
tokens.len(),
batch,
generated,
)?;
let mask = mask.map_or_else(|| vec![true; generated], <[bool]>::to_vec);
if mask.len() != generated {
return Err(RealtimeIngressError::ForcingMaskCount {
expected: generated,
actual: mask.len(),
});
}
for row in tokens.chunks_exact(generated) {
for (token, selected) in row.iter().zip(&mask) {
if *selected {
validate_token(*token, self.audio, RealtimeTokenKind::Audio)?;
}
}
}
mask
}
};
let forced_text = frame.forced_text_tokens();
if let Some(tokens) = forced_text {
validate_shape(RealtimePayloadKind::ForcedText, tokens.len(), batch, 1)?;
for &token in tokens {
validate_token(token, self.text, RealtimeTokenKind::Text)?;
}
}
Ok(ValidatedRealtimeInput {
contract: self,
frame,
forcing: RealtimeFrameForcing::new(forced_text.is_some(), forcing),
})
}
}
pub struct ValidatedRealtimeInput<'a> {
contract: &'a RealtimeIngressContract,
frame: &'a RealtimeInputFrame,
forcing: RealtimeFrameForcing,
}
impl ValidatedRealtimeInput<'_> {
pub const fn frame(&self) -> &RealtimeInputFrame {
self.frame
}
pub const fn forcing(&self) -> &RealtimeFrameForcing {
&self.forcing
}
pub fn materialize<M: RealtimeHostTokenMaterializer>(
&self,
materializer: &mut M,
) -> Result<MaterializedRealtimeInput<M::Tensor>, M::Error> {
let batch = self.frame.batch();
let input_audio = materializer.materialize_i32(
self.frame.input_audio_tokens(),
[batch, self.contract.schedule.input_audio_codebooks()],
)?;
let forced_audio = self
.frame
.forced_generated_audio_tokens()
.map(|tokens| {
materializer.materialize_i32(
tokens,
[batch, self.contract.schedule.generated_audio_codebooks()],
)
})
.transpose()?;
let forced_text = self
.frame
.forced_text_tokens()
.map(|tokens| materializer.materialize_i32(tokens, [batch, 1]))
.transpose()?;
Ok(MaterializedRealtimeInput {
schedule: self.contract.schedule.clone(),
batch,
input_audio,
forced_audio,
forced_text,
forcing: self.forcing.clone(),
retain_diagnostics: self.frame.retains_diagnostics(),
})
}
}
pub trait RealtimeHostTokenMaterializer {
type Tensor;
type Error;
fn materialize_i32(
&mut self,
values: &[i32],
shape: [usize; 2],
) -> Result<Self::Tensor, Self::Error>;
}
pub struct MaterializedRealtimeInput<T> {
schedule: RealtimeSpeechConfig,
batch: usize,
input_audio: T,
forced_audio: Option<T>,
forced_text: Option<T>,
forcing: RealtimeFrameForcing,
retain_diagnostics: bool,
}
impl<T> MaterializedRealtimeInput<T> {
pub const fn schedule(&self) -> &RealtimeSpeechConfig {
&self.schedule
}
pub const fn batch(&self) -> usize {
self.batch
}
pub const fn input_audio(&self) -> &T {
&self.input_audio
}
pub const fn forced_audio(&self) -> Option<&T> {
self.forced_audio.as_ref()
}
pub const fn forced_text(&self) -> Option<&T> {
self.forced_text.as_ref()
}
pub const fn forcing(&self) -> &RealtimeFrameForcing {
&self.forcing
}
pub const fn retains_diagnostics(&self) -> bool {
self.retain_diagnostics
}
pub fn into_parts(
self,
) -> (
RealtimeSpeechConfig,
usize,
T,
Option<T>,
Option<T>,
RealtimeFrameForcing,
bool,
) {
(
self.schedule,
self.batch,
self.input_audio,
self.forced_audio,
self.forced_text,
self.forcing,
self.retain_diagnostics,
)
}
}
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
pub enum RealtimePayloadKind {
InputAudio,
ForcedAudio,
ForcedText,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
pub enum RealtimeTokenKind {
Text,
Audio,
}
fn validate_shape(
payload: RealtimePayloadKind,
actual: usize,
rows: usize,
columns: usize,
) -> Result<(), RealtimeIngressError> {
let expected = rows
.checked_mul(columns)
.ok_or(RealtimeIngressError::ShapeOverflow { rows, columns })?;
if actual == expected {
Ok(())
} else {
Err(RealtimeIngressError::PayloadShape {
payload,
expected,
actual,
})
}
}
fn validate_token(
token: i32,
domain: TokenDomain,
kind: RealtimeTokenKind,
) -> Result<(), RealtimeIngressError> {
let token = usize::try_from(token).map_err(|_| RealtimeIngressError::TokenDomain {
kind,
token,
cardinality: domain.cardinality(),
})?;
if token < domain.cardinality() {
Ok(())
} else {
Err(RealtimeIngressError::TokenDomain {
kind,
token: i32::try_from(token).unwrap_or(i32::MAX),
cardinality: domain.cardinality(),
})
}
}
#[derive(Debug, Clone, Eq, PartialEq, thiserror::Error)]
#[non_exhaustive]
pub enum RealtimeIngressError {
#[error("realtime input batch must be positive")]
EmptyBatch,
#[error("realtime payload shape {rows}x{columns} overflowed")]
ShapeOverflow {
rows: usize,
columns: usize,
},
#[error("realtime {payload:?} payload has {actual} values, expected {expected}")]
PayloadShape {
payload: RealtimePayloadKind,
expected: usize,
actual: usize,
},
#[error("realtime generated-audio forcing mask has no payload")]
ForcingMaskWithoutPayload,
#[error("realtime forcing mask has {actual} entries, expected {expected}")]
ForcingMaskCount {
expected: usize,
actual: usize,
},
#[error("realtime {kind:?} token {token} is outside 0..{cardinality}")]
TokenDomain {
kind: RealtimeTokenKind,
token: i32,
cardinality: usize,
},
}
#[cfg(test)]
mod tests {
use std::convert::Infallible;
use eredu_core::{RealtimeFrameConvention, RealtimeInputFrame};
use super::*;
fn contract() -> RealtimeIngressContract {
RealtimeIngressContract::new(
RealtimeSpeechConfig::new(
4,
2,
2,
2,
16,
15,
RealtimeFrameConvention::FeedbackAlignedHistory,
vec![0, 1, 2, 1, 2],
)
.unwrap(),
TokenDomain::new(17),
TokenDomain::new(16),
)
.unwrap()
}
#[derive(Default)]
struct Recorder(Vec<(Vec<i32>, [usize; 2])>);
impl RealtimeHostTokenMaterializer for Recorder {
type Tensor = usize;
type Error = Infallible;
fn materialize_i32(
&mut self,
values: &[i32],
shape: [usize; 2],
) -> Result<Self::Tensor, Self::Error> {
self.0.push((values.to_vec(), shape));
Ok(self.0.len() - 1)
}
}
#[test]
fn validation_precedes_every_opaque_materialization() {
let contract = contract();
let invalid = RealtimeInputFrame::new(2, vec![1, 2, 3, 99]);
let mut materializer = Recorder::default();
assert!(matches!(
contract.validate(&invalid),
Err(RealtimeIngressError::TokenDomain { .. })
));
assert!(materializer.0.is_empty());
let valid = RealtimeInputFrame::new(2, vec![1, 2, 3, 4])
.with_partially_forced_generated_audio(vec![5, 99, 6, 99], vec![true, false])
.with_forced_text(vec![7, 8])
.with_diagnostics();
let validated = contract.validate(&valid).unwrap();
assert_eq!(validated.forcing().generated_audio(), &[true, false]);
let input = validated.materialize(&mut materializer).unwrap();
assert_eq!(materializer.0.len(), 3);
assert_eq!(materializer.0[0].1, [2, 2]);
assert_eq!(materializer.0[1].1, [2, 2]);
assert_eq!(materializer.0[2].1, [2, 1]);
assert!(input.retains_diagnostics());
}
#[test]
fn shapes_masks_and_selected_domains_fail_closed() {
let contract = contract();
assert_eq!(
contract.validate(&RealtimeInputFrame::new(0, vec![])).err(),
Some(RealtimeIngressError::EmptyBatch)
);
assert!(matches!(
contract.validate(&RealtimeInputFrame::new(2, vec![1, 2, 3])),
Err(RealtimeIngressError::PayloadShape { .. })
));
assert!(matches!(
contract.validate(
&RealtimeInputFrame::new(1, vec![1, 2])
.with_partially_forced_generated_audio(vec![3, 4], vec![true])
),
Err(RealtimeIngressError::ForcingMaskCount { .. })
));
assert!(matches!(
contract.validate(&RealtimeInputFrame::new(1, vec![1, 2]).with_forced_text(vec![17])),
Err(RealtimeIngressError::TokenDomain { .. })
));
}
}