use crc32fast::Hasher;
use regex::Regex;
use reqwest::blocking::Client;
use std::collections::HashMap;
use std::fs;
use std::path::{Path, PathBuf};
use std::time::Duration;
const MAX_IMAGE_SIZE: usize = 50 * 1024 * 1024;
pub struct ImageDownloader {
client: Client,
cache: HashMap<String, String>,
images_dir: PathBuf,
}
impl ImageDownloader {
pub fn new(output_dir: &Path) -> Self {
let client = Client::builder()
.timeout(Duration::from_secs(30))
.build()
.unwrap_or_else(|_| Client::new());
let images_dir = output_dir.join("_remote_images");
ImageDownloader {
client,
cache: HashMap::new(),
images_dir,
}
}
pub fn process_html(&mut self, html: &str) -> Result<String, Box<dyn std::error::Error>> {
let img_re = Regex::new(r#"<img\s+([^>]*?)src\s*=\s*["']((https?://[^"']+))["']([^>]*)>"#)?;
let mut result = html.to_string();
let mut replacements: Vec<(String, String)> = Vec::new();
for caps in img_re.captures_iter(html) {
let full_match = caps.get(0).unwrap().as_str();
let before_src = caps.get(1).map(|m| m.as_str()).unwrap_or("");
let url = caps.get(2).unwrap().as_str();
let after_src = caps.get(4).map(|m| m.as_str()).unwrap_or("");
if !url.starts_with("https://") && !url.starts_with("http://") {
continue;
}
match self.download_image(url) {
Ok(local_path) => {
let new_tag =
format!(r#"<img {}src="{}"{}"#, before_src, local_path, after_src);
let new_tag = if full_match.ends_with("/>") {
format!("{}/>", new_tag)
} else {
format!("{}>", new_tag)
};
replacements.push((full_match.to_string(), new_tag));
}
Err(e) => {
eprintln!(" Warning: Failed to download image {}: {}", url, e);
}
}
}
for (old, new) in replacements {
result = result.replace(&old, &new);
}
Ok(result)
}
fn download_image(&mut self, url: &str) -> Result<String, Box<dyn std::error::Error>> {
use std::io::Read;
if let Some(cached_path) = self.cache.get(url) {
return Ok(cached_path.clone());
}
fs::create_dir_all(&self.images_dir)?;
let response = self.client.get(url).send()?;
if !response.status().is_success() {
return Err(format!("HTTP {}", response.status()).into());
}
if let Some(content_length) = response.content_length() {
if content_length as usize > MAX_IMAGE_SIZE {
return Err(format!(
"Image too large (Content-Length: {:.1} MB, max {} MB)",
content_length as f64 / 1024.0 / 1024.0,
MAX_IMAGE_SIZE / 1024 / 1024
)
.into());
}
}
let mut bytes = Vec::new();
let mut total_read: usize = 0;
let mut buf = [0u8; 8192];
let mut reader = response;
loop {
let n = reader.read(&mut buf)?;
if n == 0 {
break;
}
total_read += n;
if total_read > MAX_IMAGE_SIZE {
return Err(format!(
"Image too large (>{} MB, download aborted mid-stream)",
MAX_IMAGE_SIZE / 1024 / 1024
)
.into());
}
bytes.extend_from_slice(&buf[..n]);
}
let hash = crc32_hash(url);
let ext = detect_extension(url, &bytes);
let filename = format!("{:08x}.{}", hash, ext);
let file_path = self.images_dir.join(&filename);
fs::write(&file_path, &bytes)?;
let relative_path = format!("_remote_images/{}", filename);
self.cache.insert(url.to_string(), relative_path.clone());
Ok(relative_path)
}
pub fn stats(&self) -> (usize, usize) {
(self.cache.len(), 0) }
}
fn crc32_hash(s: &str) -> u32 {
let mut hasher = Hasher::new();
hasher.update(s.as_bytes());
hasher.finalize()
}
fn detect_extension(url: &str, bytes: &[u8]) -> &'static str {
if bytes.len() >= 8 {
if bytes.starts_with(&[0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A]) {
return "png";
}
if bytes.starts_with(&[0xFF, 0xD8, 0xFF]) {
return "jpg";
}
if bytes.starts_with(b"GIF87a") || bytes.starts_with(b"GIF89a") {
return "gif";
}
if bytes.len() >= 12 && bytes.starts_with(b"RIFF") && &bytes[8..12] == b"WEBP" {
return "webp";
}
let start = String::from_utf8_lossy(&bytes[..bytes.len().min(100)]);
if start.trim_start().starts_with("<?xml") || start.trim_start().starts_with("<svg") {
return "svg";
}
if bytes.starts_with(&[0x00, 0x00, 0x01, 0x00]) {
return "ico";
}
if bytes.starts_with(b"BM") {
return "bmp";
}
}
let url_lower = url.to_lowercase();
if let Some(ext_start) = url_lower.rfind('.') {
let ext = &url_lower[ext_start + 1..];
let ext = ext.split('?').next().unwrap_or(ext);
let ext = ext.split('#').next().unwrap_or(ext);
match ext {
"png" => return "png",
"jpg" | "jpeg" => return "jpg",
"gif" => return "gif",
"webp" => return "webp",
"svg" => return "svg",
"ico" => return "ico",
"bmp" => return "bmp",
"avif" => return "avif",
_ => {}
}
}
"png"
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_crc32_hash() {
let hash1 = crc32_hash("https://example.com/image.png");
let hash2 = crc32_hash("https://example.com/image.png");
let hash3 = crc32_hash("https://example.com/other.png");
assert_eq!(hash1, hash2, "Same input should produce same hash");
assert_ne!(
hash1, hash3,
"Different input should produce different hash"
);
}
#[test]
fn test_detect_extension_from_magic_bytes() {
let png_bytes = [0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A, 0x00, 0x00];
assert_eq!(
detect_extension("http://example.com/image", &png_bytes),
"png"
);
let jpg_bytes = [0xFF, 0xD8, 0xFF, 0xE0, 0x00, 0x10, 0x4A, 0x46];
assert_eq!(
detect_extension("http://example.com/image", &jpg_bytes),
"jpg"
);
let gif_bytes = b"GIF89a\x00\x00";
assert_eq!(
detect_extension("http://example.com/image", gif_bytes),
"gif"
);
}
#[test]
fn test_detect_extension_from_url() {
let empty: &[u8] = &[];
assert_eq!(
detect_extension("https://example.com/image.png", empty),
"png"
);
assert_eq!(
detect_extension("https://example.com/image.jpg", empty),
"jpg"
);
assert_eq!(
detect_extension("https://example.com/image.jpeg", empty),
"jpg"
);
assert_eq!(
detect_extension("https://example.com/image.gif", empty),
"gif"
);
assert_eq!(
detect_extension("https://example.com/image.webp", empty),
"webp"
);
assert_eq!(
detect_extension("https://example.com/image.png?v=123", empty),
"png"
);
}
#[test]
fn test_max_image_size_is_50mb() {
assert_eq!(MAX_IMAGE_SIZE, 50 * 1024 * 1024);
}
#[test]
fn test_image_downloader_creation() {
let temp = tempfile::tempdir().unwrap();
let downloader = ImageDownloader::new(temp.path());
assert_eq!(downloader.stats(), (0, 0));
assert_eq!(downloader.images_dir, temp.path().join("_remote_images"));
}
#[test]
fn test_detect_extension_default() {
let empty: &[u8] = &[];
assert_eq!(detect_extension("https://example.com/image", empty), "png");
assert_eq!(
detect_extension("https://example.com/image.xyz", empty),
"png"
);
}
#[test]
fn test_download_rejects_oversized_content_length() {
use std::io::Write;
let listener =
std::net::TcpListener::bind("127.0.0.1:0").expect("Failed to bind test server");
let port = listener.local_addr().unwrap().port();
let url = format!("http://127.0.0.1:{}/large.png", port);
let handle = std::thread::spawn(move || {
if let Ok((mut stream, _)) = listener.accept() {
let mut buf = [0u8; 1024];
let _ = std::io::Read::read(&mut stream, &mut buf);
let fake_len = MAX_IMAGE_SIZE + 1;
let response = format!(
"HTTP/1.1 200 OK\r\nContent-Length: {}\r\nContent-Type: image/png\r\n\r\nfake",
fake_len
);
let _ = stream.write_all(response.as_bytes());
let _ = stream.flush();
}
});
let temp = tempfile::tempdir().unwrap();
let mut downloader = ImageDownloader::new(temp.path());
let result = downloader.download_image(&url);
assert!(
result.is_err(),
"Should reject image exceeding MAX_IMAGE_SIZE"
);
let err = result.unwrap_err().to_string();
assert!(
err.contains("too large"),
"Error should mention 'too large', got: {}",
err
);
handle.join().ok();
}
#[test]
fn test_download_aborts_oversized_stream() {
use std::sync::Arc;
let server = Arc::new(
tiny_http::Server::http("127.0.0.1:0").expect("Failed to start test HTTP server"),
);
let addr = server.server_addr().to_ip().unwrap();
let url = format!("http://127.0.0.1:{}/stream.bin", addr.port());
let server_clone = Arc::clone(&server);
let handle = std::thread::spawn(move || {
if let Ok(request) = server_clone.recv() {
let body = vec![0u8; MAX_IMAGE_SIZE + 8192];
let response = tiny_http::Response::from_data(body).with_status_code(200);
let _ = request.respond(response);
}
});
let temp = tempfile::tempdir().unwrap();
let mut downloader = ImageDownloader::new(temp.path());
let result = downloader.download_image(&url);
assert!(
result.is_err(),
"Should abort download when stream exceeds MAX_IMAGE_SIZE"
);
let err = result.unwrap_err().to_string();
assert!(
err.contains("too large") || err.contains("aborted"),
"Error should mention 'too large' or 'aborted', got: {}",
err
);
handle.join().ok();
}
#[test]
fn test_download_follows_redirect() {
use std::sync::Arc;
let server2 =
Arc::new(tiny_http::Server::http("127.0.0.1:0").expect("Failed to start server2"));
let addr2 = server2.server_addr().to_ip().unwrap();
let server2_clone = Arc::clone(&server2);
let handle2 = std::thread::spawn(move || {
if let Ok(request) = server2_clone.recv() {
let png: &[u8] = &[
0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A, 0x00, 0x00, 0x00, 0x0D, 0x49, 0x48, 0x44, 0x52, ];
let response = tiny_http::Response::from_data(png.to_vec()).with_status_code(200);
let _ = request.respond(response);
}
});
let server1 =
Arc::new(tiny_http::Server::http("127.0.0.1:0").expect("Failed to start server1"));
let addr1 = server1.server_addr().to_ip().unwrap();
let redirect_target = format!("http://127.0.0.1:{}/image.png", addr2.port());
let server1_clone = Arc::clone(&server1);
let handle1 = std::thread::spawn(move || {
if let Ok(request) = server1_clone.recv() {
let header =
tiny_http::Header::from_bytes(b"Location" as &[u8], redirect_target.as_bytes())
.unwrap();
let response = tiny_http::Response::from_string("Moved")
.with_status_code(301)
.with_header(header);
let _ = request.respond(response);
}
});
let temp = tempfile::tempdir().unwrap();
let mut downloader = ImageDownloader::new(temp.path());
let url = format!("http://127.0.0.1:{}/old.png", addr1.port());
let result = downloader.download_image(&url);
assert!(
result.is_ok(),
"Redirect should be followed, got: {:?}",
result.err()
);
handle1.join().ok();
handle2.join().ok();
}
}