use sha2::{Digest, Sha256};
#[derive(Debug, Clone, PartialEq, thiserror::Error)]
pub enum MediaRequestError {
#[error("multimodal input request is empty")]
EmptyRequest,
#[error("RGB8 image shape {width}x{height} requires {expected} bytes, got {actual}")]
InvalidRgbImage {
width: u32,
height: u32,
expected: usize,
actual: usize,
},
#[error("decoded audio is invalid: {0}")]
InvalidAudio(String),
#[error("decoded video is invalid: {0}")]
InvalidVideo(String),
#[error("chat media binding {index} has an empty placeholder")]
EmptyPlaceholder {
index: usize,
},
#[error(
"rendered chat contains {actual} occurrence(s) of media placeholder {placeholder:?}, but {expected} binding(s) were supplied"
)]
PlaceholderCount {
placeholder: String,
expected: usize,
actual: usize,
},
#[error(
"chat media binding {index} placeholder {placeholder:?} does not occur after the preceding binding"
)]
PlaceholderOrder {
index: usize,
placeholder: String,
},
}
#[derive(Debug, Clone, PartialEq)]
pub struct RgbImage {
pixels: Vec<u8>,
width: u32,
height: u32,
}
impl RgbImage {
pub fn new(pixels: Vec<u8>, width: u32, height: u32) -> Result<Self, MediaRequestError> {
let expected = usize::try_from(width)
.ok()
.zip(usize::try_from(height).ok())
.and_then(|(width, height)| width.checked_mul(height))
.and_then(|pixels| pixels.checked_mul(3))
.unwrap_or(usize::MAX);
if width == 0 || height == 0 || pixels.len() != expected {
return Err(MediaRequestError::InvalidRgbImage {
width,
height,
expected,
actual: pixels.len(),
});
}
Ok(Self {
pixels,
width,
height,
})
}
pub fn pixels(&self) -> &[u8] {
&self.pixels
}
pub const fn width(&self) -> u32 {
self.width
}
pub const fn height(&self) -> u32 {
self.height
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct Audio {
samples: Vec<f32>,
sample_rate: u32,
}
impl Audio {
pub fn new(samples: Vec<f32>, sample_rate: u32) -> Result<Self, MediaRequestError> {
if sample_rate == 0 {
return Err(MediaRequestError::InvalidAudio(
"sample rate must be positive".into(),
));
}
if samples.is_empty() {
return Err(MediaRequestError::InvalidAudio(
"waveform must contain at least one sample".into(),
));
}
if samples.iter().any(|sample| !sample.is_finite()) {
return Err(MediaRequestError::InvalidAudio(
"waveform samples must be finite".into(),
));
}
Ok(Self {
samples,
sample_rate,
})
}
pub fn samples(&self) -> &[f32] {
&self.samples
}
pub const fn sample_rate(&self) -> u32 {
self.sample_rate
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq)]
pub enum VideoSampling {
#[default]
ProcessorDefault,
Fps(f64),
FrameCount(usize),
All,
}
#[derive(Debug, Clone, PartialEq)]
pub struct Video {
frames: Vec<RgbImage>,
source_fps: Option<f64>,
sampling: VideoSampling,
}
impl Video {
pub fn new(
frames: Vec<RgbImage>,
source_fps: Option<f64>,
sampling: VideoSampling,
) -> Result<Self, MediaRequestError> {
if frames.is_empty() {
return Err(MediaRequestError::InvalidVideo(
"video must contain at least one frame".into(),
));
}
if source_fps.is_some_and(|fps| !fps.is_finite() || fps <= 0.0) {
return Err(MediaRequestError::InvalidVideo(
"source frame rate must be finite and positive".into(),
));
}
match sampling {
VideoSampling::Fps(fps) if !fps.is_finite() || fps <= 0.0 => {
return Err(MediaRequestError::InvalidVideo(
"sampling frame rate must be finite and positive".into(),
));
}
VideoSampling::FrameCount(0) => {
return Err(MediaRequestError::InvalidVideo(
"sampling frame count must be positive".into(),
));
}
_ => {}
}
Ok(Self {
frames,
source_fps,
sampling,
})
}
pub fn frames(&self) -> &[RgbImage] {
&self.frames
}
pub const fn source_fps(&self) -> Option<f64> {
self.source_fps
}
pub const fn sampling(&self) -> VideoSampling {
self.sampling
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum Media {
Image(RgbImage),
Video(Video),
Audio(Audio),
}
#[derive(Debug, Clone, PartialEq)]
pub enum MultimodalSegment {
Text(String),
TokenIds(Vec<u32>),
Media(Media),
}
#[derive(Debug, Clone, PartialEq)]
pub struct MultimodalRequest {
segments: Vec<MultimodalSegment>,
}
impl MultimodalRequest {
pub fn new(segments: Vec<MultimodalSegment>) -> Result<Self, MediaRequestError> {
if segments.is_empty() {
return Err(MediaRequestError::EmptyRequest);
}
Ok(Self { segments })
}
pub fn from_chat(
rendered_prompt: &str,
bindings: &[MediaBinding],
) -> Result<Self, MediaRequestError> {
validate_bindings(rendered_prompt, bindings)?;
let mut segments = Vec::with_capacity(bindings.len().saturating_mul(2) + 1);
let mut cursor = 0;
for (index, binding) in bindings.iter().enumerate() {
let remainder = &rendered_prompt[cursor..];
let relative = remainder.find(binding.placeholder()).ok_or_else(|| {
MediaRequestError::PlaceholderOrder {
index,
placeholder: binding.placeholder().into(),
}
})?;
let start = cursor + relative;
if start > cursor {
segments.push(MultimodalSegment::Text(
rendered_prompt[cursor..start].into(),
));
}
segments.push(MultimodalSegment::Media(binding.media().clone()));
cursor = start + binding.placeholder().len();
}
if cursor < rendered_prompt.len() {
segments.push(MultimodalSegment::Text(rendered_prompt[cursor..].into()));
}
Self::new(segments)
}
pub fn segments(&self) -> &[MultimodalSegment] {
&self.segments
}
pub fn tokenize<E>(
&self,
mut encode: impl FnMut(&str) -> Result<Vec<u32>, E>,
) -> Result<TokenizedMultimodalRequest, E> {
let mut segments = Vec::with_capacity(self.segments.len());
for segment in &self.segments {
segments.push(match segment {
MultimodalSegment::Text(text) => {
TokenizedMultimodalSegment::TokenIds(encode(text)?)
}
MultimodalSegment::TokenIds(ids) => {
TokenizedMultimodalSegment::TokenIds(ids.clone())
}
MultimodalSegment::Media(media) => TokenizedMultimodalSegment::Media(media.clone()),
});
}
Ok(TokenizedMultimodalRequest { segments })
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct MediaBinding {
placeholder: String,
media: Media,
}
impl MediaBinding {
pub fn new(placeholder: impl Into<String>, media: Media) -> Self {
Self {
placeholder: placeholder.into(),
media,
}
}
pub fn placeholder(&self) -> &str {
&self.placeholder
}
pub const fn media(&self) -> &Media {
&self.media
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum TokenizedMultimodalSegment {
TokenIds(Vec<u32>),
Media(Media),
}
#[derive(Debug, Clone, PartialEq)]
pub struct TokenizedMultimodalRequest {
segments: Vec<TokenizedMultimodalSegment>,
}
impl TokenizedMultimodalRequest {
pub fn segments(&self) -> &[TokenizedMultimodalSegment] {
&self.segments
}
pub fn semantic_content_fingerprint(&self) -> String {
let mut digest = Sha256::new();
digest.update(b"eredu-tokenized-multimodal-input-v1\0");
digest.update((self.segments.len() as u64).to_le_bytes());
for segment in &self.segments {
match segment {
TokenizedMultimodalSegment::TokenIds(tokens) => {
digest.update([0]);
digest.update((tokens.len() as u64).to_le_bytes());
for token in tokens {
digest.update(token.to_le_bytes());
}
}
TokenizedMultimodalSegment::Media(media) => hash_media(&mut digest, media),
}
}
let bytes = digest.finalize();
let mut encoded = String::with_capacity(71);
encoded.push_str("sha256:");
for byte in bytes {
use std::fmt::Write as _;
write!(&mut encoded, "{byte:02x}").expect("writing to a string cannot fail");
}
encoded
}
}
fn hash_image(digest: &mut Sha256, image: &RgbImage) {
digest.update(image.width.to_le_bytes());
digest.update(image.height.to_le_bytes());
digest.update((image.pixels.len() as u64).to_le_bytes());
digest.update(&image.pixels);
}
fn hash_media(digest: &mut Sha256, media: &Media) {
match media {
Media::Image(image) => {
digest.update([1]);
hash_image(digest, image);
}
Media::Video(video) => {
digest.update([2]);
digest.update((video.frames.len() as u64).to_le_bytes());
match video.source_fps {
Some(fps) => {
digest.update([1]);
digest.update(fps.to_bits().to_le_bytes());
}
None => digest.update([0]),
}
match video.sampling {
VideoSampling::ProcessorDefault => digest.update([0]),
VideoSampling::Fps(fps) => {
digest.update([1]);
digest.update(fps.to_bits().to_le_bytes());
}
VideoSampling::FrameCount(count) => {
digest.update([2]);
digest.update((count as u64).to_le_bytes());
}
VideoSampling::All => digest.update([3]),
}
for frame in &video.frames {
hash_image(digest, frame);
}
}
Media::Audio(audio) => {
digest.update([3]);
digest.update(audio.sample_rate.to_le_bytes());
digest.update((audio.samples.len() as u64).to_le_bytes());
for sample in &audio.samples {
digest.update(sample.to_bits().to_le_bytes());
}
}
}
}
fn validate_bindings(
rendered_prompt: &str,
bindings: &[MediaBinding],
) -> Result<(), MediaRequestError> {
for (index, binding) in bindings.iter().enumerate() {
if binding.placeholder.is_empty() {
return Err(MediaRequestError::EmptyPlaceholder { index });
}
if bindings[..index]
.iter()
.any(|earlier| earlier.placeholder == binding.placeholder)
{
continue;
}
let expected = bindings
.iter()
.filter(|candidate| candidate.placeholder == binding.placeholder)
.count();
let actual = rendered_prompt.matches(&binding.placeholder).count();
if actual != expected {
return Err(MediaRequestError::PlaceholderCount {
placeholder: binding.placeholder.clone(),
expected,
actual,
});
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
fn image(value: u8) -> Media {
Media::Image(RgbImage::new(vec![value; 3], 1, 1).unwrap())
}
#[test]
fn chat_composition_and_tokenization_preserve_exact_order() {
let request = MultimodalRequest::from_chat(
"before<image>middle<image>after",
&[
MediaBinding::new("<image>", image(1)),
MediaBinding::new("<image>", image(2)),
],
)
.unwrap();
let tokenized = request
.tokenize::<std::convert::Infallible>(|text| {
Ok(text
.as_bytes()
.iter()
.map(|byte| u32::from(*byte))
.collect())
})
.unwrap();
assert_eq!(tokenized.segments().len(), 5);
assert!(matches!(
&tokenized.segments()[0],
TokenizedMultimodalSegment::TokenIds(ids)
if ids == &[98, 101, 102, 111, 114, 101]
));
assert!(matches!(
&tokenized.segments()[1],
TokenizedMultimodalSegment::Media(Media::Image(image)) if image.pixels() == [1, 1, 1]
));
assert!(matches!(
&tokenized.segments()[3],
TokenizedMultimodalSegment::Media(Media::Image(image)) if image.pixels() == [2, 2, 2]
));
}
#[test]
fn validation_rejects_bad_media_and_placeholder_contracts() {
assert!(matches!(
RgbImage::new(vec![0; 2], 1, 1),
Err(MediaRequestError::InvalidRgbImage { .. })
));
assert!(Audio::new(vec![f32::NAN], 16_000).is_err());
assert!(Video::new(Vec::new(), None, VideoSampling::All).is_err());
assert!(matches!(
MultimodalRequest::from_chat(
"<image>",
&[
MediaBinding::new("<image>", image(1)),
MediaBinding::new("<image>", image(2)),
],
),
Err(MediaRequestError::PlaceholderCount { .. })
));
}
#[test]
fn request_preserves_text_only_segments() {
let request = MultimodalRequest::new(vec![
MultimodalSegment::TokenIds(vec![1, 2]),
MultimodalSegment::Text("tail".into()),
])
.unwrap();
assert_eq!(request.segments().len(), 2);
assert!(matches!(
&request.segments()[0],
MultimodalSegment::TokenIds(ids) if ids == &[1, 2]
));
}
#[test]
fn semantic_fingerprint_is_stable_and_distinguishes_equal_geometry_payloads() {
fn request(image_value: u8, audio_value: f32) -> TokenizedMultimodalRequest {
MultimodalRequest::new(vec![
MultimodalSegment::TokenIds(vec![7, 11]),
MultimodalSegment::Media(Media::Image(
RgbImage::new(vec![image_value; 12], 2, 2).unwrap(),
)),
MultimodalSegment::Media(Media::Audio(
Audio::new(vec![audio_value; 4], 16_000).unwrap(),
)),
])
.unwrap()
.tokenize::<std::convert::Infallible>(|_| unreachable!())
.unwrap()
}
let first = request(3, 0.25);
let same = request(3, 0.25);
let changed_image = request(4, 0.25);
let changed_audio = request(3, 0.5);
let fingerprint = first.semantic_content_fingerprint();
assert_eq!(fingerprint, same.semantic_content_fingerprint());
assert_eq!(fingerprint, first.semantic_content_fingerprint());
assert_ne!(fingerprint, changed_image.semantic_content_fingerprint());
assert_ne!(fingerprint, changed_audio.semantic_content_fingerprint());
assert_eq!(fingerprint.len(), 71);
assert!(fingerprint.starts_with("sha256:"));
}
}