use anyhow::{anyhow, Result};
use std::io::Read as _;
use std::path::{Path, PathBuf};
use std::time::Duration;
pub mod image_token_residual_add;
pub mod mmproj;
pub mod mmproj_weights;
pub mod pipeline;
pub mod preprocess;
pub mod vit;
pub mod vit_dump;
pub mod vit_gpu;
pub mod vit_gpu_qwen3vl;
#[allow(unused_imports)]
pub use preprocess::{preprocess_rgb_chw, PreprocessConfig, GEMMA4_VISION_CONFIG};
#[derive(Debug, Clone)]
pub struct PreprocessedImage {
pub pixel_values: Vec<f32>,
pub target_size: u32,
pub pixel_w: Option<u32>,
pub pixel_h: Option<u32>,
pub source_label: String,
}
impl PreprocessedImage {
pub fn pixel_grid(&self) -> (u32, u32) {
let w = self.pixel_w.unwrap_or(self.target_size);
let h = self.pixel_h.unwrap_or(self.target_size);
(w, h)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ImageInput {
DataUri {
mime_type: String,
payload_base64: String,
},
FilePath(PathBuf),
HttpUrl(String),
}
pub fn parse_image_url(url: &str) -> Result<ImageInput> {
if let Some(rest) = url.strip_prefix("data:") {
let (meta, payload) = rest
.split_once(",")
.ok_or_else(|| anyhow!("data URI missing comma separator"))?;
let (mime_type, encoding) = meta
.split_once(";")
.ok_or_else(|| anyhow!("data URI missing encoding section"))?;
if encoding != "base64" {
return Err(anyhow!(
"data URI encoding '{}' not supported (only 'base64')",
encoding
));
}
if !(mime_type == "image/png" || mime_type == "image/jpeg" || mime_type == "image/jpg") {
return Err(anyhow!(
"data URI mime type '{}' not supported (only image/png and image/jpeg)",
mime_type
));
}
return Ok(ImageInput::DataUri {
mime_type: mime_type.to_string(),
payload_base64: payload.to_string(),
});
}
if let Some(rest) = url.strip_prefix("file://") {
let path = if rest.starts_with('/') {
PathBuf::from(rest)
} else {
return Err(anyhow!(
"file:// URL must contain an absolute path (file:///path)"
));
};
return Ok(ImageInput::FilePath(path));
}
if url.starts_with('/') {
return Ok(ImageInput::FilePath(PathBuf::from(url)));
}
if url.starts_with("http://") || url.starts_with("https://") {
return Ok(ImageInput::HttpUrl(url.to_string()));
}
Err(anyhow!(
"unrecognized image URL scheme (expected data:, file://, or absolute path)"
))
}
pub fn load_image_bytes(input: &ImageInput) -> Result<Vec<u8>> {
match input {
ImageInput::DataUri { payload_base64, .. } => {
use base64::Engine;
let payload = base64::engine::general_purpose::STANDARD
.decode(payload_base64.trim())
.map_err(|e| anyhow!("base64 decode: {e}"))?;
Ok(payload)
}
ImageInput::FilePath(p) => read_file_bounded(p),
ImageInput::HttpUrl(url) => fetch_https_image(url),
}
}
fn read_file_bounded(p: &Path) -> Result<Vec<u8>> {
const MAX: u64 = 20 * 1024 * 1024;
let meta = std::fs::metadata(p).map_err(|e| anyhow!("stat {}: {e}", p.display()))?;
if meta.len() > MAX {
return Err(anyhow!(
"image file {} exceeds {}-byte cap (got {})",
p.display(),
MAX,
meta.len()
));
}
std::fs::read(p).map_err(|e| anyhow!("read {}: {e}", p.display()))
}
fn fetch_https_image(url: &str) -> Result<Vec<u8>> {
const MAX_BYTES: u64 = 20 * 1024 * 1024;
let fetch = || -> Result<Vec<u8>> {
let client = reqwest::blocking::Client::builder()
.timeout(Duration::from_secs(10))
.build()
.map_err(|e| anyhow!("HTTPS fetch: failed to build client: {e}"))?;
let resp = client.get(url).send().map_err(|e| {
if e.is_timeout() {
anyhow!("HTTPS fetch timed out after 10 s ({})", url)
} else {
anyhow!("HTTPS fetch network error ({}): {e}", url)
}
})?;
let status = resp.status();
if !status.is_success() {
return Err(anyhow!(
"HTTPS fetch received HTTP {} from {}",
status.as_u16(),
url
));
}
if let Some(len) = resp.content_length() {
if len > MAX_BYTES {
return Err(anyhow!(
"HTTPS fetch: Content-Length {} exceeds {}-byte cap ({})",
len,
MAX_BYTES,
url
));
}
}
let cap = MAX_BYTES + 1;
let hint = resp
.content_length()
.map(|cl| (cl as usize).min(cap as usize))
.unwrap_or(0);
let mut buf: Vec<u8> = Vec::with_capacity(hint);
resp.take(cap).read_to_end(&mut buf).map_err(|e| {
if e.kind() == std::io::ErrorKind::TimedOut {
anyhow!("HTTPS fetch body read timed out ({})", url)
} else {
anyhow!("HTTPS fetch body read error ({}): {e}", url)
}
})?;
if buf.len() as u64 >= cap {
return Err(anyhow!(
"HTTPS fetch: response body exceeds {}-byte cap ({})",
MAX_BYTES,
url
));
}
Ok(buf)
};
match tokio::runtime::Handle::try_current() {
Ok(_) => tokio::task::block_in_place(fetch),
Err(_) => fetch(),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parse_image_url_data_png() {
let got = parse_image_url("data:image/png;base64,iVBORw0K").unwrap();
match got {
ImageInput::DataUri {
mime_type,
payload_base64,
} => {
assert_eq!(mime_type, "image/png");
assert_eq!(payload_base64, "iVBORw0K");
}
other => panic!("expected DataUri, got {:?}", other),
}
}
#[test]
fn parse_image_url_data_jpeg() {
let got = parse_image_url("data:image/jpeg;base64,/9j/4AA").unwrap();
assert!(matches!(got, ImageInput::DataUri { .. }));
}
#[test]
fn parse_image_url_rejects_unsupported_mime() {
let err = parse_image_url("data:image/gif;base64,xyz").unwrap_err();
assert!(format!("{err}").contains("not supported"));
}
#[test]
fn parse_image_url_rejects_non_base64_encoding() {
let err = parse_image_url("data:image/png;utf8,hello").unwrap_err();
assert!(format!("{err}").contains("not supported"));
}
#[test]
fn parse_image_url_file_scheme() {
let got = parse_image_url("file:///tmp/cat.jpg").unwrap();
assert_eq!(got, ImageInput::FilePath(PathBuf::from("/tmp/cat.jpg")));
}
#[test]
fn parse_image_url_bare_absolute_path() {
let got = parse_image_url("/tmp/dog.png").unwrap();
assert_eq!(got, ImageInput::FilePath(PathBuf::from("/tmp/dog.png")));
}
#[test]
fn parse_image_url_http_preserved_for_fetch() {
let got = parse_image_url("https://example.com/img.jpg").unwrap();
assert!(matches!(got, ImageInput::HttpUrl(_)));
}
#[test]
fn parse_image_url_rejects_gibberish() {
let err = parse_image_url("not-a-url").unwrap_err();
assert!(format!("{err}").contains("unrecognized"));
}
#[test]
fn load_image_bytes_data_uri_round_trips_base64() {
let sig = [0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A];
use base64::Engine;
let b64 = base64::engine::general_purpose::STANDARD.encode(sig);
let url = format!("data:image/png;base64,{b64}");
let input = parse_image_url(&url).unwrap();
let bytes = load_image_bytes(&input).unwrap();
assert_eq!(bytes, sig);
}
#[test]
fn load_image_bytes_https_url_attempts_fetch() {
let input = ImageInput::HttpUrl("https://example.invalid/cat.jpg".into());
let err = load_image_bytes(&input).unwrap_err();
let msg = format!("{err}");
assert!(
!msg.contains("not yet loaded"),
"expected network error, got static rejection: {msg}"
);
assert!(
msg.contains("example.invalid")
|| msg.contains("network error")
|| msg.contains("timed out")
|| msg.contains("HTTPS fetch"),
"unexpected error message: {msg}"
);
}
#[test]
fn load_image_bytes_rejects_oversized_file() {
let err = load_image_bytes(&ImageInput::FilePath(PathBuf::from(
"/tmp/does-not-exist-xyz-42.png",
)))
.unwrap_err();
let msg = format!("{err}");
assert!(
msg.contains("stat") || msg.contains("No such"),
"unexpected error: {msg}"
);
}
#[test]
fn fetch_https_image_cap_fires_without_content_length() {
use std::io::Write as _;
use std::net::TcpListener;
use std::sync::{
atomic::{AtomicU64, Ordering},
Arc,
};
use std::thread;
const MAX_BYTES: u64 = 20 * 1024 * 1024;
let listener = TcpListener::bind("127.0.0.1:0").expect("bind ephemeral port");
let addr = listener.local_addr().expect("local_addr");
let bytes_accepted = Arc::new(AtomicU64::new(0));
let bytes_accepted_srv = Arc::clone(&bytes_accepted);
let srv = thread::spawn(move || {
let (mut stream, _peer) = listener.accept().expect("accept");
drain_http_request(&mut stream);
let header = b"HTTP/1.0 200 OK\r\nContent-Type: application/octet-stream\r\n\r\n";
if stream.write_all(header).is_err() {
return;
}
let chunk = vec![0u8; 64 * 1024];
loop {
match stream.write_all(&chunk) {
Ok(_) => {
bytes_accepted_srv.fetch_add(chunk.len() as u64, Ordering::Relaxed);
if bytes_accepted_srv.load(Ordering::Relaxed)
> MAX_BYTES + chunk.len() as u64
{
break;
}
}
Err(_) => break, }
}
});
let url = format!("http://127.0.0.1:{}/oversized.bin", addr.port());
let err = fetch_https_image(&url)
.expect_err("expected fetch_https_image to return an error for oversized body");
let msg = format!("{err}");
assert!(
msg.contains("cap") || msg.contains("exceed"),
"error message should mention the cap; got: {msg}"
);
let _ = srv.join();
let written = bytes_accepted.load(Ordering::Relaxed);
let slop = 2 * 1024 * 1024u64; assert!(
written <= MAX_BYTES + slop,
"server wrote {written} bytes before client closed — \
cap is {} bytes; too much slop suggests client buffered full body",
MAX_BYTES
);
}
fn drain_http_request(stream: &mut std::net::TcpStream) {
use std::io::BufRead as _;
let mut reader = std::io::BufReader::new(stream.try_clone().expect("clone"));
let mut line = String::new();
loop {
line.clear();
match reader.read_line(&mut line) {
Ok(0) | Err(_) => break,
Ok(_) => {
if line == "\r\n" || line == "\n" {
break;
}
}
}
}
}
}