tokio_prompt_orchestrator/
multi_modal.rs1use serde::{Deserialize, Serialize};
9
10#[derive(Serialize, Deserialize)]
13struct TextBody<T> {
14 text: T,
15}
16
17fn ser_text<S: serde::Serializer>(text: &str, s: S) -> Result<S::Ok, S::Error> {
18 TextBody { text }.serialize(s)
19}
20
21fn de_text<'de, D: serde::Deserializer<'de>>(d: D) -> Result<String, D::Error> {
22 TextBody::<String>::deserialize(d).map(|b| b.text)
23}
24
25#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
27#[serde(tag = "type", rename_all = "snake_case")]
28pub enum ContentPart {
29 #[serde(serialize_with = "ser_text", deserialize_with = "de_text")]
31 Text(String),
32
33 ImageUrl {
35 url: String,
37 alt_text: Option<String>,
39 width: Option<u32>,
41 height: Option<u32>,
43 },
44
45 CodeSnippet {
47 language: String,
49 code: String,
51 filename: Option<String>,
53 },
54
55 DataTable {
57 headers: Vec<String>,
59 rows: Vec<Vec<String>>,
61 },
62
63 AudioTranscript {
65 text: String,
67 language: String,
69 confidence: f64,
71 },
72}
73
74#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
76pub struct MultiModalContent {
77 pub parts: Vec<ContentPart>,
79}
80
81impl MultiModalContent {
82 pub fn new() -> Self {
84 Self::default()
85 }
86
87 pub fn text_only(&self) -> String {
90 self.parts
91 .iter()
92 .filter_map(|p| match p {
93 ContentPart::Text(t) => Some(t.as_str()),
94 ContentPart::AudioTranscript { text, .. } => Some(text.as_str()),
95 _ => None,
96 })
97 .collect::<Vec<_>>()
98 .join(" ")
99 }
100
101 pub fn token_estimate(&self) -> usize {
108 self.parts.iter().map(|p| match p {
109 ContentPart::Text(t) => {
110 let words = t.split_whitespace().count();
111 ((words as f64) / 0.75).ceil() as usize
112 }
113 ContentPart::AudioTranscript { text, .. } => {
114 let words = text.split_whitespace().count();
115 ((words as f64) / 0.75).ceil() as usize
116 }
117 ContentPart::ImageUrl { .. } => 85,
118 ContentPart::CodeSnippet { code, .. } => (code.len() / 4).max(1),
119 ContentPart::DataTable { headers, rows } => {
120 let header_chars: usize = headers.iter().map(|h| h.len()).sum();
121 let row_chars: usize = rows.iter().flat_map(|r| r.iter()).map(|c| c.len()).sum();
122 ((header_chars + row_chars) / 4).max(1)
123 }
124 }).sum()
125 }
126
127 pub fn has_images(&self) -> bool {
129 self.parts.iter().any(|p| matches!(p, ContentPart::ImageUrl { .. }))
130 }
131
132 pub fn has_code(&self) -> bool {
134 self.parts.iter().any(|p| matches!(p, ContentPart::CodeSnippet { .. }))
135 }
136
137 pub fn tables(&self) -> Vec<(&Vec<String>, &Vec<Vec<String>>)> {
139 self.parts
140 .iter()
141 .filter_map(|p| match p {
142 ContentPart::DataTable { headers, rows } => Some((headers, rows)),
143 _ => None,
144 })
145 .collect()
146 }
147}
148
149#[derive(Debug, Default)]
163pub struct ContentBuilder {
164 parts: Vec<ContentPart>,
165}
166
167impl ContentBuilder {
168 pub fn new() -> Self {
170 Self::default()
171 }
172
173 pub fn add_text(mut self, text: impl Into<String>) -> Self {
175 self.parts.push(ContentPart::Text(text.into()));
176 self
177 }
178
179 pub fn add_image(
181 mut self,
182 url: impl Into<String>,
183 alt_text: Option<impl Into<String>>,
184 width: Option<u32>,
185 height: Option<u32>,
186 ) -> Self {
187 self.parts.push(ContentPart::ImageUrl {
188 url: url.into(),
189 alt_text: alt_text.map(|a| a.into()),
190 width,
191 height,
192 });
193 self
194 }
195
196 pub fn add_code(
198 mut self,
199 language: impl Into<String>,
200 code: impl Into<String>,
201 filename: Option<impl Into<String>>,
202 ) -> Self {
203 self.parts.push(ContentPart::CodeSnippet {
204 language: language.into(),
205 code: code.into(),
206 filename: filename.map(|f| f.into()),
207 });
208 self
209 }
210
211 pub fn add_table(
213 mut self,
214 headers: Vec<String>,
215 rows: Vec<Vec<String>>,
216 ) -> Self {
217 self.parts.push(ContentPart::DataTable { headers, rows });
218 self
219 }
220
221 pub fn add_audio(
223 mut self,
224 text: impl Into<String>,
225 language: impl Into<String>,
226 confidence: f64,
227 ) -> Self {
228 self.parts.push(ContentPart::AudioTranscript {
229 text: text.into(),
230 language: language.into(),
231 confidence,
232 });
233 self
234 }
235
236 pub fn build(self) -> MultiModalContent {
238 MultiModalContent { parts: self.parts }
239 }
240}
241
242pub struct ContentSerializer;
246
247impl ContentSerializer {
248 pub fn serialize(content: &MultiModalContent) -> Result<String, String> {
253 serde_json::to_string(content).map_err(|e| e.to_string())
254 }
255
256 pub fn deserialize(json: &str) -> Result<MultiModalContent, String> {
259 serde_json::from_str(json).map_err(|e| e.to_string())
260 }
261}
262
263#[cfg(test)]
264mod tests {
265 use super::*;
266
267 #[test]
270 fn text_only_collects_text_and_audio() {
271 let content = ContentBuilder::new()
272 .add_text("Hello")
273 .add_image("https://x.com/img.png", None::<String>, None, None)
274 .add_audio("world", "en-US", 0.99)
275 .build();
276 assert_eq!(content.text_only(), "Hello world");
277 }
278
279 #[test]
280 fn text_only_empty_when_no_text_parts() {
281 let content = ContentBuilder::new()
282 .add_image("https://x.com/img.png", None::<String>, None, None)
283 .build();
284 assert_eq!(content.text_only(), "");
285 }
286
287 #[test]
290 fn has_images_true_when_image_present() {
291 let c = ContentBuilder::new()
292 .add_image("https://x.com/a.png", None::<String>, None, None)
293 .build();
294 assert!(c.has_images());
295 }
296
297 #[test]
298 fn has_images_false_when_no_images() {
299 let c = ContentBuilder::new().add_text("hello").build();
300 assert!(!c.has_images());
301 }
302
303 #[test]
304 fn has_code_true_when_code_present() {
305 let c = ContentBuilder::new()
306 .add_code("rust", "fn main() {}", None::<String>)
307 .build();
308 assert!(c.has_code());
309 }
310
311 #[test]
312 fn has_code_false_when_no_code() {
313 let c = ContentBuilder::new().add_text("hello").build();
314 assert!(!c.has_code());
315 }
316
317 #[test]
320 fn tables_returns_all_data_tables() {
321 let c = ContentBuilder::new()
322 .add_text("intro")
323 .add_table(
324 vec!["Name".into(), "Age".into()],
325 vec![vec!["Alice".into(), "30".into()]],
326 )
327 .add_table(
328 vec!["X".into()],
329 vec![],
330 )
331 .build();
332 assert_eq!(c.tables().len(), 2);
333 }
334
335 #[test]
336 fn tables_empty_when_no_tables() {
337 let c = ContentBuilder::new().add_text("hello").build();
338 assert!(c.tables().is_empty());
339 }
340
341 #[test]
344 fn token_estimate_text() {
345 let c = ContentBuilder::new()
347 .add_text("one two three four")
348 .build();
349 assert_eq!(c.token_estimate(), 6);
350 }
351
352 #[test]
353 fn token_estimate_image() {
354 let c = ContentBuilder::new()
355 .add_image("https://x.com/img.png", None::<String>, None, None)
356 .build();
357 assert_eq!(c.token_estimate(), 85);
358 }
359
360 #[test]
361 fn token_estimate_code() {
362 let c = ContentBuilder::new()
364 .add_code("python", "12345678", None::<String>)
365 .build();
366 assert_eq!(c.token_estimate(), 2);
367 }
368
369 #[test]
370 fn token_estimate_combined() {
371 let c = ContentBuilder::new()
372 .add_text("hello world") .add_image("https://x.com/a.png", None::<String>, None, None) .build();
375 assert_eq!(c.token_estimate(), 3 + 85);
376 }
377
378 #[test]
381 fn builder_produces_correct_part_count() {
382 let c = ContentBuilder::new()
383 .add_text("t")
384 .add_image("u", None::<String>, Some(800), Some(600))
385 .add_code("rs", "fn f() {}", Some("lib.rs"))
386 .add_table(vec!["H".into()], vec![])
387 .add_audio("transcript", "en", 0.9)
388 .build();
389 assert_eq!(c.parts.len(), 5);
390 }
391
392 #[test]
393 fn builder_image_stores_dimensions() {
394 let c = ContentBuilder::new()
395 .add_image("https://x.com/a.png", Some("alt"), Some(1920), Some(1080))
396 .build();
397 if let ContentPart::ImageUrl { width, height, alt_text, .. } = &c.parts[0] {
398 assert_eq!(*width, Some(1920));
399 assert_eq!(*height, Some(1080));
400 assert_eq!(alt_text.as_deref(), Some("alt"));
401 } else {
402 panic!("expected ImageUrl part");
403 }
404 }
405
406 #[test]
409 fn serialize_deserialize_roundtrip() {
410 let original = ContentBuilder::new()
411 .add_text("Hello world")
412 .add_image("https://x.com/img.png", Some("desc"), Some(640), Some(480))
413 .add_code("rust", "fn main() {}", Some("main.rs"))
414 .add_table(
415 vec!["Col1".into(), "Col2".into()],
416 vec![vec!["A".into(), "B".into()]],
417 )
418 .add_audio("audio text", "en-GB", 0.95)
419 .build();
420
421 let json = ContentSerializer::serialize(&original).expect("serialize must succeed");
422 let recovered = ContentSerializer::deserialize(&json).expect("deserialize must succeed");
423 assert_eq!(original, recovered);
424 }
425
426 #[test]
427 fn deserialize_invalid_json_returns_error() {
428 assert!(ContentSerializer::deserialize("{not valid json}").is_err());
429 }
430
431 #[test]
432 fn serialize_empty_content() {
433 let c = MultiModalContent::new();
434 let json = ContentSerializer::serialize(&c).expect("serialize must succeed");
435 let back = ContentSerializer::deserialize(&json).expect("deserialize must succeed");
436 assert_eq!(c, back);
437 }
438}