Skip to main content

tokio_prompt_orchestrator/
multi_modal.rs

1//! Multi-modal content handling for prompts and responses.
2//!
3//! [`MultiModalContent`] holds an ordered sequence of [`ContentPart`] values
4//! that can represent text, images, code snippets, data tables, and audio
5//! transcripts. A fluent [`ContentBuilder`] and a JSON
6//! [`ContentSerializer`] are also provided.
7
8use serde::{Deserialize, Serialize};
9
10/// JSON body of a [`ContentPart::Text`] part. Internally tagged enums cannot
11/// hold a bare string, so the text is wrapped in an object.
12#[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/// A single piece of content within a multi-modal message.
26#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
27#[serde(tag = "type", rename_all = "snake_case")]
28pub enum ContentPart {
29    /// Plain text. Serialized as `{"type":"text","text":"..."}`.
30    #[serde(serialize_with = "ser_text", deserialize_with = "de_text")]
31    Text(String),
32
33    /// An image referenced by URL.
34    ImageUrl {
35        /// Fully-qualified image URL.
36        url: String,
37        /// Optional alt-text description.
38        alt_text: Option<String>,
39        /// Optional pixel width.
40        width: Option<u32>,
41        /// Optional pixel height.
42        height: Option<u32>,
43    },
44
45    /// A fenced code block with optional filename metadata.
46    CodeSnippet {
47        /// Programming language identifier (e.g. `"rust"`, `"python"`).
48        language: String,
49        /// The source code.
50        code: String,
51        /// Optional source file name.
52        filename: Option<String>,
53    },
54
55    /// A tabular dataset with headers and rows.
56    DataTable {
57        /// Column header labels.
58        headers: Vec<String>,
59        /// Data rows; each row is a vector of string cells.
60        rows: Vec<Vec<String>>,
61    },
62
63    /// A transcribed audio segment.
64    AudioTranscript {
65        /// Transcribed text.
66        text: String,
67        /// BCP-47 language tag (e.g. `"en-US"`).
68        language: String,
69        /// Transcription confidence in `[0.0, 1.0]`.
70        confidence: f64,
71    },
72}
73
74/// An ordered collection of [`ContentPart`] values.
75#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
76pub struct MultiModalContent {
77    /// Ordered parts comprising this message.
78    pub parts: Vec<ContentPart>,
79}
80
81impl MultiModalContent {
82    /// Create an empty content object.
83    pub fn new() -> Self {
84        Self::default()
85    }
86
87    /// Concatenate all [`ContentPart::Text`] and [`ContentPart::AudioTranscript`]
88    /// parts into a single string separated by spaces.
89    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    /// Rough token estimate for the entire content.
102    ///
103    /// - **Text / AudioTranscript**: `word_count / 0.75` (≈ 1.33 tokens/word)
104    /// - **ImageUrl**: 85 tokens (fixed OpenAI estimate)
105    /// - **CodeSnippet**: `code.len() / 4` (≈ 4 chars/token)
106    /// - **DataTable**: sum of all cell lengths divided by 4
107    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    /// Returns `true` if any part is an [`ContentPart::ImageUrl`].
128    pub fn has_images(&self) -> bool {
129        self.parts.iter().any(|p| matches!(p, ContentPart::ImageUrl { .. }))
130    }
131
132    /// Returns `true` if any part is a [`ContentPart::CodeSnippet`].
133    pub fn has_code(&self) -> bool {
134        self.parts.iter().any(|p| matches!(p, ContentPart::CodeSnippet { .. }))
135    }
136
137    /// Return references to all [`ContentPart::DataTable`] parts.
138    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/// Fluent builder for [`MultiModalContent`].
150///
151/// # Example
152/// ```
153/// use tokio_prompt_orchestrator::multi_modal::ContentBuilder;
154///
155/// let content = ContentBuilder::new()
156///     .add_text("Hello!")
157///     .add_image("https://example.com/img.png", Some("a cat"), None, None)
158///     .build();
159///
160/// assert!(content.has_images());
161/// ```
162#[derive(Debug, Default)]
163pub struct ContentBuilder {
164    parts: Vec<ContentPart>,
165}
166
167impl ContentBuilder {
168    /// Create a new empty builder.
169    pub fn new() -> Self {
170        Self::default()
171    }
172
173    /// Append a plain-text part.
174    pub fn add_text(mut self, text: impl Into<String>) -> Self {
175        self.parts.push(ContentPart::Text(text.into()));
176        self
177    }
178
179    /// Append an image-URL part.
180    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    /// Append a code-snippet part.
197    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    /// Append a data-table part.
212    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    /// Append an audio-transcript part.
222    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    /// Consume the builder and return the [`MultiModalContent`].
237    pub fn build(self) -> MultiModalContent {
238        MultiModalContent { parts: self.parts }
239    }
240}
241
242/// Serialize and deserialize [`MultiModalContent`] to/from JSON strings.
243///
244/// Uses `serde_json` which is already a dependency of this crate.
245pub struct ContentSerializer;
246
247impl ContentSerializer {
248    /// Serialize `content` to a compact JSON string.
249    ///
250    /// Returns an error string if serialization fails (this should not happen
251    /// for well-formed content).
252    pub fn serialize(content: &MultiModalContent) -> Result<String, String> {
253        serde_json::to_string(content).map_err(|e| e.to_string())
254    }
255
256    /// Deserialize a [`MultiModalContent`] from a JSON string previously
257    /// produced by [`ContentSerializer::serialize`].
258    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    // ── ContentPart variants ──────────────────────────────────────────────────
268
269    #[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    // ── has_images / has_code ─────────────────────────────────────────────────
288
289    #[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    // ── tables ────────────────────────────────────────────────────────────────
318
319    #[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    // ── token_estimate ────────────────────────────────────────────────────────
342
343    #[test]
344    fn token_estimate_text() {
345        // "one two three four" → 4 words → ceil(4 / 0.75) = 6
346        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        // code with 8 chars → 8 / 4 = 2 tokens
363        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")  // 2 words → ceil(2/0.75) = 3
373            .add_image("https://x.com/a.png", None::<String>, None, None) // 85
374            .build();
375        assert_eq!(c.token_estimate(), 3 + 85);
376    }
377
378    // ── ContentBuilder ────────────────────────────────────────────────────────
379
380    #[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    // ── ContentSerializer ─────────────────────────────────────────────────────
407
408    #[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}