use self::image_data::{ImageData, ImageDataBase64};
use self::pair::Pair;
use self::request::{ImageRequest, Progress};
use crate::download_progress::DownloadProgress;
use crate::error::ImageDownloadError;
use crate::util::Util;
use crate::{constants, dom};
use base64::Engine;
use dom_query::{Document, NodeRef};
use futures::channel::mpsc::{self, Sender};
use futures::{SinkExt, StreamExt};
use image::codecs::jpeg::JpegEncoder;
use image::{ImageFormat, ImageReader};
use reqwest::Client;
use std::io::Cursor;
mod image_data;
mod pair;
mod request;
pub struct ImageDownloader {
max_size: (u32, u32),
}
impl ImageDownloader {
pub fn new(max_size: (u32, u32)) -> Self {
ImageDownloader { max_size }
}
pub async fn single_from_url(
url: &str,
client: &Client,
mut progress: Option<Sender<DownloadProgress>>,
) -> Result<Vec<u8>, ImageDownloadError> {
let response = client.get(url).send().await?;
let content_type = Util::get_content_type(&response);
let content_length = Util::get_content_length(&response);
if content_type.is_err()
&& let Ok(content_length) = content_length
&& content_length > constants::UNKNOWN_CONTENT_SIZE_LIMIT
{
return Err(ImageDownloadError::ContentType);
}
let content_type = content_type?;
let content_length = content_length.unwrap_or(0);
if !content_type.contains("image") {
return Err(ImageDownloadError::ContentType);
}
if content_length > constants::MAX_IMAGE_SIZE {
tracing::warn!(%url, content_length, "Image is too large");
return Err(ImageDownloadError::TooLarge);
}
let mut stream = response.bytes_stream();
let mut downloaded_bytes = 0;
let mut result = Vec::with_capacity(content_length);
while let Some(item) = stream.next().await {
let chunk = item?;
downloaded_bytes += chunk.len();
if downloaded_bytes > constants::MAX_IMAGE_SIZE {
tracing::warn!(%url, downloaded_bytes, "Image is too large");
return Err(ImageDownloadError::TooLarge);
}
if let Some(sender) = progress.as_mut() {
_ = sender
.send(DownloadProgress {
total_size: content_length,
downloaded: downloaded_bytes,
})
.await;
}
result.extend_from_slice(&chunk);
}
Ok(result)
}
pub async fn download_images_from_string(
&self,
html: &str,
client: &Client,
mut progress: Option<Sender<DownloadProgress>>,
) -> Result<String, ImageDownloadError> {
let image_urls = Self::harvest_image_urls_from_html(html)?;
let (tx, mut rx) = mpsc::channel::<Progress>(2);
tokio::spawn(async move {
let mut total_size = 0_usize;
let mut downloaded = 0_usize;
while let Some(update) = rx.next().await {
match update {
Progress::Expected(size) => total_size += size,
Progress::Downloaded(size) => downloaded += size,
}
if let Some(progress) = progress.as_mut() {
_ = progress
.send(DownloadProgress {
total_size,
downloaded,
})
.await;
}
}
});
let downloaded_images = futures::stream::iter(image_urls)
.map(|image_url| self.download_image(image_url, client, tx.clone()))
.buffer_unordered(constants::MAX_PARALLEL_IMAGE_DOWNLOADS)
.filter_map(|result| async move { result.ok() })
.collect::<Vec<_>>()
.await;
Self::replace_downloaded_images(html, downloaded_images)
}
fn replace_downloaded_images(
html: &str,
downloaded_images: Vec<Pair<ImageDataBase64>>,
) -> Result<String, ImageDownloadError> {
let doc = Document::from(html);
let img_nodes = dom::select(&doc, "img");
for downloaded_image_pair in downloaded_images {
let url = &downloaded_image_pair.value.url;
let node = img_nodes
.iter()
.find(|node| dom::attr(node, "src").as_deref() == Some(url.as_str()));
if let Some(node) = node {
node.set_attr("src", &downloaded_image_pair.value.data);
if let Some(parent_data) = downloaded_image_pair.parent_value {
node.set_attr("big-src", &parent_data.data)
}
}
}
Ok(Self::serialize_fragment(&doc))
}
fn serialize_fragment(doc: &Document) -> String {
match doc.body() {
Some(body) => body.inner_html().to_string(),
None => doc.html().to_string(),
}
}
fn harvest_image_urls_from_html(html: &str) -> Result<Vec<Pair<String>>, ImageDownloadError> {
let doc = Document::from(html);
let mut image_urls = Vec::new();
for node in dom::select(&doc, "img") {
if let Ok(url) = Self::harvest_image_urls_from_node(&node) {
image_urls.push(url);
}
}
Ok(image_urls)
}
fn harvest_image_urls_from_node(node: &NodeRef) -> Result<Pair<String>, ImageDownloadError> {
let src = match dom::attr(node, "src") {
Some(src) => {
if src.starts_with("data:") {
return Err(ImageDownloadError::Unknown);
} else {
src
}
}
None => {
return Err(ImageDownloadError::Unknown);
}
};
let parent_url = Self::check_image_parent(node).ok();
let image_url = Pair {
value: src,
parent_value: parent_url,
};
Ok(image_url)
}
async fn download_image(
&self,
image_url: Pair<String>,
client: &Client,
mut tx: Sender<Progress>,
) -> Result<Pair<ImageDataBase64>, ImageDownloadError> {
let mut image = ImageRequest::new(image_url.value, client)
.await?
.download(&mut tx)
.await?;
let mut parent_image = None;
if let Some(parent_url) = image_url.parent_value
&& let Ok(parent_request) = ImageRequest::new(parent_url, client).await
{
parent_image = parent_request.download(&mut tx).await.ok();
}
if image.content_type != "image/svg+xml" && image.content_type != "image/gif" {
let max_size = self.max_size;
let data = std::mem::take(&mut image.data);
let (data, resized) = tokio::task::spawn_blocking(move || {
let resized = Self::scale_image(&data, max_size);
(data, resized)
})
.await
.map_err(|_| ImageDownloadError::ImageScale)?;
image.data = data;
if let Some((resized_data, content_type)) = resized {
let resized_image = ImageData {
url: image.url.clone(),
data: resized_data,
content_type: content_type.into(),
};
if parent_image.is_none() {
parent_image = Some(image);
}
image = resized_image;
}
}
Ok(Pair {
value: Self::to_data_url(image),
parent_value: parent_image.map(Self::to_data_url),
})
}
fn to_data_url(image: ImageData) -> ImageDataBase64 {
let base64 = base64::engine::general_purpose::STANDARD.encode(&image.data);
ImageDataBase64 {
url: image.url,
data: format!("data:{};base64,{base64}", image.content_type),
}
}
fn scale_image(
image_buffer: &[u8],
max_dimensions: (u32, u32),
) -> Option<(Vec<u8>, &'static str)> {
let reader = ImageReader::new(Cursor::new(image_buffer))
.with_guessed_format()
.ok()?;
let source_format = reader.format();
let image = match reader.decode() {
Err(error) => {
tracing::error!(%error, "Failed to open image to resize");
return None;
}
Ok(image) => image,
};
if image.width() <= max_dimensions.0 && image.height() <= max_dimensions.1 {
return None;
}
let image = image.resize(
max_dimensions.0,
max_dimensions.1,
image::imageops::FilterType::Lanczos3,
);
let format = match source_format {
Some(ImageFormat::Jpeg) => ImageFormat::Jpeg,
Some(ImageFormat::WebP) if image.color().has_alpha() => ImageFormat::WebP,
Some(ImageFormat::WebP) => ImageFormat::Jpeg,
_ => ImageFormat::Png,
};
let mut resized_buf: Vec<u8> = Vec::new();
let result = if format == ImageFormat::Jpeg {
image.write_with_encoder(JpegEncoder::new_with_quality(
&mut resized_buf,
constants::RESIZED_JPEG_QUALITY,
))
} else {
image.write_to(&mut Cursor::new(&mut resized_buf), format)
};
if let Err(error) = result {
tracing::error!(%error, "Failed to save resized image");
return None;
}
Some((resized_buf, format.to_mime_type()))
}
fn check_image_parent(node: &NodeRef) -> Result<String, ImageDownloadError> {
let parent = match node.parent() {
Some(parent) => parent,
None => {
tracing::debug!("No parent node");
return Err(ImageDownloadError::ParentDownload);
}
};
if !dom::tag_name_is(&parent, "a") {
tracing::debug!("parent is not an <a> node");
return Err(ImageDownloadError::ParentDownload);
}
let href = match dom::attr(&parent, "href") {
Some(href) => href,
None => {
tracing::debug!("Parent doesn't have href prop");
return Err(ImageDownloadError::ParentDownload);
}
};
Ok(href)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_server;
use crate::test_util::assert_html_eq;
use reqwest::Client;
use std::fs;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
#[test]
fn replace_downloaded_images_with_quote_in_url() {
let html = r#"<p><img src="https://example.com/it's.jpg"><img src="https://example.com/b.jpg"><img src="https://example.com/b.jpg"></p>"#;
let image = |url: &str, data: &str| Pair {
value: ImageDataBase64 {
url: url.into(),
data: data.into(),
},
parent_value: None,
};
let downloaded = vec![
image(
"https://example.com/it's.jpg",
"data:image/jpeg;base64,QQ==",
),
image("https://example.com/b.jpg", "data:image/jpeg;base64,Qg=="),
image("https://example.com/b.jpg", "data:image/jpeg;base64,Qw=="),
];
let result = ImageDownloader::replace_downloaded_images(html, downloaded).unwrap();
let document = dom_query::Document::from(result.as_str());
let sources = document
.select("img")
.iter()
.map(|img| img.attr("src").unwrap_or_default().to_string())
.collect::<Vec<_>>();
assert_eq!(
sources,
[
"data:image/jpeg;base64,QQ==",
"data:image/jpeg;base64,Qg==",
"data:image/jpeg;base64,Qw=="
]
);
}
#[tokio::test]
async fn oversized_images_are_rejected() {
let base = test_server::serve(3, |request_line, stream| {
if request_line.contains("/huge-header ") {
test_server::write_huge_content_length(stream, "image/png");
} else {
test_server::write_endless_body(stream, "image/png");
}
});
let client = Client::new();
for path in ["/huge-header", "/endless"] {
let url = base.join(path).unwrap();
let result = ImageDownloader::single_from_url(url.as_str(), &client, None).await;
assert!(
matches!(result, Err(ImageDownloadError::TooLarge)),
"{path}: {result:?}"
);
}
let url = base.join("/huge-header").unwrap();
let result = ImageRequest::new(url.into(), &client).await;
assert!(
matches!(result, Err(ImageDownloadError::TooLarge)),
"{result:?}"
);
}
fn encode(image: &image::DynamicImage, format: ImageFormat) -> Vec<u8> {
let mut buf = Vec::new();
image.write_to(&mut Cursor::new(&mut buf), format).unwrap();
buf
}
fn data_url(image: &[u8], content_type: &str) -> String {
let base64 = base64::engine::general_purpose::STANDARD.encode(image);
format!("data:{content_type};base64,{base64}")
}
#[test]
fn scale_image_output_format() {
let rgb = image::DynamicImage::ImageRgb8(image::RgbImage::from_fn(300, 200, |x, y| {
image::Rgb([x as u8, y as u8, 128])
}));
let rgba = image::DynamicImage::ImageRgba8(rgb.to_rgba8());
for (image, source, expected) in [
(&rgb, ImageFormat::Jpeg, "image/jpeg"),
(&rgb, ImageFormat::WebP, "image/jpeg"),
(&rgba, ImageFormat::WebP, "image/webp"),
(&rgba, ImageFormat::Png, "image/png"),
(&rgb, ImageFormat::Bmp, "image/png"),
] {
let (data, content_type) =
ImageDownloader::scale_image(&encode(image, source), (150, 150)).unwrap();
assert_eq!(content_type, expected, "{source:?}");
let scaled = image::load_from_memory(&data).unwrap();
assert_eq!((scaled.width(), scaled.height()), (150, 100), "{source:?}");
}
assert!(
ImageDownloader::scale_image(&encode(&rgb, ImageFormat::Jpeg), (300, 200)).is_none()
);
}
#[tokio::test]
async fn download_images_from_string_limits_connections() {
const IMAGES: usize = 20;
let big_jpeg = encode(
&image::DynamicImage::ImageRgb8(image::RgbImage::new(100, 100)),
ImageFormat::Jpeg,
);
let open = Arc::new(AtomicUsize::new(0));
let most_open = Arc::new(AtomicUsize::new(0));
let base = {
let (open, most_open, big_jpeg) = (open.clone(), most_open.clone(), big_jpeg.clone());
test_server::serve_parallel(IMAGES + 2, move |request_line, stream| {
let now_open = open.fetch_add(1, Ordering::SeqCst) + 1;
most_open.fetch_max(now_open, Ordering::SeqCst);
std::thread::sleep(Duration::from_millis(50));
if request_line.contains("/chunked ") {
test_server::write_chunked_response(stream, "image/png", b"chunked");
} else if request_line.contains("/big ") {
test_server::write_bytes_response(stream, "image/jpeg", &big_jpeg);
} else {
test_server::write_bytes_response(stream, "image/png", b"small");
}
open.fetch_sub(1, Ordering::SeqCst);
})
};
let mut html = format!(r#"<p><img src="{base}chunked"><img src="{base}big">"#);
for i in 0..IMAGES {
html.push_str(&format!(r#"<img src="{base}{i}">"#));
}
html.push_str("</p>");
let (tx, mut rx) = mpsc::channel(1);
let progress = tokio::spawn(async move {
let mut last = None;
while let Some(progress) = rx.next().await {
last = Some(progress);
}
last
});
let result = ImageDownloader::new((50, 50))
.download_images_from_string(&html, &Client::new(), Some(tx))
.await
.unwrap();
let most_open = most_open.load(Ordering::SeqCst);
assert!(
most_open <= constants::MAX_PARALLEL_IMAGE_DOWNLOADS,
"{most_open} connections at once"
);
let document = Document::from(result.as_str());
let images = dom::select(&document, "img");
assert_eq!(images.len(), IMAGES + 2);
assert_eq!(
dom::attr(&images[0], "src").unwrap(),
data_url(b"chunked", "image/png")
);
let src = dom::attr(&images[1], "src").unwrap();
assert!(src.starts_with("data:image/jpeg;base64,"), "{src}");
assert_eq!(
dom::attr(&images[1], "big-src").unwrap(),
data_url(&big_jpeg, "image/jpeg")
);
for image in &images[2..] {
assert_eq!(
dom::attr(image, "src").unwrap(),
data_url(b"small", "image/png")
);
}
let progress = progress.await.unwrap().unwrap();
let expected_size = "chunked".len() + big_jpeg.len() + IMAGES * "small".len();
assert_eq!(progress.total_size, expected_size);
assert_eq!(progress.downloaded, expected_size);
}
#[tokio::test]
#[ignore = "downloads content from the web"]
async fn fedora31() {
let image_dowloader = ImageDownloader::new((2048, 2048));
let html = fs::read_to_string(r"./resources/tests/images/planet_gnome/source.html")
.expect("Failed to read HTML");
let result = image_dowloader
.download_images_from_string(&html, &Client::new(), None)
.await
.expect("Failed to downalod images");
let expected = fs::read_to_string(r"./resources/tests/images/planet_gnome/expected.html")
.expect("Failed to create output file");
assert_html_eq(&expected, &result);
}
}