1use crate::vision::{V_MERGE, V_PATCH, V_PATCH_IN, V_TEMPORAL};
16use base64::Engine as _;
17use image::RgbImage;
18use image::imageops::FilterType;
19
20pub const MIN_PIXELS: usize = 65536;
22pub const MAX_PIXELS: usize = 16_777_216;
23const FACTOR: usize = V_PATCH * V_MERGE; pub struct PreppedImage {
26 pub patches: Vec<f32>,
28 pub gh: usize,
29 pub gw: usize,
30}
31
32impl PreppedImage {
33 pub fn n_tokens(&self) -> usize {
35 self.gh * self.gw / (V_MERGE * V_MERGE)
36 }
37}
38
39fn round_half_even(x: f64) -> f64 {
44 let r = x.round();
45 if (x - x.trunc()).abs() == 0.5 && r % 2.0 != 0.0 {
46 r - x.signum()
47 } else {
48 r
49 }
50}
51
52pub fn smart_resize(h: usize, w: usize) -> Result<(usize, usize), String> {
55 if h < 2 || w < 2 {
56 return Err(format!("image too small: {w}x{h}"));
57 }
58 let ar = h.max(w) as f64 / h.min(w) as f64;
59 if ar > 200.0 {
60 return Err(format!("aspect ratio {ar:.0} exceeds 200"));
61 }
62 let f = FACTOR as f64;
63 let (hf, wf) = (h as f64, w as f64);
64 let mut h_bar = (round_half_even(hf / f) * f).max(f);
65 let mut w_bar = (round_half_even(wf / f) * f).max(f);
66 if h_bar * w_bar > MAX_PIXELS as f64 {
67 let beta = (hf * wf / MAX_PIXELS as f64).sqrt();
68 h_bar = ((hf / beta / f).floor() * f).max(f);
69 w_bar = ((wf / beta / f).floor() * f).max(f);
70 } else if h_bar * w_bar < MIN_PIXELS as f64 {
71 let beta = (MIN_PIXELS as f64 / (hf * wf)).sqrt();
72 h_bar = (hf * beta / f).ceil() * f;
73 w_bar = (wf * beta / f).ceil() * f;
74 }
75 Ok((h_bar as usize, w_bar as usize))
76}
77
78fn decode_frame(bytes: &[u8]) -> Result<(RgbImage, usize, usize), String> {
81 let img = image::load_from_memory(bytes).map_err(|e| format!("image decode: {e}"))?;
82 let rgb = img.to_rgb8();
83 let (w, h) = (rgb.width() as usize, rgb.height() as usize);
84 let (rh, rw) = smart_resize(h, w)?;
85 let resized = image::imageops::resize(&rgb, rw as u32, rh as u32, FilterType::CatmullRom);
86 Ok((resized, rh / V_PATCH, rw / V_PATCH))
87}
88
89fn fill_slot(rows: &mut [f32], frame: &RgbImage, gh: usize, gw: usize, t: usize) {
92 let inv = 1.0f32 / 127.5;
93 for py in 0..gh {
94 for px in 0..gw {
95 let row = &mut rows[(py * gw + px) * V_PATCH_IN..(py * gw + px + 1) * V_PATCH_IN];
96 for c in 0..3 {
97 let base = c * V_TEMPORAL * V_PATCH * V_PATCH + t * V_PATCH * V_PATCH;
98 for ph in 0..V_PATCH {
99 for pw in 0..V_PATCH {
100 let p =
101 frame.get_pixel((px * V_PATCH + pw) as u32, (py * V_PATCH + ph) as u32);
102 row[base + ph * V_PATCH + pw] = p.0[c] as f32 * inv - 1.0;
103 }
104 }
105 }
106 }
107 }
108}
109
110pub fn prep_image_bytes(bytes: &[u8]) -> Result<PreppedImage, String> {
113 let (frame, gh, gw) = decode_frame(bytes)?;
114 let mut patches = vec![0f32; gh * gw * V_PATCH_IN];
115 for t in 0..V_TEMPORAL {
116 fill_slot(&mut patches, &frame, gh, gw, t);
117 }
118 Ok(PreppedImage { patches, gh, gw })
119}
120
121pub fn prep_data_uri(uri: &str) -> Result<PreppedImage, String> {
123 let bytes = decode_data_uri(uri)?;
124 prep_image_bytes(&bytes)
125}
126
127pub struct VisionUnit {
131 pub prep: PreppedImage,
132 pub video: Option<usize>,
134}
135
136pub struct PreppedVideo {
141 pub groups: Vec<PreppedImage>,
142 pub timestamps: Vec<f32>,
143}
144
145pub fn video_max_pixels() -> usize {
148 std::env::var("MEMRA_VIDEO_MAX_PIXELS")
149 .ok()
150 .and_then(|v| v.parse().ok())
151 .unwrap_or(2_097_152)
152}
153pub const VID_MIN_PIXELS: usize = 4096;
154pub const VID_MAX_FRAMES: usize = 32;
156
157pub const GIF_MAX_FRAMES: usize = 512;
169pub const GIF_MAX_TOTAL_PIXELS: usize = 1 << 26; fn smart_resize_video(frames: usize, h: usize, w: usize) -> Result<(usize, usize), String> {
173 if h < 2 || w < 2 {
174 return Err(format!("frame too small: {w}x{h}"));
175 }
176 let ar = h.max(w) as f64 / h.min(w) as f64;
177 if ar > 200.0 {
178 return Err(format!("aspect ratio {ar:.0} exceeds 200"));
179 }
180 let f = FACTOR as f64;
181 let (hf, wf) = (h as f64, w as f64);
182 let t_bar = ((frames as f64 / V_TEMPORAL as f64).round() * V_TEMPORAL as f64).max(2.0);
183 let mut h_bar = (round_half_even(hf / f) * f).max(f);
184 let mut w_bar = (round_half_even(wf / f) * f).max(f);
185 let (min_px, max_px) = (VID_MIN_PIXELS as f64, video_max_pixels() as f64);
186 if t_bar * h_bar * w_bar > max_px {
187 let beta = (frames as f64 * hf * wf / max_px).sqrt();
188 h_bar = ((hf / beta / f).floor() * f).max(f);
189 w_bar = ((wf / beta / f).floor() * f).max(f);
190 } else if t_bar * h_bar * w_bar < min_px {
191 let beta = (min_px / (frames as f64 * hf * wf)).sqrt();
192 h_bar = (hf * beta / f).ceil() * f;
193 w_bar = (wf * beta / f).ceil() * f;
194 }
195 Ok((h_bar as usize, w_bar as usize))
196}
197
198pub fn prep_video_gif(bytes: &[u8]) -> Result<PreppedVideo, String> {
203 use image::AnimationDecoder;
204 use image::ImageDecoder as _;
205 let dec = image::codecs::gif::GifDecoder::new(std::io::Cursor::new(bytes))
206 .map_err(|e| format!("gif decode: {e}"))?;
207 let (cw, ch) = dec.dimensions();
213 let canvas_px = (cw as usize) * (ch as usize);
214 if canvas_px == 0 {
215 return Err("gif has an empty canvas".into());
216 }
217 let max_frames = GIF_MAX_FRAMES.min(GIF_MAX_TOTAL_PIXELS / canvas_px);
218 if max_frames == 0 {
219 return Err(format!(
220 "gif canvas {cw}x{ch} exceeds the decode budget ({GIF_MAX_TOTAL_PIXELS} px)"
221 ));
222 }
223 let mut frames: Vec<(RgbImage, f32)> = Vec::new(); let mut t = 0f32;
225 for fr in dec.into_frames() {
226 if frames.len() >= max_frames {
227 return Err(format!(
228 "gif exceeds the decode budget: more than {max_frames} frames at {cw}x{ch} \
229 (ceiling {GIF_MAX_FRAMES} frames / {GIF_MAX_TOTAL_PIXELS} total px)"
230 ));
231 }
232 let fr = fr.map_err(|e| format!("gif frame: {e}"))?;
233 let (num, den) = fr.delay().numer_denom_ms();
234 let dt = if den == 0 {
235 100.0
236 } else {
237 num as f32 / den as f32
238 } / 1000.0;
239 frames.push((
240 image::DynamicImage::ImageRgba8(fr.into_buffer()).to_rgb8(),
241 t,
242 ));
243 t += dt.max(0.01);
244 }
245 if frames.is_empty() {
246 return Err("gif has no frames".into());
247 }
248 if frames.len() == 1 {
250 let f0 = frames[0].clone();
251 frames.push((f0.0, f0.1));
252 }
253 let total = frames.len();
255 let take = total.min(VID_MAX_FRAMES) & !1;
256 let picked: Vec<usize> = (0..take)
257 .map(|i| i * total / take) .collect();
259 let (h, w) = (frames[0].0.height() as usize, frames[0].0.width() as usize);
260 let (rh, rw) = smart_resize_video(take, h, w)?;
261 let (gh, gw) = (rh / V_PATCH, rw / V_PATCH);
262 let mut groups = Vec::with_capacity(take / 2);
263 let mut timestamps = Vec::with_capacity(take / 2);
264 for g in 0..take / 2 {
265 let (a, b) = (picked[2 * g], picked[2 * g + 1]);
266 let mut patches = vec![0f32; gh * gw * V_PATCH_IN];
267 for (slot, idx) in [(0usize, a), (1usize, b)] {
268 let resized = image::imageops::resize(
269 &frames[idx].0,
270 rw as u32,
271 rh as u32,
272 FilterType::CatmullRom,
273 );
274 fill_slot(&mut patches, &resized, gh, gw, slot);
275 }
276 groups.push(PreppedImage { patches, gh, gw });
277 timestamps.push(frames[a].1);
278 }
279 Ok(PreppedVideo { groups, timestamps })
280}
281
282pub fn decode_data_uri(uri: &str) -> Result<Vec<u8>, String> {
284 let rest = uri
285 .strip_prefix("data:")
286 .ok_or_else(|| "expected data: URI (http fetch requires MEMRA_FETCH_URLS=1)".to_string())?;
287 let (meta, payload) = rest
288 .split_once(',')
289 .ok_or_else(|| "malformed data URI: no comma".to_string())?;
290 if !meta.ends_with(";base64") {
291 return Err("data URI must be base64-encoded".into());
292 }
293 base64::engine::general_purpose::STANDARD
294 .decode(payload.trim())
295 .map_err(|e| format!("base64 decode: {e}"))
296}
297
298#[cfg(test)]
299mod tests {
300 use super::*;
301
302 #[test]
303 fn smart_resize_multiples_and_budget() {
304 let (h, w) = smart_resize(1080, 1920).unwrap();
306 assert_eq!(h % 32, 0);
307 assert_eq!(w % 32, 0);
308 assert!(h * w >= MIN_PIXELS && h * w <= MAX_PIXELS);
309 let (h, w) = smart_resize(64, 64).unwrap();
311 assert!(h * w >= MIN_PIXELS);
312 let (h, w) = smart_resize(8000, 12000).unwrap();
314 assert!(h * w <= MAX_PIXELS);
315 assert!(smart_resize(10, 4000).is_err()); }
317
318 #[test]
319 fn patchify_shape_and_order() {
320 let mut img = RgbImage::new(64, 64);
322 for (x, y, p) in img.enumerate_pixels_mut() {
323 *p = image::Rgb([x as u8, y as u8, 200]);
324 }
325 let mut buf = std::io::Cursor::new(Vec::new());
326 img.write_to(&mut buf, image::ImageFormat::Png).unwrap();
327 let prep = prep_image_bytes(buf.get_ref()).unwrap();
328 assert_eq!(prep.patches.len(), prep.gh * prep.gw * V_PATCH_IN);
329 assert_eq!(prep.gh % V_MERGE, 0);
330 assert_eq!(prep.gw % V_MERGE, 0);
331 let row = &prep.patches[0..V_PATCH_IN];
333 let slot = V_PATCH * V_PATCH;
334 for c in 0..3 {
335 let b = c * V_TEMPORAL * slot;
336 assert_eq!(row[b..b + slot], row[b + slot..b + 2 * slot]);
337 }
338 assert!(prep.patches.iter().all(|v| (-1.0..=1.0).contains(v)));
340 }
341
342 fn crafted_gif(w: u16, h: u16, frames: usize) -> Vec<u8> {
347 let mut b = Vec::new();
348 b.extend_from_slice(b"GIF89a");
349 b.extend_from_slice(&w.to_le_bytes());
350 b.extend_from_slice(&h.to_le_bytes());
351 b.push(0x80); b.push(0); b.push(0); b.extend_from_slice(&[0, 0, 0, 0xFF, 0xFF, 0xFF]); for _ in 0..frames {
356 b.push(0x2C); b.extend_from_slice(&0u16.to_le_bytes()); b.extend_from_slice(&0u16.to_le_bytes()); b.extend_from_slice(&1u16.to_le_bytes()); b.extend_from_slice(&1u16.to_le_bytes()); b.push(0); b.push(0x02); b.extend_from_slice(&[0x02, 0x44, 0x01]); b.push(0x00); }
366 b.push(0x3B); b
368 }
369
370 #[test]
371 fn gif_decode_bomb_is_refused_before_full_expansion() {
372 fn expect_err(bytes: &[u8]) -> String {
376 match prep_video_gif(bytes) {
377 Err(e) => e,
378 Ok(_) => panic!("decode-bomb GIF was accepted"),
379 }
380 }
381 let bomb = crafted_gif(2000, 2000, 64);
382 assert!(bomb.len() < 2048, "the bomb itself is tiny on the wire");
383 let err = expect_err(&bomb);
384 assert!(err.contains("decode budget"), "{err}");
385
386 let ok = crafted_gif(2000, 2000, 4);
388 let vid = prep_video_gif(&ok).unwrap();
389 assert_eq!(vid.groups.len(), 2); let err = expect_err(&crafted_gif(8, 8, GIF_MAX_FRAMES + 8));
393 assert!(err.contains("decode budget"), "{err}");
394
395 let err = expect_err(&crafted_gif(0xFFFF, 0xFFFF, 1));
397 assert!(err.contains("exceeds the decode budget"), "{err}");
398 }
399
400 #[test]
401 fn data_uri_roundtrip() {
402 let png = {
403 let img = RgbImage::new(32, 32);
404 let mut buf = std::io::Cursor::new(Vec::new());
405 img.write_to(&mut buf, image::ImageFormat::Png).unwrap();
406 buf.into_inner()
407 };
408 let uri = format!(
409 "data:image/png;base64,{}",
410 base64::engine::general_purpose::STANDARD.encode(&png)
411 );
412 let prep = prep_data_uri(&uri).unwrap();
413 assert_eq!(prep.n_tokens(), prep.gh * prep.gw / 4);
414 assert!(decode_data_uri("http://x/y.png").is_err());
415 }
416}