Skip to main content

robit_agent/
media.rs

1//! Media handling utilities: download, encode, etc.
2
3use std::path::{Path, PathBuf};
4
5use base64::{engine::general_purpose, Engine as _};
6use thiserror::Error;
7
8/// Errors that can occur while handling media.
9#[derive(Debug, Error)]
10pub enum MediaError {
11    #[error("HTTP error: {0}")]
12    Http(#[from] reqwest::Error),
13    #[error("IO error: {0}")]
14    Io(#[from] std::io::Error),
15    #[error("Invalid media content: empty or corrupted")]
16    InvalidContent,
17    #[error("Image processing error: {0}")]
18    Image(String),
19}
20
21/// Download media from URL and save to the specified directory.
22///
23/// Returns the path to the saved file.
24pub async fn download_media(
25    url: &str,
26    filename: Option<&str>,
27    save_dir: &PathBuf,
28) -> Result<PathBuf, MediaError> {
29    // Create directory if it doesn't exist
30    tokio::fs::create_dir_all(save_dir).await?;
31
32    // Determine filename
33    let save_filename = match filename {
34        Some(s) => s.to_string(),
35        None => uuid::Uuid::new_v4().to_string(),
36    };
37    let save_path = save_dir.join(save_filename);
38
39    // Download
40    let client = reqwest::Client::new();
41    let response = client.get(url).send().await?;
42    let bytes = response.bytes().await?;
43
44    if bytes.is_empty() {
45        return Err(MediaError::InvalidContent);
46    }
47
48    tokio::fs::write(&save_path, &bytes).await?;
49
50    Ok(save_path)
51}
52
53/// An image (or other media) encoded as a base64 data URL, plus info about
54/// any compression applied. `compression.bytes` is emptied after the data
55/// URL is built to avoid holding a second copy of the payload.
56#[derive(Debug)]
57pub struct EncodedImage {
58    pub data_url: String,
59    pub compression: CompressedImage,
60}
61
62/// Download media from URL and encode as base64 data URL, compressing
63/// images first (see [`compress_image_bytes`]).
64pub async fn download_and_encode_base64(
65    url: &str,
66    content_type: &str,
67    max_image_dim: u32,
68) -> Result<EncodedImage, MediaError> {
69    let client = reqwest::Client::new();
70    let bytes = client.get(url).send().await?.bytes().await?;
71
72    if bytes.is_empty() {
73        return Err(MediaError::InvalidContent);
74    }
75
76    encode_image_bytes(&bytes, content_type, max_image_dim).await
77}
78
79/// Read a local file and encode it as a base64 data URL, compressing images
80/// first (see [`compress_image_bytes`]). The MIME type is inferred from the
81/// file extension.
82pub async fn encode_file_base64(
83    path: &Path,
84    max_image_dim: u32,
85) -> Result<EncodedImage, MediaError> {
86    let bytes = tokio::fs::read(path).await?;
87
88    if bytes.is_empty() {
89        return Err(MediaError::InvalidContent);
90    }
91
92    let mime_type = mime_from_extension(path);
93    encode_image_bytes(&bytes, mime_type, max_image_dim).await
94}
95
96/// Compress (if an image and enabled) and base64-encode raw media bytes.
97async fn encode_image_bytes(
98    bytes: &[u8],
99    mime: &str,
100    max_image_dim: u32,
101) -> Result<EncodedImage, MediaError> {
102    let mut compression = if mime.starts_with("image/") {
103        let owned = bytes.to_vec();
104        let mime = mime.to_string();
105        // Decoding/resizing/encoding is pure CPU work on multi-MB payloads —
106        // keep it off the async runtime threads.
107        tokio::task::spawn_blocking(move || compress_image_bytes(&owned, &mime, max_image_dim))
108            .await
109            .map_err(|e| MediaError::Image(format!("compression task failed: {}", e)))??
110    } else {
111        CompressedImage {
112            bytes: bytes.to_vec(),
113            mime: mime.to_string(),
114            orig_dims: (0, 0),
115            new_dims: (0, 0),
116            kept_original: true,
117        }
118    };
119
120    let data_url = format!(
121        "data:{};base64,{}",
122        compression.mime,
123        general_purpose::STANDARD.encode(&compression.bytes)
124    );
125    // The bytes now only exist inside the data URL; drop the copy.
126    compression.bytes = Vec::new();
127    Ok(EncodedImage {
128        data_url,
129        compression,
130    })
131}
132
133/// Infer a MIME type from the file extension.
134///
135/// Falls back to `application/octet-stream` for unknown extensions.
136fn mime_from_extension(path: &Path) -> &'static str {
137    match path
138        .extension()
139        .and_then(|e| e.to_str())
140        .map(|e| e.to_ascii_lowercase())
141        .as_deref()
142    {
143        Some("png") => "image/png",
144        Some("jpg") | Some("jpeg") => "image/jpeg",
145        Some("gif") => "image/gif",
146        Some("webp") => "image/webp",
147        _ => "application/octet-stream",
148    }
149}
150
151/// Result of preparing an image for the vision-model context.
152#[derive(Debug)]
153pub struct CompressedImage {
154    /// Encoded image bytes (possibly the untouched original).
155    pub bytes: Vec<u8>,
156    /// MIME type of `bytes`.
157    pub mime: String,
158    /// Original (width, height); (0, 0) when unknown (passthrough).
159    pub orig_dims: (u32, u32),
160    /// Final (width, height); (0, 0) when unknown (passthrough).
161    pub new_dims: (u32, u32),
162    /// True when the original bytes were kept (GIF, compression disabled,
163    /// or re-encoding would have grown the file).
164    pub kept_original: bool,
165}
166
167/// Downscale and re-encode an image for the vision-model context.
168///
169/// Base64 image payloads are re-sent with every LLM call, so multi-MB 2K
170/// images quickly blow past provider-gateway request-body limits (HTTP 413).
171/// This function proportionally downscales the image so its longest side is
172/// at most `max_dim`, flattens any alpha channel onto white, and re-encodes
173/// as JPEG (quality 85) — typically a 10x size reduction for generated 2K
174/// PNGs with no practical loss for vision models, which ingest images at
175/// roughly this resolution anyway.
176///
177/// Pass-through cases (original bytes kept unchanged):
178/// - `max_dim == 0` (compression disabled)
179/// - GIF (re-encoding would keep only the first frame of an animation)
180/// - re-encoding produced a LARGER file (image was already well optimized)
181/// Decode guards: images whose width or height exceeds this are rejected
182/// before any pixel buffer is allocated. File size says nothing about the
183/// decoded size — a tiny, highly compressible PNG can decompress to hundreds
184/// of MB (decode bomb), enough to OOM a resident bot process.
185const MAX_DECODE_DIMENSION: u32 = 16384;
186/// Hard cap on total allocation during decode (256MB).
187const MAX_DECODE_ALLOC: u64 = 256 * 1024 * 1024;
188
189fn compress_image_bytes(bytes: &[u8], mime: &str, max_dim: u32) -> Result<CompressedImage, MediaError> {
190    if max_dim == 0 || mime == "image/gif" {
191        return Ok(CompressedImage {
192            bytes: bytes.to_vec(),
193            mime: mime.to_string(),
194            orig_dims: (0, 0),
195            new_dims: (0, 0),
196            kept_original: true,
197        });
198    }
199
200    let mut limits = image::Limits::default();
201    limits.max_image_width = Some(MAX_DECODE_DIMENSION);
202    limits.max_image_height = Some(MAX_DECODE_DIMENSION);
203    limits.max_alloc = Some(MAX_DECODE_ALLOC);
204
205    let mut reader = image::ImageReader::new(std::io::Cursor::new(bytes))
206        .with_guessed_format()
207        .map_err(|e| MediaError::Image(format!("format detection failed: {}", e)))?;
208    reader.limits(limits);
209    let img = reader
210        .decode()
211        .map_err(|e| MediaError::Image(format!("decode failed: {}", e)))?;
212    let orig_dims = (img.width(), img.height());
213
214    // JPEG has no alpha channel — flatten transparency onto white.
215    let rgb = flatten_alpha_to_white(img);
216
217    let final_img = if orig_dims.0.max(orig_dims.1) > max_dim {
218        let longest = orig_dims.0.max(orig_dims.1) as f32;
219        let scale = max_dim as f32 / longest;
220        let nw = ((orig_dims.0 as f32 * scale).round() as u32).max(1);
221        let nh = ((orig_dims.1 as f32 * scale).round() as u32).max(1);
222        image::imageops::resize(&rgb, nw, nh, image::imageops::FilterType::Lanczos3)
223    } else {
224        rgb
225    };
226
227    let mut jpeg = Vec::new();
228    {
229        use image::ImageEncoder as _;
230        let encoder = image::codecs::jpeg::JpegEncoder::new_with_quality(&mut jpeg, 85);
231        encoder
232            .write_image(
233                final_img.as_raw(),
234                final_img.width(),
235                final_img.height(),
236                image::ExtendedColorType::Rgb8,
237            )
238            .map_err(|e| MediaError::Image(format!("JPEG encode failed: {}", e)))?;
239    }
240
241    // Keep the original when re-encoding grew the file.
242    if jpeg.len() >= bytes.len() {
243        return Ok(CompressedImage {
244            bytes: bytes.to_vec(),
245            mime: mime.to_string(),
246            orig_dims,
247            new_dims: orig_dims,
248            kept_original: true,
249        });
250    }
251
252    Ok(CompressedImage {
253        bytes: jpeg,
254        mime: "image/jpeg".to_string(),
255        orig_dims,
256        new_dims: (final_img.width(), final_img.height()),
257        kept_original: false,
258    })
259}
260
261/// Composite an image's alpha channel onto a white background, returning RGB.
262fn flatten_alpha_to_white(img: image::DynamicImage) -> image::RgbImage {
263    use image::GenericImageView;
264    let (w, h) = img.dimensions();
265    let mut out = image::RgbImage::new(w, h);
266    for (x, y, p) in img.pixels() {
267        let a = p.0[3] as f32 / 255.0;
268        let blend = |c: u8| ((c as f32 * a) + (255.0 * (1.0 - a))).round() as u8;
269        out.put_pixel(x, y, image::Rgb([blend(p.0[0]), blend(p.0[1]), blend(p.0[2])]));
270    }
271    out
272}
273
274#[cfg(test)]
275mod tests {
276    use super::*;
277
278    /// Deterministic pseudo-random RGB PNG. Noise barely compresses in PNG,
279    /// so the downscaled JPEG is guaranteed to be much smaller — this keeps
280    /// the size assertions meaningful and reproducible.
281    fn noise_png(w: u32, h: u32) -> Vec<u8> {
282        let mut img = image::RgbImage::new(w, h);
283        let mut seed: u32 = 0x1234_5678;
284        for (_, _, p) in img.enumerate_pixels_mut() {
285            seed = seed.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
286            *p = image::Rgb([(seed >> 16) as u8, (seed >> 8) as u8, seed as u8]);
287        }
288        let mut buf = std::io::Cursor::new(Vec::new());
289        image::DynamicImage::ImageRgb8(img)
290            .write_to(&mut buf, image::ImageFormat::Png)
291            .unwrap();
292        buf.into_inner()
293    }
294
295    fn solid_png(w: u32, h: u32) -> Vec<u8> {
296        let img = image::RgbImage::from_pixel(w, h, image::Rgb([200, 30, 30]));
297        let mut buf = std::io::Cursor::new(Vec::new());
298        image::DynamicImage::ImageRgb8(img)
299            .write_to(&mut buf, image::ImageFormat::Png)
300            .unwrap();
301        buf.into_inner()
302    }
303
304    fn noise_gif(w: u32, h: u32) -> Vec<u8> {
305        let mut img = image::RgbImage::new(w, h);
306        let mut seed: u32 = 0xDEAD_BEEF;
307        for (_, _, p) in img.enumerate_pixels_mut() {
308            seed = seed.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
309            *p = image::Rgb([(seed >> 16) as u8, (seed >> 8) as u8, seed as u8]);
310        }
311        let mut buf = std::io::Cursor::new(Vec::new());
312        image::DynamicImage::ImageRgb8(img)
313            .write_to(&mut buf, image::ImageFormat::Gif)
314            .unwrap();
315        buf.into_inner()
316    }
317
318    #[test]
319    fn big_image_is_downscaled_and_reencoded_as_jpeg() {
320        let orig = noise_png(2048, 2048);
321        let out = compress_image_bytes(&orig, "image/png", 1024).unwrap();
322        assert_eq!(out.mime, "image/jpeg");
323        assert_eq!(out.orig_dims, (2048, 2048));
324        assert_eq!(out.new_dims, (1024, 1024));
325        assert!(!out.kept_original);
326        // JPEG magic bytes
327        assert_eq!(&out.bytes[0..2], &[0xFF, 0xD8]);
328        assert!(
329            out.bytes.len() * 10 < orig.len(),
330            "2048px noise PNG should shrink >10x as a 1024px JPEG: {} -> {} bytes",
331            orig.len(),
332            out.bytes.len()
333        );
334    }
335
336    #[test]
337    fn small_image_keeps_dimensions_still_reencodes() {
338        let orig = noise_png(800, 600);
339        let out = compress_image_bytes(&orig, "image/png", 1024).unwrap();
340        assert_eq!(out.new_dims, (800, 600));
341        assert_eq!(out.mime, "image/jpeg");
342        assert!(!out.kept_original);
343    }
344
345    #[test]
346    fn gif_passes_through_unchanged() {
347        // Re-encoding a GIF keeps only the first frame, so GIFs must never
348        // go through the compression path.
349        let orig = noise_gif(32, 32);
350        let out = compress_image_bytes(&orig, "image/gif", 1024).unwrap();
351        assert!(out.kept_original);
352        assert_eq!(out.bytes, orig);
353        assert_eq!(out.mime, "image/gif");
354    }
355
356    #[test]
357    fn tiny_image_keeps_original_when_reencoding_would_grow() {
358        let orig = solid_png(8, 8);
359        let out = compress_image_bytes(&orig, "image/png", 1024).unwrap();
360        assert!(out.kept_original, "tiny PNG should not be replaced by a larger JPEG");
361        assert_eq!(out.bytes, orig);
362        assert_eq!(out.mime, "image/png");
363    }
364
365    #[test]
366    fn transparent_pixels_flatten_to_white() {
367        // 256px RGBA noise under full transparency: the PNG is large (noise)
368        // while the flattened JPEG is tiny, so the compression path runs.
369        let mut img = image::RgbaImage::new(256, 256);
370        let mut seed: u32 = 0x0BAD_C0DE;
371        for (_, _, p) in img.enumerate_pixels_mut() {
372            seed = seed.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
373            *p = image::Rgba([(seed >> 16) as u8, (seed >> 8) as u8, seed as u8, 0]);
374        }
375        let mut buf = std::io::Cursor::new(Vec::new());
376        image::DynamicImage::ImageRgba8(img)
377            .write_to(&mut buf, image::ImageFormat::Png)
378            .unwrap();
379        let out = compress_image_bytes(&buf.get_ref(), "image/png", 1024).unwrap();
380        assert_eq!(out.mime, "image/jpeg");
381        assert!(!out.kept_original);
382        let decoded = image::load_from_memory(&out.bytes).unwrap();
383        use image::GenericImageView as _;
384        let p = decoded.get_pixel(0, 0);
385        assert!(
386            p.0[0] >= 250 && p.0[1] >= 250 && p.0[2] >= 250,
387            "transparent pixels should flatten to white, got {:?}",
388            p
389        );
390    }
391
392    #[test]
393    fn zero_max_dimension_disables_compression() {
394        let orig = noise_png(2048, 2048);
395        let out = compress_image_bytes(&orig, "image/png", 0).unwrap();
396        assert!(out.kept_original);
397        assert_eq!(out.bytes, orig);
398        assert_eq!(out.mime, "image/png");
399    }
400
401    #[test]
402    fn non_square_images_scale_by_longest_side() {
403        // Portrait: 600×2000 → longest side 2000 → 307×1024
404        let out = compress_image_bytes(&noise_png(600, 2000), "image/png", 1024).unwrap();
405        assert_eq!(out.new_dims, (307, 1024));
406        assert_eq!(out.orig_dims, (600, 2000));
407
408        // Extreme aspect: 10000×10 → 1024×1
409        let out = compress_image_bytes(&noise_png(10000, 10), "image/png", 1024).unwrap();
410        assert_eq!(out.new_dims, (1024, 1));
411    }
412
413    #[test]
414    fn oversized_dimensions_are_rejected_before_decoding() {
415        // A highly compressible 20000×10 PNG is a tiny file, but decoding
416        // must refuse it (decode bomb guard) instead of trusting file size.
417        let bomb = solid_png(20000, 10);
418        let result = compress_image_bytes(&bomb, "image/png", 1024);
419        assert!(
420            result.is_err(),
421            "images wider/taller than the decode cap must be rejected, got Ok"
422        );
423    }
424}