#[derive(Debug, Clone, PartialEq, thiserror::Error)]
pub enum MediaRequestError {
#[error("multimodal input request is empty")]
EmptyRequest,
#[error("multimodal input request contains no media")]
MissingMedia,
#[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);
}
if !segments
.iter()
.any(|segment| matches!(segment, MultimodalSegment::Media(_)))
{
return Err(MediaRequestError::MissingMedia);
}
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
}
}
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 { .. })
));
}
}